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:
Connor Carpenter
2026-08-04 18:53:45 -07:00
committed by GitHub
co-authored by ishandhanani Alex Nails
parent 29831d58ef
commit a0b04dbe4c
10 changed files with 982 additions and 86 deletions
+2 -2
View File
@@ -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
+235 -17
View File
@@ -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());
}
}
}