From 122c25033629424f1cd125b8b41f7aaa6cbab20b Mon Sep 17 00:00:00 2001 From: Simo Lin Date: Sun, 21 Dec 2025 15:12:36 -1000 Subject: [PATCH] [model-gateway] add retry and circuit breaker support to gRPC routers (#15585) --- .../grpc/common/stages/request_execution.rs | 19 ++- sgl-model-gateway/src/routers/grpc/context.rs | 19 +++ .../src/routers/grpc/pd_router.rs | 112 +++++++++++++++--- sgl-model-gateway/src/routers/grpc/router.rs | 101 +++++++++++++--- 4 files changed, 211 insertions(+), 40 deletions(-) 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 593b0aa5b..8c2052e92 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}, + context::{ClientSelection, ExecutionResult, LoadGuards, RequestContext, WorkerSelection}, proto_wrapper::{ProtoGenerateRequest, ProtoStream}, }, }; @@ -96,9 +96,10 @@ impl PipelineStage for RequestExecutionStage { let result = async { match self.mode { - ExecutionMode::Single => self.execute_single(proto_request, clients).await, + ExecutionMode::Single => self.execute_single(proto_request, clients, workers).await, ExecutionMode::DualDispatch => { - self.execute_dual_dispatch(proto_request, clients).await + self.execute_dual_dispatch(proto_request, clients, workers) + .await } } } @@ -120,6 +121,7 @@ impl RequestExecutionStage { &self, proto_request: ProtoGenerateRequest, clients: &mut ClientSelection, + workers: &WorkerSelection, ) -> Result { let client = clients.single_mut().ok_or_else(|| { error!( @@ -132,7 +134,12 @@ impl RequestExecutionStage { ) })?; - let stream = client.generate(proto_request).await.map_err(|e| { + let result = client.generate(proto_request).await; + + // Record circuit breaker outcome + workers.record_outcome(result.is_ok()); + + let stream = result.map_err(|e| { error!( function = "execute_single", error = %e, @@ -151,6 +158,7 @@ impl RequestExecutionStage { &self, proto_request: ProtoGenerateRequest, clients: &mut ClientSelection, + workers: &WorkerSelection, ) -> Result { let (prefill_client, decode_client) = clients.dual_mut().ok_or_else(|| { error!( @@ -171,6 +179,9 @@ impl RequestExecutionStage { decode_client.generate(decode_request) ); + // Record circuit breaker outcomes for each worker individually + workers.record_dual_outcomes(prefill_result.is_ok(), decode_result.is_ok()); + // Handle prefill result let prefill_stream = match prefill_result { Ok(s) => s, diff --git a/sgl-model-gateway/src/routers/grpc/context.rs b/sgl-model-gateway/src/routers/grpc/context.rs index abcb17b88..c90339943 100644 --- a/sgl-model-gateway/src/routers/grpc/context.rs +++ b/sgl-model-gateway/src/routers/grpc/context.rs @@ -368,6 +368,25 @@ impl WorkerSelection { } } + /// Record circuit breaker outcome for all workers + pub fn record_outcome(&self, success: bool) { + match self { + Self::Single { worker } => worker.record_outcome(success), + Self::Dual { prefill, decode } => { + prefill.record_outcome(success); + decode.record_outcome(success); + } + } + } + + /// Record circuit breaker outcomes for dual dispatch (individual tracking) + pub fn record_dual_outcomes(&self, prefill_success: bool, decode_success: bool) { + if let Self::Dual { prefill, decode } = self { + prefill.record_outcome(prefill_success); + decode.record_outcome(decode_success); + } + } + #[allow(clippy::type_complexity)] pub fn dual(&self) -> Option<(&Arc, &Arc)> { match self { diff --git a/sgl-model-gateway/src/routers/grpc/pd_router.rs b/sgl-model-gateway/src/routers/grpc/pd_router.rs index 35dd222df..5f9b081ad 100644 --- a/sgl-model-gateway/src/routers/grpc/pd_router.rs +++ b/sgl-model-gateway/src/routers/grpc/pd_router.rs @@ -7,7 +7,9 @@ use tracing::debug; use super::{context::SharedComponents, pipeline::RequestPipeline}; use crate::{ app_context::AppContext, - core::{ConnectionMode, WorkerRegistry, WorkerType}, + config::types::RetryConfig, + core::{is_retryable_status, ConnectionMode, RetryExecutor, WorkerRegistry, WorkerType}, + observability::metrics::{metrics_labels, Metrics}, protocols::{chat::ChatCompletionRequest, generate::GenerateRequest}, routers::RouterTrait, }; @@ -18,6 +20,7 @@ pub struct GrpcPDRouter { worker_registry: Arc, pipeline: RequestPipeline, shared_components: Arc, + retry_config: RetryConfig, } impl GrpcPDRouter { @@ -66,6 +69,7 @@ impl GrpcPDRouter { worker_registry, pipeline, shared_components, + retry_config: ctx.router_config.effective_retry_config(), }) } @@ -81,15 +85,50 @@ impl GrpcPDRouter { model_id ); - // Use pipeline for ALL requests (streaming and non-streaming) - self.pipeline - .execute_generate( - Arc::new(body.clone()), - headers.cloned(), - model_id.map(|s| s.to_string()), - self.shared_components.clone(), - ) - .await + // Clone values needed for retry closure + let request = Arc::new(body.clone()); + let headers_cloned = headers.cloned(); + let model_id_cloned = model_id.map(|s| s.to_string()); + let components = self.shared_components.clone(); + let pipeline = &self.pipeline; + + RetryExecutor::execute_response_with_retry( + &self.retry_config, + |_attempt| { + let request = Arc::clone(&request); + let headers = headers_cloned.clone(); + let model_id = model_id_cloned.clone(); + let components = Arc::clone(&components); + async move { + pipeline + .execute_generate(request, headers, model_id, components) + .await + } + }, + |res, _attempt| is_retryable_status(res.status()), + |delay, attempt| { + Metrics::record_worker_retry( + metrics_labels::WORKER_PREFILL, + metrics_labels::ENDPOINT_GENERATE, + ); + Metrics::record_worker_retry( + metrics_labels::WORKER_DECODE, + metrics_labels::ENDPOINT_GENERATE, + ); + Metrics::record_worker_retry_backoff(attempt, delay); + }, + || { + Metrics::record_worker_retries_exhausted( + metrics_labels::WORKER_PREFILL, + metrics_labels::ENDPOINT_GENERATE, + ); + Metrics::record_worker_retries_exhausted( + metrics_labels::WORKER_DECODE, + metrics_labels::ENDPOINT_GENERATE, + ); + }, + ) + .await } /// Main route_chat implementation with PD dual dispatch @@ -104,15 +143,50 @@ impl GrpcPDRouter { model_id ); - // Use pipeline for ALL requests (streaming and non-streaming) - self.pipeline - .execute_chat( - Arc::new(body.clone()), - headers.cloned(), - model_id.map(|s| s.to_string()), - self.shared_components.clone(), - ) - .await + // Clone values needed for retry closure + let request = Arc::new(body.clone()); + let headers_cloned = headers.cloned(); + let model_id_cloned = model_id.map(|s| s.to_string()); + let components = self.shared_components.clone(); + let pipeline = &self.pipeline; + + RetryExecutor::execute_response_with_retry( + &self.retry_config, + |_attempt| { + let request = Arc::clone(&request); + let headers = headers_cloned.clone(); + let model_id = model_id_cloned.clone(); + let components = Arc::clone(&components); + async move { + pipeline + .execute_chat(request, headers, model_id, components) + .await + } + }, + |res, _attempt| is_retryable_status(res.status()), + |delay, attempt| { + Metrics::record_worker_retry( + metrics_labels::WORKER_PREFILL, + metrics_labels::ENDPOINT_CHAT, + ); + Metrics::record_worker_retry( + metrics_labels::WORKER_DECODE, + metrics_labels::ENDPOINT_CHAT, + ); + Metrics::record_worker_retry_backoff(attempt, delay); + }, + || { + Metrics::record_worker_retries_exhausted( + metrics_labels::WORKER_PREFILL, + metrics_labels::ENDPOINT_CHAT, + ); + Metrics::record_worker_retries_exhausted( + metrics_labels::WORKER_DECODE, + metrics_labels::ENDPOINT_CHAT, + ); + }, + ) + .await } } diff --git a/sgl-model-gateway/src/routers/grpc/router.rs b/sgl-model-gateway/src/routers/grpc/router.rs index 9c5af2a47..7e8b5db2a 100644 --- a/sgl-model-gateway/src/routers/grpc/router.rs +++ b/sgl-model-gateway/src/routers/grpc/router.rs @@ -22,7 +22,9 @@ use super::{ }; use crate::{ app_context::AppContext, - core::WorkerRegistry, + config::types::RetryConfig, + core::{is_retryable_status, RetryExecutor, WorkerRegistry}, + observability::metrics::{metrics_labels, Metrics}, protocols::{ chat::ChatCompletionRequest, generate::GenerateRequest, @@ -41,6 +43,7 @@ pub struct GrpcRouter { shared_components: Arc, responses_context: responses::ResponsesContext, harmony_responses_context: responses::ResponsesContext, + retry_config: RetryConfig, } impl GrpcRouter { @@ -126,6 +129,7 @@ impl GrpcRouter { shared_components, responses_context, harmony_responses_context, + retry_config: ctx.router_config.effective_retry_config(), }) } @@ -151,14 +155,45 @@ impl GrpcRouter { &self.pipeline }; - pipeline - .execute_chat( - Arc::new(body.clone()), - headers.cloned(), - model_id.map(|s| s.to_string()), - self.shared_components.clone(), - ) - .await + // Clone values needed for retry closure + let request = Arc::new(body.clone()); + let headers_cloned = headers.cloned(); + let model_id_cloned = model_id.map(|s| s.to_string()); + let components = self.shared_components.clone(); + + RetryExecutor::execute_response_with_retry( + &self.retry_config, + // Operation: execute pipeline (creates fresh context each attempt) + |_attempt| { + let request = Arc::clone(&request); + let headers = headers_cloned.clone(); + let model_id = model_id_cloned.clone(); + let components = Arc::clone(&components); + async move { + pipeline + .execute_chat(request, headers, model_id, components) + .await + } + }, + // Should retry: check if status is retryable + |res, _attempt| is_retryable_status(res.status()), + // On backoff: record retry metrics + |delay, attempt| { + Metrics::record_worker_retry( + metrics_labels::WORKER_REGULAR, + metrics_labels::ENDPOINT_CHAT, + ); + Metrics::record_worker_retry_backoff(attempt, delay); + }, + // On exhausted: record exhaustion + || { + Metrics::record_worker_retries_exhausted( + metrics_labels::WORKER_REGULAR, + metrics_labels::ENDPOINT_CHAT, + ); + }, + ) + .await } /// Main route_generate implementation @@ -170,14 +205,46 @@ impl GrpcRouter { ) -> Response { debug!("Processing generate request for model: {:?}", model_id); - self.pipeline - .execute_generate( - Arc::new(body.clone()), - headers.cloned(), - model_id.map(|s| s.to_string()), - self.shared_components.clone(), - ) - .await + // Clone values needed for retry closure + let request = Arc::new(body.clone()); + let headers_cloned = headers.cloned(); + let model_id_cloned = model_id.map(|s| s.to_string()); + let components = self.shared_components.clone(); + let pipeline = &self.pipeline; + + RetryExecutor::execute_response_with_retry( + &self.retry_config, + // Operation: execute pipeline (creates fresh context each attempt) + |_attempt| { + let request = Arc::clone(&request); + let headers = headers_cloned.clone(); + let model_id = model_id_cloned.clone(); + let components = Arc::clone(&components); + async move { + pipeline + .execute_generate(request, headers, model_id, components) + .await + } + }, + // Should retry: check if status is retryable + |res, _attempt| is_retryable_status(res.status()), + // On backoff: record retry metrics + |delay, attempt| { + Metrics::record_worker_retry( + metrics_labels::WORKER_REGULAR, + metrics_labels::ENDPOINT_GENERATE, + ); + Metrics::record_worker_retry_backoff(attempt, delay); + }, + // On exhausted: record exhaustion + || { + Metrics::record_worker_retries_exhausted( + metrics_labels::WORKER_REGULAR, + metrics_labels::ENDPOINT_GENERATE, + ); + }, + ) + .await } /// Main route_responses implementation