diff --git a/model_gateway/src/policies/cache_aware.rs b/model_gateway/src/policies/cache_aware.rs index 3e30f683fc..f2c869c8b0 100644 --- a/model_gateway/src/policies/cache_aware.rs +++ b/model_gateway/src/policies/cache_aware.rs @@ -483,7 +483,7 @@ impl CacheAwarePolicy { // Use shortest queue when imbalanced let min_load_idx = healthy_indices .iter() - .min_by_key(|&&idx| workers[idx].load()) + .min_by_key(|&&idx| (workers[idx].load(), workers[idx].processed_requests(), idx)) .copied()?; let worker_url = workers[min_load_idx].url(); @@ -868,7 +868,7 @@ impl CacheAwarePolicy { // No cache overlap — min-load fallback (no token tree involved) let min_idx = healthy_indices .iter() - .min_by_key(|&&idx| workers[idx].load()) + .min_by_key(|&&idx| (workers[idx].load(), workers[idx].processed_requests(), idx)) .copied()?; debug!( worker = workers[min_idx].url(), @@ -982,7 +982,9 @@ impl CacheAwarePolicy { } else { healthy_indices .iter() - .min_by_key(|&&idx| workers[idx].load()) + .min_by_key(|&&idx| { + (workers[idx].load(), workers[idx].processed_requests(), idx) + }) .copied() }; @@ -1064,7 +1066,9 @@ impl CacheAwarePolicy { } else { healthy_indices .iter() - .min_by_key(|&&idx| workers[idx].load()) + .min_by_key(|&&idx| { + (workers[idx].load(), workers[idx].processed_requests(), idx) + }) .copied() }; @@ -2276,24 +2280,28 @@ mod tests { // Empty indexer → has_event_indexer returns false → falls through to token tree assert!(!policy.has_event_indexer("unknown")); - // Route a request — should use token tree, not event-driven min-load + // Tokens must be >= PAGE_SIZE (16) to populate the tree; shorter + // sequences are uncacheable and fall through to min-load. + let tokens: Vec = (1..=16).collect(); + + // First request populates the token tree for the selected worker. let idx = policy .select_worker( &workers, &SelectWorkerInfo { - tokens: Some(&[1, 2, 3, 4]), + tokens: Some(&tokens), ..Default::default() }, ) .unwrap(); assert!(idx < 2); // valid worker via token tree - // Route the same tokens again — token tree should route to same worker (cache hit) + // Same tokens again — token-tree cache hit routes to the same worker. let idx2 = policy .select_worker( &workers, &SelectWorkerInfo { - tokens: Some(&[1, 2, 3, 4]), + tokens: Some(&tokens), ..Default::default() }, ) diff --git a/model_gateway/src/routers/http/pd_router.rs b/model_gateway/src/routers/http/pd_router.rs index 8707d3e398..b4af3db7c4 100644 --- a/model_gateway/src/routers/http/pd_router.rs +++ b/model_gateway/src/routers/http/pd_router.rs @@ -489,8 +489,8 @@ impl PDRouter { &self, res: reqwest::Response, context: &PDRequestContext<'_>, - prefill: Arc, decode: Arc, + load_guards: Vec, ) -> Response { let status = res.status(); @@ -524,8 +524,7 @@ impl PDRouter { context.return_logprob, Some(decode_url), Some(response_headers), - prefill, - decode, + load_guards, ) } else { // Handle non-streaming error response @@ -609,12 +608,10 @@ impl PDRouter { prefill: Arc, decode: Arc, ) -> Response { - // For non-streaming: use guard for automatic load management - // For streaming: load will be managed in create_streaming_response - let _prefill_guard = - (!context.is_stream).then(|| WorkerLoadGuard::new(prefill.clone(), headers)); - let _decode_guard = - (!context.is_stream).then(|| WorkerLoadGuard::new(decode.clone(), headers)); + let load_guards = vec![ + WorkerLoadGuard::new(prefill.clone(), headers), + WorkerLoadGuard::new(decode.clone(), headers), + ]; let mut headers_with_trace = headers.cloned().unwrap_or_default(); inject_trace_context_http(&mut headers_with_trace); @@ -680,7 +677,7 @@ impl PDRouter { ); return self - .handle_decode_error_response(decode_response, &context, prefill, decode) + .handle_decode_error_response(decode_response, &context, decode, load_guards) .await; } @@ -714,8 +711,7 @@ impl PDRouter { context.return_logprob, None, Some(response_headers), - prefill, - decode, + load_guards, ) } else { // Non-streaming response @@ -902,8 +898,7 @@ impl PDRouter { return_logprob: bool, decode_url: Option, headers: Option, - prefill: Arc, - decode: Arc, + load_guards: Vec, ) -> Response { use crate::worker::AttachedBody; @@ -949,11 +944,6 @@ impl PDRouter { let stream = UnboundedReceiverStream::new(rx); let body = Body::from_stream(stream); - let guards = vec![ - WorkerLoadGuard::new(prefill, headers.as_ref()), - WorkerLoadGuard::new(decode, headers.as_ref()), - ]; - let mut response = Response::new(body); *response.status_mut() = status; @@ -961,7 +951,7 @@ impl PDRouter { response_headers.insert(CONTENT_TYPE, HeaderValue::from_static("text/event-stream")); *response.headers_mut() = response_headers; - AttachedBody::wrap_response(response, guards) + AttachedBody::wrap_response(response, load_guards) } // Helper to process non-streaming decode response with logprob merging @@ -1585,8 +1575,13 @@ mod tests { headers: None, }; + let load_guards = vec![ + WorkerLoadGuard::new(prefill.clone(), None), + WorkerLoadGuard::new(decode.clone(), None), + ]; + let response = router - .handle_decode_error_response(decode_response, &context, prefill, decode) + .handle_decode_error_response(decode_response, &context, decode, load_guards) .await; let body = axum::body::to_bytes(response.into_body(), usize::MAX) @@ -1638,6 +1633,14 @@ mod tests { let stream = UnboundedReceiverStream::new(rx); { + let guards = vec![ + WorkerLoadGuard::new(prefill_ref.clone(), None), + WorkerLoadGuard::new(decode_ref.clone(), None), + ]; + + assert_eq!(prefill_ref.load(), 1); + assert_eq!(decode_ref.load(), 1); + let response = router.create_streaming_response( stream.map(Ok), StatusCode::OK, @@ -1645,8 +1648,7 @@ mod tests { false, None, None, - prefill_ref.clone(), - decode_ref.clone(), + guards, ); // Guards are now attached to response body, so load should be 1