use std::collections::BTreeSet;
use std::sync::Arc;
use std::time::SystemTime;
use super::CompactOutput;
use super::CompactRequest;
use super::Model;
use super::ModelEventSink;
use super::ModelOutput;
use super::ModelPricing;
use super::ModelRequest;
use super::PromptCacheMode;
use super::ToolDefinition;
use super::has_prompt_cache_breakpoint;
use super::mark_prompt_cache_breakpoint;
use crate::Error;
use crate::Result;
use crate::protocol::ModelChoice;
use crate::protocol::ModelStepDiagnostics;
use crate::protocol::PromptCacheDiagnostics;
use crate::protocol::TokenUsage;
use crate::protocol::ToolDiscoveryMode;
pub struct ModelRouter {
default: String,
routes: Vec<ModelRoute>,
}
struct ModelRoute {
choice: ModelChoice,
provider: Arc<dyn Model>,
credential: ModelCredentialLifetime,
}
impl ModelRouter {
pub fn new(id: impl Into<String>, provider: Arc<dyn Model>) -> Self {
let id = id.into();
let choice = inferred_choice(&id, provider.as_ref());
Self {
default: id,
routes: vec![ModelRoute {
choice,
provider,
credential: ModelCredentialLifetime::default(),
}],
}
}
pub fn register(&mut self, id: impl Into<String>, provider: Arc<dyn Model>) -> Result<()> {
let id = id.into();
if self.routes.iter().any(|route| route.choice.route == id) {
return Err(Error::Duplicate(format!("model provider `{id}`")));
}
self.routes.push(ModelRoute {
choice: inferred_choice(&id, provider.as_ref()),
provider,
credential: ModelCredentialLifetime::default(),
});
Ok(())
}
pub fn set_credential_lifetime(
&mut self,
id: &str,
credential: ModelCredentialLifetime,
) -> Result<()> {
let route = self
.routes
.iter_mut()
.find(|route| route.choice.route == id)
.ok_or_else(|| Error::Unknown(format!("model provider `{id}`")))?;
route.credential = credential;
Ok(())
}
#[must_use]
pub fn choices(
&self,
) -> impl DoubleEndedIterator<Item = &ModelChoice> + ExactSizeIterator + Clone {
self.routes.iter().map(|route| &route.choice)
}
pub fn resolve_choice(
&self,
route: &str,
reasoning_effort: Option<&str>,
) -> Result<&ModelChoice> {
let choice = self
.choices()
.find(|choice| choice.route == route)
.ok_or_else(|| Error::Unknown(format!("model route `{route}`")))?;
let Some(reasoning_effort) = reasoning_effort else {
return Ok(choice);
};
self.choices()
.find(|candidate| {
candidate.group == choice.group
&& candidate.reasoning_effort.as_deref() == Some(reasoning_effort)
})
.ok_or_else(|| {
Error::Unknown(format!(
"reasoning effort `{reasoning_effort}` for model route `{route}`"
))
})
}
pub fn configure_choice(&mut self, mut choice: ModelChoice) -> Result<()> {
if choice.group.trim().is_empty() || choice.model.trim().is_empty() {
return Err(Error::Config(
"model choice group and model cannot be empty".into(),
));
}
if choice.context_window.is_some_and(|window| window <= 0) {
return Err(Error::Config(
"model choice context window must be positive".into(),
));
}
let current = self
.routes
.iter_mut()
.find(|current| current.choice.route == choice.route)
.ok_or_else(|| Error::Unknown(format!("model route `{}`", choice.route)))?;
choice.supports_image_input = current.provider.supports_image_input();
choice.supports_realtime_voice = current.provider.supports_realtime_voice();
choice.tool_discovery = current.provider.tool_discovery();
current.choice = choice;
Ok(())
}
#[must_use]
pub fn default_provider(&self) -> &str {
&self.default
}
pub async fn respond(
&self,
provider: &str,
request: ModelRequest<'_>,
events: ModelEventSink,
) -> Result<ModelOutput> {
let route = self.route(provider)?;
while_valid(&route.credential, || {
route.provider.respond(request, events)
})
.await
}
pub fn compaction_endpoint(&self, provider: &str) -> Result<bool> {
Ok(self.provider(provider)?.compaction_endpoint())
}
pub fn supports_image_input(&self, provider: &str) -> Result<bool> {
Ok(self.provider(provider)?.supports_image_input())
}
pub fn supports_realtime_voice(&self, provider: &str) -> Result<bool> {
Ok(self.provider(provider)?.supports_realtime_voice())
}
pub async fn start_realtime_voice(
&self,
provider: &str,
request: super::RealtimeVoiceRequest,
) -> Result<super::RealtimeVoiceCall> {
let route = self.route(provider)?;
let mut credential = route.credential.clone();
credential.expires_at = credential
.expires_at
.map(super::RealtimeVoiceCall::cleanup_deadline);
let mut call =
while_valid(&credential, || route.provider.start_realtime_voice(request)).await?;
call.limit_credential(credential);
Ok(call)
}
pub fn tool_discovery(&self, provider: &str) -> Result<ToolDiscoveryMode> {
Ok(self.provider(provider)?.tool_discovery())
}
pub(crate) fn prepare_tool_definitions(
&self,
provider: &str,
mut direct: Vec<ToolDefinition>,
deferred: Vec<ToolDefinition>,
materialized: &BTreeSet<String>,
) -> Result<(Vec<ToolDefinition>, Vec<ToolDefinition>)> {
match self.provider(provider)?.tool_discovery() {
ToolDiscoveryMode::Native => Ok((direct, deferred)),
ToolDiscoveryMode::Rebuild => {
direct.extend(
deferred
.iter()
.filter(|tool| materialized.contains(&tool.name))
.cloned(),
);
Ok((direct, Vec::new()))
}
}
}
pub(crate) fn prepare_turn_input(
&self,
context: &[serde_json::Value],
input: &mut serde_json::Value,
) {
if !has_prompt_cache_breakpoint(context) {
let _ = mark_prompt_cache_breakpoint(input);
}
}
pub fn prompt_cache_capability(&self, provider: &str) -> Result<PromptCacheMode> {
Ok(self.provider(provider)?.prompt_cache_capability())
}
pub fn pricing(&self, provider: &str) -> Result<Option<ModelPricing>> {
Ok(self.provider(provider)?.pricing())
}
pub fn estimated_cost_microusd(
&self,
provider: &str,
usage: &TokenUsage,
) -> Result<Option<u64>> {
Ok(self
.provider(provider)?
.pricing()
.and_then(|pricing| pricing.estimate_microusd(usage)))
}
pub(crate) fn model_step_diagnostics(
&self,
provider: &str,
context_epoch: u64,
rewrite_reasons: Vec<String>,
usage: &TokenUsage,
) -> Result<ModelStepDiagnostics> {
let model = self.provider(provider)?;
let capability = model.prompt_cache_capability();
Ok(ModelStepDiagnostics {
provider: provider.into(),
prompt_cache: PromptCacheDiagnostics {
capability,
context_epoch,
outcome: capability.outcome(usage, !rewrite_reasons.is_empty()),
rewrite_reasons,
},
estimated_cost_microusd: model
.pricing()
.and_then(|pricing| pricing.estimate_microusd(usage)),
})
}
pub async fn compact(
&self,
provider: &str,
request: CompactRequest<'_>,
) -> Result<CompactOutput> {
let route = self.route(provider)?;
while_valid(&route.credential, || route.provider.compact(request)).await
}
fn provider(&self, id: &str) -> Result<&dyn Model> {
Ok(self.route(id)?.provider.as_ref())
}
fn route(&self, id: &str) -> Result<&ModelRoute> {
self.routes
.iter()
.find(|route| route.choice.route == id)
.ok_or_else(|| Error::Unknown(format!("model provider `{id}`")))
}
}
#[derive(Clone, Default)]
pub struct ModelCredentialLifetime {
pub expires_at: Option<SystemTime>,
pub revoked: Option<tokio::sync::watch::Receiver<()>>,
}
impl ModelCredentialLifetime {
pub(super) async fn ended(mut self) {
let expiry = async {
match self.expires_at {
Some(deadline) => {
tokio::time::sleep(
deadline
.duration_since(SystemTime::now())
.unwrap_or_default(),
)
.await
}
None => std::future::pending().await,
}
};
let revoked = async {
match self.revoked.as_mut() {
Some(revoked) => {
let _ = revoked.changed().await;
}
None => std::future::pending().await,
}
};
tokio::select! { _ = expiry => {}, _ = revoked => {} }
}
fn is_valid(&self) -> bool {
self.expires_at
.is_none_or(|deadline| deadline > SystemTime::now())
&& self
.revoked
.as_ref()
.is_none_or(|revoked| matches!(revoked.has_changed(), Ok(false)))
}
}
async fn while_valid<T, F: Future<Output = Result<T>>>(
credential: &ModelCredentialLifetime,
operation: impl FnOnce() -> F,
) -> Result<T> {
if !credential.is_valid() {
return Err(expired_credential());
}
tokio::select! {
biased;
_ = credential.clone().ended() => Err(expired_credential()),
result = async { operation().await } => result,
}
}
fn expired_credential() -> Error {
Error::Provider(crate::ProviderError::http(
"model credential has expired or been revoked",
401,
None,
))
}
fn inferred_choice(route: &str, provider: &dyn Model) -> ModelChoice {
let mut info = provider.info();
if info.model.is_empty() {
info.model = route.to_string();
}
ModelChoice {
route: route.to_string(),
group: route.to_string(),
model: info.model,
reasoning_effort: info.reasoning_effort,
context_window: None,
supports_image_input: provider.supports_image_input(),
supports_realtime_voice: provider.supports_realtime_voice(),
tool_discovery: provider.tool_discovery(),
}
}