diff --git a/proto/sglang/runtime/v1/sglang.proto b/proto/sglang/runtime/v1/sglang.proto index 4452ec6ea..980506957 100644 --- a/proto/sglang/runtime/v1/sglang.proto +++ b/proto/sglang/runtime/v1/sglang.proto @@ -34,6 +34,20 @@ service SglangService { rpc UpdateWeightsFromDisk(UpdateWeightsRequest) returns (UpdateWeightsResponse); } +// Disaggregated prefill/decode rendezvous parameters. +// +// Carried on Generate/TextGenerate requests when the worker is part of a +// disaggregated deployment. The prefill leg writes its KV cache to the +// bootstrap server identified by (bootstrap_host, bootstrap_port) under +// the room id; the decode leg fetches it from the same room. Both legs +// receive the same `bootstrap_room` from the router so they rendezvous +// on the KV transfer. +message DisaggregatedParams { + string bootstrap_host = 1; + int32 bootstrap_port = 2; + int64 bootstrap_room = 3; +} + // Sampling parameters shared across text and tokenized RPCs. message SamplingParams { optional float temperature = 1; @@ -69,6 +83,7 @@ message TextGenerateRequest { optional int32 routed_dp_rank = 11; map trace_headers = 12; optional string session_id = 13; + optional DisaggregatedParams disaggregated_params = 14; } message TextGenerateResponse { @@ -92,6 +107,7 @@ message GenerateRequest { optional int32 routed_dp_rank = 10; map trace_headers = 11; optional string session_id = 12; + optional DisaggregatedParams disaggregated_params = 13; } message GenerateResponse { diff --git a/rust/sglang-grpc/src/utils/request_utils.rs b/rust/sglang-grpc/src/utils/request_utils.rs index 1d12e7b2e..4685a8209 100644 --- a/rust/sglang-grpc/src/utils/request_utils.rs +++ b/rust/sglang-grpc/src/utils/request_utils.rs @@ -66,6 +66,26 @@ fn trace_headers_to_json(headers: &HashMap) -> Option, + params: &Option, +) { + 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")); + } + } }