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:
co-authored by
Ishan Dhanani
parent
ca8f15cd70
commit
cc7d6659fd
@@ -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"));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user