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

This commit is contained in:
Simo Lin
2025-12-21 15:12:36 -10:00
committed by GitHub
parent a3a552232d
commit 122c250336
4 changed files with 211 additions and 40 deletions

View File

@@ -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<ExecutionResult, Response> {
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<ExecutionResult, Response> {
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,

View File

@@ -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<dyn Worker>, &Arc<dyn Worker>)> {
match self {

View File

@@ -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<WorkerRegistry>,
pipeline: RequestPipeline,
shared_components: Arc<SharedComponents>,
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
}
}

View File

@@ -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<SharedComponents>,
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