Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 16 additions & 8 deletions model_gateway/src/policies/cache_aware.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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()?;
Comment thread
coderabbitai[bot] marked this conversation as resolved.

let worker_url = workers[min_load_idx].url();
Expand Down Expand Up @@ -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(),
Expand Down Expand Up @@ -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()
};

Expand Down Expand Up @@ -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)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Reserve worker before tree insertion

When several cold cache-miss requests are selected concurrently on different Tokio worker threads, this processed_requests() tie-breaker is not made visible until after match_and_insert_with finishes and increment_processed() runs below, while the WorkerLoadGuard is only created after select_pd_pair returns. For long prompts, multiple requests can therefore observe identical (load, processed_requests) values, all choose the lowest index, and insert for that same worker—the imbalance scenario this change is trying to avoid. Reserve or increment the chosen worker before the tree insertion, and apply the same ordering to the token/min-load paths.

Useful? React with 👍 / 👎.

})
.copied()
};

Expand Down Expand Up @@ -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<u32> = (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()
},
)
Expand Down
48 changes: 25 additions & 23 deletions model_gateway/src/routers/http/pd_router.rs
Original file line number Diff line number Diff line change
Expand Up @@ -489,8 +489,8 @@ impl PDRouter {
&self,
res: reqwest::Response,
context: &PDRequestContext<'_>,
prefill: Arc<dyn Worker>,
decode: Arc<dyn Worker>,
load_guards: Vec<WorkerLoadGuard>,
Comment thread
SYChen123 marked this conversation as resolved.
) -> Response {
let status = res.status();

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -609,12 +608,10 @@ impl PDRouter {
prefill: Arc<dyn Worker>,
decode: Arc<dyn Worker>,
) -> 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);
Expand Down Expand Up @@ -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;
}

Expand Down Expand Up @@ -714,8 +711,7 @@ impl PDRouter {
context.return_logprob,
None,
Some(response_headers),
prefill,
decode,
load_guards,
)
} else {
// Non-streaming response
Expand Down Expand Up @@ -902,8 +898,7 @@ impl PDRouter {
return_logprob: bool,
decode_url: Option<String>,
headers: Option<HeaderMap>,
prefill: Arc<dyn Worker>,
decode: Arc<dyn Worker>,
load_guards: Vec<WorkerLoadGuard>,
) -> Response {
use crate::worker::AttachedBody;

Expand Down Expand Up @@ -949,19 +944,14 @@ 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;

let mut response_headers = headers.unwrap_or_default();
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
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -1638,15 +1633,22 @@ 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,
None,
false,
None,
None,
prefill_ref.clone(),
decode_ref.clone(),
guards,
);

// Guards are now attached to response body, so load should be 1
Expand Down
Loading