feat(grpc): support disaggregated generation requests (#30440)

Signed-off-by: Connor Carpenter <connorc@nvidia.com>
Co-authored-by: Ishan Dhanani <ishandhanani@gmail.com>
This commit is contained in:
Connor Carpenter
2026-07-08 14:53:56 -07:00
committed by GitHub
co-authored by Ishan Dhanani
parent ca8f15cd70
commit cc7d6659fd
2 changed files with 86 additions and 0 deletions
@@ -66,6 +66,26 @@ fn trace_headers_to_json(headers: &HashMap<String, String>) -> Option<serde_json
}
}
fn insert_disaggregated_params(
request: &mut HashMap<String, serde_json::Value>,
params: &Option<proto::DisaggregatedParams>,
) {
if let Some(params) = params {
request.insert(
"bootstrap_host".into(),
serde_json::json!(params.bootstrap_host),
);
request.insert(
"bootstrap_port".into(),
serde_json::json!(params.bootstrap_port),
);
request.insert(
"bootstrap_room".into(),
serde_json::json!(params.bootstrap_room),
);
}
}
fn now_timestamp() -> f64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
@@ -131,6 +151,7 @@ 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_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);
}
@@ -178,6 +199,7 @@ 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_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);
}
@@ -269,4 +291,52 @@ mod tests {
Some(&serde_json::json!("session-1"))
);
}
#[test]
fn generate_dicts_include_disaggregated_params() {
let disaggregated_params = Some(proto::DisaggregatedParams {
bootstrap_host: "10.0.0.1".to_string(),
bootstrap_port: 8998,
bootstrap_room: i64::MAX,
});
let text_req = proto::TextGenerateRequest {
disaggregated_params: disaggregated_params.clone(),
..Default::default()
};
let token_req = proto::GenerateRequest {
disaggregated_params,
..Default::default()
};
for request in [
build_text_generate_dict("request-1", &text_req),
build_generate_dict("request-2", &token_req),
] {
assert_eq!(
request.get("bootstrap_host"),
Some(&serde_json::json!("10.0.0.1"))
);
assert_eq!(
request.get("bootstrap_port"),
Some(&serde_json::json!(8998))
);
assert_eq!(
request.get("bootstrap_room"),
Some(&serde_json::json!(i64::MAX))
);
}
}
#[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());
for request in [text_request, token_request] {
assert!(!request.contains_key("bootstrap_host"));
assert!(!request.contains_key("bootstrap_port"));
assert!(!request.contains_key("bootstrap_room"));
}
}
}