feat(grpc): add generation request semantics (#32588)
Signed-off-by: Connor Carpenter <connorc@nvidia.com> Co-authored-by: ishandhanani <82981111+ishandhanani@users.noreply.github.com> Co-authored-by: Alex Nails <alex.nails@radixark.ai>
This commit is contained in:
co-authored by
ishandhanani
Alex Nails
parent
29831d58ef
commit
a0b04dbe4c
@@ -229,7 +229,7 @@ impl proto::sglang_service_server::SglangService for SglangServiceImpl {
|
||||
.rid
|
||||
.clone()
|
||||
.unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
|
||||
let req_dict = build_text_generate_dict(&rid, &req);
|
||||
let req_dict = build_text_generate_dict(&rid, &req).map_err(Status::invalid_argument)?;
|
||||
|
||||
let mut receiver = self
|
||||
.bridge
|
||||
@@ -298,7 +298,7 @@ impl proto::sglang_service_server::SglangService for SglangServiceImpl {
|
||||
.rid
|
||||
.clone()
|
||||
.unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
|
||||
let req_dict = build_generate_dict(&rid, &req);
|
||||
let req_dict = build_generate_dict(&rid, &req).map_err(Status::invalid_argument)?;
|
||||
|
||||
let mut receiver = self
|
||||
.bridge
|
||||
|
||||
@@ -2,8 +2,25 @@ use std::collections::HashMap;
|
||||
|
||||
use crate::proto;
|
||||
|
||||
fn regex_escape_literal(value: &str) -> String {
|
||||
let mut escaped = String::with_capacity(value.len());
|
||||
for character in value.chars() {
|
||||
if matches!(
|
||||
character,
|
||||
'.' | '+' | '*' | '?' | '^' | '$' | '(' | ')' | '[' | ']' | '{' | '}' | '|' | '\\'
|
||||
) {
|
||||
escaped.push('\\');
|
||||
}
|
||||
escaped.push(character);
|
||||
}
|
||||
escaped
|
||||
}
|
||||
|
||||
/// Convert proto SamplingParams to a serde_json map (used as Python dict via PyO3).
|
||||
fn sampling_params_to_map(params: &Option<proto::SamplingParams>) -> serde_json::Value {
|
||||
#[allow(deprecated)]
|
||||
fn sampling_params_to_map(
|
||||
params: &Option<proto::SamplingParams>,
|
||||
) -> Result<serde_json::Value, String> {
|
||||
match params {
|
||||
Some(p) => {
|
||||
let mut map = serde_json::Map::new();
|
||||
@@ -46,15 +63,90 @@ fn sampling_params_to_map(params: &Option<proto::SamplingParams>) -> serde_json:
|
||||
if let Some(v) = p.n {
|
||||
map.insert("n".into(), serde_json::json!(v));
|
||||
}
|
||||
if let Some(ref v) = p.json_schema {
|
||||
map.insert("json_schema".into(), serde_json::json!(v));
|
||||
if let Some(v) = p.seed {
|
||||
map.insert("sampling_seed".into(), serde_json::json!(v));
|
||||
}
|
||||
if let Some(ref v) = p.regex {
|
||||
map.insert("regex".into(), serde_json::json!(v));
|
||||
if p.guided_decoding.is_some() && (p.json_schema.is_some() || p.regex.is_some()) {
|
||||
return Err(
|
||||
"legacy json_schema/regex cannot be combined with guided_decoding".into(),
|
||||
);
|
||||
}
|
||||
serde_json::Value::Object(map)
|
||||
if let Some(guided) = p.guided_decoding.as_ref() {
|
||||
use proto::guided_decoding::Constraint;
|
||||
match guided.constraint.as_ref() {
|
||||
Some(Constraint::JsonSchema(value)) if !value.is_empty() => {
|
||||
map.insert("json_schema".into(), serde_json::json!(value));
|
||||
}
|
||||
Some(Constraint::Regex(value)) if !value.is_empty() => {
|
||||
map.insert("regex".into(), serde_json::json!(value));
|
||||
}
|
||||
Some(Constraint::Ebnf(value)) if !value.is_empty() => {
|
||||
map.insert("ebnf".into(), serde_json::json!(value));
|
||||
}
|
||||
Some(Constraint::Choice(choice))
|
||||
if !choice.values.is_empty()
|
||||
&& choice.values.iter().all(|value| !value.is_empty()) =>
|
||||
{
|
||||
let alternatives = choice
|
||||
.values
|
||||
.iter()
|
||||
.map(|value| regex_escape_literal(value))
|
||||
.collect::<Vec<_>>()
|
||||
.join("|");
|
||||
map.insert(
|
||||
"regex".into(),
|
||||
serde_json::json!(format!("(?:{alternatives})")),
|
||||
);
|
||||
}
|
||||
Some(Constraint::StructuralTag(value)) if !value.is_empty() => {
|
||||
map.insert("structural_tag".into(), serde_json::json!(value));
|
||||
}
|
||||
Some(Constraint::Choice(_)) => {
|
||||
return Err("guided choice must contain only non-empty values".into());
|
||||
}
|
||||
Some(_) => return Err("guided decoding constraint must not be empty".into()),
|
||||
None => return Err("guided decoding constraint must be specified".into()),
|
||||
}
|
||||
} else {
|
||||
if let Some(value) = p.json_schema.as_ref() {
|
||||
if value.is_empty() {
|
||||
return Err("legacy json_schema must not be empty".into());
|
||||
}
|
||||
map.insert("json_schema".into(), serde_json::json!(value));
|
||||
}
|
||||
if let Some(value) = p.regex.as_ref() {
|
||||
if value.is_empty() {
|
||||
return Err("legacy regex must not be empty".into());
|
||||
}
|
||||
map.insert("regex".into(), serde_json::json!(value));
|
||||
}
|
||||
}
|
||||
Ok(serde_json::Value::Object(map))
|
||||
}
|
||||
None => serde_json::Value::Object(serde_json::Map::new()),
|
||||
None => Ok(serde_json::Value::Object(serde_json::Map::new())),
|
||||
}
|
||||
}
|
||||
|
||||
fn insert_generation_controls(
|
||||
d: &mut HashMap<String, serde_json::Value>,
|
||||
priority: Option<i32>,
|
||||
require_reasoning: Option<bool>,
|
||||
max_thinking_tokens: Option<u32>,
|
||||
) {
|
||||
if let Some(priority) = priority {
|
||||
d.insert("priority".into(), serde_json::json!(priority));
|
||||
}
|
||||
if let Some(require_reasoning) = require_reasoning {
|
||||
d.insert(
|
||||
"require_reasoning".into(),
|
||||
serde_json::json!(require_reasoning),
|
||||
);
|
||||
}
|
||||
if let Some(max_thinking_tokens) = max_thinking_tokens {
|
||||
d.insert(
|
||||
"max_thinking_tokens".into(),
|
||||
serde_json::json!(max_thinking_tokens),
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -111,13 +203,13 @@ pub(crate) fn extract_model_path(json_info: &str) -> String {
|
||||
pub(crate) fn build_text_generate_dict(
|
||||
rid: &str,
|
||||
req: &proto::TextGenerateRequest,
|
||||
) -> HashMap<String, serde_json::Value> {
|
||||
) -> Result<HashMap<String, serde_json::Value>, String> {
|
||||
let mut d = HashMap::new();
|
||||
d.insert("rid".into(), serde_json::json!(rid));
|
||||
d.insert("text".into(), serde_json::json!(req.text));
|
||||
d.insert(
|
||||
"sampling_params".into(),
|
||||
sampling_params_to_map(&req.sampling_params),
|
||||
sampling_params_to_map(&req.sampling_params)?,
|
||||
);
|
||||
d.insert(
|
||||
"stream".into(),
|
||||
@@ -151,25 +243,31 @@ pub(crate) fn build_text_generate_dict(
|
||||
if let Some(ref session_id) = req.session_id {
|
||||
d.insert("session_id".into(), serde_json::json!(session_id));
|
||||
}
|
||||
insert_generation_controls(
|
||||
&mut d,
|
||||
req.priority,
|
||||
req.require_reasoning,
|
||||
req.max_thinking_tokens,
|
||||
);
|
||||
insert_disaggregated_params(&mut d, &req.disaggregated_params);
|
||||
if let Some(trace) = trace_headers_to_json(&req.trace_headers) {
|
||||
d.insert("external_trace_header".into(), trace);
|
||||
}
|
||||
d.insert("received_time".into(), serde_json::json!(now_timestamp()));
|
||||
d
|
||||
Ok(d)
|
||||
}
|
||||
|
||||
/// Build a request dict for GenerateReqInput from proto GenerateRequest (tokenized).
|
||||
pub(crate) fn build_generate_dict(
|
||||
rid: &str,
|
||||
req: &proto::GenerateRequest,
|
||||
) -> HashMap<String, serde_json::Value> {
|
||||
) -> Result<HashMap<String, serde_json::Value>, String> {
|
||||
let mut d = HashMap::new();
|
||||
d.insert("rid".into(), serde_json::json!(rid));
|
||||
d.insert("input_ids".into(), serde_json::json!(req.input_ids));
|
||||
d.insert(
|
||||
"sampling_params".into(),
|
||||
sampling_params_to_map(&req.sampling_params),
|
||||
sampling_params_to_map(&req.sampling_params)?,
|
||||
);
|
||||
d.insert(
|
||||
"stream".into(),
|
||||
@@ -199,12 +297,18 @@ pub(crate) fn build_generate_dict(
|
||||
if let Some(ref session_id) = req.session_id {
|
||||
d.insert("session_id".into(), serde_json::json!(session_id));
|
||||
}
|
||||
insert_generation_controls(
|
||||
&mut d,
|
||||
req.priority,
|
||||
req.require_reasoning,
|
||||
req.max_thinking_tokens,
|
||||
);
|
||||
insert_disaggregated_params(&mut d, &req.disaggregated_params);
|
||||
if let Some(trace) = trace_headers_to_json(&req.trace_headers) {
|
||||
d.insert("external_trace_header".into(), trace);
|
||||
}
|
||||
d.insert("received_time".into(), serde_json::json!(now_timestamp()));
|
||||
d
|
||||
Ok(d)
|
||||
}
|
||||
|
||||
/// Build a request dict for EmbeddingReqInput from proto TextEmbedRequest.
|
||||
@@ -267,6 +371,7 @@ pub(crate) fn build_classify_dict(
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
#[allow(deprecated)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
@@ -283,11 +388,15 @@ mod tests {
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
build_text_generate_dict("request-1", &text_req).get("session_id"),
|
||||
build_text_generate_dict("request-1", &text_req)
|
||||
.unwrap()
|
||||
.get("session_id"),
|
||||
Some(&serde_json::json!("session-1"))
|
||||
);
|
||||
assert_eq!(
|
||||
build_generate_dict("request-2", &token_req).get("session_id"),
|
||||
build_generate_dict("request-2", &token_req)
|
||||
.unwrap()
|
||||
.get("session_id"),
|
||||
Some(&serde_json::json!("session-1"))
|
||||
);
|
||||
}
|
||||
@@ -312,6 +421,7 @@ mod tests {
|
||||
build_text_generate_dict("request-1", &text_req),
|
||||
build_generate_dict("request-2", &token_req),
|
||||
] {
|
||||
let request = request.unwrap();
|
||||
assert_eq!(
|
||||
request.get("bootstrap_host"),
|
||||
Some(&serde_json::json!("10.0.0.1"))
|
||||
@@ -330,8 +440,9 @@ mod tests {
|
||||
#[test]
|
||||
fn generate_dicts_omit_disaggregated_params_when_absent() {
|
||||
let text_request =
|
||||
build_text_generate_dict("request-1", &proto::TextGenerateRequest::default());
|
||||
let token_request = build_generate_dict("request-2", &proto::GenerateRequest::default());
|
||||
build_text_generate_dict("request-1", &proto::TextGenerateRequest::default()).unwrap();
|
||||
let token_request =
|
||||
build_generate_dict("request-2", &proto::GenerateRequest::default()).unwrap();
|
||||
|
||||
for request in [text_request, token_request] {
|
||||
assert!(!request.contains_key("bootstrap_host"));
|
||||
@@ -339,4 +450,111 @@ mod tests {
|
||||
assert!(!request.contains_key("bootstrap_room"));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generate_dicts_preserve_optional_generation_controls() {
|
||||
let sampling_params = proto::SamplingParams {
|
||||
seed: Some(42),
|
||||
..Default::default()
|
||||
};
|
||||
let text_request = proto::TextGenerateRequest {
|
||||
sampling_params: Some(sampling_params.clone()),
|
||||
priority: Some(3),
|
||||
require_reasoning: Some(false),
|
||||
max_thinking_tokens: Some(128),
|
||||
..Default::default()
|
||||
};
|
||||
let token_request = proto::GenerateRequest {
|
||||
sampling_params: Some(proto::SamplingParams {
|
||||
seed: Some(42),
|
||||
..Default::default()
|
||||
}),
|
||||
priority: Some(3),
|
||||
require_reasoning: Some(false),
|
||||
max_thinking_tokens: Some(128),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
for mapped in [
|
||||
build_text_generate_dict("text-request", &text_request).unwrap(),
|
||||
build_generate_dict("token-request", &token_request).unwrap(),
|
||||
] {
|
||||
assert_eq!(mapped["priority"], serde_json::json!(3));
|
||||
assert_eq!(mapped["require_reasoning"], serde_json::json!(false));
|
||||
assert_eq!(mapped["max_thinking_tokens"], serde_json::json!(128));
|
||||
assert_eq!(
|
||||
mapped["sampling_params"]["sampling_seed"],
|
||||
serde_json::json!(42)
|
||||
);
|
||||
}
|
||||
|
||||
for mapped in [
|
||||
build_text_generate_dict("text-request", &Default::default()).unwrap(),
|
||||
build_generate_dict("token-request", &Default::default()).unwrap(),
|
||||
] {
|
||||
assert!(!mapped.contains_key("priority"));
|
||||
assert!(!mapped.contains_key("require_reasoning"));
|
||||
assert!(!mapped.contains_key("max_thinking_tokens"));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn guided_choice_maps_to_escaped_regex() {
|
||||
let request = proto::GenerateRequest {
|
||||
sampling_params: Some(proto::SamplingParams {
|
||||
guided_decoding: Some(proto::GuidedDecoding {
|
||||
constraint: Some(proto::guided_decoding::Constraint::Choice(
|
||||
proto::ChoiceConstraint {
|
||||
values: vec!["a+b".into(), "x.y".into()],
|
||||
},
|
||||
)),
|
||||
}),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
let mapped = build_generate_dict("request", &request).unwrap();
|
||||
assert_eq!(
|
||||
mapped["sampling_params"]["regex"],
|
||||
serde_json::json!("(?:a\\+b|x\\.y)")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_guidance_combinations_are_rejected() {
|
||||
let conflicting = proto::GenerateRequest {
|
||||
sampling_params: Some(proto::SamplingParams {
|
||||
regex: Some("[a-z]+".into()),
|
||||
guided_decoding: Some(proto::GuidedDecoding {
|
||||
constraint: Some(proto::guided_decoding::Constraint::Regex("[0-9]+".into())),
|
||||
}),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let empty_choice = proto::GenerateRequest {
|
||||
sampling_params: Some(proto::SamplingParams {
|
||||
guided_decoding: Some(proto::GuidedDecoding {
|
||||
constraint: Some(proto::guided_decoding::Constraint::Choice(
|
||||
proto::ChoiceConstraint { values: vec![] },
|
||||
)),
|
||||
}),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let empty_legacy_regex = proto::GenerateRequest {
|
||||
sampling_params: Some(proto::SamplingParams {
|
||||
regex: Some(String::new()),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
for request in [conflicting, empty_choice, empty_legacy_regex] {
|
||||
assert!(build_generate_dict("request", &request).is_err());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user