Skip to main content

tea_model/
router.rs

1use std::collections::BTreeMap;
2use std::sync::Arc;
3
4use tea_protocol::{ModelRef, ProviderId};
5use thiserror::Error;
6
7use crate::{ModelProvider, ModelSpec};
8
9/// Object-safe lookup port that routes provider-qualified model identities.
10pub trait ModelRouter: std::fmt::Debug + Send + Sync {
11    /// Returns the provider registered under the canonical identity.
12    fn provider(&self, provider_id: &ProviderId) -> Option<&dyn ModelProvider>;
13
14    /// Returns all advertised models in deterministic provider/model order.
15    fn models(&self) -> &[ModelSpec];
16
17    /// Resolves one complete model identity.
18    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/// Immutable reference registry for a fixed runtime provider generation.
36#[derive(Debug)]
37pub struct ModelRegistry {
38    providers: BTreeMap<ProviderId, Arc<dyn ModelProvider>>,
39    models: Vec<ModelSpec>,
40}
41
42impl ModelRegistry {
43    /// Builds a validated registry from one immutable provider generation.
44    ///
45    /// # Errors
46    ///
47    /// Returns an error for duplicate provider identities or a provider that
48    /// advertises a model owned by another provider.
49    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    /// Returns the registered provider count.
81    #[must_use]
82    pub fn provider_count(&self) -> usize {
83        self.providers.len()
84    }
85
86    /// Returns registered provider identities in canonical order.
87    #[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/// Invalid immutable provider-registry composition.
104#[derive(Debug, Clone, PartialEq, Eq, Error)]
105pub enum ModelRegistryError {
106    /// At least one provider is required.
107    #[error("model registry requires at least one provider")]
108    Empty,
109    /// Two adapters claimed the same provider identity.
110    #[error("model provider {0} is registered more than once")]
111    DuplicateProvider(ProviderId),
112    /// An adapter advertised a model owned by a different provider.
113    #[error("model provider {0} advertises a model owned by another provider")]
114    ProviderCatalogMismatch(ProviderId),
115}