1use std::collections::BTreeMap;
2use std::sync::Arc;
3
4use tea_protocol::{ModelRef, ProviderId};
5use thiserror::Error;
6
7use crate::{ModelProvider, ModelSpec};
8
9pub trait ModelRouter: std::fmt::Debug + Send + Sync {
11 fn provider(&self, provider_id: &ProviderId) -> Option<&dyn ModelProvider>;
13
14 fn models(&self) -> &[ModelSpec];
16
17 fn model(&self, model_ref: &ModelRef) -> Option<&ModelSpec> {
19 self.provider(model_ref.provider_id())?
20 .model(model_ref.model_id())
21 .filter(|model| model.model_ref() == model_ref)
22 }
23}
24
25impl<T: ModelProvider> ModelRouter for T {
26 fn provider(&self, provider_id: &ProviderId) -> Option<&dyn ModelProvider> {
27 (self.provider_id() == provider_id).then_some(self)
28 }
29
30 fn models(&self) -> &[ModelSpec] {
31 ModelProvider::models(self)
32 }
33}
34
35#[derive(Debug)]
37pub struct ModelRegistry {
38 providers: BTreeMap<ProviderId, Arc<dyn ModelProvider>>,
39 models: Vec<ModelSpec>,
40}
41
42impl ModelRegistry {
43 pub fn new(
50 providers: impl IntoIterator<Item = Arc<dyn ModelProvider>>,
51 ) -> Result<Self, ModelRegistryError> {
52 let mut by_id = BTreeMap::new();
53 for provider in providers {
54 let provider_id = provider.provider_id().clone();
55 if provider
56 .models()
57 .iter()
58 .any(|model| model.provider_id() != &provider_id)
59 {
60 return Err(ModelRegistryError::ProviderCatalogMismatch(provider_id));
61 }
62 if by_id.insert(provider_id.clone(), provider).is_some() {
63 return Err(ModelRegistryError::DuplicateProvider(provider_id));
64 }
65 }
66 if by_id.is_empty() {
67 return Err(ModelRegistryError::Empty);
68 }
69 let mut models = by_id
70 .values()
71 .flat_map(|provider| provider.models().iter().cloned())
72 .collect::<Vec<_>>();
73 models.sort_by(|left, right| left.model_ref().cmp(right.model_ref()));
74 Ok(Self {
75 providers: by_id,
76 models,
77 })
78 }
79
80 #[must_use]
82 pub fn provider_count(&self) -> usize {
83 self.providers.len()
84 }
85
86 #[must_use]
88 pub fn provider_ids(&self) -> Vec<ProviderId> {
89 self.providers.keys().cloned().collect()
90 }
91}
92
93impl ModelRouter for ModelRegistry {
94 fn provider(&self, provider_id: &ProviderId) -> Option<&dyn ModelProvider> {
95 self.providers.get(provider_id).map(AsRef::as_ref)
96 }
97
98 fn models(&self) -> &[ModelSpec] {
99 &self.models
100 }
101}
102
103#[derive(Debug, Clone, PartialEq, Eq, Error)]
105pub enum ModelRegistryError {
106 #[error("model registry requires at least one provider")]
108 Empty,
109 #[error("model provider {0} is registered more than once")]
111 DuplicateProvider(ProviderId),
112 #[error("model provider {0} advertises a model owned by another provider")]
114 ProviderCatalogMismatch(ProviderId),
115}