use crate::skins::{OpenAIErrorHandler, SkinErrorHandler};
use crate::{router::Router, service::ProviderManager};
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
use tokio_util::sync::CancellationToken;
#[derive(Clone)]
pub struct SkinContext {
pub router: Arc<Router>,
pub model_resolver: Arc<RwLock<ModelResolver>>,
pub provider_manager: Arc<RwLock<ProviderManager>>,
pub cancel_tokens: Arc<CancellationToken>,
pub error_handler: Arc<dyn SkinErrorHandler + Send + Sync>,
}
impl SkinContext {
pub fn new(router: Router) -> Self {
Self {
router: Arc::new(router),
model_resolver: Arc::new(RwLock::new(ModelResolver::new())),
provider_manager: Arc::new(RwLock::new(ProviderManager::new())),
cancel_tokens: Arc::new(CancellationToken::new()),
error_handler: Arc::new(OpenAIErrorHandler),
}
}
pub fn with_provider_manager(
router: Router,
provider_manager: Arc<RwLock<ProviderManager>>,
) -> Self {
Self {
router: Arc::new(router),
model_resolver: Arc::new(RwLock::new(ModelResolver::new())),
provider_manager,
cancel_tokens: Arc::new(CancellationToken::new()),
error_handler: Arc::new(OpenAIErrorHandler),
}
}
pub fn with_error_handler(
router: Router,
provider_manager: Arc<RwLock<ProviderManager>>,
error_handler: Arc<dyn SkinErrorHandler + Send + Sync>,
) -> Self {
Self {
router: Arc::new(router),
model_resolver: Arc::new(RwLock::new(ModelResolver::new())),
provider_manager,
cancel_tokens: Arc::new(CancellationToken::new()),
error_handler,
}
}
}
pub fn determine_skin_from_path(path: &str) -> Arc<dyn SkinErrorHandler + Send + Sync> {
if path.starts_with("/api/openai/v1/") || path.starts_with("/api/openai-compatible/v1/") {
Arc::new(OpenAIErrorHandler)
} else if path.starts_with("/api/anthropic/v1/") {
Arc::new(OpenAIErrorHandler) } else {
Arc::new(OpenAIErrorHandler)
}
}
pub struct ModelResolver {
models: HashMap<String, crate::types::ModelRef>,
}
impl Default for ModelResolver {
fn default() -> Self {
Self::new()
}
}
impl ModelResolver {
pub fn new() -> Self {
Self {
models: HashMap::new(),
}
}
pub fn register(&mut self, model_ref: crate::types::ModelRef) {
self.models.insert(model_ref.alias.clone(), model_ref);
}
pub fn resolve(&self, alias: &str) -> Option<&crate::types::ModelRef> {
self.models.get(alias)
}
}
impl SkinContext {
pub async fn resolve_model_ref(&self, model: &str) -> Option<crate::types::ModelRef> {
let mgr = self.provider_manager.read().await;
let discovered = if let Some(m) = mgr.get_model(model) {
Some(m.clone())
} else {
let mut candidate: Option<crate::types::DiscoveredModel> = None;
if let Some((prefix, rest)) = model.split_once('/') {
let prefix_lower = prefix.to_lowercase();
candidate = mgr
.list_models()
.into_iter()
.find(|m| m.provider_name == prefix_lower && m.name == rest)
.cloned();
if candidate.is_none() {
use crate::types::ProviderKind as PK;
let kind_hint = match prefix_lower.as_str() {
"openai-compat" => Some(PK::OpenAICompat),
"openai" => Some(PK::OpenAI),
"openrouter" => Some(PK::OpenRouter),
_ => None,
};
if let Some(k) = kind_hint {
candidate = mgr
.list_models()
.into_iter()
.find(|m| m.name == rest && m.provider_kind == k)
.cloned();
}
}
}
candidate.or_else(|| {
mgr.list_models()
.into_iter()
.find(|m| m.name == model)
.cloned()
})
}?;
let provider = mgr
.list_providers()
.into_iter()
.find(|p| p.name.to_lowercase() == discovered.provider_name)
.or_else(|| {
mgr.list_providers()
.into_iter()
.find(|p| p.endpoint.kind == discovered.provider_kind)
})?;
Some(crate::types::ModelRef {
alias: discovered.id.clone(),
provider: provider.clone(),
model_id: discovered.name.clone(),
input_modalities: discovered.input_modalities.clone(),
output_modalities: discovered.output_modalities.clone(),
})
}
}