diff --git a/docs/docs/advanced_features/pd_disaggregation.mdx b/docs/docs/advanced_features/pd_disaggregation.mdx index e7a1a5754..1cb97e88d 100644 --- a/docs/docs/advanced_features/pd_disaggregation.mdx +++ b/docs/docs/advanced_features/pd_disaggregation.mdx @@ -28,6 +28,25 @@ When you need to profile prefill or decode workers in PD disaggregation mode, pl For deploying PD disaggregation at scale with load balancing and fault tolerance, SGLang provides a router. The router can distribute requests between prefill and decode instances using various routing policies. For detailed information on setting up routing with PD disaggregation, including configuration options and deployment patterns, see the [SGLang Model Gateway (former Router)](./sgl_model_gateway#prefill-decode-disaggregation). +### Responses API under HTTP PD + +HTTP PD supports foreground generation through `POST /v1/responses`, including streaming generation with `stream=true` through the HTTP PD router. This requires both the router's Responses dispatch support and serving workers that accept Responses PD routing metadata. + +Store-backed and stateful Responses features are **not supported under PD**: + +- Response persistence (`store`) and response history. +- Conversation chaining with `previous_response_id`. +- Background Responses (`background=true`), including background streams. +- Retrieval by response ID (`GET /v1/responses/{id}`). +- Cancel-by-ID (`POST /v1/responses/{id}/cancel`). +- Other workflows requiring shared process-local Responses history. + +Prefill and decode run in separate serving processes. Their process-local Responses stores cannot provide coherent shared state: both legs need the same conversation history to reconstruct the prompt, and prefill's stored response is not available to decode. Send explicit conversation history in foreground requests instead of relying on stored response IDs. + +Built-in tools (`web_search`, `code_interpreter`) are also unsupported under PD: each tool call requires another generation, but the router dispatches exactly one prefill/decode pair per HTTP request. + +Serving-side capability enforcement is tracked in [#39122](https://github.com/sgl-project/sglang/pull/39122). That change introduces opt-in standalone storage with `--enable-response-store`, disallows it under PD, and rejects stateful requests with HTTP 400 before generation. These checks require a serving version containing that change; older versions must not be assumed to reject every unsupported workflow at admission. The HTTP PD router also rejects detached background requests because it does not implement their retrieval/cancel lifecycle. + ## Mooncake ### Requirements diff --git a/sgl-model-gateway/src/routers/http/pd_router.rs b/sgl-model-gateway/src/routers/http/pd_router.rs index ba5dd0560..d5ca6557f 100644 --- a/sgl-model-gateway/src/routers/http/pd_router.rs +++ b/sgl-model-gateway/src/routers/http/pd_router.rs @@ -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::(&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>, diff --git a/sgl-model-gateway/tests/routing/pd_routing_test.rs b/sgl-model-gateway/tests/routing/pd_routing_test.rs index 649ee9133..6ad31453d 100644 --- a/sgl-model-gateway/tests/routing/pd_routing_test.rs +++ b/sgl-model-gateway/tests/routing/pd_routing_test.rs @@ -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>>, + 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| { + 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::(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::(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) -> 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::().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") + ); + } +}