Skip to content
Closed
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
311 changes: 297 additions & 14 deletions model_gateway/src/routers/common/worker_selection.rs
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,10 @@ use crate::{
common::header_utils::{apply_provider_headers, extract_auth_header},
error,
},
worker::{ConnectionMode, ProviderType, RuntimeType, Worker, WorkerRegistry, WorkerType},
worker::{
ConnectionMode, ProviderType, RuntimeType, Worker, WorkerRegistry, WorkerType,
UNKNOWN_MODEL_ID,
},
};

/// Holds references to shared infrastructure needed for worker selection.
Expand Down Expand Up @@ -102,23 +105,82 @@ impl<'a> WorkerSelector<'a> {
})
}

/// Available workers passing every request filter, ready for load-based
/// selection.
///
/// When the model is known, this uses the registry's bounded,
/// wildcard-safe per-model lookup (`get_candidates_for_model`) instead
/// of scanning the whole fleet — the O(total fleet) scan was the hot-path
/// cost this addresses. If that bounded set turns up no available
/// candidates, it falls back to the full scan so behavior is never worse
/// (the bounded index can briefly lag a concurrent registration, and we
/// must never regress a servable model into "no worker found").
fn get_candidates(&self, req: &SelectWorkerRequest<'_>) -> Vec<Arc<dyn Worker>> {
let workers = self.registry.get_workers_filtered(
None, // model_id index lookup not used — we filter via supports_model
req.worker_type,
req.connection_mode,
req.runtime_type,
false, // we filter availability ourselves for consistent behavior
);
if req.model_id != UNKNOWN_MODEL_ID {
let bounded =
Self::filter_candidates(self.registry.get_candidates_for_model(req.model_id), req);
if !bounded.is_empty() {
return bounded;
Comment on lines +119 to +123

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

⚠️ Potential issue | 🔴 Critical | ⚡ Quick win

Keep provider-scoped requests on the full provider-aware path.

The bounded lookup shrinks the candidate set before filter_by_provider decides whether a multi-provider filter is needed. In a fleet with an OpenAI worker for the requested model and an Anthropic worker for another model, an Anthropic request can see only the OpenAI worker in bounded, so filter_by_provider treats it as a single-provider set and allows routing to the wrong provider. That violates the SelectWorkerRequest::provider credential-leakage guard used by the Anthropic router. A conservative fix is to skip the bounded path whenever req.provider.is_some() unless provider diversity is tracked registry-wide.

🛡️ Conservative fix
-        if req.model_id != UNKNOWN_MODEL_ID {
+        if req.model_id != UNKNOWN_MODEL_ID && req.provider.is_none() {
             let bounded =
                 Self::filter_candidates(self.registry.get_candidates_for_model(req.model_id), req);
             if !bounded.is_empty() {
                 return bounded;
             }
@@
-        if req.model_id != UNKNOWN_MODEL_ID {
+        if req.model_id != UNKNOWN_MODEL_ID && req.provider.is_none() {
             let bounded = Self::healthy_supporting_candidates(
                 self.registry.get_candidates_for_model(req.model_id),
                 req,
             );

Also applies to: 194-199

🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

In `@model_gateway/src/routers/common/worker_selection.rs` around lines 119 - 123,
The bounded lookup optimization that filters candidates by model ID before
provider filtering can cause provider-scoped requests to be routed incorrectly
by hiding workers from other providers in the candidate set. When a provider is
specified in the request, the bounded path should be skipped to allow the full
provider-aware filtering logic to work correctly. Modify the condition that
checks `req.model_id != UNKNOWN_MODEL_ID` to also require
`req.provider.is_none()` before taking the bounded path. Apply this same fix to
both occurrences of this pattern (the one shown at startLine 119 and the other
mentioned around line 194-199).

}
Comment on lines +119 to +124

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

🟡 Nit: The PR description notes that a stale wildcard (one that has since called set_models but stays in wildcard_workers until the next mutation) is "harmless" because supports_model re-filters. That's true downstream in find_best_worker, but the interaction with this early return is subtle: if the bounded set is non-empty only because of stale wildcards (all later rejected by supports_model), and the requested model is reachable only via alias (not indexed), the fallback to full scan never runs and the alias-only worker is never found.

The scenario is narrow (stale wildcard + alias-only model + no indexed workers for the model), but it's the one case where the "never worse" fallback guarantee doesn't hold. Worth noting in the doc comment or adding a supports_model pre-check to the emptiness test:

Suggested change
if req.model_id != UNKNOWN_MODEL_ID {
let bounded =
Self::filter_candidates(self.registry.get_candidates_for_model(req.model_id), req);
if !bounded.is_empty() {
return bounded;
}
if req.model_id != UNKNOWN_MODEL_ID {
let bounded =
Self::filter_candidates(self.registry.get_candidates_for_model(req.model_id), req);
if bounded.iter().any(|w| w.supports_model(req.model_id)) {
return bounded;
}
}

This way the fallback triggers when the bounded set has only stale wildcards, at the cost of one extra supports_model pass on the bounded set (which is cheap and already done in find_best_worker).

Comment on lines +120 to +124

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

⚠️ Potential issue | 🟠 Major | ⚡ Quick win

Fall back when bounded candidates do not support the requested model.

get_candidates returns any non-empty bounded set before find_best_worker applies supports_model. Because stale wildcard entries are explicitly allowed, a stale-but-available wildcard can make bounded non-empty, then get rejected later by supports_model, skipping the full-scan alias fallback and missing a healthy alias-only worker.

🐛 Proposed fix
             let bounded =
                 Self::filter_candidates(self.registry.get_candidates_for_model(req.model_id), req);
-            if !bounded.is_empty() {
+            if bounded.iter().any(|w| w.supports_model(req.model_id)) {
                 return bounded;
             }
📝 Committable suggestion

‼️ IMPORTANT
Carefully review the code before committing. Ensure that it accurately replaces the highlighted code, contains no missing lines, and has no issues with indentation. Thoroughly test & benchmark the code to ensure it meets the requirements.

Suggested change
let bounded =
Self::filter_candidates(self.registry.get_candidates_for_model(req.model_id), req);
if !bounded.is_empty() {
return bounded;
}
let bounded =
Self::filter_candidates(self.registry.get_candidates_for_model(req.model_id), req);
if bounded.iter().any(|w| w.supports_model(req.model_id)) {
return bounded;
}
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

In `@model_gateway/src/routers/common/worker_selection.rs` around lines 120 - 124,
The early return in the bounded candidates check does not verify that the
returned candidates actually support the requested model. Before returning the
bounded set when it is non-empty, add a check to ensure at least one candidate
in bounded supports the model by filtering or checking against supports_model.
If no candidates in the bounded set support the requested model, allow execution
to continue to the full-scan alias fallback mechanism instead of returning an
unsupported candidate set. This prevents stale wildcard entries from causing
premature returns and skipping the fallback that might find a healthy alias-only
worker.

}

// Unknown model (caller wants any worker) or empty bounded set:
// fall back to the full scan, preserving the original behavior.
Self::filter_candidates(
self.registry.get_workers_filtered(
None, // full scan — wildcard-safe via supports_model below
req.worker_type,
req.connection_mode,
req.runtime_type,
false, // we filter availability ourselves for consistent behavior
),
req,
)
}

let candidates: Vec<_> = workers.into_iter().filter(|w| w.is_available()).collect();
/// Apply the request's worker_type/connection_mode/runtime_type,
/// availability, and provider filters to a candidate set.
///
/// `get_workers_filtered` already applies the type/mode/runtime filters,
/// so re-applying them to that path is a cheap no-op; the bounded
/// per-model path relies on them here.
fn filter_candidates(
workers: Vec<Arc<dyn Worker>>,
req: &SelectWorkerRequest<'_>,
) -> Vec<Arc<dyn Worker>> {
let candidates: Vec<_> = workers
.into_iter()
.filter(|w| Self::matches_filters(w, req) && w.is_available())
.collect();

match &req.provider {
Some(provider) => filter_by_provider(candidates, provider),
None => candidates,
}
}

/// Per-worker worker_type/connection_mode/runtime_type filter, mirroring
/// `WorkerRegistry::get_workers_filtered`. Applied to the bounded
/// per-model candidate set (which is not pre-filtered).
fn matches_filters(worker: &Arc<dyn Worker>, req: &SelectWorkerRequest<'_>) -> bool {
if let Some(ref wtype) = req.worker_type {
if *worker.worker_type() != *wtype {
return false;
}
}
if let Some(ref conn) = req.connection_mode {
if worker.connection_mode() != conn {
return false;
}
}
if let Some(ref rt) = req.runtime_type {
if worker.metadata().spec.runtime_type != *rt {
return false;
}
}
true
}

fn find_best_worker(&self, req: &SelectWorkerRequest<'_>) -> Option<Arc<dyn Worker>> {
self.get_candidates(req)
.into_iter()
Expand All @@ -129,18 +191,46 @@ impl<'a> WorkerSelector<'a> {
/// Check if any healthy worker supports the model (regardless of circuit breaker).
/// Used to distinguish "model not found" from "all workers circuit-broken".
fn any_worker_supports_model(&self, req: &SelectWorkerRequest<'_>) -> bool {
if req.model_id != UNKNOWN_MODEL_ID {
let bounded = Self::healthy_supporting_candidates(
self.registry.get_candidates_for_model(req.model_id),
req,
);
if bounded.iter().any(|w| w.supports_model(req.model_id)) {
return true;
}
// Bounded set yielded no supporting worker — fall through to the
// full scan so we never falsely report "model not found" if the
// per-model index briefly lags a registration.
}

let workers = self.registry.get_workers_filtered(
None,
req.worker_type,
req.connection_mode,
req.runtime_type,
true, // healthy only — model exists even if circuit-broken
);
let candidates = match &req.provider {
Some(p) => filter_by_provider(workers, p),
None => workers,
};
candidates.iter().any(|w| w.supports_model(req.model_id))
Self::healthy_supporting_candidates(workers, req)
.iter()
.any(|w| w.supports_model(req.model_id))
}

/// Filter a candidate set to healthy workers passing the request's
/// type/mode/runtime and provider filters (no circuit-breaker check —
/// the model exists even if every worker is circuit-broken).
fn healthy_supporting_candidates(
workers: Vec<Arc<dyn Worker>>,
req: &SelectWorkerRequest<'_>,
) -> Vec<Arc<dyn Worker>> {
let candidates: Vec<_> = workers
.into_iter()
.filter(|w| Self::matches_filters(w, req) && w.is_healthy())
.collect();
match &req.provider {
Some(p) => filter_by_provider(candidates, p),
None => candidates,
}
}

/// Refresh model lists for healthy external workers in parallel.
Expand Down Expand Up @@ -274,3 +364,196 @@ async fn refresh_worker_models(
}
}
}

#[cfg(test)]
mod tests {
use openai_protocol::{model_card::ModelCard, worker::HealthCheckConfig};

use super::*;
use crate::worker::BasicWorkerBuilder;

fn ready_worker(builder: BasicWorkerBuilder) -> Arc<dyn Worker> {
// disable_health_check makes the worker start Ready (and thus
// is_available, since a fresh circuit breaker permits execution).
let worker: Arc<dyn Worker> = Arc::new(
builder
.health_config(HealthCheckConfig {
disable_health_check: true,
..Default::default()
})
.build(),
);
assert!(worker.is_available(), "test worker must be routable");
worker
}

fn selector_request(model_id: &str) -> SelectWorkerRequest<'_> {
SelectWorkerRequest {
model_id,
..Default::default()
}
}

fn client() -> reqwest::Client {
reqwest::Client::new()
}

/// The core regression: a wildcard worker (no models declared) must be
/// selectable for an arbitrary model via the bounded candidate path,
/// even though it is not in `get_by_model(arbitrary_model)`.
#[tokio::test]
async fn wildcard_worker_selected_for_arbitrary_model_via_bounded_path() {
let registry = WorkerRegistry::new();
registry
.register(ready_worker(BasicWorkerBuilder::new(
"http://wildcard:8080",
)))
.unwrap();

// Sanity: not in the per-model index for this arbitrary model.
assert!(registry.get_by_model("totally-made-up-model").is_empty());

let client = client();
let selector = WorkerSelector::new(&registry, &client);
let chosen = selector
.select_worker(&selector_request("totally-made-up-model"))
.await
.expect("wildcard worker should serve any model");
assert_eq!(chosen.url(), "http://wildcard:8080");
}

/// Per-model bounding returns the right worker for a normal model and
/// excludes workers serving other models.
#[tokio::test]
async fn bounded_path_selects_correct_model_and_excludes_others() {
let registry = WorkerRegistry::new();
registry
.register(ready_worker(
BasicWorkerBuilder::new("http://a:8080").model(ModelCard::new("model-a")),
))
.unwrap();
registry
.register(ready_worker(
BasicWorkerBuilder::new("http://b:8080").model(ModelCard::new("model-b")),
))
.unwrap();

let client = client();
let selector = WorkerSelector::new(&registry, &client);

let chosen = selector
.select_worker(&selector_request("model-a"))
.await
.expect("model-a worker exists");
assert_eq!(
chosen.url(),
"http://a:8080",
"must not pick model-b worker"
);
}

/// Empty-bounded-set fallback: a worker reachable only by an *alias*
/// (not its indexed id, and not a wildcard) is found via the full-scan
/// fallback so we never regress a servable model into model_not_found.
#[tokio::test]
async fn empty_bounded_set_falls_back_to_full_scan() {
let registry = WorkerRegistry::new();
// Indexed under id "gpt-4"; supports "gpt-4-latest" only via alias.
registry
.register(ready_worker(
BasicWorkerBuilder::new("http://aliased:8080")
.model(ModelCard::new("gpt-4").with_alias("gpt-4-latest")),
))
.unwrap();

// The bounded lookup keys on the literal id and finds nothing for the
// alias (no wildcard workers either) — proving the fallback is what
// surfaces the worker.
assert!(
registry.get_candidates_for_model("gpt-4-latest").is_empty(),
"alias is not an index key, so the bounded set is empty"
);

let client = client();
let selector = WorkerSelector::new(&registry, &client);
let chosen = selector
.select_worker(&selector_request("gpt-4-latest"))
.await
.expect("alias-only worker must be found via fallback");
assert_eq!(chosen.url(), "http://aliased:8080");
}

/// A genuinely unknown model with no wildcard workers yields
/// model_not_found (the fallback does not invent workers).
#[tokio::test]
async fn unknown_model_without_wildcard_is_not_found() {
let registry = WorkerRegistry::new();
registry
.register(ready_worker(
BasicWorkerBuilder::new("http://a:8080").model(ModelCard::new("model-a")),
))
.unwrap();

let client = client();
let selector = WorkerSelector::new(&registry, &client);
let result = selector
.select_worker(&selector_request("nonexistent"))
.await;
assert!(result.is_err(), "no worker and no wildcard → error");
}

/// With UNKNOWN_MODEL_ID the common/HTTP path takes the full-scan branch
/// (no per-model index lookup). A wildcard worker — which `supports_model`
/// accepts for the sentinel — is selectable, matching the pre-bounding
/// behavior. (A *specific* worker does not `supports_model` the sentinel,
/// so this path only resolves to wildcards, unchanged by this PR.)
#[tokio::test]
async fn unknown_model_id_selects_wildcard_via_full_scan() {
let registry = WorkerRegistry::new();
registry
.register(ready_worker(BasicWorkerBuilder::new(
"http://wildcard:8080",
)))
.unwrap();

let client = client();
let selector = WorkerSelector::new(&registry, &client);
let chosen = selector
.select_worker(&selector_request(UNKNOWN_MODEL_ID))
.await
.expect("wildcard worker for unknown model id");
assert_eq!(chosen.url(), "http://wildcard:8080");
}

/// The bounded path still honors the worker_type filter.
#[tokio::test]
async fn bounded_path_honors_worker_type_filter() {
let registry = WorkerRegistry::new();
registry
.register(ready_worker(
BasicWorkerBuilder::new("http://regular:8080")
.model(ModelCard::new("m"))
.worker_type(WorkerType::Regular),
))
.unwrap();
registry
.register(ready_worker(
BasicWorkerBuilder::new("http://prefill:8080")
.model(ModelCard::new("m"))
.worker_type(WorkerType::Prefill),
))
.unwrap();

let client = client();
let selector = WorkerSelector::new(&registry, &client);
let chosen = selector
.select_worker(&SelectWorkerRequest {
model_id: "m",
worker_type: Some(WorkerType::Prefill),
..Default::default()
})
.await
.expect("prefill worker exists for model m");
assert_eq!(chosen.url(), "http://prefill:8080");
}
}
Loading
Loading