[model-gateway] add retry and circuit breaker support to gRPC routers (#15585)

This commit is contained in:
Simo Lin
2025-12-21 17:12:36 -08:00
committed by GitHub
parent a3a552232d
commit 122c250336
4 changed files with 211 additions and 40 deletions
@@ -8,7 +8,7 @@ use super::PipelineStage;
use crate::routers::{ use crate::routers::{
error, error,
grpc::{ grpc::{
context::{ClientSelection, ExecutionResult, LoadGuards, RequestContext}, context::{ClientSelection, ExecutionResult, LoadGuards, RequestContext, WorkerSelection},
proto_wrapper::{ProtoGenerateRequest, ProtoStream}, proto_wrapper::{ProtoGenerateRequest, ProtoStream},
}, },
}; };
@@ -96,9 +96,10 @@ impl PipelineStage for RequestExecutionStage {
let result = async { let result = async {
match self.mode { match self.mode {
ExecutionMode::Single => self.execute_single(proto_request, clients).await, ExecutionMode::Single => self.execute_single(proto_request, clients, workers).await,
ExecutionMode::DualDispatch => { 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, &self,
proto_request: ProtoGenerateRequest, proto_request: ProtoGenerateRequest,
clients: &mut ClientSelection, clients: &mut ClientSelection,
workers: &WorkerSelection,
) -> Result<ExecutionResult, Response> { ) -> Result<ExecutionResult, Response> {
let client = clients.single_mut().ok_or_else(|| { let client = clients.single_mut().ok_or_else(|| {
error!( 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!( error!(
function = "execute_single", function = "execute_single",
error = %e, error = %e,
@@ -151,6 +158,7 @@ impl RequestExecutionStage {
&self, &self,
proto_request: ProtoGenerateRequest, proto_request: ProtoGenerateRequest,
clients: &mut ClientSelection, clients: &mut ClientSelection,
workers: &WorkerSelection,
) -> Result<ExecutionResult, Response> { ) -> Result<ExecutionResult, Response> {
let (prefill_client, decode_client) = clients.dual_mut().ok_or_else(|| { let (prefill_client, decode_client) = clients.dual_mut().ok_or_else(|| {
error!( error!(
@@ -171,6 +179,9 @@ impl RequestExecutionStage {
decode_client.generate(decode_request) 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 // Handle prefill result
let prefill_stream = match prefill_result { let prefill_stream = match prefill_result {
Ok(s) => s, Ok(s) => s,
@@ -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)] #[allow(clippy::type_complexity)]
pub fn dual(&self) -> Option<(&Arc<dyn Worker>, &Arc<dyn Worker>)> { pub fn dual(&self) -> Option<(&Arc<dyn Worker>, &Arc<dyn Worker>)> {
match self { match self {
+89 -15
View File
@@ -7,7 +7,9 @@ use tracing::debug;
use super::{context::SharedComponents, pipeline::RequestPipeline}; use super::{context::SharedComponents, pipeline::RequestPipeline};
use crate::{ use crate::{
app_context::AppContext, 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}, protocols::{chat::ChatCompletionRequest, generate::GenerateRequest},
routers::RouterTrait, routers::RouterTrait,
}; };
@@ -18,6 +20,7 @@ pub struct GrpcPDRouter {
worker_registry: Arc<WorkerRegistry>, worker_registry: Arc<WorkerRegistry>,
pipeline: RequestPipeline, pipeline: RequestPipeline,
shared_components: Arc<SharedComponents>, shared_components: Arc<SharedComponents>,
retry_config: RetryConfig,
} }
impl GrpcPDRouter { impl GrpcPDRouter {
@@ -66,6 +69,7 @@ impl GrpcPDRouter {
worker_registry, worker_registry,
pipeline, pipeline,
shared_components, shared_components,
retry_config: ctx.router_config.effective_retry_config(),
}) })
} }
@@ -81,13 +85,48 @@ impl GrpcPDRouter {
model_id model_id
); );
// Use pipeline for ALL requests (streaming and non-streaming) // Clone values needed for retry closure
self.pipeline let request = Arc::new(body.clone());
.execute_generate( let headers_cloned = headers.cloned();
Arc::new(body.clone()), let model_id_cloned = model_id.map(|s| s.to_string());
headers.cloned(), let components = self.shared_components.clone();
model_id.map(|s| s.to_string()), let pipeline = &self.pipeline;
self.shared_components.clone(),
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 .await
} }
@@ -104,13 +143,48 @@ impl GrpcPDRouter {
model_id model_id
); );
// Use pipeline for ALL requests (streaming and non-streaming) // Clone values needed for retry closure
self.pipeline let request = Arc::new(body.clone());
.execute_chat( let headers_cloned = headers.cloned();
Arc::new(body.clone()), let model_id_cloned = model_id.map(|s| s.to_string());
headers.cloned(), let components = self.shared_components.clone();
model_id.map(|s| s.to_string()), let pipeline = &self.pipeline;
self.shared_components.clone(),
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 .await
} }
+79 -12
View File
@@ -22,7 +22,9 @@ use super::{
}; };
use crate::{ use crate::{
app_context::AppContext, app_context::AppContext,
core::WorkerRegistry, config::types::RetryConfig,
core::{is_retryable_status, RetryExecutor, WorkerRegistry},
observability::metrics::{metrics_labels, Metrics},
protocols::{ protocols::{
chat::ChatCompletionRequest, chat::ChatCompletionRequest,
generate::GenerateRequest, generate::GenerateRequest,
@@ -41,6 +43,7 @@ pub struct GrpcRouter {
shared_components: Arc<SharedComponents>, shared_components: Arc<SharedComponents>,
responses_context: responses::ResponsesContext, responses_context: responses::ResponsesContext,
harmony_responses_context: responses::ResponsesContext, harmony_responses_context: responses::ResponsesContext,
retry_config: RetryConfig,
} }
impl GrpcRouter { impl GrpcRouter {
@@ -126,6 +129,7 @@ impl GrpcRouter {
shared_components, shared_components,
responses_context, responses_context,
harmony_responses_context, harmony_responses_context,
retry_config: ctx.router_config.effective_retry_config(),
}) })
} }
@@ -151,12 +155,43 @@ impl GrpcRouter {
&self.pipeline &self.pipeline
}; };
// 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 pipeline
.execute_chat( .execute_chat(request, headers, model_id, components)
Arc::new(body.clone()), .await
headers.cloned(), }
model_id.map(|s| s.to_string()), },
self.shared_components.clone(), // 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 .await
} }
@@ -170,12 +205,44 @@ impl GrpcRouter {
) -> Response { ) -> Response {
debug!("Processing generate request for model: {:?}", model_id); debug!("Processing generate request for model: {:?}", model_id);
self.pipeline // Clone values needed for retry closure
.execute_generate( let request = Arc::new(body.clone());
Arc::new(body.clone()), let headers_cloned = headers.cloned();
headers.cloned(), let model_id_cloned = model_id.map(|s| s.to_string());
model_id.map(|s| s.to_string()), let components = self.shared_components.clone();
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 .await
} }