diff --git a/experimental/sgl-router/src/proxy/abort.rs b/experimental/sgl-router/src/proxy/abort.rs new file mode 100644 index 000000000000..86618922588c --- /dev/null +++ b/experimental/sgl-router/src/proxy/abort.rs @@ -0,0 +1,46 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 The SGLang Authors +// SPDX-License-Identifier: Apache-2.0 + +use axum::http::{header::AUTHORIZATION, HeaderMap}; +use reqwest::{Client, RequestBuilder, Url}; +use std::time::Duration; + +/// Cancels unfinished engine work without delaying request cleanup. +pub(super) struct AbortOnDrop(Option); + +impl AbortOnDrop { + pub(super) fn new( + client: &Client, + worker: &Url, + headers: &HeaderMap, + rid: Option<&str>, + ) -> Self { + Self(rid.filter(|rid| !rid.is_empty()).map(|rid| { + let mut request = client + .post(worker.join("/abort_request").expect("validated worker URL")) + .json(&serde_json::json!({"rid": rid, "abort_all": false})) + .timeout(Duration::from_secs(5)); + if let Some(auth) = headers.get(AUTHORIZATION) { + request = request.header(AUTHORIZATION, auth); + } + request + })) + } + + pub(super) fn disarm(&mut self) { + self.0 = None; + } +} + +impl Drop for AbortOnDrop { + fn drop(&mut self) { + let Some(request) = self.0.take() else { return }; + if let Ok(runtime) = tokio::runtime::Handle::try_current() { + runtime.spawn(async move { + if let Err(error) = request.send().await.and_then(|r| r.error_for_status()) { + tracing::warn!(%error, "engine abort failed"); + } + }); + } + } +} diff --git a/experimental/sgl-router/src/proxy/mod.rs b/experimental/sgl-router/src/proxy/mod.rs index e23d9063fd2b..bc47ecfe1a8b 100644 --- a/experimental/sgl-router/src/proxy/mod.rs +++ b/experimental/sgl-router/src/proxy/mod.rs @@ -3,8 +3,11 @@ //! HTTP proxy — forwards requests to the upstream SGLang worker. +mod abort; pub mod sse; +use abort::AbortOnDrop; + use crate::health::circuit_breaker::CircuitBreaker; use crate::server::error::ApiError; use crate::server::header_utils::should_forward_request_header; @@ -203,6 +206,7 @@ impl Proxy { /// path concatenation (no double-slash) and pass a typed URL to the /// split error variants (`UpstreamUnreachable` / `UpstreamTimeout` / /// `UpstreamStatus`). + #[allow(clippy::too_many_arguments)] pub async fn forward_json_to( &self, worker_url: &str, @@ -211,6 +215,7 @@ impl Proxy { path: &str, headers: &HeaderMap, body: Bytes, + request_id: Option<&str>, ) -> Result, ApiError> { let permit = breaker.acquire().ok_or_else(|| ApiError::BreakerOpen { worker: worker_url.to_string(), @@ -228,6 +233,8 @@ impl Proxy { req = req .header("content-type", "application/json") .timeout(self.request_timeout); + let mut abort = + AbortOnDrop::new(self.client_for(protocol), &worker_url, headers, request_id); let resp = req.send().await.map_err(|e| { breaker.record_failure(); Self::classify_reqwest_error_for(worker_url.clone(), e, path) @@ -257,6 +264,7 @@ impl Proxy { return Err(ApiError::UpstreamStatus { status }); } }; + abort.disarm(); match breaker_outcome(status) { BreakerOutcome::Failure => breaker.record_failure(), BreakerOutcome::Success => breaker.record_success(), @@ -301,6 +309,7 @@ impl Proxy { path: &str, headers: &HeaderMap, body: Bytes, + request_id: Option<&str>, stream_guards: Option>, on_first_byte: Option>, on_stream_end: Option>, @@ -322,11 +331,16 @@ impl Proxy { req = req .header("content-type", "application/json") .header("accept", "text/event-stream"); + let mut abort = + AbortOnDrop::new(self.client_for(protocol), &worker_url, headers, request_id); let resp = req.send().await.map_err(|e| { breaker.record_failure(); Self::classify_reqwest_error_for(worker_url.clone(), e, path) })?; let status = resp.status(); + if !status.is_success() { + abort.disarm(); + } let upstream_ct = resp .headers() .get(reqwest::header::CONTENT_TYPE) @@ -385,6 +399,15 @@ impl Proxy { None }; permit.disarm(); + let on_complete: Option> = + Some(Box::new(move |end| { + if end.reason == sse::StreamEndReason::Completed { + abort.disarm(); + } + if let Some(hook) = on_complete { + hook(end); + } + })); let body = sse::bytes_stream_to_body( resp.bytes_stream(), stream_guards, @@ -465,6 +488,7 @@ mod tests { "/chat", &headers, Bytes::new(), + None, ) .now_or_never() .is_none()); @@ -481,6 +505,7 @@ mod tests { None, None, None, + None, ) .now_or_never() .is_none()); @@ -583,6 +608,7 @@ mod tests { None, None, None, + None, expiration, ) .await @@ -694,6 +720,7 @@ mod tests { "/v1/chat/completions", &headers, Bytes::from_static(b"{}"), + None, ) .await .expect("dispatch should reach the worker (breaker must stay closed)"); @@ -737,6 +764,7 @@ mod tests { "/v1/chat/completions", &headers, Bytes::from_static(b"{}"), + None, ) .await; } @@ -780,6 +808,7 @@ mod tests { "/v1/chat/completions", &headers, Bytes::from_static(b"{}"), + None, ) .await .expect("the half-open probe must be admitted and reach the worker"); @@ -819,6 +848,7 @@ mod tests { None, None, None, + None, ) .await .expect("streaming dispatch should reach the worker"); diff --git a/experimental/sgl-router/src/server/routes/chat/forward.rs b/experimental/sgl-router/src/server/routes/chat/forward.rs index 562310ab3a90..7dc195f7f677 100644 --- a/experimental/sgl-router/src/server/routes/chat/forward.rs +++ b/experimental/sgl-router/src/server/routes/chat/forward.rs @@ -81,7 +81,8 @@ pub(super) async fn forward_chat_request( }; (decode, bootstrap) }); - let body = request.into_outgoing_body(ctx, pd.as_ref().map(|(_, bootstrap)| bootstrap))?; + let (body, request_id) = + request.into_outgoing_body(ctx, pd.as_ref().map(|(_, bootstrap)| bootstrap), &headers)?; let prefill_load_guards = (worker_load_guard, active_request_guard); // In PD mode, prefill runs independently and decode supplies the client response. @@ -112,6 +113,7 @@ pub(super) async fn forward_chat_request( &response_worker, &headers, body, + request_id.as_deref(), response_load_guards, &metrics, expiration_token.clone(), @@ -171,6 +173,7 @@ fn spawn_prefill_request( CHAT_PATH, &headers, body, + None, ) .await { @@ -188,11 +191,13 @@ fn spawn_prefill_request( }); } +#[allow(clippy::too_many_arguments)] async fn forward_to_response_worker( ctx: &AppContext, worker: &Worker, headers: &HeaderMap, body: Bytes, + request_id: Option<&str>, load_guards: LoadGuards, metrics: &DispatchMetrics, expiration: CancellationToken, @@ -209,6 +214,7 @@ async fn forward_to_response_worker( CHAT_PATH, headers, body, + request_id, Some(stream_guards), Some(metrics.first_byte_callback()), Some(metrics.stream_end_callback(worker.url.clone())), @@ -226,6 +232,7 @@ async fn forward_to_response_worker( CHAT_PATH, headers, body, + request_id, ) .await } diff --git a/experimental/sgl-router/src/server/routes/chat/preparation.rs b/experimental/sgl-router/src/server/routes/chat/preparation.rs index 8d9150087d54..b70ccbf0db72 100644 --- a/experimental/sgl-router/src/server/routes/chat/preparation.rs +++ b/experimental/sgl-router/src/server/routes/chat/preparation.rs @@ -23,6 +23,7 @@ pub(super) struct PreparedChatRequest { pub(super) streaming: bool, pub(super) max_output_tokens: Option, pub(super) body: Bytes, + rid: Option, pub(super) tokens: Option, /// Token count for routing/load accounting; estimated from body size when unavailable. pub(super) input_token_count: usize, @@ -67,6 +68,7 @@ impl PreparedChatRequest { streaming: fields.stream.unwrap_or(false), max_output_tokens: fields.requested_max_output_tokens(), body, + rid: fields.rid, tokens, input_token_count, can_forward_input_ids, @@ -79,7 +81,8 @@ impl PreparedChatRequest { self, ctx: &AppContext, bootstrap: Option<&BootstrapFields>, - ) -> Result { + headers: &axum::http::HeaderMap, + ) -> Result<(Bytes, Option), ApiError> { // Routing tokens can replace engine tokenization only for supported chat templates. let input_ids = match (self.tokens.as_ref(), self.parsed_body.as_ref()) { (Some(tokens), Some(parsed_body)) @@ -98,13 +101,40 @@ impl PreparedChatRequest { ) { ctx.metrics.record_ingress_tokenize_error(&self.model.0); } - build_outgoing_body( + let body = build_outgoing_body( &self.body, self.parsed_body, input_ids, bootstrap, &self.sampling_defaults, - ) + )?; + // PD prefill must outlive the client to finish KV transfer. + if bootstrap.is_some() { + return Ok((body, None)); + } + if let Some(rid) = self.rid { + return Ok(( + body, + rid.as_str().filter(|s| !s.is_empty()).map(str::to_owned), + )); + } + let rid = headers + .get("x-request-id") + .and_then(|v| v.to_str().ok()) + .filter(|v| !v.is_empty()) + .map(str::to_owned) + .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()); + // Splice into the validated object, preserving unrelated JSON values verbatim. + let close = body + .iter() + .rposition(|&b| b == b'}') + .ok_or_else(invalid_request)?; + let mut output = Vec::with_capacity(body.len() + rid.len() + 16); + output.extend_from_slice(&body[..close]); + output.extend_from_slice(b",\"rid\":"); + serde_json::to_writer(&mut output, &rid).expect("serialize request ID"); + output.extend_from_slice(&body[close..]); + Ok((Bytes::from(output), Some(rid))) } } @@ -115,6 +145,7 @@ pub(super) struct RoutingFields { pub(super) model: Option, max_tokens: Option, max_completion_tokens: Option, + rid: Option, sampling: [SamplingValue; SamplingField::ALL.len()], } @@ -223,6 +254,7 @@ enum RoutingKey { Model, MaxTokens, MaxCompletionTokens, + Rid, } impl RoutingKey { @@ -232,6 +264,7 @@ impl RoutingKey { Self::Model => 1 << 1, Self::MaxTokens => 1 << 2, Self::MaxCompletionTokens => 1 << 3, + Self::Rid => 1 << 4, } } @@ -241,6 +274,7 @@ impl RoutingKey { Self::Model => "model", Self::MaxTokens => "max_tokens", Self::MaxCompletionTokens => "max_completion_tokens", + Self::Rid => "rid", } } } @@ -263,6 +297,7 @@ impl<'de> Deserialize<'de> for RequestKey { fn visit_str(self, v: &str) -> Result { Ok(match v { + "rid" => RequestKey::Routing(RoutingKey::Rid), "stream" => RequestKey::Routing(RoutingKey::Stream), "model" => RequestKey::Routing(RoutingKey::Model), "max_tokens" => RequestKey::Routing(RoutingKey::MaxTokens), @@ -307,6 +342,7 @@ impl<'de> serde::de::Visitor<'de> for RoutingFieldsVisitor { // Track keys separately so even a repeated null is rejected. seen_routing_keys |= field.bit(); match field { + RoutingKey::Rid => fields.rid = map.next_value()?, RoutingKey::Stream => fields.stream = map.next_value()?, RoutingKey::Model => fields.model = map.next_value()?, RoutingKey::MaxTokens => fields.max_tokens = map.next_value()?, diff --git a/experimental/sgl-router/tests/proxy/cache_aware_input_ids.rs b/experimental/sgl-router/tests/proxy/cache_aware_input_ids.rs index 9e21e044b284..5b900822d94b 100644 --- a/experimental/sgl-router/tests/proxy/cache_aware_input_ids.rs +++ b/experimental/sgl-router/tests/proxy/cache_aware_input_ids.rs @@ -198,7 +198,15 @@ async fn caller_input_ids_are_used_for_routing_and_preserved() { send(Arc::clone(&ctx), request.clone()).await, StatusCode::OK ); - assert_eq!(captured(&mock), request, "body must be forwarded untouched"); + let mut forwarded = captured(&mock); + assert!(forwarded + .as_object_mut() + .unwrap() + .remove("rid") + .unwrap() + .as_str() + .is_some()); + assert_eq!(forwarded, request, "payload must be forwarded untouched"); } // Bypasses are not rendering failures. assert!(!ctx diff --git a/experimental/sgl-router/tests/proxy/chat_routing.rs b/experimental/sgl-router/tests/proxy/chat_routing.rs index 1f92bb685cfa..604ec5d2cc43 100644 --- a/experimental/sgl-router/tests/proxy/chat_routing.rs +++ b/experimental/sgl-router/tests/proxy/chat_routing.rs @@ -21,6 +21,7 @@ use std::sync::Arc; use std::time::Duration; use tower::ServiceExt; +mod cancellation; mod reorg; const TEST_TIMEOUT: Duration = Duration::from_secs(5); @@ -1033,6 +1034,7 @@ async fn forward_json_to_records_failure_on_body_drop() { "/v1/chat/completions", &headers, body, + None, ) .await; assert!(res.is_err(), "body drop should surface as ApiError"); @@ -1090,6 +1092,7 @@ async fn forward_json_to_records_success_only_after_body_completes() { "/v1/chat/completions", &headers, bytes::Bytes::from_static(b"{}"), + None, ) .await; assert!(res.is_ok(), "clean OK call must succeed: {res:?}"); @@ -1145,6 +1148,7 @@ async fn forward_streaming_to_records_failure_on_mid_stream_drop() { None, None, None, + None, ) .await; @@ -1248,6 +1252,7 @@ async fn forward_json_to_records_failure_on_5xx() { "/v1/chat/completions", &headers, body, + None, ) .await; @@ -1282,6 +1287,7 @@ async fn forward_json_to_rejects_when_breaker_open() { "/v1/chat/completions", &headers, body, + None, ) .await; @@ -1320,6 +1326,7 @@ async fn forward_json_to_malformed_url_returns_worker_misconfigured_and_trips_br "/v1/chat/completions", &headers, body, + None, ) .await; diff --git a/experimental/sgl-router/tests/proxy/chat_routing/cancellation.rs b/experimental/sgl-router/tests/proxy/chat_routing/cancellation.rs new file mode 100644 index 000000000000..61d7cbb96ee6 --- /dev/null +++ b/experimental/sgl-router/tests/proxy/chat_routing/cancellation.rs @@ -0,0 +1,105 @@ +use super::*; +use axum::{extract::State, http::HeaderMap, routing::post, Json, Router}; +use serde_json::{json, Value}; +use tokio::sync::mpsc; + +type Events = mpsc::UnboundedSender<(&'static str, Value)>; + +async fn chat(State(events): State, Json(body): Json) -> Body { + events.send(("chat", body.clone())).unwrap(); + if body["hold"] == true { + if body["stream"] == true { + return Body::from_stream(futures::stream::pending::< + Result, + >()); + } + std::future::pending::<()>().await; + } + Body::from(if body["stream"] == true { + "data: [DONE]\n\n" + } else { + "{}" + }) +} + +async fn abort(State(events): State, headers: HeaderMap, Json(body): Json) { + assert_eq!(headers["authorization"], "Bearer test"); + events.send(("abort", body)).unwrap(); +} + +#[tokio::test] +async fn cancellation_aborts_only_unfinished_requests() { + let (events, mut rx) = mpsc::unbounded_channel(); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let url = format!("http://{}", listener.local_addr().unwrap()); + let server = tokio::spawn(async move { + axum::serve( + listener, + Router::new() + .route("/v1/chat/completions", post(chat)) + .route("/abort_request", post(abort)) + .with_state(events), + ) + .await + .unwrap(); + }); + for streaming in [false, true] { + for hold in [false, true] { + let app = build_router(build_ctx_with_worker(&url)); + let mut body = json!({"model": "tiny", "stream": streaming, "hold": hold}); + if streaming { + body["rid"] = json!("caller-rid"); + } + let request = Request::post("/v1/chat/completions") + .header("content-type", "application/json") + .header("authorization", "Bearer test") + .header("x-request-id", "gateway-rid") + .body(Body::from(body.to_string())) + .unwrap(); + let task = tokio::spawn(app.oneshot(request)); + let (event, forwarded) = tokio::time::timeout(TEST_TIMEOUT, rx.recv()) + .await + .unwrap() + .unwrap(); + assert_eq!(event, "chat"); + assert_eq!( + forwarded["rid"], + if streaming { + "caller-rid" + } else { + "gateway-rid" + } + ); + if hold && !streaming { + task.abort(); + assert!(task.await.unwrap_err().is_cancelled()); + } else { + let response = tokio::time::timeout(TEST_TIMEOUT, task) + .await + .unwrap() + .unwrap() + .unwrap(); + if hold { + drop(response); + } else { + response.into_body().collect().await.unwrap(); + } + } + if hold { + let (event, aborted) = tokio::time::timeout(TEST_TIMEOUT, rx.recv()) + .await + .unwrap() + .unwrap(); + assert_eq!(event, "abort"); + assert_eq!( + aborted, + json!({"rid": forwarded["rid"], "abort_all": false}) + ); + } + assert!(tokio::time::timeout(Duration::from_millis(20), rx.recv()) + .await + .is_err()); + } + } + server.abort(); +} diff --git a/experimental/sgl-router/tests/proxy/h2c_forward.rs b/experimental/sgl-router/tests/proxy/h2c_forward.rs index 983fce2b4601..cb7eaa0a5ff3 100644 --- a/experimental/sgl-router/tests/proxy/h2c_forward.rs +++ b/experimental/sgl-router/tests/proxy/h2c_forward.rs @@ -72,6 +72,7 @@ async fn h2c_client_reaches_http2_only_worker() { "/v1/chat/completions", &axum::http::HeaderMap::new(), Bytes::from_static(b"{}"), + None, ) .await .expect("h2c client must reach an HTTP/2-only worker"); @@ -96,6 +97,7 @@ async fn http1_client_cannot_reach_http2_only_worker() { "/v1/chat/completions", &axum::http::HeaderMap::new(), Bytes::from_static(b"{}"), + None, ) .await; assert!( @@ -164,6 +166,7 @@ async fn h2c_client_streams_sse_from_http2_only_worker() { &axum::http::HeaderMap::new(), Bytes::from_static(b"{}"), None, + None, Some(Box::new(move || { flag.store(true, std::sync::atomic::Ordering::SeqCst); })), diff --git a/experimental/sgl-router/tests/proxy/roundrobin_input_ids.rs b/experimental/sgl-router/tests/proxy/roundrobin_input_ids.rs index 729061421db6..77d60e8e320e 100644 --- a/experimental/sgl-router/tests/proxy/roundrobin_input_ids.rs +++ b/experimental/sgl-router/tests/proxy/roundrobin_input_ids.rs @@ -134,7 +134,15 @@ fn captured(mock: &MockWorker) -> Value { .last_body .clone() .expect("worker captured a request body"); - serde_json::from_slice(&b).expect("captured body is valid JSON") + let mut body: Value = serde_json::from_slice(&b).expect("captured body is valid JSON"); + assert!(body + .as_object_mut() + .unwrap() + .remove("rid") + .unwrap() + .as_str() + .is_some()); + body } /// A round-robin (load-only) policy still forwards `input_ids` on a