[PD] Add /v1/responses support to the HTTP PD router (#36141)

Co-authored-by: Shangming Cai <csmthu@gmail.com>
Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com>
This commit is contained in:
lz
2026-09-13 22:48:13 +08:00
committed by GitHub
co-authored by Shangming Cai Xinyuan Tong
parent 7f1f8c706a
commit 6220f45d8e
3 changed files with 564 additions and 6 deletions
@@ -4,7 +4,10 @@ use async_trait::async_trait;
use axum::{
body::Body,
extract::Request,
http::{header::CONTENT_TYPE, HeaderMap, HeaderValue, StatusCode},
http::{
header::{CONTENT_LENGTH, CONTENT_TYPE},
HeaderMap, HeaderValue, StatusCode,
},
response::{IntoResponse, Response},
};
use futures_util::StreamExt;
@@ -36,6 +39,7 @@ use crate::{
embedding::EmbeddingRequest,
generate::GenerateRequest,
rerank::RerankRequest,
responses::ResponsesRequest,
},
routers::{
error,
@@ -419,6 +423,10 @@ impl PDRouter {
Ok(v) => v,
Err(e) => return Self::handle_serialization_error(e),
};
// ResponsesRequest serializes an absent stream as null, which SRT rejects.
if context.route == "/v1/responses" {
json_request["stream"] = Value::Bool(context.is_stream);
}
json_request = match Self::inject_bootstrap_into_value(
json_request,
@@ -530,7 +538,8 @@ impl PDRouter {
if context.is_stream {
// Handle streaming error response
let response_headers = header_utils::preserve_response_headers(res.headers());
let mut response_headers = header_utils::preserve_response_headers(res.headers());
response_headers.remove(CONTENT_LENGTH);
let error_payload = match res.bytes().await {
Ok(error_body) => match serde_json::from_slice::<Value>(&error_body) {
Ok(error_json) => {
@@ -555,10 +564,7 @@ impl PDRouter {
}
};
let sse_data = format!(
"data: {{'error': {}}}",
serde_json::to_string(&error_payload).unwrap_or_default()
);
let sse_data = format!("data: {}\n\n", json!({ "error": error_payload }));
let error_stream = tokio_stream::once(Ok(axum::body::Bytes::from(sse_data)));
self.create_streaming_response(
@@ -1654,6 +1660,50 @@ impl RouterTrait for PDRouter {
self.execute_dual_dispatch(headers, body, context).await
}
async fn route_responses(
&self,
headers: Option<&HeaderMap>,
body: &ResponsesRequest,
model_id: Option<&str>,
) -> Response {
let is_stream = body.is_stream();
// Reject detached requests even when workers lack response-store
// admission checks: the PD router cannot complete their retrieval /
// cancel lifecycle. Attached requests still undergo serving-side
// capability validation, including rejection of background streams
// when response storage is unavailable.
if body.background.unwrap_or(false) && !is_stream {
warn!("PD mode does not support detached background responses; returning bad request");
return error::bad_request(
"pd_unsupported_background_responses",
"PD mode does not support background responses without streaming",
);
}
let request_text = if self.policies_need_request_text() {
let text = body.extract_text_for_routing();
(!text.is_empty()).then_some(text)
} else {
None
};
let context = PDRequestContext {
route: "/v1/responses",
// The Responses API carries one logical response per request.
batch_size: None,
is_stream,
// The PD logprob merging expects /generate-style meta_info,
// which the Responses API schema does not carry.
return_logprob: false,
request_text,
model_id,
headers: headers.cloned(),
};
self.execute_dual_dispatch(headers, body, context).await
}
async fn route_rerank(
&self,
headers: Option<&HeaderMap>,
@@ -220,3 +220,492 @@ mod pd_routing_tests {
ctx.shutdown().await;
}
}
#[cfg(test)]
mod pd_responses_routing_tests {
use std::{sync::Arc, time::Duration};
use axum::routing::post;
use http_body_util::BodyExt;
use smg::{
config::{PolicyConfig, RetryConfig},
core::{
BasicWorkerBuilder, DPAwareWorkerBuilder, Worker, WorkerRegistry,
WorkerType as CoreWorkerType,
},
policies::PolicyRegistry,
protocols::responses::ResponsesRequest,
routers::{pd_router::PDRouter, RouterTrait},
};
use tokio::{
io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader},
net::{TcpListener, TcpStream},
sync::{oneshot, Mutex},
task::JoinHandle,
time::timeout,
};
use super::*;
const STORE_ERROR: &str = r#"{"error":{"message":"Response store is disabled. Stateful Responses require --enable-response-store on a standalone server; response storage is unavailable in PD mode.","type":"invalid_request_error","param":"previous_response_id","code":400}}"#;
const DECODE_RESPONSE: &str =
r#"{"id":"resp_decode","object":"response","status":"completed"}"#;
const DECODE_STREAM: &str = "event: response.completed\ndata: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_decode\",\"status\":\"completed\"}}\n\n";
#[derive(Clone, Copy)]
enum WorkerReply {
Json(StatusCode, &'static str),
Stream(&'static str),
}
/// Regression test: `/v1/responses` must be routed through the PD
/// dual-dispatch path instead of falling back to the `RouterTrait`
/// default `501 NOT_IMPLEMENTED` implementation.
#[tokio::test]
async fn test_pd_mode_responses_routing() {
let config = RouterConfig::builder()
.prefill_decode_mode(
vec![("http://127.0.0.1:19830".to_string(), None)],
vec!["http://127.0.0.1:19831".to_string()],
)
.round_robin_policy()
.host("127.0.0.1")
.port(3803)
.max_payload_size(256 * 1024 * 1024)
.request_timeout_secs(600)
.worker_startup_timeout_secs(5)
.worker_startup_check_interval_secs(1)
.max_concurrent_requests(64)
.queue_timeout_secs(60)
.build_unchecked();
let ctx = AppTestContext::new_with_config(
config,
vec![
TestWorkerConfig::prefill(19830),
TestWorkerConfig::decode(19831),
],
)
.await;
let app = ctx.create_app().await;
let payload = json!({
"model": "mock-model",
"input": "PD mode responses request",
"stream": false
});
let req = Request::builder()
.method("POST")
.uri("/v1/responses")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(serde_json::to_string(&payload).unwrap()))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::OK,
"PD mode /v1/responses request should be dual-dispatched, not 501"
);
ctx.shutdown().await;
}
fn make_pd_router() -> PDRouter {
PDRouter {
worker_registry: Arc::new(WorkerRegistry::new()),
policy_registry: Arc::new(PolicyRegistry::new(PolicyConfig::RoundRobin)),
client: reqwest::Client::new(),
retry_config: RetryConfig::default(),
api_key: None,
enable_igw: false,
}
}
fn register_workers(router: &PDRouter, prefill_url: String, decode_url: String) {
let prefill = DPAwareWorkerBuilder::new(prefill_url, 2, 4)
.worker_type(CoreWorkerType::Prefill {
bootstrap_port: Some(8998),
})
.build();
prefill.set_healthy(true);
let decode = BasicWorkerBuilder::new(decode_url)
.worker_type(CoreWorkerType::Decode)
.build();
decode.set_healthy(true);
router.worker_registry.register(Arc::new(prefill));
router.worker_registry.register(Arc::new(decode));
}
/// Spawn a local worker that records every JSON body POSTed to
/// `/v1/responses` and returns either a JSON or an SSE response.
async fn spawn_capture_worker(
captured: Arc<Mutex<Vec<serde_json::Value>>>,
reply: WorkerReply,
) -> (String, JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let app = axum::Router::new().route(
"/v1/responses",
post(move |axum::Json(body): axum::Json<serde_json::Value>| {
let captured = captured.clone();
async move {
// SRT forwards Responses.stream to ChatCompletionRequest's non-nullable bool.
let stream = body.get("stream").cloned().unwrap_or(json!(false));
if serde_json::from_value::<bool>(stream).is_err() {
return axum::response::Response::builder()
.status(StatusCode::BAD_REQUEST)
.header(CONTENT_TYPE, "application/json")
.body(Body::from(
r#"{"error":{"message":"stream must be a boolean"}}"#,
))
.unwrap();
}
captured.lock().await.push(body);
let (status, content_type, body) = match reply {
WorkerReply::Json(status, body) => (status, "application/json", body),
WorkerReply::Stream(body) => (StatusCode::OK, "text/event-stream", body),
};
axum::response::Response::builder()
.status(status)
.header(CONTENT_TYPE, content_type)
.header("content-length", body.len())
.body(Body::from(body))
.unwrap()
}
}),
);
let task = tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
(format!("http://{addr}"), task)
}
fn responses_request(payload: serde_json::Value) -> ResponsesRequest {
serde_json::from_value(payload).expect("valid responses request")
}
/// PDRouter must inject the PD bootstrap metadata into both worker
/// requests, forward `disagg_prefill_dp_rank` to decode for a DP-aware
/// prefill worker, target `/v1/responses` on both workers, and return the
/// decode response.
#[tokio::test]
async fn test_pd_responses_default_buffered_request_injects_bootstrap_metadata() {
let prefill_bodies = Arc::new(Mutex::new(Vec::new()));
let decode_bodies = Arc::new(Mutex::new(Vec::new()));
let (prefill_url, prefill_task) = spawn_capture_worker(
prefill_bodies.clone(),
WorkerReply::Json(StatusCode::OK, r#"{"id":"resp_prefill"}"#),
)
.await;
let (decode_url, decode_task) = spawn_capture_worker(
decode_bodies.clone(),
WorkerReply::Json(StatusCode::OK, DECODE_RESPONSE),
)
.await;
let router = make_pd_router();
register_workers(&router, prefill_url, decode_url);
let request = responses_request(json!({
"model": "mock-model", "input": "Hello PD responses"
}));
assert_eq!(request.stream, None);
let body = timeout(Duration::from_secs(5), async {
let response = router.route_responses(None, &request, None).await;
assert_eq!(response.status(), StatusCode::OK);
response.into_body().collect().await.unwrap().to_bytes()
})
.await
.expect("PD response must complete");
assert_eq!(body.as_ref(), DECODE_RESPONSE.as_bytes());
let prefill_bodies = prefill_bodies.lock().await;
let decode_bodies = decode_bodies.lock().await;
assert_eq!(prefill_bodies.len(), 1);
assert_eq!(decode_bodies.len(), 1);
let prefill_body = &prefill_bodies[0];
let decode_body = &decode_bodies[0];
for body in [prefill_body, decode_body] {
assert_eq!(body["stream"], false);
assert_eq!(body["input"], "Hello PD responses");
assert_eq!(body["model"], "mock-model");
assert_eq!(body["bootstrap_host"], "127.0.0.1");
assert_eq!(body["bootstrap_port"], 8998);
assert!(body["bootstrap_room"].is_u64());
}
assert_eq!(
decode_body["bootstrap_room"],
prefill_body["bootstrap_room"]
);
assert_eq!(prefill_body["data_parallel_rank"], 2);
assert!(prefill_body.get("disagg_prefill_dp_rank").is_none());
assert_eq!(decode_body["disagg_prefill_dp_rank"], 2);
prefill_task.abort();
decode_task.abort();
}
/// Streaming Responses requests must flow through the PD dual-dispatch
/// path and return the decode SSE stream unchanged, with bootstrap
/// metadata still injected into both worker requests.
#[tokio::test]
async fn test_pd_responses_streaming_passthrough() {
let prefill_bodies = Arc::new(Mutex::new(Vec::new()));
let decode_bodies = Arc::new(Mutex::new(Vec::new()));
let (prefill_url, prefill_task) = spawn_capture_worker(
prefill_bodies.clone(),
WorkerReply::Stream("data: {\"id\":\"resp_prefill\"}\n\n"),
)
.await;
let (decode_url, decode_task) =
spawn_capture_worker(decode_bodies.clone(), WorkerReply::Stream(DECODE_STREAM)).await;
let router = make_pd_router();
register_workers(&router, prefill_url, decode_url);
let request = responses_request(json!({
"model": "mock-model", "input": "Hello PD streaming", "stream": true
}));
let body = timeout(Duration::from_secs(5), async {
let response = router.route_responses(None, &request, None).await;
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(response.headers()[CONTENT_TYPE], "text/event-stream");
response.into_body().collect().await.unwrap().to_bytes()
})
.await
.expect("PD stream must complete without a [DONE] sentinel");
assert_eq!(body.as_ref(), DECODE_STREAM.as_bytes());
let prefill_bodies = prefill_bodies.lock().await;
let decode_bodies = decode_bodies.lock().await;
assert_eq!(prefill_bodies.len(), 1);
assert_eq!(decode_bodies.len(), 1);
assert_eq!(decode_bodies[0]["stream"], true);
assert!(prefill_bodies[0]["bootstrap_room"].is_u64());
assert_eq!(
decode_bodies[0]["bootstrap_room"],
prefill_bodies[0]["bootstrap_room"]
);
prefill_task.abort();
decode_task.abort();
}
#[tokio::test]
async fn test_pd_responses_decode_error_is_valid_sse() {
let (prefill_url, prefill_task) = spawn_capture_worker(
Arc::new(Mutex::new(Vec::new())),
WorkerReply::Json(StatusCode::OK, "{}"),
)
.await;
let (decode_url, decode_task) = spawn_capture_worker(
Arc::new(Mutex::new(Vec::new())),
WorkerReply::Json(StatusCode::BAD_REQUEST, STORE_ERROR),
)
.await;
let router = make_pd_router();
register_workers(&router, prefill_url, decode_url);
let request = responses_request(json!({
"model":"mock-model", "input":"hi", "stream":true
}));
let body = timeout(Duration::from_secs(5), async {
let response = router.route_responses(None, &request, None).await;
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
assert_eq!(response.headers()[CONTENT_TYPE], "text/event-stream");
assert!(!response.headers().contains_key("content-length"));
response.into_body().collect().await.unwrap().to_bytes()
})
.await
.expect("decode error must complete");
let body = std::str::from_utf8(&body).unwrap();
let data = body
.strip_suffix("\n\n")
.expect("SSE event delimiter")
.strip_prefix("data: ")
.unwrap();
let data: serde_json::Value = serde_json::from_str(data).unwrap();
assert_eq!(data["error"]["status"], 400);
assert_eq!(
data["error"]["message"],
serde_json::from_str::<serde_json::Value>(STORE_ERROR).unwrap()
);
prefill_task.abort();
decode_task.abort();
}
/// Read the gateway's fixed-length JSON POST without an HTTP server that
/// could hide disconnects by keeping a pending handler alive.
async fn read_responses_post(socket: &mut BufReader<TcpStream>) -> serde_json::Value {
let mut line = String::new();
socket.read_line(&mut line).await.unwrap();
assert_eq!(line, "POST /v1/responses HTTP/1.1\r\n");
let mut content_length = None;
loop {
line.clear();
assert_ne!(socket.read_line(&mut line).await.unwrap(), 0);
if line == "\r\n" {
break;
}
let (name, value) = line.split_once(':').unwrap();
if name.eq_ignore_ascii_case("content-length") {
content_length = Some(value.trim().parse::<usize>().unwrap());
}
}
let mut body = vec![0; content_length.expect("JSON POST must have Content-Length")];
socket.read_exact(&mut body).await.unwrap();
serde_json::from_slice(&body).unwrap()
}
/// Model #39122's pre-generation admission error without depending on its
/// serving implementation. Prefill waits until decode has received the
/// request, then rejects it without producing KV. Decode sends no headers
/// or KV-timeout response: only a router disconnect can finish its task.
#[tokio::test]
async fn test_pd_responses_stateful_400_cancels_pending_decode() {
for (param, value, stream) in [
("previous_response_id", json!("resp_previous"), false),
("previous_response_id", json!("resp_previous"), true),
("background", json!(true), true),
] {
let prefill_listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let decode_listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let prefill_url = format!("http://{}", prefill_listener.local_addr().unwrap());
let decode_url = format!("http://{}", decode_listener.local_addr().unwrap());
let (decode_started_tx, decode_started_rx) = oneshot::channel();
let admission_error = json!({
"error": {
"message": "Response store is disabled. Stateful Responses require --enable-response-store on a standalone server; response storage is unavailable in PD mode.",
"type": "BadRequestError",
"param": param,
"code": 400
}
});
let error_body = admission_error.to_string();
let mut prefill_task = tokio::spawn(async move {
let (socket, _) = prefill_listener.accept().await.unwrap();
let mut socket = BufReader::new(socket);
let body = read_responses_post(&mut socket).await;
decode_started_rx.await.unwrap();
let response = format!(
"HTTP/1.1 400 Bad Request\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
error_body.len(), error_body
);
socket
.get_mut()
.write_all(response.as_bytes())
.await
.unwrap();
body
});
let mut decode_task = tokio::spawn(async move {
let (socket, _) = decode_listener.accept().await.unwrap();
let mut socket = BufReader::new(socket);
let body = read_responses_post(&mut socket).await;
decode_started_tx.send(()).unwrap();
let mut byte = [0];
assert_eq!(
socket.read(&mut byte).await.unwrap(),
0,
"router must close the pending decode connection after prefill rejects"
);
body
});
let router = make_pd_router();
let prefill = Arc::new(
BasicWorkerBuilder::new(prefill_url)
.worker_type(CoreWorkerType::Prefill {
bootstrap_port: Some(9001),
})
.build(),
);
let decode = Arc::new(
BasicWorkerBuilder::new(decode_url)
.worker_type(CoreWorkerType::Decode)
.build(),
);
prefill.set_healthy(true);
decode.set_healthy(true);
router.worker_registry.register(prefill.clone());
router.worker_registry.register(decode.clone());
let mut payload = json!({
"model": "mock-model",
"input": "Continue the conversation",
"stream": stream
});
payload[param] = value.clone();
let request = responses_request(payload);
// The timeout is only a regression guard. No worker timer, client
// timeout, test teardown, or router drop can release decode here.
let result = timeout(Duration::from_secs(5), async {
let response = router.route_responses(None, &request, None).await;
let status = response.status();
let error_code = response.headers().get("x-smg-error-code").cloned();
let body = response.into_body().collect().await.unwrap().to_bytes();
let (prefill_body, decode_body) = tokio::join!(&mut prefill_task, &mut decode_task);
(
status,
error_code,
body,
prefill_body.unwrap(),
decode_body.unwrap(),
)
})
.await;
if result.is_err() {
// join! may already have consumed a completed handle. Only
// abort and await tasks that are still pending on failure.
for task in [&mut prefill_task, &mut decode_task] {
if !task.is_finished() {
task.abort();
let _ = task.await;
}
}
}
let (status, error_code, body, prefill_body, decode_body) = result.expect(
"prefill admission error must finish routing and disconnect decode promptly",
);
assert_eq!(status, StatusCode::BAD_REQUEST, "{param}, stream={stream}");
assert_eq!(error_code.unwrap(), "prefill_bad_request");
let body: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert!(body["error"]["message"]
.as_str()
.unwrap()
.contains(&admission_error.to_string()));
assert_eq!(prefill_body[param], value);
assert_eq!(decode_body[param], value);
assert!(prefill_body["bootstrap_room"].is_u64());
assert_eq!(
prefill_body["bootstrap_room"],
decode_body["bootstrap_room"]
);
assert_eq!(prefill.load(), 0);
assert_eq!(decode.load(), 0);
}
}
/// Detached background responses (`background=true, stream=false`) are
/// retrieved through the /v1/responses/{id} endpoints, which the PD router
/// does not implement. They must be rejected deterministically instead of
/// dual-dispatched.
#[tokio::test]
async fn test_pd_responses_detached_background_rejected() {
let router = make_pd_router();
let request = responses_request(json!({
"model": "mock-model",
"input": "Hello background",
"background": true,
"stream": false
}));
let response = router.route_responses(None, &request, None).await;
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
assert_eq!(
response
.headers()
.get("x-smg-error-code")
.and_then(|v| v.to_str().ok()),
Some("pd_unsupported_background_responses")
);
}
}