Fix cache aware wrong routing caused by incorrect load tracking (#15101)

This commit is contained in:
fzyzcjy
2025-12-15 10:03:54 +08:00
committed by GitHub
parent 1cd0c3bfb7
commit 62b3fdae43
4 changed files with 35 additions and 111 deletions
@@ -95,8 +95,10 @@ impl CircuitBreaker {
/// Create a new circuit breaker with custom configuration and metric label /// Create a new circuit breaker with custom configuration and metric label
pub fn with_config_and_label(config: CircuitBreakerConfig, metric_label: String) -> Self { pub fn with_config_and_label(config: CircuitBreakerConfig, metric_label: String) -> Self {
let init_state = CircuitState::Closed;
RouterMetrics::set_cb_state(&metric_label, init_state.to_int());
Self { Self {
state: Arc::new(RwLock::new(CircuitState::Closed)), state: Arc::new(RwLock::new(init_state)),
consecutive_failures: Arc::new(AtomicU32::new(0)), consecutive_failures: Arc::new(AtomicU32::new(0)),
consecutive_successes: Arc::new(AtomicU32::new(0)), consecutive_successes: Arc::new(AtomicU32::new(0)),
total_failures: Arc::new(AtomicU64::new(0)), total_failures: Arc::new(AtomicU64::new(0)),
+1 -1
View File
@@ -33,7 +33,7 @@ pub use model_type::{Endpoint, ModelType};
pub use retry::{is_retryable_status, BackoffCalculator, RetryError, RetryExecutor}; pub use retry::{is_retryable_status, BackoffCalculator, RetryError, RetryExecutor};
pub use worker::{ pub use worker::{
worker_to_info, BasicWorker, ConnectionMode, DPAwareWorker, HealthChecker, HealthConfig, worker_to_info, BasicWorker, ConnectionMode, DPAwareWorker, HealthChecker, HealthConfig,
RuntimeType, Worker, WorkerFactory, WorkerLoadGuard, WorkerType, RuntimeType, Worker, WorkerFactory, WorkerLoadGuard, WorkerLoadGuardV2, WorkerType,
}; };
pub use worker_builder::{BasicWorkerBuilder, DPAwareWorkerBuilder}; pub use worker_builder::{BasicWorkerBuilder, DPAwareWorkerBuilder};
pub use worker_manager::{LoadMonitor, WorkerManager}; pub use worker_manager::{LoadMonitor, WorkerManager};
+18 -20
View File
@@ -130,26 +130,6 @@ pub trait Worker: Send + Sync + fmt::Debug {
/// Record the outcome of a request to this worker /// Record the outcome of a request to this worker
fn record_outcome(&self, success: bool) { fn record_outcome(&self, success: bool) {
self.circuit_breaker().record_outcome(success); self.circuit_breaker().record_outcome(success);
let after = self.circuit_breaker().state();
if before != after {
let from = before.as_str();
let to = after.as_str();
RouterMetrics::record_cb_state_transition(self.url(), from, to);
}
let state_code = self.circuit_breaker().state().to_int();
RouterMetrics::set_cb_state(self.url(), state_code);
// Update consecutive failures/successes gauges
RouterMetrics::set_cb_consecutive_failures(
self.url(),
self.circuit_breaker().failure_count(),
);
RouterMetrics::set_cb_consecutive_successes(
self.url(),
self.circuit_breaker().success_count(),
);
} }
/// Check if this worker is DP-aware /// Check if this worker is DP-aware
@@ -1054,6 +1034,24 @@ pub fn workers_to_urls(workers: &[Box<dyn Worker>]) -> Vec<String> {
workers.iter().map(|w| w.url().to_string()).collect() workers.iter().map(|w| w.url().to_string()).collect()
} }
// TODO migrate code to V2 (and then remove this name suffix)
pub struct WorkerLoadGuardV2 {
worker: Arc<dyn Worker>,
}
impl WorkerLoadGuardV2 {
pub fn new(worker: Arc<dyn Worker>) -> Self {
worker.increment_load();
Self { worker }
}
}
impl Drop for WorkerLoadGuardV2 {
fn drop(&mut self) {
self.worker.decrement_load();
}
}
/// RAII guard for worker load management /// RAII guard for worker load management
pub struct WorkerLoadGuard<'a> { pub struct WorkerLoadGuard<'a> {
workers: Vec<&'a dyn Worker>, workers: Vec<&'a dyn Worker>,
+13 -89
View File
@@ -19,7 +19,8 @@ use tracing::{debug, error};
use crate::{ use crate::{
config::types::RetryConfig, config::types::RetryConfig,
core::{ core::{
is_retryable_status, ConnectionMode, RetryExecutor, Worker, WorkerRegistry, WorkerType, is_retryable_status, ConnectionMode, RetryExecutor, Worker, WorkerLoadGuardV2,
WorkerRegistry, WorkerType,
}, },
observability::{ observability::{
events::{self, Event}, events::{self, Event},
@@ -265,19 +266,8 @@ impl Router {
None => self.policy_registry.get_default_policy(), None => self.policy_registry.get_default_policy(),
}; };
let load_incremented = if policy.name() == "cache_aware" { let load_guard =
worker.increment_load(); (policy.name() == "cache_aware").then(|| WorkerLoadGuardV2::new(worker.clone()));
true
} else {
false
};
// Keep a clone for potential cleanup on retry
let worker_for_cleanup = if load_incremented {
Some(worker.clone())
} else {
None
};
events::RequestSentEvent { events::RequestSentEvent {
url: worker.url().to_string(), url: worker.url().to_string(),
@@ -294,7 +284,7 @@ impl Router {
route, route,
worker.url(), worker.url(),
is_stream, is_stream,
load_incremented, load_guard,
) )
.await; .await;
@@ -302,14 +292,6 @@ impl Router {
worker.record_outcome(response.status().is_success()); worker.record_outcome(response.status().is_success());
// For retryable failures, we need to decrement load since send_typed_request
// won't have done it (it only decrements on success or non-retryable failures)
if is_retryable_status(response.status()) && load_incremented {
if let Some(cleanup_worker) = worker_for_cleanup {
cleanup_worker.decrement_load();
}
}
response response
} }
@@ -453,7 +435,7 @@ impl Router {
route: &'static str, route: &'static str,
worker_url: &str, worker_url: &str,
is_stream: bool, is_stream: bool,
load_incremented: bool, // Whether load was incremented for this request mut load_guard: Option<WorkerLoadGuardV2>,
) -> Response { ) -> Response {
// Get the worker once and reuse for API key and load tracking // Get the worker once and reuse for API key and load tracking
let worker = self.worker_registry.get_by_url(worker_url); let worker = self.worker_registry.get_by_url(worker_url);
@@ -536,13 +518,6 @@ impl Router {
worker_url, route, e worker_url, route, e
); );
// Decrement load on error if it was incremented
if load_incremented {
if let Some(ref w) = worker {
w.decrement_load();
}
}
return convert_reqwest_error(e); return convert_reqwest_error(e);
} }
}; };
@@ -567,19 +542,9 @@ impl Router {
} }
}; };
// Decrement load counter for non-streaming requests if it was incremented drop(load_guard);
if load_incremented {
if let Some(ref w) = worker {
w.decrement_load();
}
}
response response
} else if load_incremented { } else {
// For streaming with load tracking, we need to manually decrement when done
// Clone the worker Arc for the async block instead of looking it up again
let stream_worker = worker.clone();
// Preserve headers for streaming response // Preserve headers for streaming response
let mut response_headers = header_utils::preserve_response_headers(res.headers()); let mut response_headers = header_utils::preserve_response_headers(res.headers());
// Ensure we set the correct content-type for SSE // Ensure we set the correct content-type for SSE
@@ -591,16 +556,14 @@ impl Router {
// Spawn task to forward stream and detect completion // Spawn task to forward stream and detect completion
tokio::spawn(async move { tokio::spawn(async move {
let mut stream = stream; let mut stream = stream;
let mut decremented = false;
while let Some(chunk) = stream.next().await { while let Some(chunk) = stream.next().await {
match chunk { match chunk {
Ok(bytes) => { Ok(bytes) => {
// Check for stream end marker using memmem for efficiency // Check for stream end marker using memmem for efficiency
if memmem::find(&bytes, b"data: [DONE]").is_some() { if load_guard.is_some()
if let Some(ref w) = stream_worker { && memmem::find(&bytes, b"data: [DONE]").is_some()
w.decrement_load(); {
decremented = true; load_guard = None;
}
} }
if tx.send(Ok(bytes)).is_err() { if tx.send(Ok(bytes)).is_err() {
break; break;
@@ -612,46 +575,7 @@ impl Router {
} }
} }
} }
if !decremented { drop(load_guard);
if let Some(ref w) = stream_worker {
w.decrement_load();
}
}
});
let stream = UnboundedReceiverStream::new(rx);
let body = Body::from_stream(stream);
let mut response = Response::new(body);
*response.status_mut() = status;
*response.headers_mut() = response_headers;
response
} else {
// For requests without load tracking, just stream
// Preserve headers for streaming response
let mut response_headers = header_utils::preserve_response_headers(res.headers());
// Ensure we set the correct content-type for SSE
response_headers.insert(CONTENT_TYPE, HeaderValue::from_static("text/event-stream"));
let stream = res.bytes_stream();
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
// Spawn task to forward stream
tokio::spawn(async move {
let mut stream = stream;
while let Some(chunk) = stream.next().await {
match chunk {
Ok(bytes) => {
if tx.send(Ok(bytes)).is_err() {
break;
}
}
Err(e) => {
let _ = tx.send(Err(format!("Stream error: {}", e)));
break;
}
}
}
}); });
let stream = UnboundedReceiverStream::new(rx); let stream = UnboundedReceiverStream::new(rx);