diff --git a/sgl-model-gateway/src/core/mod.rs b/sgl-model-gateway/src/core/mod.rs index 388c75404..01a5b6d7a 100644 --- a/sgl-model-gateway/src/core/mod.rs +++ b/sgl-model-gateway/src/core/mod.rs @@ -32,8 +32,8 @@ pub use model_card::{ModelCard, ProviderType}; pub use model_type::{Endpoint, ModelType}; pub use retry::{is_retryable_status, BackoffCalculator, RetryError, RetryExecutor}; pub use worker::{ - worker_to_info, BasicWorker, ConnectionMode, DPAwareWorker, HealthChecker, HealthConfig, - RuntimeType, Worker, WorkerFactory, WorkerLoadGuard, WorkerLoadGuardV2, WorkerType, + attach_guards_to_response, worker_to_info, BasicWorker, ConnectionMode, DPAwareWorker, + HealthChecker, HealthConfig, RuntimeType, Worker, WorkerFactory, WorkerLoadGuard, WorkerType, }; pub use worker_builder::{BasicWorkerBuilder, DPAwareWorkerBuilder}; pub use worker_manager::{LoadMonitor, WorkerManager}; diff --git a/sgl-model-gateway/src/core/worker.rs b/sgl-model-gateway/src/core/worker.rs index 74422c033..29f2da960 100644 --- a/sgl-model-gateway/src/core/worker.rs +++ b/sgl-model-gateway/src/core/worker.rs @@ -8,6 +8,7 @@ use std::{ }; use async_trait::async_trait; +use axum::body::Body; use serde::{Deserialize, Serialize}; use serde_json; use tokio::{sync::OnceCell, time}; @@ -1052,54 +1053,118 @@ pub fn workers_to_urls(workers: &[Box]) -> Vec { workers.iter().map(|w| w.url().to_string()).collect() } -// TODO migrate code to V2 (and then remove this name suffix) -pub struct WorkerLoadGuardV2 { +/// RAII guard for worker load management +/// +/// Automatically decrements worker load when dropped. Can be attached to +/// an axum Response to tie the guard's lifetime to the response body, +/// which is essential for streaming responses where the function returns +/// immediately but the stream continues in the background. +pub struct WorkerLoadGuard { worker: Arc, } -impl WorkerLoadGuardV2 { +impl WorkerLoadGuard { pub fn new(worker: Arc) -> Self { worker.increment_load(); Self { worker } } + + /// Attach this guard to a Response, tying the guard's lifetime to the response body. + /// + /// When the response body is fully consumed or dropped (e.g., client disconnects), + /// the guard is dropped and worker load is decremented automatically. + /// + /// This is the proper RAII pattern for SSE/streaming responses where the handler + /// returns immediately but the stream continues in a background task. + pub fn attach_to_response( + self, + response: axum::response::Response, + ) -> axum::response::Response { + let (parts, body) = response.into_parts(); + + // Wrap body with guard - guard drops when body drops + let guarded_body = GuardedBody { + inner: body, + _guard: self, + }; + + axum::response::Response::from_parts(parts, Body::new(guarded_body)) + } } -impl Drop for WorkerLoadGuardV2 { +impl Drop for WorkerLoadGuard { fn drop(&mut self) { self.worker.decrement_load(); } } -/// RAII guard for worker load management -pub struct WorkerLoadGuard<'a> { - workers: Vec<&'a dyn Worker>, +/// Attach multiple guards to a Response (for dual prefill/decode workers) +pub fn attach_guards_to_response( + guards: Vec, + response: axum::response::Response, +) -> axum::response::Response { + let (parts, body) = response.into_parts(); + + let guarded_body = MultiGuardedBody { + inner: body, + _guards: guards, + }; + + axum::response::Response::from_parts(parts, Body::new(guarded_body)) } -impl<'a> WorkerLoadGuard<'a> { - /// Create a new load guard for a single worker - pub fn new(worker: &'a dyn Worker) -> Self { - worker.increment_load(); - Self { - workers: vec![worker], - } +/// Body wrapper that holds a WorkerLoadGuard +/// +/// When this body is dropped (stream ends or client disconnects), +/// the guard is dropped, decrementing worker load. +struct GuardedBody { + inner: Body, + _guard: WorkerLoadGuard, +} + +/// Body wrapper that holds multiple WorkerLoadGuards (for dual prefill/decode) +struct MultiGuardedBody { + inner: Body, + _guards: Vec, +} + +impl http_body::Body for GuardedBody { + type Data = bytes::Bytes; + type Error = axum::Error; + + fn poll_frame( + mut self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll, Self::Error>>> { + std::pin::Pin::new(&mut self.inner).poll_frame(cx) } - /// Create a new load guard for multiple workers - pub fn new_multi(workers: Vec<&'a dyn Worker>) -> Self { - // Increment load counters for all workers - for worker in &workers { - worker.increment_load(); - } - Self { workers } + fn is_end_stream(&self) -> bool { + self.inner.is_end_stream() + } + + fn size_hint(&self) -> http_body::SizeHint { + self.inner.size_hint() } } -impl<'a> Drop for WorkerLoadGuard<'a> { - fn drop(&mut self) { - // Decrement load counters for all workers - for worker in &self.workers { - worker.decrement_load(); - } +impl http_body::Body for MultiGuardedBody { + type Data = bytes::Bytes; + type Error = axum::Error; + + fn poll_frame( + mut self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll, Self::Error>>> { + std::pin::Pin::new(&mut self.inner).poll_frame(cx) + } + + fn is_end_stream(&self) -> bool { + self.inner.is_end_stream() + } + + fn size_hint(&self) -> http_body::SizeHint { + self.inner.size_hint() } } @@ -1537,81 +1602,6 @@ mod tests { assert_eq!(worker.worker_type(), &WorkerType::Decode); } - #[test] - fn test_load_guard_single_worker() { - use crate::core::BasicWorkerBuilder; - let worker = BasicWorkerBuilder::new("http://test:8080") - .worker_type(WorkerType::Regular) - .build(); - assert_eq!(worker.load(), 0); - - { - let _guard = WorkerLoadGuard::new(&worker); - assert_eq!(worker.load(), 1); - } - - assert_eq!(worker.load(), 0); - } - - #[test] - fn test_load_guard_multiple_workers() { - let workers: Vec> = vec![ - Box::new( - BasicWorkerBuilder::new("http://w1:8080") - .worker_type(WorkerType::Regular) - .build(), - ), - Box::new( - BasicWorkerBuilder::new("http://w2:8080") - .worker_type(WorkerType::Regular) - .build(), - ), - Box::new( - BasicWorkerBuilder::new("http://w3:8080") - .worker_type(WorkerType::Regular) - .build(), - ), - ]; - - let worker_refs: Vec<&dyn Worker> = workers.iter().map(|w| w.as_ref()).collect(); - - { - let _guard = WorkerLoadGuard::new_multi(worker_refs); - assert_eq!(workers[0].load(), 1); - assert_eq!(workers[1].load(), 1); - assert_eq!(workers[2].load(), 1); - } - - assert_eq!(workers[0].load(), 0); - assert_eq!(workers[1].load(), 0); - assert_eq!(workers[2].load(), 0); - } - - #[test] - fn test_load_guard_panic_safety() { - use crate::core::BasicWorkerBuilder; - let worker = Arc::new( - BasicWorkerBuilder::new("http://test:8080") - .worker_type(WorkerType::Regular) - .build(), - ); - assert_eq!(worker.load(), 0); - - let worker_clone = Arc::clone(&worker); - - use std::panic::AssertUnwindSafe; - - let result = std::panic::catch_unwind(AssertUnwindSafe(|| { - let _guard = WorkerLoadGuard::new(worker_clone.as_ref()); - assert_eq!(worker_clone.load(), 1); - panic!("Test panic"); - })); - - assert!(result.is_err()); - - assert_eq!(worker.load(), 0); - } - #[test] fn test_urls_to_workers() { let urls = vec!["http://w1:8080".to_string(), "http://w2:8080".to_string()]; diff --git a/sgl-model-gateway/src/routers/grpc/common/stages/request_execution.rs b/sgl-model-gateway/src/routers/grpc/common/stages/request_execution.rs index e976bd145..593b0aa5b 100644 --- a/sgl-model-gateway/src/routers/grpc/common/stages/request_execution.rs +++ b/sgl-model-gateway/src/routers/grpc/common/stages/request_execution.rs @@ -8,7 +8,7 @@ use super::PipelineStage; use crate::routers::{ error, grpc::{ - context::{ClientSelection, ExecutionResult, LoadGuards, RequestContext, WorkerSelection}, + context::{ClientSelection, ExecutionResult, LoadGuards, RequestContext}, proto_wrapper::{ProtoGenerateRequest, ProtoStream}, }, }; @@ -69,16 +69,7 @@ impl PipelineStage for RequestExecutionStage { ) })?; - let load_guards = match workers { - WorkerSelection::Single { worker } => { - LoadGuards::Single(crate::core::WorkerLoadGuardV2::new(worker.clone())) - } - WorkerSelection::Dual { prefill, decode } => LoadGuards::Dual { - prefill: crate::core::WorkerLoadGuardV2::new(prefill.clone()), - decode: crate::core::WorkerLoadGuardV2::new(decode.clone()), - }, - }; - ctx.state.load_guards = Some(load_guards); + ctx.state.load_guards = Some(LoadGuards::from(workers)); // Extract dispatch metadata for tracing span let request_id = ctx diff --git a/sgl-model-gateway/src/routers/grpc/context.rs b/sgl-model-gateway/src/routers/grpc/context.rs index 36f25673f..abcb17b88 100644 --- a/sgl-model-gateway/src/routers/grpc/context.rs +++ b/sgl-model-gateway/src/routers/grpc/context.rs @@ -14,7 +14,7 @@ use super::{ proto_wrapper::{ProtoGenerateComplete, ProtoGenerateRequest, ProtoStream}, }; use crate::{ - core::{Worker, WorkerLoadGuardV2}, + core::{attach_guards_to_response, Worker, WorkerLoadGuard}, protocols::{ chat::{ChatCompletionRequest, ChatCompletionResponse}, generate::{GenerateRequest, GenerateResponse}, @@ -149,13 +149,47 @@ pub struct DispatchMetadata { /// Load guards for worker load tracking /// Automatically decrements load when dropped pub enum LoadGuards { - Single(WorkerLoadGuardV2), + Single(WorkerLoadGuard), Dual { - prefill: WorkerLoadGuardV2, - decode: WorkerLoadGuardV2, + prefill: WorkerLoadGuard, + decode: WorkerLoadGuard, }, } +impl From<&WorkerSelection> for LoadGuards { + fn from(selection: &WorkerSelection) -> Self { + match selection { + WorkerSelection::Single { worker } => { + LoadGuards::Single(WorkerLoadGuard::new(worker.clone())) + } + WorkerSelection::Dual { prefill, decode } => LoadGuards::Dual { + prefill: WorkerLoadGuard::new(prefill.clone()), + decode: WorkerLoadGuard::new(decode.clone()), + }, + } + } +} + +impl LoadGuards { + /// Attach these load guards to a Response, tying their lifetime to the response body. + /// + /// When the response body is fully consumed or dropped (e.g., client disconnects), + /// the guards are dropped and worker load is decremented automatically. + /// + /// This is the proper RAII pattern for SSE/streaming responses. + pub fn attach_to_response( + self, + response: axum::response::Response, + ) -> axum::response::Response { + let guards = match self { + LoadGuards::Single(guard) => vec![guard], + LoadGuards::Dual { prefill, decode } => vec![prefill, decode], + }; + + attach_guards_to_response(guards, response) + } +} + /// Response processing state (Step 6) #[derive(Default)] pub struct ResponseState { diff --git a/sgl-model-gateway/src/routers/grpc/harmony/responses.rs b/sgl-model-gateway/src/routers/grpc/harmony/responses.rs index da39bed03..1f6549729 100644 --- a/sgl-model-gateway/src/routers/grpc/harmony/responses.rs +++ b/sgl-model-gateway/src/routers/grpc/harmony/responses.rs @@ -787,8 +787,8 @@ async fn execute_mcp_tool_loop_streaming( "Harmony Responses streaming iteration" ); - // Execute pipeline and get stream - let execution_result = match ctx + // Execute pipeline and get stream + load guards + let (execution_result, _load_guards) = match ctx .pipeline .execute_harmony_responses_streaming(¤t_request, ctx) .await @@ -805,6 +805,7 @@ async fn execute_mcp_tool_loop_streaming( }; // Process stream with token-level streaming (mixed tools - emits correct events per tool type) + // Load guards are held during processing and dropped when iteration completes let iteration_result = match HarmonyStreamingProcessor::process_responses_iteration_stream( execution_result, emitter, @@ -999,8 +1000,8 @@ async fn execute_without_mcp_streaming( ) { debug!("No MCP tools - executing single iteration"); - // Execute pipeline and get stream - let execution_result = match ctx + // Execute pipeline and get stream + load guards + let (execution_result, _load_guards) = match ctx .pipeline .execute_harmony_responses_streaming(current_request, ctx) .await @@ -1018,6 +1019,7 @@ async fn execute_without_mcp_streaming( // Process stream (emits all output items during streaming - function tool path emits function_call_arguments.* events) // Pass empty HashSet so all tools are treated as function tools (per-tool detection) + // Load guards are held during processing and dropped when iteration completes let empty_mcp_tools = std::collections::HashSet::new(); let iteration_result = match HarmonyStreamingProcessor::process_responses_iteration_stream( execution_result, @@ -1033,6 +1035,7 @@ async fn execute_without_mcp_streaming( return; } }; + // _load_guards dropped here after iteration completes // Extract usage from iteration result let usage = match iteration_result { diff --git a/sgl-model-gateway/src/routers/grpc/harmony/stages/response_processing.rs b/sgl-model-gateway/src/routers/grpc/harmony/stages/response_processing.rs index eacc8c7bd..ccf43430c 100644 --- a/sgl-model-gateway/src/routers/grpc/harmony/stages/response_processing.rs +++ b/sgl-model-gateway/src/routers/grpc/harmony/stages/response_processing.rs @@ -70,15 +70,22 @@ impl PipelineStage for HarmonyResponseProcessingStage { // For streaming, delegate to streaming processor and return SSE response if is_streaming { - return Ok(Some( - self.streaming_processor - .clone() - .process_streaming_chat_response( - execution_result, - ctx.chat_request_arc(), - dispatch, - ), - )); + let response = self + .streaming_processor + .clone() + .process_streaming_chat_response( + execution_result, + ctx.chat_request_arc(), + dispatch, + ); + + // Attach load guards to response body for proper RAII lifecycle + let response = match ctx.state.load_guards.take() { + Some(guards) => guards.attach_to_response(response), + None => response, + }; + + return Ok(Some(response)); } // For non-streaming, delegate to Harmony response processor to build ChatCompletionResponse diff --git a/sgl-model-gateway/src/routers/grpc/harmony/streaming.rs b/sgl-model-gateway/src/routers/grpc/harmony/streaming.rs index 14a53ec06..55406dfbe 100644 --- a/sgl-model-gateway/src/routers/grpc/harmony/streaming.rs +++ b/sgl-model-gateway/src/routers/grpc/harmony/streaming.rs @@ -116,6 +116,9 @@ impl HarmonyStreamingProcessor { /// Process a streaming Harmony Chat Completion response /// /// Returns an SSE response with streaming token updates. + /// + /// Note: Caller should attach load guards to the returned response using + /// `WorkerLoadGuard::attach_to_response()` for proper RAII lifecycle management. pub fn process_streaming_chat_response( self: Arc, execution_result: context::ExecutionResult, diff --git a/sgl-model-gateway/src/routers/grpc/pipeline.rs b/sgl-model-gateway/src/routers/grpc/pipeline.rs index a0df33540..183ab2d91 100644 --- a/sgl-model-gateway/src/routers/grpc/pipeline.rs +++ b/sgl-model-gateway/src/routers/grpc/pipeline.rs @@ -548,12 +548,13 @@ impl RequestPipeline { /// Execute Harmony Responses pipeline iteration with streaming support /// /// This version executes the pipeline up to the dispatch stage and returns - /// the raw ExecutionResult (with stream) for token-level streaming processing. + /// the raw ExecutionResult (with stream) and LoadGuards for token-level streaming processing. + /// The caller is responsible for keeping load_guards alive until stream processing completes. pub async fn execute_harmony_responses_streaming( &self, request: &crate::protocols::responses::ResponsesRequest, harmony_ctx: &harmony::responses::HarmonyResponsesContext, - ) -> Result { + ) -> Result<(ExecutionResult, Option), Response> { // Create RequestContext for this Responses request let mut ctx = RequestContext::for_responses( Arc::new(request.clone()), @@ -585,8 +586,8 @@ impl RequestPipeline { } } - // Extract execution_result (the raw stream from workers) - ctx.state.response.execution_result.take().ok_or_else(|| { + // Extract execution_result (the raw stream from workers) and load_guards + let execution_result = ctx.state.response.execution_result.take().ok_or_else(|| { error!( function = "execute_harmony_responses_streaming", "No ExecutionResult produced by pipeline" @@ -595,6 +596,10 @@ impl RequestPipeline { "no_execution_result_produced", "No ExecutionResult produced by pipeline", ) - }) + })?; + + let load_guards = ctx.state.load_guards.take(); + + Ok((execution_result, load_guards)) } } diff --git a/sgl-model-gateway/src/routers/grpc/regular/stages/chat/response_processing.rs b/sgl-model-gateway/src/routers/grpc/regular/stages/chat/response_processing.rs index b88fa6144..5c970c2ad 100644 --- a/sgl-model-gateway/src/routers/grpc/regular/stages/chat/response_processing.rs +++ b/sgl-model-gateway/src/routers/grpc/regular/stages/chat/response_processing.rs @@ -80,14 +80,20 @@ impl ChatResponseProcessingStage { .clone(); if is_streaming { - // Streaming: Use StreamingProcessor and return SSE response (done) - return Ok(Some( - self.streaming_processor.clone().process_streaming_response( - execution_result, - ctx.chat_request_arc(), // Cheap Arc clone (8 bytes) - dispatch, - ), - )); + // Streaming: Use StreamingProcessor and return SSE response + let response = self.streaming_processor.clone().process_streaming_response( + execution_result, + ctx.chat_request_arc(), // Cheap Arc clone (8 bytes) + dispatch, + ); + + // Attach load guards to response body for proper RAII lifecycle + let response = match ctx.state.load_guards.take() { + Some(guards) => guards.attach_to_response(response), + None => response, + }; + + return Ok(Some(response)); } // Non-streaming: Delegate to ResponseProcessor diff --git a/sgl-model-gateway/src/routers/grpc/regular/stages/generate/response_processing.rs b/sgl-model-gateway/src/routers/grpc/regular/stages/generate/response_processing.rs index 1ae502894..e386b4d9a 100644 --- a/sgl-model-gateway/src/routers/grpc/regular/stages/generate/response_processing.rs +++ b/sgl-model-gateway/src/routers/grpc/regular/stages/generate/response_processing.rs @@ -78,14 +78,20 @@ impl GenerateResponseProcessingStage { .clone(); if is_streaming { - // Streaming: Use StreamingProcessor and return SSE response (done) - return Ok(Some( - self.streaming_processor.clone().process_streaming_generate( - execution_result, - ctx.generate_request_arc(), // Cheap Arc clone (8 bytes) - dispatch, - ), - )); + // Streaming: Use StreamingProcessor and return SSE response + let response = self.streaming_processor.clone().process_streaming_generate( + execution_result, + ctx.generate_request_arc(), // Cheap Arc clone (8 bytes) + dispatch, + ); + + // Attach load guards to response body for proper RAII lifecycle + let response = match ctx.state.load_guards.take() { + Some(guards) => guards.attach_to_response(response), + None => response, + }; + + return Ok(Some(response)); } // Non-streaming: Delegate to ResponseProcessor diff --git a/sgl-model-gateway/src/routers/grpc/regular/streaming.rs b/sgl-model-gateway/src/routers/grpc/regular/streaming.rs index 5875ea130..e62bdee53 100644 --- a/sgl-model-gateway/src/routers/grpc/regular/streaming.rs +++ b/sgl-model-gateway/src/routers/grpc/regular/streaming.rs @@ -81,6 +81,9 @@ impl StreamingProcessor { /// - Channel creation /// - Background task spawning /// - SSE response building + /// + /// Note: Caller should attach load guards to the returned response using + /// `WorkerLoadGuard::attach_to_response()` for proper RAII lifecycle management. pub fn process_streaming_response( self: Arc, execution_result: context::ExecutionResult, @@ -633,6 +636,9 @@ impl StreamingProcessor { /// Process streaming generate response and return SSE response /// /// Simpler than chat - no tool/reasoning parsing, just text accumulation + /// + /// Note: Caller should attach load guards to the returned response using + /// `WorkerLoadGuard::attach_to_response()` for proper RAII lifecycle management. pub fn process_streaming_generate( self: Arc, execution_result: context::ExecutionResult, diff --git a/sgl-model-gateway/src/routers/http/pd_router.rs b/sgl-model-gateway/src/routers/http/pd_router.rs index 08762c6eb..6494814b1 100644 --- a/sgl-model-gateway/src/routers/http/pd_router.rs +++ b/sgl-model-gateway/src/routers/http/pd_router.rs @@ -8,6 +8,7 @@ use axum::{ response::{IntoResponse, Response}, }; use futures_util::StreamExt; +use memchr::memmem; use reqwest::Client; use serde::Serialize; use serde_json::{json, Value}; @@ -337,8 +338,8 @@ impl PDRouter { headers, json_request, context, - prefill.as_ref(), - decode.as_ref(), + Arc::clone(&prefill), + Arc::clone(&decode), start_time, ) .await; @@ -410,8 +411,8 @@ impl PDRouter { &self, res: reqwest::Response, context: &PDRequestContext<'_>, - prefill: &dyn Worker, - decode: &dyn Worker, + prefill: Arc, + decode: Arc, ) -> Response { let status = res.status(); @@ -526,17 +527,14 @@ impl PDRouter { headers: Option<&HeaderMap>, json_request: Value, context: PDRequestContext<'_>, - prefill: &dyn Worker, - decode: &dyn Worker, + prefill: Arc, + decode: Arc, _start_time: Instant, ) -> Response { // For non-streaming: use guard for automatic load management // For streaming: load will be managed in create_streaming_response - let _guard = if !context.is_stream { - Some(WorkerLoadGuard::new_multi(vec![prefill, decode])) - } else { - None - }; + let _prefill_guard = (!context.is_stream).then(|| WorkerLoadGuard::new(prefill.clone())); + let _decode_guard = (!context.is_stream).then(|| WorkerLoadGuard::new(decode.clone())); let mut headers_with_trace = headers.cloned().unwrap_or_default(); inject_trace_context_http(&mut headers_with_trace); @@ -807,30 +805,19 @@ impl PDRouter { return_logprob: bool, decode_url: Option, headers: Option, - prefill: &dyn Worker, - decode: &dyn Worker, + prefill: Arc, + decode: Arc, ) -> Response { - prefill.increment_load(); - decode.increment_load(); - - let prefill_url = prefill.url().to_string(); - let decode_url_str = decode.url().to_string(); + use crate::core::attach_guards_to_response; let (tx, rx) = tokio::sync::mpsc::unbounded_channel(); - let registry = self.worker_registry.clone(); - tokio::spawn(async move { - let mut stream_completed = false; - futures_util::pin_mut!(stream); while let Some(chunk_result) = stream.next().await { match chunk_result { Ok(chunk) => { - let is_done = chunk - .as_ref() - .windows(12) - .any(|window| window == b"data: [DONE]"); + let is_done = memmem::find(&chunk, b"data: [DONE]").is_some(); let result = if return_logprob && prefill_logprobs.is_some() { Self::merge_streaming_logprobs(prefill_logprobs.clone(), &chunk) @@ -844,7 +831,6 @@ impl PDRouter { } if is_done { - stream_completed = true; break; } } @@ -857,22 +843,6 @@ impl PDRouter { } } } - - if let Some(worker) = registry.get_by_url(&prefill_url) { - worker.decrement_load(); - debug!( - "Decremented load for prefill worker: {} (stream_completed: {})", - prefill_url, stream_completed - ); - } - - if let Some(worker) = registry.get_by_url(&decode_url_str) { - worker.decrement_load(); - debug!( - "Decremented load for decode worker: {} (stream_completed: {})", - decode_url_str, stream_completed - ); - } }); let stream = UnboundedReceiverStream::new(rx); @@ -885,7 +855,10 @@ impl PDRouter { headers.insert(CONTENT_TYPE, HeaderValue::from_static("text/event-stream")); *response.headers_mut() = headers; - response + // Attach load guards to response body for proper RAII lifecycle + // Guards are dropped when response body is consumed or client disconnects + let guards = vec![WorkerLoadGuard::new(prefill), WorkerLoadGuard::new(decode)]; + attach_guards_to_response(guards, response) } // Helper to process non-streaming decode response with logprob merging @@ -1453,23 +1426,27 @@ mod tests { #[test] fn test_worker_load_metrics() { - let prefill_worker = create_test_worker( + let prefill_worker: Arc = Arc::from(create_test_worker( "http://prefill".to_string(), WorkerType::Prefill { bootstrap_port: None, }, true, - ); - let decode_worker = - create_test_worker("http://decode".to_string(), WorkerType::Decode, true); + )); + let decode_worker: Arc = Arc::from(create_test_worker( + "http://decode".to_string(), + WorkerType::Decode, + true, + )); - let _guard = - WorkerLoadGuard::new_multi(vec![prefill_worker.as_ref(), decode_worker.as_ref()]); + let _prefill_guard = WorkerLoadGuard::new(prefill_worker.clone()); + let _decode_guard = WorkerLoadGuard::new(decode_worker.clone()); assert_eq!(prefill_worker.load(), 1); assert_eq!(decode_worker.load(), 1); - drop(_guard); + drop(_prefill_guard); + drop(_decode_guard); assert_eq!(prefill_worker.load(), 0); assert_eq!(decode_worker.load(), 0); @@ -1507,31 +1484,37 @@ mod tests { let (tx, rx) = tokio::sync::mpsc::unbounded_channel(); let stream = UnboundedReceiverStream::new(rx); - let _response = router.create_streaming_response( - stream.map(Ok), - StatusCode::OK, - None, - false, - None, - None, - prefill_ref.as_ref(), - decode_ref.as_ref(), - ); + { + let response = router.create_streaming_response( + stream.map(Ok), + StatusCode::OK, + None, + false, + None, + None, + prefill_ref.clone(), + decode_ref.clone(), + ); - assert_eq!(prefill_ref.load(), 1); - assert_eq!(decode_ref.load(), 1); + // Guards are now attached to response body, so load should be 1 + assert_eq!(prefill_ref.load(), 1); + assert_eq!(decode_ref.load(), 1); - tx.send(bytes::Bytes::from("test data")).unwrap(); + tx.send(bytes::Bytes::from("test data")).unwrap(); - sleep(Duration::from_millis(10)).await; + sleep(Duration::from_millis(10)).await; - assert_eq!(prefill_ref.load(), 1); - assert_eq!(decode_ref.load(), 1); + // Load still 1 while response body exists + assert_eq!(prefill_ref.load(), 1); + assert_eq!(decode_ref.load(), 1); - drop(tx); + drop(tx); - sleep(Duration::from_millis(100)).await; + // Response (and its body with guards) dropped here + drop(response); + } + // Guards dropped when response dropped assert_eq!(prefill_ref.load(), 0); assert_eq!(decode_ref.load(), 0); } diff --git a/sgl-model-gateway/src/routers/http/router.rs b/sgl-model-gateway/src/routers/http/router.rs index adde3f8de..e39783567 100644 --- a/sgl-model-gateway/src/routers/http/router.rs +++ b/sgl-model-gateway/src/routers/http/router.rs @@ -11,7 +11,6 @@ use axum::{ Json, }; use futures_util::StreamExt; -use memchr::memmem; use reqwest::Client; use tokio_stream::wrappers::UnboundedReceiverStream; use tracing::{debug, error}; @@ -19,7 +18,7 @@ use tracing::{debug, error}; use crate::{ config::types::RetryConfig, core::{ - is_retryable_status, ConnectionMode, RetryExecutor, Worker, WorkerLoadGuardV2, + is_retryable_status, ConnectionMode, RetryExecutor, Worker, WorkerLoadGuard, WorkerRegistry, WorkerType, }, observability::{ @@ -265,7 +264,7 @@ impl Router { }; let load_guard = - (policy.name() == "cache_aware").then(|| WorkerLoadGuardV2::new(worker.clone())); + (policy.name() == "cache_aware").then(|| WorkerLoadGuard::new(worker.clone())); events::RequestSentEvent { url: worker.url().to_string(), @@ -443,7 +442,7 @@ impl Router { route: &'static str, worker_url: &str, is_stream: bool, - mut load_guard: Option, + load_guard: Option, ) -> Response { // Get the worker once and reuse for API key and load tracking let worker = self.worker_registry.get_by_url(worker_url); @@ -550,7 +549,7 @@ impl Router { } }; - drop(load_guard); + // load_guard dropped here automatically after response body is read response } else { // Preserve headers for streaming response @@ -561,18 +560,12 @@ impl Router { let stream = res.bytes_stream(); let (tx, rx) = tokio::sync::mpsc::unbounded_channel(); - // Spawn task to forward stream and detect completion + // Spawn task to forward stream tokio::spawn(async move { let mut stream = stream; while let Some(chunk) = stream.next().await { match chunk { Ok(bytes) => { - // Check for stream end marker using memmem for efficiency - if load_guard.is_some() - && memmem::find(&bytes, b"data: [DONE]").is_some() - { - load_guard = None; - } if tx.send(Ok(bytes)).is_err() { break; } @@ -583,7 +576,6 @@ impl Router { } } } - drop(load_guard); }); let stream = UnboundedReceiverStream::new(rx); @@ -592,6 +584,12 @@ impl Router { let mut response = Response::new(body); *response.status_mut() = status; *response.headers_mut() = response_headers; + + // Attach load guard to response body for proper RAII lifecycle + // Guard is dropped when response body is consumed or client disconnects + if let Some(guard) = load_guard { + response = guard.attach_to_response(response); + } response } } diff --git a/sgl-model-gateway/tests/load_guard_raii_test.rs b/sgl-model-gateway/tests/load_guard_raii_test.rs new file mode 100644 index 000000000..ea54e0d48 --- /dev/null +++ b/sgl-model-gateway/tests/load_guard_raii_test.rs @@ -0,0 +1,215 @@ +//! Tests for WorkerLoadGuard RAII pattern with response body attachment +//! +//! These tests verify that load guards properly decrement worker load when: +//! - Response body is fully consumed +//! - Response body is dropped (client disconnect simulation) +//! - Multiple guards are attached (dual prefill/decode workers) + +use std::sync::Arc; + +use axum::{body::Body, response::Response}; +use bytes::Bytes; +use futures_util::StreamExt; +use http_body_util::BodyExt; +use sgl_model_gateway::core::{ + attach_guards_to_response, BasicWorkerBuilder, Worker, WorkerLoadGuard, +}; +use tokio::sync::mpsc; +use tokio_stream::wrappers::UnboundedReceiverStream; + +/// Helper to create an SSE streaming response +fn create_sse_response(rx: mpsc::UnboundedReceiver) -> Response { + let stream = UnboundedReceiverStream::new(rx).map(Ok::<_, std::io::Error>); + let body = Body::from_stream(stream); + Response::new(body) +} + +/// Helper to create a test worker +fn create_test_worker() -> Arc { + Arc::new(BasicWorkerBuilder::new("http://localhost:8000").build()) +} + +#[tokio::test] +async fn test_guard_dropped_when_response_body_consumed() { + let worker = create_test_worker(); + assert_eq!(worker.load(), 0); + + // Create a simple response with some data + let body = Body::from("Hello, World!"); + let response = Response::new(body); + + // Attach guard + let guard = WorkerLoadGuard::new(worker.clone()); + assert_eq!(worker.load(), 1); + + let guarded_response = guard.attach_to_response(response); + + // Load should still be 1 (guard is in the body) + assert_eq!(worker.load(), 1); + + // Consume the response body + let body = guarded_response.into_body(); + let _bytes = body.collect().await.unwrap().to_bytes(); + + // After consuming, guard should be dropped, load should be 0 + assert_eq!(worker.load(), 0); +} + +#[tokio::test] +async fn test_guard_dropped_when_response_dropped_without_consumption() { + let worker = create_test_worker(); + assert_eq!(worker.load(), 0); + + { + let body = Body::from("Hello, World!"); + let response = Response::new(body); + + let guard = WorkerLoadGuard::new(worker.clone()); + assert_eq!(worker.load(), 1); + + let _guarded_response = guard.attach_to_response(response); + + // Load is still 1 + assert_eq!(worker.load(), 1); + + // Response goes out of scope here + } + + // After response is dropped, guard should be dropped, load should be 0 + assert_eq!(worker.load(), 0); +} + +#[tokio::test] +async fn test_streaming_guard_dropped_when_stream_ends() { + let worker = create_test_worker(); + assert_eq!(worker.load(), 0); + + // Create a channel for SSE streaming + let (tx, rx) = mpsc::unbounded_channel::(); + + let response = create_sse_response(rx); + let guard = WorkerLoadGuard::new(worker.clone()); + assert_eq!(worker.load(), 1); + + let guarded_response = guard.attach_to_response(response); + + // Spawn a task to consume the response + let worker_clone = worker.clone(); + let consume_task = tokio::spawn(async move { + { + let mut body = guarded_response.into_body(); + while let Some(result) = body.frame().await { + if result.is_err() { + break; + } + } + // Body is still in scope here, guard not dropped yet + } + // Body dropped here, guard should be dropped + assert_eq!(worker_clone.load(), 0); + }); + + // Send some data + tx.send(Bytes::from("data: chunk1\n\n")).unwrap(); + tx.send(Bytes::from("data: chunk2\n\n")).unwrap(); + + // Load should still be 1 while streaming + tokio::time::sleep(tokio::time::Duration::from_millis(10)).await; + assert_eq!(worker.load(), 1); + + // Close the sender to end the stream + drop(tx); + + // Wait for consumer to finish + consume_task.await.unwrap(); + + // Load should now be 0 + assert_eq!(worker.load(), 0); +} + +#[tokio::test] +async fn test_streaming_guard_dropped_on_client_disconnect() { + let worker = create_test_worker(); + assert_eq!(worker.load(), 0); + + let (tx, rx) = mpsc::unbounded_channel::(); + + let response = create_sse_response(rx); + let guard = WorkerLoadGuard::new(worker.clone()); + assert_eq!(worker.load(), 1); + + let guarded_response = guard.attach_to_response(response); + + // Start consuming but drop early (simulate client disconnect) + { + let mut body = guarded_response.into_body(); + + // Read one frame + tx.send(Bytes::from("data: chunk1\n\n")).unwrap(); + let _ = body.frame().await; + + // Load still 1 + assert_eq!(worker.load(), 1); + + // Body dropped here (simulating client disconnect) + } + + // Guard should be dropped when body is dropped + assert_eq!(worker.load(), 0); + + // tx is still open but no one is listening + drop(tx); +} + +#[tokio::test] +async fn test_multiple_guards_all_dropped() { + let worker1 = create_test_worker(); + let worker2 = create_test_worker(); + assert_eq!(worker1.load(), 0); + assert_eq!(worker2.load(), 0); + + { + let body = Body::from("Hello"); + let response = Response::new(body); + + // Create guards for both workers (simulates dual prefill/decode) + let guard1 = WorkerLoadGuard::new(worker1.clone()); + let guard2 = WorkerLoadGuard::new(worker2.clone()); + assert_eq!(worker1.load(), 1); + assert_eq!(worker2.load(), 1); + + // Attach both guards using attach_guards_to_response + let _response = attach_guards_to_response(vec![guard1, guard2], response); + + // Both loads are 1 + assert_eq!(worker1.load(), 1); + assert_eq!(worker2.load(), 1); + } + + // Both guards dropped when response goes out of scope + assert_eq!(worker1.load(), 0); + assert_eq!(worker2.load(), 0); +} + +#[tokio::test] +async fn test_guard_with_empty_body() { + let worker = create_test_worker(); + assert_eq!(worker.load(), 0); + + { + let body = Body::empty(); + let response = Response::new(body); + + let guard = WorkerLoadGuard::new(worker.clone()); + assert_eq!(worker.load(), 1); + + let guarded_response = guard.attach_to_response(response); + + // Consume empty body + let body = guarded_response.into_body(); + let bytes = body.collect().await.unwrap().to_bytes(); + assert!(bytes.is_empty()); + } + + assert_eq!(worker.load(), 0); +}