Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 46 additions & 0 deletions experimental/sgl-router/src/proxy/abort.rs
Original file line number Diff line number Diff line change
@@ -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<RequestBuilder>);

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");
}
});
}
}
}
30 changes: 30 additions & 0 deletions experimental/sgl-router/src/proxy/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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,
Expand All @@ -211,6 +215,7 @@ impl Proxy {
path: &str,
headers: &HeaderMap,
body: Bytes,
request_id: Option<&str>,
) -> Result<Response<Body>, ApiError> {
let permit = breaker.acquire().ok_or_else(|| ApiError::BreakerOpen {
worker: worker_url.to_string(),
Expand All @@ -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)
Expand Down Expand Up @@ -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(),
Expand Down Expand Up @@ -301,6 +309,7 @@ impl Proxy {
path: &str,
headers: &HeaderMap,
body: Bytes,
request_id: Option<&str>,
stream_guards: Option<Box<dyn Send + 'static>>,
on_first_byte: Option<Box<dyn FnOnce() + Send + 'static>>,
on_stream_end: Option<Box<dyn FnOnce(sse::StreamEnd) + Send + 'static>>,
Expand All @@ -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)
Expand Down Expand Up @@ -385,6 +399,15 @@ impl Proxy {
None
};
permit.disarm();
let on_complete: Option<Box<dyn FnOnce(sse::StreamEnd) + Send>> =
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,
Expand Down Expand Up @@ -465,6 +488,7 @@ mod tests {
"/chat",
&headers,
Bytes::new(),
None,
)
.now_or_never()
.is_none());
Expand All @@ -481,6 +505,7 @@ mod tests {
None,
None,
None,
None,
)
.now_or_never()
.is_none());
Expand Down Expand Up @@ -583,6 +608,7 @@ mod tests {
None,
None,
None,
None,
expiration,
)
.await
Expand Down Expand Up @@ -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)");
Expand Down Expand Up @@ -737,6 +764,7 @@ mod tests {
"/v1/chat/completions",
&headers,
Bytes::from_static(b"{}"),
None,
)
.await;
}
Expand Down Expand Up @@ -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");
Expand Down Expand Up @@ -819,6 +848,7 @@ mod tests {
None,
None,
None,
None,
)
.await
.expect("streaming dispatch should reach the worker");
Expand Down
9 changes: 8 additions & 1 deletion experimental/sgl-router/src/server/routes/chat/forward.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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(),
Expand Down Expand Up @@ -171,6 +173,7 @@ fn spawn_prefill_request(
CHAT_PATH,
&headers,
body,
None,
)
.await
{
Expand All @@ -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,
Expand All @@ -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())),
Expand All @@ -226,6 +232,7 @@ async fn forward_to_response_worker(
CHAT_PATH,
headers,
body,
request_id,
)
.await
}
Expand Down
42 changes: 39 additions & 3 deletions experimental/sgl-router/src/server/routes/chat/preparation.rs
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@ pub(super) struct PreparedChatRequest {
pub(super) streaming: bool,
pub(super) max_output_tokens: Option<u64>,
pub(super) body: Bytes,
rid: Option<Value>,
pub(super) tokens: Option<RequestTokens>,
/// Token count for routing/load accounting; estimated from body size when unavailable.
pub(super) input_token_count: usize,
Expand Down Expand Up @@ -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,
Expand All @@ -79,7 +81,8 @@ impl PreparedChatRequest {
self,
ctx: &AppContext,
bootstrap: Option<&BootstrapFields>,
) -> Result<Bytes, ApiError> {
headers: &axum::http::HeaderMap,
) -> Result<(Bytes, Option<String>), 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))
Expand All @@ -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)))
}
}

Expand All @@ -115,6 +145,7 @@ pub(super) struct RoutingFields {
pub(super) model: Option<String>,
max_tokens: Option<u64>,
max_completion_tokens: Option<u64>,
rid: Option<Value>,
sampling: [SamplingValue; SamplingField::ALL.len()],
}

Expand Down Expand Up @@ -223,6 +254,7 @@ enum RoutingKey {
Model,
MaxTokens,
MaxCompletionTokens,
Rid,
}

impl RoutingKey {
Expand All @@ -232,6 +264,7 @@ impl RoutingKey {
Self::Model => 1 << 1,
Self::MaxTokens => 1 << 2,
Self::MaxCompletionTokens => 1 << 3,
Self::Rid => 1 << 4,
}
}

Expand All @@ -241,6 +274,7 @@ impl RoutingKey {
Self::Model => "model",
Self::MaxTokens => "max_tokens",
Self::MaxCompletionTokens => "max_completion_tokens",
Self::Rid => "rid",
}
}
}
Expand All @@ -263,6 +297,7 @@ impl<'de> Deserialize<'de> for RequestKey {

fn visit_str<E>(self, v: &str) -> Result<RequestKey, E> {
Ok(match v {
"rid" => RequestKey::Routing(RoutingKey::Rid),
"stream" => RequestKey::Routing(RoutingKey::Stream),
"model" => RequestKey::Routing(RoutingKey::Model),
"max_tokens" => RequestKey::Routing(RoutingKey::MaxTokens),
Expand Down Expand Up @@ -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()?,
Expand Down
10 changes: 9 additions & 1 deletion experimental/sgl-router/tests/proxy/cache_aware_input_ids.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading
Loading