use axum::body::Body;
use axum::extract::{Request, State};
use axum::http::{HeaderMap, StatusCode};
use axum::response::{IntoResponse, Response};
use serde_json::{Value, json};
use crate::app_state::AppState;
use crate::config::UpstreamProvider;
use crate::model_catalog::ModelCatalogCache;
use crate::subscription::{SubscriptionProvider, SubscriptionReader};
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ModelRouteError {
ModelRequired,
NotFound(String),
Ambiguous(String),
}
#[path = "model_routing_snapshot.rs"]
pub(crate) mod snapshot;
pub(crate) use snapshot::{
RoutedState, ValidatedSubscription, route_pinned_subscription, route_subscription_model,
};
#[path = "model_routing_catalog_snapshot.rs"]
mod catalog_snapshot;
pub(crate) use catalog_snapshot::ConfiguredCatalogSnapshot;
impl std::fmt::Display for ModelRouteError {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::ModelRequired => {
formatter.write_str("model is required when UPSTREAM_PROVIDER=auto")
}
Self::NotFound(message) | Self::Ambiguous(message) => formatter.write_str(message),
}
}
}
pub(crate) fn model_route_error_response(error: &ModelRouteError) -> Response {
let (status, error_type) = match error {
ModelRouteError::ModelRequired | ModelRouteError::Ambiguous(_) => {
(StatusCode::BAD_REQUEST, "invalid_request_error")
}
ModelRouteError::NotFound(_) => (StatusCode::NOT_FOUND, "not_found_error"),
};
crate::proxy::error_response(status, error_type, &error.to_string())
}
pub(crate) fn model_not_found_response(model: &str, catalog: &[String]) -> Response {
let detail = if catalog.is_empty() {
String::new()
} else {
format!("; this deployment advertises: {}", advertised_list(catalog))
};
model_route_error_response(&ModelRouteError::NotFound(format!(
"model '{model}' is not available{detail}"
)))
}
const fn provider_owner(provider: SubscriptionProvider) -> &'static str {
match provider {
SubscriptionProvider::Claude => "anthropic",
SubscriptionProvider::Codex => "openai",
SubscriptionProvider::Gemini => "google",
SubscriptionProvider::Qwen => "qwen",
}
}
fn provider_hint(model: &str) -> Option<SubscriptionProvider> {
if model.starts_with("claude-") {
Some(SubscriptionProvider::Claude)
} else if model.starts_with("gpt-")
|| model.starts_with("codex-")
|| model
.strip_prefix('o')
.and_then(|suffix| suffix.chars().next())
.is_some_and(|character| character.is_ascii_digit())
{
Some(SubscriptionProvider::Codex)
} else if model.starts_with("gemini-") {
Some(SubscriptionProvider::Gemini)
} else if model.starts_with("qwen-") {
Some(SubscriptionProvider::Qwen)
} else {
None
}
}
fn providers_for_model(model: &str, catalogs: &ModelCatalogCache) -> Vec<SubscriptionProvider> {
SubscriptionProvider::ALL
.into_iter()
.filter(|provider| catalogs.models(*provider).iter().any(|id| id == model))
.collect()
}
#[must_use]
pub fn provider_for_model(
model: &str,
catalogs: &ModelCatalogCache,
) -> Option<SubscriptionProvider> {
let providers = providers_for_model(model, catalogs);
if providers.len() == 1 {
return providers.first().copied();
}
provider_hint(model).filter(|provider| providers.contains(provider))
}
fn credential_state(
provider: SubscriptionProvider,
catalogs: &ModelCatalogCache,
) -> Option<String> {
if !catalogs.provider_is_degraded(provider) {
return None;
}
let status = catalogs.status(provider);
Some(match (status.discovered, status.last_error) {
(true, _) => format!(
"the {provider} catalog is retained for diagnostics but its credential is not usable"
),
(false, Some(_)) => {
format!("{provider} has never completed a live catalog discovery")
}
(false, None) => format!("no {provider} credential has been read yet"),
})
}
fn credential_states(model: &str, catalogs: &ModelCatalogCache) -> Vec<String> {
provider_hint(model).map_or_else(
|| {
SubscriptionProvider::ALL
.into_iter()
.filter(|provider| catalogs.provider_has_observation(*provider))
.filter_map(|provider| credential_state(provider, catalogs))
.collect()
},
|provider| credential_state(provider, catalogs).into_iter().collect(),
)
}
const ADVERTISED_IN_ERRORS: usize = 24;
fn advertised_detail(available: &[SubscriptionProvider], catalogs: &ModelCatalogCache) -> String {
let mut ids = available
.iter()
.flat_map(|provider| catalogs.models(*provider))
.collect::<Vec<_>>();
ids.sort_unstable();
ids.dedup();
if ids.is_empty() {
return String::new();
}
format!("; this deployment advertises: {}", advertised_list(&ids))
}
fn advertised_list(ids: &[String]) -> String {
if ids.len() > ADVERTISED_IN_ERRORS {
let shown = ids[..ADVERTISED_IN_ERRORS].join(", ");
let rest = ids.len() - ADVERTISED_IN_ERRORS;
return format!("{shown} and {rest} more");
}
ids.join(", ")
}
pub fn available_provider_for_model(
model: &str,
available: &[SubscriptionProvider],
catalogs: &ModelCatalogCache,
) -> Result<SubscriptionProvider, ModelRouteError> {
let advertised = providers_for_model(model, catalogs);
if advertised.is_empty() {
let causes = credential_states(model, catalogs);
let detail = if causes.is_empty() {
advertised_detail(available, catalogs)
} else {
format!(": {}", causes.join("; "))
};
return Err(ModelRouteError::NotFound(format!(
"model '{model}' is not advertised by any subscription{detail}"
)));
}
let provider = provider_hint(model)
.filter(|provider| advertised.contains(provider))
.or_else(|| {
let healthy = advertised
.iter()
.copied()
.filter(|provider| available.contains(provider))
.collect::<Vec<_>>();
(healthy.len() == 1).then(|| healthy[0])
})
.or_else(|| (advertised.len() == 1).then(|| advertised[0]))
.ok_or_else(|| {
let providers = advertised
.iter()
.map(ToString::to_string)
.collect::<Vec<_>>()
.join(", ");
ModelRouteError::Ambiguous(format!(
"model '{model}' is advertised by multiple subscriptions ({providers}); pin \
UPSTREAM_PROVIDER to disambiguate"
))
})?;
available
.contains(&provider)
.then_some(provider)
.ok_or_else(|| {
let cause = credential_state(provider, catalogs).unwrap_or_else(|| {
format!(
"the last credential check found no usable {provider} credential (missing or \
rejected upstream)"
)
});
ModelRouteError::NotFound(format!(
"model '{model}' has no healthy {provider} credential: {cause}"
))
})
}
pub async fn healthy_providers(
client: &reqwest::Client,
readers: &[SubscriptionReader],
token_cache: &crate::refresh::TokenCache,
now_ms: i64,
) -> Vec<SubscriptionProvider> {
token_cache.register_readers(crate::credential_recovery_store::PRIMARY_ACCOUNT, readers);
let checks = SubscriptionProvider::ALL
.into_iter()
.map(|provider| async move {
readers
.iter()
.find(|reader| reader.provider() == provider)?;
let token = token_cache
.get_fresh_registered(
client,
provider,
crate::credential_recovery_store::PRIMARY_ACCOUNT,
now_ms,
)
.await
.ok()?;
if token_cache.evidence(provider) == Some(crate::refresh::CredentialEvidence::Rejected)
{
tracing::debug!("{provider} credential was rejected upstream; not routable");
return None;
}
if !token.is_expired(now_ms) {
return Some(provider);
}
tracing::debug!(
"{provider} credential is stamped expired and could not be refreshed; keeping it \
routable until an upstream rejects it"
);
Some(provider)
});
futures_util::future::join_all(checks)
.await
.into_iter()
.flatten()
.collect()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ProviderHealthState {
Starting,
Healthy,
Degraded,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ProviderHealth {
pub provider: SubscriptionProvider,
pub healthy: bool,
pub reason: Option<String>,
pub summary: Option<&'static str>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct ProviderHealthReport {
pub provider: SubscriptionProvider,
pub state: ProviderHealthState,
pub reason: Option<String>,
pub summary: Option<&'static str>,
}
impl ProviderHealthReport {
#[must_use]
pub const fn is_degraded(&self) -> bool {
matches!(self.state, ProviderHealthState::Degraded)
}
#[must_use]
pub const fn is_serving(&self) -> bool {
!self.is_degraded()
}
}
#[must_use]
pub fn configured_provider_health(
readers: &[SubscriptionReader],
token_cache: &crate::refresh::TokenCache,
catalogs: &ModelCatalogCache,
) -> Vec<ProviderHealth> {
SubscriptionProvider::ALL
.into_iter()
.filter(|provider| readers.iter().any(|reader| reader.provider() == *provider))
.map(|provider| {
let rejected = token_cache.evidence(provider)
== Some(crate::refresh::CredentialEvidence::Rejected);
let status = catalogs.status(provider);
let reason =
if rejected {
Some(token_cache.last_refresh_error(provider).unwrap_or_else(|| {
format!("the {provider} credential was rejected upstream")
}))
} else if status.discovered && !status.credential_healthy {
credential_state(provider, catalogs)
} else {
None
};
ProviderHealth {
provider,
healthy: reason.is_none(),
reason,
summary: rejected
.then_some("the credential was rejected upstream and needs re-authentication"),
}
})
.collect()
}
fn account_health(
provider: SubscriptionProvider,
account: &str,
credential: Result<Option<crate::subscription::SubscriptionToken>, String>,
token_cache: &crate::refresh::TokenCache,
catalogs: &ModelCatalogCache,
) -> Option<ProviderHealthReport> {
let token = match credential {
Ok(Some(token)) => token,
Ok(None) => return None,
Err(error) => {
return Some(ProviderHealthReport {
provider,
state: ProviderHealthState::Degraded,
reason: Some(error),
summary: Some("the configured credential could not be read"),
});
}
};
let rejected = token_cache.evidence_for(provider, account)
== Some(crate::refresh::CredentialEvidence::Rejected);
let status = catalogs.status_for(provider, account);
let account_matches = match (token.account_id.as_deref(), status.account.as_deref()) {
(Some(current), Some(discovered)) => current == discovered,
(None, None) => true,
_ => false,
};
let (state, reason, summary) = if rejected {
(
ProviderHealthState::Degraded,
Some(
token_cache
.last_refresh_error_for(provider, account)
.unwrap_or_else(|| format!("the {provider} credential was rejected upstream")),
),
Some("the credential was rejected upstream and needs re-authentication"),
)
} else if !account_matches {
(ProviderHealthState::Starting, None, None)
} else if status.discovered && !status.credential_healthy {
(
ProviderHealthState::Degraded,
Some(format!("the {provider} catalog credential is not usable")),
Some("the credential was rejected upstream and needs re-authentication"),
)
} else if status.discovered {
(ProviderHealthState::Healthy, None, None)
} else {
(ProviderHealthState::Starting, None, None)
};
Some(ProviderHealthReport {
provider,
state,
reason,
summary,
})
}
async fn configured_account_health_report(state: &AppState) -> Vec<(String, ProviderHealthReport)> {
let checks = SubscriptionProvider::ALL
.into_iter()
.map(|provider| async move {
let accounts = state
.account_router
.as_ref()
.filter(|router| router.provider() == provider)
.map_or_else(
|| {
state
.subscription_readers
.iter()
.find(|reader| reader.provider() == provider)
.map(|reader| {
state.subscription_cache.register_reader(
crate::credential_recovery_store::PRIMARY_ACCOUNT,
reader,
);
vec![crate::credential_recovery_store::PRIMARY_ACCOUNT.to_string()]
})
.unwrap_or_default()
},
|router| {
router
.subscription_readers()
.into_iter()
.map(|(account, _)| account)
.collect()
},
);
let mut reports = Vec::new();
for account in accounts {
let credential = state
.subscription_cache
.load_authoritative(provider, &account)
.await;
if let Some(report) = account_health(
provider,
&account,
credential,
&state.subscription_cache,
&state.model_catalogs,
) {
reports.push((account, report));
}
}
reports
});
futures_util::future::join_all(checks)
.await
.into_iter()
.flatten()
.collect()
}
fn aggregate_provider_health(
accounts: &[(String, ProviderHealthReport)],
) -> Vec<ProviderHealthReport> {
SubscriptionProvider::ALL
.into_iter()
.filter_map(|provider| {
let reports = accounts
.iter()
.map(|(_, report)| report)
.filter(|report| report.provider == provider)
.collect::<Vec<_>>();
if reports.is_empty() {
return None;
}
let state = if reports
.iter()
.any(|report| report.state == ProviderHealthState::Healthy)
{
ProviderHealthState::Healthy
} else if reports
.iter()
.any(|report| report.state == ProviderHealthState::Starting)
{
ProviderHealthState::Starting
} else {
ProviderHealthState::Degraded
};
let degraded = reports
.iter()
.find(|report| report.state == ProviderHealthState::Degraded);
Some(ProviderHealthReport {
provider,
state,
reason: degraded.and_then(|report| report.reason.clone()),
summary: degraded.and_then(|report| report.summary),
})
})
.collect()
}
pub(crate) async fn configured_provider_health_report(
state: &AppState,
) -> Vec<ProviderHealthReport> {
aggregate_provider_health(&configured_account_health_report(state).await)
}
pub(crate) async fn configured_catalog_snapshot(state: &AppState) -> ConfiguredCatalogSnapshot {
let accounts = configured_account_health_report(state).await;
let models = SubscriptionProvider::ALL
.into_iter()
.map(|provider| {
let healthy_accounts = accounts
.iter()
.filter(|(_, report)| {
report.provider == provider && report.state == ProviderHealthState::Healthy
})
.map(|(account, _)| account.clone())
.collect::<Vec<_>>();
(
provider,
state
.model_catalogs
.models_for_accounts(provider, &healthy_accounts),
)
})
.collect();
ConfiguredCatalogSnapshot {
health: aggregate_provider_health(&accounts),
models,
}
}
#[must_use]
pub fn model_catalog(providers: &[SubscriptionProvider], catalogs: &ModelCatalogCache) -> Value {
model_catalog_with(providers, catalogs, |provider| catalogs.models(provider))
}
fn model_catalog_with(
providers: &[SubscriptionProvider],
catalogs: &ModelCatalogCache,
models: impl Fn(SubscriptionProvider) -> Vec<String>,
) -> Value {
let now = chrono::Utc::now().timestamp();
let degraded = providers
.iter()
.filter(|provider| catalogs.provider_is_degraded(**provider))
.map(|provider| provider.as_str())
.collect::<Vec<_>>();
let healthy_providers = providers
.iter()
.filter(|provider| !catalogs.provider_is_degraded(**provider))
.map(|provider| provider.as_str())
.collect::<Vec<_>>();
let data = providers
.iter()
.flat_map(|provider| {
let owner = provider_owner(*provider);
models(*provider).into_iter().map(move |id| {
json!({
"id": id,
"object": "model",
"created": now,
"owned_by": owner,
})
})
})
.collect::<Vec<_>>();
json!({
"object": "list",
"data": data,
"using_fallback": false,
"degraded_providers": degraded,
"healthy_providers": healthy_providers,
})
}
#[must_use]
pub async fn pinned_model_catalog(state: &AppState, provider: SubscriptionProvider) -> Value {
let snapshot = configured_catalog_snapshot(state).await;
let health = snapshot
.health()
.iter()
.filter(|entry| entry.provider == provider)
.cloned()
.collect::<Vec<_>>();
let mut catalog = if health
.first()
.is_some_and(|entry| entry.state == ProviderHealthState::Healthy)
{
model_catalog_with(&[provider], &state.model_catalogs, |provider| {
snapshot.models(provider)
})
} else {
model_catalog(&[], &state.model_catalogs)
};
merge_configured_degradation(&health, &mut catalog);
catalog
}
fn merge_configured_degradation(health: &[ProviderHealthReport], catalog: &mut Value) {
let mut degraded = catalog
.get("degraded_providers")
.and_then(Value::as_array)
.cloned()
.unwrap_or_default();
let mut reasons = serde_json::Map::new();
for entry in health.iter().filter(|entry| entry.is_degraded()) {
let name = Value::from(entry.provider.as_str());
if !degraded.contains(&name) {
degraded.push(name);
}
if let Some(summary) = entry.summary {
reasons.insert(entry.provider.as_str().to_string(), Value::from(summary));
}
}
if let Some(object) = catalog.as_object_mut() {
object.insert("degraded_providers".into(), Value::Array(degraded));
object.insert("degraded_reasons".into(), Value::Object(reasons));
}
}
pub async fn models(State(state): State<AppState>, headers: HeaderMap) -> Response {
if let Err(response) = crate::proxy::authenticate_client(&state, &headers) {
return *response;
}
let models = match state.upstream_provider {
UpstreamProvider::Auto => {
let snapshot = configured_catalog_snapshot(&state).await;
let healthy = snapshot.healthy_providers();
let mut catalog = model_catalog_with(&healthy, &state.model_catalogs, |provider| {
snapshot.models(provider)
});
merge_configured_degradation(snapshot.health(), &mut catalog);
append_stored_provider_models(&state, &mut catalog);
catalog
}
UpstreamProvider::Anthropic => {
pinned_model_catalog(&state, SubscriptionProvider::Claude).await
}
UpstreamProvider::Gonka => state.gonka.as_ref().map_or_else(
|| crate::gonka::list_models(&crate::config::default_gonka_model()),
|gonka| crate::gonka::list_models(&gonka.model),
),
UpstreamProvider::Crater => crate::crater::list_models(),
UpstreamProvider::Codex => pinned_model_catalog(&state, SubscriptionProvider::Codex).await,
UpstreamProvider::Qwen => pinned_model_catalog(&state, SubscriptionProvider::Qwen).await,
UpstreamProvider::Gemini => {
pinned_model_catalog(&state, SubscriptionProvider::Gemini).await
}
UpstreamProvider::OpenAICompatible => {
crate::provider_proxy::openai_compatible_models(&state)
}
};
(StatusCode::OK, axum::Json(models)).into_response()
}
pub async fn route_anthropic_request(
state: &AppState,
request: Request,
) -> Result<(AppState, Request), Response> {
route_anthropic_request_with_subscription(state, request)
.await
.map(|(routed, request)| (routed.state, request))
}
pub(crate) async fn route_anthropic_request_with_subscription(
state: &AppState,
request: Request,
) -> Result<(RoutedState, Request), Response> {
let path = request.uri().path().to_string();
let (parts, body) = request.into_parts();
let body_bytes = axum::body::to_bytes(body, state.max_proxy_request_bytes)
.await
.map_err(|error| {
crate::proxy::error_response(
StatusCode::PAYLOAD_TOO_LARGE,
"invalid_request_error",
&format!(
"request body exceeds the {} byte proxy limit: {error}",
state.max_proxy_request_bytes
),
)
})?;
let routing_body = serde_json::from_slice(&body_bytes).map_err(|error| {
crate::proxy::error_response(
StatusCode::BAD_REQUEST,
"invalid_request_error",
&format!("Failed to parse request body as JSON: {error}"),
)
})?;
let routed = if path.ends_with("/messages") || path.ends_with("/messages/count_tokens") {
route_state_with_subscription(state, &routing_body)
.await
.map_err(|error| model_route_error_response(&error))?
} else {
route_pinned_subscription(state, SubscriptionProvider::Claude)
.await
.map_err(|error| {
crate::proxy::error_response(
StatusCode::BAD_REQUEST,
"invalid_request_error",
&error.to_string(),
)
})?
};
Ok((routed, Request::from_parts(parts, Body::from(body_bytes))))
}
pub async fn route_provider(
state: &AppState,
provider: SubscriptionProvider,
) -> Result<AppState, String> {
let healthy = healthy_providers(
&state.client,
&state.subscription_readers,
&state.subscription_cache,
chrono::Utc::now().timestamp_millis(),
)
.await;
let reader = state
.subscription_readers
.iter()
.find(|reader| reader.provider() == provider)
.filter(|_| healthy.contains(&provider))
.cloned()
.ok_or_else(|| format!("no healthy {provider} credential is available"))?;
let mut routed = state.clone();
routed.upstream_provider = match provider {
SubscriptionProvider::Claude => UpstreamProvider::Anthropic,
SubscriptionProvider::Codex => UpstreamProvider::Codex,
SubscriptionProvider::Gemini => UpstreamProvider::Gemini,
SubscriptionProvider::Qwen => UpstreamProvider::Qwen,
};
if provider != SubscriptionProvider::Claude {
routed.account_router = None;
routed.subscription_reader = Some(reader);
}
Ok(routed)
}
fn append_stored_provider_models(state: &AppState, catalog: &mut Value) {
let Ok(providers) = state.provider_store.list() else {
return;
};
let Some(data) = catalog.get_mut("data").and_then(Value::as_array_mut) else {
return;
};
for provider in providers.iter().filter(|record| record.enabled) {
for model in &provider.models {
if data
.iter()
.any(|entry| entry.get("id").and_then(Value::as_str) == Some(model.as_str()))
{
data.push(json!({
"id": format!("{}/{}", provider.name, model),
"object": "model",
"owned_by": provider.name,
}));
continue;
}
data.push(json!({
"id": model,
"object": "model",
"owned_by": provider.name,
}));
}
}
}
fn stored_provider_for_model(
state: &AppState,
model: &str,
) -> Result<Option<crate::providers::ResolvedProvider>, ModelRouteError> {
if let Some((name, bare)) = model.split_once('/') {
return match state.provider_store.resolve(name) {
Ok(Some(provider)) if provider.declares(bare) => Ok(Some(provider)),
Ok(Some(_)) => Err(ModelRouteError::NotFound(format!(
"provider '{name}' does not advertise model '{bare}'"
))),
_ => Ok(None),
};
}
let Ok(providers) = state.provider_store.list() else {
return Ok(None);
};
let mut declaring = providers
.into_iter()
.filter(|record| record.enabled && record.models.iter().any(|id| id == model))
.map(|record| record.name);
let Some(first) = declaring.next() else {
return Ok(None);
};
if let Some(second) = declaring.next() {
return Err(ModelRouteError::Ambiguous(format!(
"model '{model}' is declared by multiple stored providers ({first}, {second}); name \
one as '<provider>/{model}' to disambiguate"
)));
}
Ok(state.provider_store.resolve(&first).ok().flatten())
}
fn route_stored_provider(
state: &AppState,
provider: &crate::providers::ResolvedProvider,
model: &str,
) -> AppState {
let mut routed = state.clone();
routed.upstream_provider = UpstreamProvider::OpenAICompatible;
routed
.openai_compatible
.provider_name
.clone_from(&provider.name);
routed.bridge_model = Some(bare_model_id(model).to_string());
routed
}
#[must_use]
pub fn bare_model_id(model: &str) -> &str {
model.split_once('/').map_or(model, |(_, bare)| bare)
}
pub(crate) async fn route_state_with_subscription(
state: &AppState,
body: &Value,
) -> Result<RoutedState, ModelRouteError> {
if state.upstream_provider != UpstreamProvider::Auto {
return Ok(RoutedState {
state: state.clone(),
subscription: None,
});
}
let model = body
.get("model")
.and_then(Value::as_str)
.filter(|model| !model.is_empty())
.ok_or(ModelRouteError::ModelRequired)?;
if let Some(stored) = stored_provider_for_model(state, model)? {
return Ok(RoutedState {
state: route_stored_provider(state, &stored, model),
subscription: None,
});
}
route_subscription_model(state, model).await
}
pub async fn route_state(state: &AppState, body: &Value) -> Result<AppState, ModelRouteError> {
route_state_with_subscription(state, body)
.await
.map(|routed| routed.state)
}
#[cfg(test)]
#[path = "model_routing_tests.rs"]
mod tests;
#[cfg(test)]
#[path = "model_routing_health_tests.rs"]
mod health_tests;
#[cfg(test)]
#[path = "model_routing_snapshot_tests.rs"]
mod snapshot_tests;
#[cfg(test)]
#[path = "model_routing_pool_tests.rs"]
mod pool_tests;
#[cfg(test)]
#[path = "model_routing_recovery_tests.rs"]
mod recovery_tests;
#[cfg(test)]
#[path = "model_routing_evidence_tests.rs"]
mod evidence_tests;