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
216 changes: 209 additions & 7 deletions crates/protocols/src/models.rs
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,12 @@ use serde_json::Value;

use super::{model_card::ModelCard, worker::ProviderType};

#[derive(Debug, Clone)]
struct UpstreamModelInfo {
id: String,
created: Option<u64>,
}

/// A single model entry in the `/v1/models` response.
#[derive(Debug, Clone, Serialize, Deserialize, schemars::JsonSchema)]
pub struct ModelObject {
Expand Down Expand Up @@ -58,13 +64,209 @@ impl ListModelsResponse {
let Some(data) = json.get("data").and_then(|d| d.as_array()) else {
return Vec::new();
};
data.iter()
.filter_map(|m| m.get("id").and_then(|id| id.as_str()))
.map(|id| {
let mut card = ModelCard::new(id);
card.provider.clone_from(&provider);
card
let models: Vec<_> = data
.iter()
.filter_map(|model| {
let id = model.get("id").and_then(|id| id.as_str())?;
let created = model.get("created").and_then(Value::as_u64);
Some(UpstreamModelInfo {
id: id.to_string(),
created,
})
})
.collect();

group_upstream_models_into_cards(models, provider)
}
}

fn is_ascii_digits(value: &str, len: usize) -> bool {
value.len() == len && value.bytes().all(|b| b.is_ascii_digit())
}

fn strip_date_suffix(id: &str) -> Option<String> {
let parts: Vec<_> = id.split('-').collect();

if parts.len() >= 4
&& is_ascii_digits(parts[parts.len() - 3], 4)
&& is_ascii_digits(parts[parts.len() - 2], 2)
&& is_ascii_digits(parts[parts.len() - 1], 2)
{
return Some(parts[..parts.len() - 3].join("-"));
}

if parts.len() >= 3
&& is_ascii_digits(parts[parts.len() - 2], 4)
&& is_ascii_digits(parts[parts.len() - 1], 2)
{
return Some(parts[..parts.len() - 2].join("-"));
}

None
}

fn xai_group_key_and_rank(id: &str) -> Option<(String, u16)> {
let parts: Vec<_> = id.split('-').collect();
if parts.len() == 3
&& parts[0] == "grok"
&& !parts[1].is_empty()
&& parts[1].bytes().all(|b| b.is_ascii_digit())
&& is_ascii_digits(parts[2], 4)
{
return parts[2]
.parse::<u16>()
.ok()
.map(|rank| (format!("{}-{}", parts[0], parts[1]), rank));
}

None
}

fn alias_group_key(id: &str) -> String {
if let Some((group_key, _)) = xai_group_key_and_rank(id) {
return group_key;
}

strip_date_suffix(id).unwrap_or_else(|| id.to_string())
}

fn xai_revision_rank(id: &str) -> Option<u16> {
xai_group_key_and_rank(id).map(|(_, rank)| rank)
}

fn select_primary_and_aliases(
group_key: &str,
variants: &[UpstreamModelInfo],
) -> (String, Vec<String>) {
let primary_id = if variants.iter().any(|variant| variant.id == group_key) {
group_key.to_string()
} else if variants
.iter()
.all(|variant| xai_revision_rank(&variant.id).is_some())
{
variants
.iter()
.max_by_key(|variant| {
(
variant.created.unwrap_or(0),
xai_revision_rank(&variant.id).unwrap_or(0),
variant.id.as_str(),
)
})
.collect()
.map(|variant| variant.id.clone())
.unwrap_or_else(|| group_key.to_string())
} else {
variants
.iter()
.map(|variant| variant.id.as_str())
.min_by(|a, b| a.len().cmp(&b.len()).then_with(|| a.cmp(b)))
.map(ToOwned::to_owned)
.unwrap_or_else(|| group_key.to_string())
};

let mut aliases: Vec<String> = variants
.iter()
.filter_map(|variant| (variant.id != primary_id).then(|| variant.id.clone()))
.collect();

if group_key != primary_id && !aliases.iter().any(|alias| alias == group_key) {
aliases.push(group_key.to_string());
}

aliases.sort();
aliases.dedup();

(primary_id, aliases)
}

fn group_upstream_models_into_cards(
models: Vec<UpstreamModelInfo>,
provider: Option<ProviderType>,
) -> Vec<ModelCard> {
let mut groups = std::collections::BTreeMap::<String, Vec<UpstreamModelInfo>>::new();
for model in models {
groups
.entry(alias_group_key(&model.id))
.or_default()
.push(model);
}

groups
.into_iter()
.map(|(group_key, variants)| {
let (primary_id, aliases) = select_primary_and_aliases(&group_key, &variants);
let mut card = ModelCard::new(primary_id).with_aliases(aliases);
card.provider.clone_from(&provider);
card
})
.collect()
}

#[cfg(test)]
mod tests {
use serde_json::json;

use super::ListModelsResponse;
use crate::worker::ProviderType;

#[test]
fn parse_upstream_groups_openai_date_variants_under_stable_name() {
let json = json!({
"object": "list",
"data": [
{"id": "gpt-4o", "object": "model"},
{"id": "gpt-4o-2024-08-06", "object": "model"},
{"id": "gpt-4o-2024-11-20", "object": "model"}
]
});

let cards = ListModelsResponse::parse_upstream(&json, Some(ProviderType::OpenAI));

assert_eq!(cards.len(), 1);
assert_eq!(cards[0].id, "gpt-4o");
assert!(cards[0]
.aliases
.iter()
.any(|alias| alias == "gpt-4o-2024-08-06"));
assert!(cards[0]
.aliases
.iter()
.any(|alias| alias == "gpt-4o-2024-11-20"));
assert_eq!(cards[0].provider, Some(ProviderType::OpenAI));
}

#[test]
fn parse_upstream_adds_xai_family_alias_for_revisioned_model() {
let json = json!({
"object": "list",
"data": [
{"id": "grok-4-0709", "object": "model", "created": 100}
]
});

let cards = ListModelsResponse::parse_upstream(&json, Some(ProviderType::XAI));

assert_eq!(cards.len(), 1);
assert_eq!(cards[0].id, "grok-4-0709");
assert!(cards[0].aliases.iter().any(|alias| alias == "grok-4"));
assert_eq!(cards[0].provider, Some(ProviderType::XAI));
}

#[test]
fn parse_upstream_picks_latest_xai_revision_as_primary_model() {
let json = json!({
"object": "list",
"data": [
{"id": "grok-4-0709", "object": "model", "created": 100},
{"id": "grok-4-0812", "object": "model", "created": 200}
]
});

let cards = ListModelsResponse::parse_upstream(&json, Some(ProviderType::XAI));

assert_eq!(cards.len(), 1);
assert_eq!(cards[0].id, "grok-4-0812");
assert!(cards[0].aliases.iter().any(|alias| alias == "grok-4"));
assert!(cards[0].aliases.iter().any(|alias| alias == "grok-4-0709"));
}
}
11 changes: 6 additions & 5 deletions model_gateway/src/routers/openai/chat.rs
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@ use tokio_stream::wrappers::UnboundedReceiverStream;
use super::{
context::{ComponentRefs, PayloadState, RequestContext, SharedComponents, WorkerSelection},
provider::ProviderRegistry,
router::resolve_provider,
router::{request_provider, resolve_effective_model, resolve_provider},
};
use crate::{
config::types::RetryConfig,
Expand All @@ -26,7 +26,7 @@ use crate::{
header_utils::{apply_provider_headers, extract_auth_header},
worker_selection::{SelectWorkerRequest, WorkerSelector},
},
worker::{is_retryable_status, Endpoint, ProviderType, RetryExecutor, WorkerRegistry},
worker::{is_retryable_status, Endpoint, RetryExecutor, WorkerRegistry},
};

/// Shared context passed to chat routing functions.
Expand Down Expand Up @@ -62,7 +62,7 @@ pub(super) async fn route_chat(
.select_worker(&SelectWorkerRequest {
model_id: model,
headers,
provider: Some(ProviderType::OpenAI),
provider: Some(request_provider(model)),
..Default::default()
})
.await
Expand All @@ -80,6 +80,7 @@ pub(super) async fn route_chat(
return response;
}
};
let effective_model = resolve_effective_model(worker.as_ref(), model);

let mut payload = match to_value(body) {
Ok(v) => v,
Expand All @@ -100,9 +101,9 @@ pub(super) async fn route_chat(
};

// Patch the serialized payload to use the effective model consistently.
payload["model"] = serde_json::Value::String(model.to_owned());
payload["model"] = serde_json::Value::String(effective_model.clone());

let provider = resolve_provider(deps.provider_registry, worker.as_ref(), model);
let provider = resolve_provider(deps.provider_registry, worker.as_ref(), &effective_model);
if let Err(e) = provider.transform_request(&mut payload, Endpoint::Chat) {
Metrics::record_router_error(
metrics_labels::ROUTER_OPENAI,
Expand Down
11 changes: 6 additions & 5 deletions model_gateway/src/routers/openai/responses/route.rs
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@ use super::{
ComponentRefs, PayloadState, RequestContext, ResponsesComponents, WorkerSelection,
},
provider::ProviderRegistry,
router::resolve_provider,
router::{request_provider, resolve_effective_model, resolve_provider},
},
handle_non_streaming_response, handle_streaming_response,
};
Expand All @@ -26,7 +26,7 @@ use crate::{
error,
worker_selection::{SelectWorkerRequest, WorkerSelector},
},
worker::{Endpoint, ProviderType, WorkerRegistry},
worker::{Endpoint, WorkerRegistry},
};

/// Shared context passed to responses routing functions.
Expand Down Expand Up @@ -63,7 +63,7 @@ pub(in crate::routers::openai) async fn route_responses(
.select_worker(&SelectWorkerRequest {
model_id: model,
headers,
provider: Some(ProviderType::OpenAI),
provider: Some(request_provider(model)),
..Default::default()
Comment on lines 63 to 67

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

This provider hint is being applied too early.

Like the chat path, this filters on request_provider(model) before alias resolution. If the client sends an alias that only exists in discovered model metadata, the selector can exclude the real worker before resolve_effective_model(...) ever runs.

🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.

In `@model_gateway/src/routers/openai/responses/route.rs` around lines 63 - 67,
The provider hint is being applied too early on the SelectWorkerRequest causing
premature filtering; remove or omit the provider: Some(request_provider(model))
from the .select_worker(...) call (the SelectWorkerRequest construction) so
alias-based models in discovered metadata aren't excluded before alias
resolution, and instead pass the provider hint only after
resolve_effective_model(...) has run and returned the resolved model/provider.

})
.await
Expand All @@ -81,6 +81,7 @@ pub(in crate::routers::openai) async fn route_responses(
return response;
}
};
let effective_model = resolve_effective_model(worker.as_ref(), model);

// Validate mutual exclusivity of conversation and previous_response_id
// Treat empty strings as unset to match other metadata paths
Expand All @@ -105,7 +106,7 @@ pub(in crate::routers::openai) async fn route_responses(
}

let mut request_body = body.clone();
request_body.model = model_id.to_string();
request_body.model = effective_model.clone();
request_body.conversation = None;

let original_previous_response_id = match super::history::load_input_history(
Expand Down Expand Up @@ -143,7 +144,7 @@ pub(in crate::routers::openai) async fn route_responses(
}
};

let provider = resolve_provider(deps.provider_registry, worker.as_ref(), model);
let provider = resolve_provider(deps.provider_registry, worker.as_ref(), &effective_model);
if let Err(e) = provider.transform_request(&mut payload, Endpoint::Responses) {
Metrics::record_router_error(
metrics_labels::ROUTER_OPENAI,
Expand Down
24 changes: 23 additions & 1 deletion model_gateway/src/routers/openai/router.rs
Original file line number Diff line number Diff line change
Expand Up @@ -36,11 +36,33 @@ use crate::{
///
/// Checks (in order): worker's per-model provider, model name heuristic,
/// then falls back to the default provider.
pub(super) fn request_provider(model: &str) -> ProviderType {
ProviderType::from_model_name(model).unwrap_or(ProviderType::OpenAI)
}

/// Resolve the canonical upstream model ID from the worker's live model view.
pub(super) fn resolve_effective_model(worker: &dyn Worker, model: &str) -> String {
worker
.models()
.into_iter()
.find(|candidate| candidate.matches(model))
.map(|candidate| candidate.id)
Comment on lines +45 to +49

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

high

Calling worker.models().into_iter().find(...) on every request is a major performance bottleneck. worker.models() clones the entire list of model cards, which can be large for external providers. This overhead will significantly impact request latency. The Worker trait should be extended with methods to resolve canonical IDs and providers efficiently to avoid full list clones on the hot path. This optimization should be extracted into a shared mechanism to ensure consistency across all routing paths.

References
  1. If an optimization is applicable to multiple code paths, extract it into a shared helper function to ensure consistency and avoid code duplication.

.unwrap_or_else(|| worker.canonical_model_id(model).to_string())
Comment on lines +43 to +50

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

Canonicalize realtime model IDs before proxying.

This helper only fixes the call sites that use it. In this file, the realtime REST/WS paths still pass the raw model into forward_realtime_rest / handle_realtime_ws (Lines 206-214, 226-234, 245-253, and 272-278), so grok-4 can now select the xAI worker but still gets forwarded upstream as grok-4 instead of the worker’s canonical ID. That leaves alias routing broken for the realtime endpoints.

🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.

In `@model_gateway/src/routers/openai/router.rs` around lines 43 - 50, The
realtime REST/WS handlers still pass the raw model name into upstream proxies,
so alias routing breaks; update each call site that forwards realtime requests
(calls to forward_realtime_rest and handle_realtime_ws) to first canonicalize
the model using resolve_effective_model(worker, model) and pass that returned
String into the proxy functions instead of the original model variable; ensure
both REST and WS branches (the locations that currently call
forward_realtime_rest(...) and handle_realtime_ws(...)) are changed so the
canonical ID is used for upstream forwarding.

}

pub(super) fn resolve_provider(
registry: &ProviderRegistry,
worker: &dyn Worker,
model: &str,
) -> Arc<dyn super::provider::Provider> {
if let Some(pt) = worker
.models()
.into_iter()
.find(|candidate| candidate.matches(model))
.and_then(|candidate| candidate.provider)
{
Comment on lines +58 to +63

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

high

Similar to resolve_effective_model, this block calls worker.models(), causing expensive full-list clones on every request. This logic should be moved into a more efficient method on the Worker trait or BasicWorker implementation that can query the provider for a model without cloning the entire model list. Applying this optimization consistently across code paths aligns with repository guidelines.

References
  1. If an optimization is applicable to multiple code paths, extract it into a shared helper function to ensure consistency and avoid code duplication.

return registry.get_arc(&pt);
}
if let Some(pt) = worker.provider_for_model(model) {
return registry.get_arc(pt);
}
Expand Down Expand Up @@ -117,7 +139,7 @@ impl OpenAIRouter {
.select_worker(&SelectWorkerRequest {
model_id,
headers,
provider: Some(ProviderType::OpenAI),
provider: Some(request_provider(model_id)),
..Default::default()
})
.await
Expand Down
Loading