use super::candidate::{
PricingSku, ReadyRouteCandidate, ResolvedAuthSource, ResolvedEndpoint, ValidationReport,
};
use super::descriptor::ProviderDescriptor;
use super::errors::RouteError;
use super::ids::{LogicalModelRef, ModelId, ProviderId, WireModelId};
use super::offering::{ProviderModelOffering, RouteLimits, bundled_offerings};
use crate::ProviderKind;
use crate::catalog::{CatalogOffering, bundled_catalog_offerings};
#[derive(Debug, Clone, Default)]
pub struct RouteRequest {
pub explicit_provider: Option<ProviderKind>,
pub model_selector: Option<LogicalModelRef>,
pub saved_provider_model: Option<WireModelId>,
pub base_url_override: Option<String>,
}
#[derive(Debug, Clone)]
pub struct RouteResolver {
offerings: Vec<ProviderModelOffering>,
}
impl Default for RouteResolver {
fn default() -> Self {
Self::new()
}
}
impl RouteResolver {
#[must_use]
pub fn new() -> Self {
Self::from_offerings(default_offerings())
}
#[must_use]
pub fn from_offerings(offerings: Vec<ProviderModelOffering>) -> Self {
Self { offerings }
}
pub fn resolve(&self, req: &RouteRequest) -> Result<ReadyRouteCandidate, RouteError> {
let provider_kind = req.explicit_provider.unwrap_or_default();
let descriptor = ProviderDescriptor::for_kind(provider_kind);
let provider_id = descriptor.id();
let default_offering = self.default_offering(&provider_id);
let logical_model = match &req.model_selector {
Some(selector) => selector.clone(),
None => {
let raw = req
.saved_provider_model
.as_ref()
.map(|w| w.as_str().to_string())
.unwrap_or_else(|| {
default_offering.map_or_else(
|| descriptor.default_wire_model().as_str().to_string(),
|offering| offering.wire_model_id.as_str().to_string(),
)
});
LogicalModelRef::from(raw)
}
};
if logical_model.raw().is_empty() {
return Err(RouteError::EmptyModel);
}
let is_auto = logical_model.is_auto();
let class = if request_uses_custom_endpoint(&descriptor, req.base_url_override.as_deref()) {
ProviderClass::LocalOrCustom
} else {
classify(provider_kind)
};
let (wire_model_id, canonical_model, endpoint_key, limits, pricing) = if is_auto {
default_offering.map_or_else(
|| {
(
descriptor.default_wire_model(),
None,
"chat".to_string(),
RouteLimits::default(),
PricingSku::UnknownOrStale,
)
},
|offering| {
(
offering.wire_model_id.clone(),
offering.canonical_model.clone(),
offering.endpoint_key.clone(),
offering.limits,
offering.pricing.clone(),
)
},
)
} else {
self.scope_selector(provider_kind, &provider_id, &logical_model, class)?
};
let endpoint = ResolvedEndpoint {
base_url: req
.base_url_override
.clone()
.unwrap_or_else(|| descriptor.default_base_url().to_string()),
endpoint_key,
protocol: descriptor.protocol(),
};
let mut messages = Vec::new();
if endpoint_uses_insecure_http(&endpoint.base_url) {
messages
.push("endpoint uses insecure http:// (credentials sent in plaintext)".to_string());
}
let validation = ValidationReport { ok: true, messages };
Ok(ReadyRouteCandidate::new(
provider_id,
provider_kind,
logical_model,
canonical_model,
wire_model_id,
endpoint,
ResolvedAuthSource::Missing,
descriptor.protocol(),
limits,
Some(pricing),
validation,
))
}
fn scope_selector(
&self,
provider_kind: ProviderKind,
provider_id: &ProviderId,
logical_model: &LogicalModelRef,
class: ProviderClass,
) -> Result<
(
WireModelId,
Option<ModelId>,
String,
RouteLimits,
PricingSku,
),
RouteError,
> {
let raw = logical_model.raw();
for offering in &self.offerings {
if offering.provider != *provider_id {
continue;
}
let matches_canonical = offering
.canonical_model
.as_ref()
.is_some_and(|m| m.as_str() == raw);
let matches_wire = offering.wire_model_id.as_str() == raw;
if matches_canonical || matches_wire {
return Ok((
offering.wire_model_id.clone(),
offering.canonical_model.clone(),
offering.endpoint_key.clone(),
offering.limits,
offering.pricing.clone(),
));
}
}
match class {
ProviderClass::StrictDirect => {
if self.selector_matches_other_provider_offering(provider_id, raw) {
return Err(RouteError::ForeignModelForDirectProvider {
provider: provider_id.clone(),
model: raw.to_string(),
});
}
if logical_model.namespace_hint().is_some() {
return Err(RouteError::ForeignModelForDirectProvider {
provider: provider_id.clone(),
model: raw.to_string(),
});
}
Ok((
WireModelId::from(raw),
None,
"chat".to_string(),
RouteLimits::default(),
PricingSku::UnknownOrStale,
))
}
ProviderClass::Aggregator | ProviderClass::LocalOrCustom => {
let _ = provider_kind;
Ok((
WireModelId::from(raw),
None,
"chat".to_string(),
RouteLimits::default(),
PricingSku::UnknownOrStale,
))
}
}
}
fn default_offering(&self, provider_id: &ProviderId) -> Option<&ProviderModelOffering> {
self.offerings
.iter()
.find(|offering| offering.provider == *provider_id && offering.default_for_provider)
}
fn selector_matches_other_provider_offering(
&self,
provider_id: &ProviderId,
raw: &str,
) -> bool {
self.offerings.iter().any(|offering| {
offering.provider != *provider_id
&& (offering.wire_model_id.as_str() == raw
|| offering
.canonical_model
.as_ref()
.is_some_and(|model| model.as_str() == raw))
})
}
}
fn default_offerings() -> Vec<ProviderModelOffering> {
let mut seen: std::collections::HashSet<(String, String)> = std::collections::HashSet::new();
let mut out = Vec::new();
let asset_rows = bundled_catalog_offerings()
.iter()
.map(CatalogOffering::to_offering)
.collect::<Vec<_>>();
for offering in bundled_offerings().into_iter().chain(asset_rows) {
let key = (
offering.provider.as_str().to_string(),
offering.wire_model_id.as_str().to_string(),
);
if seen.insert(key) {
out.push(offering);
}
}
out
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum ProviderClass {
StrictDirect,
Aggregator,
LocalOrCustom,
}
fn classify(kind: ProviderKind) -> ProviderClass {
match kind {
ProviderKind::Deepseek | ProviderKind::Zai => ProviderClass::StrictDirect,
ProviderKind::Ollama | ProviderKind::Vllm | ProviderKind::Sglang | ProviderKind::Openai => {
ProviderClass::LocalOrCustom
}
_ => ProviderClass::Aggregator,
}
}
fn request_uses_custom_endpoint(
descriptor: &ProviderDescriptor,
base_url_override: Option<&str>,
) -> bool {
base_url_override.is_some_and(|base_url| {
normalize_route_base_url(base_url)
!= normalize_route_base_url(descriptor.default_base_url())
})
}
fn normalize_route_base_url(base_url: &str) -> String {
let trimmed = base_url.trim().trim_end_matches('/');
let deepseek_domains = ["api.deepseek.com", "api.deepseeki.com"];
if deepseek_domains
.iter()
.any(|domain| trimmed.to_ascii_lowercase().contains(domain))
{
return trimmed.trim_end_matches("/v1").to_string();
}
if let Some(idx) = trimmed.find("://") {
let (scheme, rest) = trimmed.split_at(idx);
let scheme = scheme.to_ascii_lowercase();
let rest = &rest[3..];
let (authority, path) = match rest.find('/') {
Some(p) => (&rest[..p], &rest[p..]),
None => (rest, ""),
};
return format!("{scheme}://{}{path}", authority.to_ascii_lowercase());
}
trimmed.to_ascii_lowercase()
}
fn endpoint_uses_insecure_http(base_url: &str) -> bool {
let trimmed = base_url.trim();
let Some(rest) = strip_http_scheme(trimmed) else {
return false;
};
!is_loopback_host(host_of_authority(rest))
}
fn strip_http_scheme(base_url: &str) -> Option<&str> {
let idx = base_url.find("://")?;
let (scheme, rest) = base_url.split_at(idx);
if scheme.eq_ignore_ascii_case("http") {
Some(&rest[3..])
} else {
None
}
}
fn host_of_authority(rest: &str) -> &str {
let authority = rest.split('/').next().unwrap_or(rest);
let authority = authority.rsplit('@').next().unwrap_or(authority);
if let Some(inner) = authority.strip_prefix('[') {
return inner.split(']').next().unwrap_or(inner);
}
authority.split(':').next().unwrap_or(authority)
}
fn is_loopback_host(host: &str) -> bool {
let host = host.trim().trim_matches(|c| c == '[' || c == ']');
host.eq_ignore_ascii_case("localhost")
|| host == "127.0.0.1"
|| host == "::1"
|| host
.strip_prefix("127.")
.is_some_and(|_| host.split('.').count() == 4)
}