Skip to main content

ferrin_core/registry/
provider_registry.rs

1//! Registry of providers addressed as `provider:model`.
2
3use std::collections::BTreeMap;
4use std::fmt;
5use std::sync::Arc;
6
7use ferrin_spec::EmbeddingModelRef;
8use ferrin_spec::ImageModelRef;
9use ferrin_spec::LanguageModelRef;
10use ferrin_spec::ModelKind;
11use ferrin_spec::NoSuchModelError;
12use ferrin_spec::Provider;
13use ferrin_spec::ProviderId;
14use ferrin_spec::ProviderRef;
15use ferrin_spec::RealtimeModelRef;
16use ferrin_spec::RerankingModelRef;
17use ferrin_spec::SpeechModelRef;
18use ferrin_spec::SpeechTranslationModelRef;
19use ferrin_spec::TranscriptionModelRef;
20use ferrin_spec::VideoModelRef;
21
22use crate::error::Error;
23use crate::error::NoSuchProviderDetails;
24use crate::middleware::EmbeddingModelMiddleware;
25use crate::middleware::ImageModelMiddleware;
26use crate::middleware::LanguageModelMiddleware;
27use crate::middleware::wrap_embedding_model;
28use crate::middleware::wrap_image_model;
29use crate::middleware::wrap_language_model;
30
31/// Resolves model ids of the form `<provider><separator><model>`.
32#[derive(Clone)]
33pub struct ProviderRegistry {
34    providers: BTreeMap<String, ProviderRef>,
35    separator: String,
36    language_model_middleware: Vec<Arc<dyn LanguageModelMiddleware>>,
37    embedding_model_middleware: Vec<Arc<dyn EmbeddingModelMiddleware>>,
38    image_model_middleware: Vec<Arc<dyn ImageModelMiddleware>>,
39    id: ProviderId,
40}
41
42impl fmt::Debug for ProviderRegistry {
43    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
44        f.debug_struct("ProviderRegistry")
45            .field("providers", &self.providers.keys().collect::<Vec<_>>())
46            .field("separator", &self.separator)
47            .field(
48                "language_model_middleware",
49                &self.language_model_middleware.len(),
50            )
51            .field(
52                "embedding_model_middleware",
53                &self.embedding_model_middleware.len(),
54            )
55            .field("image_model_middleware", &self.image_model_middleware.len())
56            .finish()
57    }
58}
59
60/// Creates a registry from `(id, provider)` pairs with the default
61/// separator `:`.
62pub fn create_provider_registry(
63    providers: impl IntoIterator<Item = (impl Into<String>, ProviderRef)>,
64) -> ProviderRegistry {
65    let mut builder = ProviderRegistry::builder();
66    for (id, provider) in providers {
67        builder = builder.provider(id, provider);
68    }
69    builder.build()
70}
71
72impl ProviderRegistry {
73    /// Starts building a registry.
74    #[must_use]
75    pub fn builder() -> ProviderRegistryBuilder {
76        ProviderRegistryBuilder::default()
77    }
78
79    /// The registered provider ids.
80    pub fn provider_ids(&self) -> impl Iterator<Item = &str> + '_ {
81        self.providers.keys().map(String::as_str)
82    }
83
84    /// Looks up a provider by id.
85    #[must_use]
86    pub fn provider(&self, id: &str) -> Option<&ProviderRef> {
87        self.providers.get(id)
88    }
89
90    fn split<'a>(&self, id: &'a str, kind: ModelKind) -> Result<(&ProviderRef, &'a str), Error> {
91        let not_found = |provider_id: &str| {
92            Error::no_such_provider(NoSuchProviderDetails {
93                provider_id: ProviderId::new(provider_id),
94                available_providers: self.providers.keys().map(ProviderId::new).collect(),
95                model_id: id.to_owned(),
96                model_kind: kind,
97            })
98        };
99        let Some((provider_id, model_id)) = id.split_once(self.separator.as_str()) else {
100            return Err(not_found(id));
101        };
102        let provider = self
103            .providers
104            .get(provider_id)
105            .ok_or_else(|| not_found(provider_id))?;
106        Ok((provider, model_id))
107    }
108
109    /// Resolves a language model, applying the registry middleware.
110    ///
111    /// # Errors
112    ///
113    /// [`Error::NoSuchProvider`] for unknown providers or malformed ids,
114    /// [`Error::Provider`] (`NoSuchModel`) from the provider.
115    pub fn language_model(&self, id: &str) -> Result<LanguageModelRef, Error> {
116        let (provider, model_id) = self.split(id, ModelKind::Language)?;
117        let model = provider.language_model(model_id)?;
118        if self.language_model_middleware.is_empty() {
119            return Ok(model);
120        }
121        let inner = model
122            .into_model()
123            .map_err(|id| Error::NoDefaultRegistry { model_id: id })?;
124        Ok(LanguageModelRef::from_arc(wrap_language_model(
125            inner,
126            self.language_model_middleware.iter().cloned(),
127        )))
128    }
129
130    /// Resolves an embedding model, applying the registry middleware.
131    ///
132    /// # Errors
133    ///
134    /// See [`ProviderRegistry::language_model`].
135    pub fn embedding_model(&self, id: &str) -> Result<EmbeddingModelRef, Error> {
136        let (provider, model_id) = self.split(id, ModelKind::Embedding)?;
137        let model = provider.embedding_model(model_id)?;
138        if self.embedding_model_middleware.is_empty() {
139            return Ok(model);
140        }
141        let inner = model
142            .into_model()
143            .map_err(|id| Error::NoDefaultRegistry { model_id: id })?;
144        Ok(EmbeddingModelRef::from_arc(wrap_embedding_model(
145            inner,
146            self.embedding_model_middleware.iter().cloned(),
147        )))
148    }
149
150    /// Resolves an image model, applying the registry middleware.
151    ///
152    /// # Errors
153    ///
154    /// See [`ProviderRegistry::language_model`].
155    pub fn image_model(&self, id: &str) -> Result<ImageModelRef, Error> {
156        let (provider, model_id) = self.split(id, ModelKind::Image)?;
157        let model = provider.image_model(model_id)?;
158        if self.image_model_middleware.is_empty() {
159            return Ok(model);
160        }
161        let inner = model
162            .into_model()
163            .map_err(|id| Error::NoDefaultRegistry { model_id: id })?;
164        Ok(ImageModelRef::from_arc(wrap_image_model(
165            inner,
166            self.image_model_middleware.iter().cloned(),
167        )))
168    }
169
170    /// Resolves a transcription model.
171    ///
172    /// # Errors
173    ///
174    /// See [`ProviderRegistry::language_model`].
175    pub fn transcription_model(&self, id: &str) -> Result<TranscriptionModelRef, Error> {
176        let (provider, model_id) = self.split(id, ModelKind::Transcription)?;
177        Ok(provider.transcription_model(model_id)?)
178    }
179
180    /// Resolves a speech model.
181    ///
182    /// # Errors
183    ///
184    /// See [`ProviderRegistry::language_model`].
185    pub fn speech_model(&self, id: &str) -> Result<SpeechModelRef, Error> {
186        let (provider, model_id) = self.split(id, ModelKind::Speech)?;
187        Ok(provider.speech_model(model_id)?)
188    }
189
190    /// Resolves a reranking model.
191    ///
192    /// # Errors
193    ///
194    /// See [`ProviderRegistry::language_model`].
195    pub fn reranking_model(&self, id: &str) -> Result<RerankingModelRef, Error> {
196        let (provider, model_id) = self.split(id, ModelKind::Reranking)?;
197        Ok(provider.reranking_model(model_id)?)
198    }
199
200    /// Resolves a video model.
201    ///
202    /// # Errors
203    ///
204    /// See [`ProviderRegistry::language_model`].
205    pub fn video_model(&self, id: &str) -> Result<VideoModelRef, Error> {
206        let (provider, model_id) = self.split(id, ModelKind::Video)?;
207        Ok(provider.video_model(model_id)?)
208    }
209
210    /// Resolves a speech translation model.
211    ///
212    /// # Errors
213    ///
214    /// See [`ProviderRegistry::language_model`].
215    pub fn speech_translation_model(&self, id: &str) -> Result<SpeechTranslationModelRef, Error> {
216        let (provider, model_id) = self.split(id, ModelKind::SpeechTranslation)?;
217        Ok(provider.speech_translation_model(model_id)?)
218    }
219
220    /// Resolves a realtime model through the provider's realtime factory.
221    ///
222    /// # Errors
223    ///
224    /// See [`ProviderRegistry::language_model`]; providers without realtime
225    /// support yield a `NoSuchModel` error.
226    pub fn realtime_model(&self, id: &str) -> Result<RealtimeModelRef, Error> {
227        let (provider, model_id) = self.split(id, ModelKind::Realtime)?;
228        let factory = provider.realtime().ok_or_else(|| {
229            Error::from(
230                NoSuchModelError::new(model_id, ModelKind::Realtime)
231                    .with_provider(provider.provider_id())
232                    .with_message(format!(
233                        "provider `{}` does not support realtime sessions",
234                        provider.provider_id()
235                    )),
236            )
237        })?;
238        Ok(factory.model(model_id)?)
239    }
240}
241
242fn to_no_such_model(error: Error, id: &str, kind: ModelKind) -> NoSuchModelError {
243    match error {
244        Error::Provider(provider) => match *provider {
245            ferrin_spec::ProviderError::NoSuchModel(inner) => *inner,
246            other => NoSuchModelError::new(id, kind).with_message(other.to_string()),
247        },
248        other => NoSuchModelError::new(id, kind).with_message(other.to_string()),
249    }
250}
251
252impl Provider for ProviderRegistry {
253    fn provider_id(&self) -> &ProviderId {
254        &self.id
255    }
256
257    fn language_model(&self, model_id: &str) -> Result<LanguageModelRef, NoSuchModelError> {
258        ProviderRegistry::language_model(self, model_id)
259            .map_err(|error| to_no_such_model(error, model_id, ModelKind::Language))
260    }
261
262    fn embedding_model(&self, model_id: &str) -> Result<EmbeddingModelRef, NoSuchModelError> {
263        ProviderRegistry::embedding_model(self, model_id)
264            .map_err(|error| to_no_such_model(error, model_id, ModelKind::Embedding))
265    }
266
267    fn image_model(&self, model_id: &str) -> Result<ImageModelRef, NoSuchModelError> {
268        ProviderRegistry::image_model(self, model_id)
269            .map_err(|error| to_no_such_model(error, model_id, ModelKind::Image))
270    }
271
272    fn transcription_model(
273        &self,
274        model_id: &str,
275    ) -> Result<TranscriptionModelRef, NoSuchModelError> {
276        ProviderRegistry::transcription_model(self, model_id)
277            .map_err(|error| to_no_such_model(error, model_id, ModelKind::Transcription))
278    }
279
280    fn speech_model(&self, model_id: &str) -> Result<SpeechModelRef, NoSuchModelError> {
281        ProviderRegistry::speech_model(self, model_id)
282            .map_err(|error| to_no_such_model(error, model_id, ModelKind::Speech))
283    }
284
285    fn reranking_model(&self, model_id: &str) -> Result<RerankingModelRef, NoSuchModelError> {
286        ProviderRegistry::reranking_model(self, model_id)
287            .map_err(|error| to_no_such_model(error, model_id, ModelKind::Reranking))
288    }
289
290    fn video_model(&self, model_id: &str) -> Result<VideoModelRef, NoSuchModelError> {
291        ProviderRegistry::video_model(self, model_id)
292            .map_err(|error| to_no_such_model(error, model_id, ModelKind::Video))
293    }
294
295    fn speech_translation_model(
296        &self,
297        model_id: &str,
298    ) -> Result<SpeechTranslationModelRef, NoSuchModelError> {
299        ProviderRegistry::speech_translation_model(self, model_id)
300            .map_err(|error| to_no_such_model(error, model_id, ModelKind::SpeechTranslation))
301    }
302}
303
304/// Builder of a [`ProviderRegistry`].
305pub struct ProviderRegistryBuilder {
306    providers: BTreeMap<String, ProviderRef>,
307    separator: String,
308    language_model_middleware: Vec<Arc<dyn LanguageModelMiddleware>>,
309    embedding_model_middleware: Vec<Arc<dyn EmbeddingModelMiddleware>>,
310    image_model_middleware: Vec<Arc<dyn ImageModelMiddleware>>,
311}
312
313impl Default for ProviderRegistryBuilder {
314    fn default() -> Self {
315        Self {
316            providers: BTreeMap::new(),
317            separator: ":".to_owned(),
318            language_model_middleware: Vec::new(),
319            embedding_model_middleware: Vec::new(),
320            image_model_middleware: Vec::new(),
321        }
322    }
323}
324
325impl fmt::Debug for ProviderRegistryBuilder {
326    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
327        f.debug_struct("ProviderRegistryBuilder")
328            .field("providers", &self.providers.keys().collect::<Vec<_>>())
329            .field("separator", &self.separator)
330            .field(
331                "language_model_middleware",
332                &self.language_model_middleware.len(),
333            )
334            .field(
335                "embedding_model_middleware",
336                &self.embedding_model_middleware.len(),
337            )
338            .field("image_model_middleware", &self.image_model_middleware.len())
339            .finish()
340    }
341}
342
343impl ProviderRegistryBuilder {
344    /// Registers `provider` under `id` (replacing an existing entry).
345    #[must_use]
346    pub fn provider(mut self, id: impl Into<String>, provider: ProviderRef) -> Self {
347        self.providers.insert(id.into(), provider);
348        self
349    }
350
351    /// Sets the separator between provider and model id (default `:`).
352    #[must_use]
353    pub fn separator(mut self, separator: impl Into<String>) -> Self {
354        self.separator = separator.into();
355        self
356    }
357
358    /// Applies `middleware` to every resolved language model.
359    #[must_use]
360    pub fn language_model_middleware(
361        mut self,
362        middleware: Arc<dyn LanguageModelMiddleware>,
363    ) -> Self {
364        self.language_model_middleware.push(middleware);
365        self
366    }
367
368    /// Applies `middleware` to every resolved embedding model.
369    #[must_use]
370    pub fn embedding_model_middleware(
371        mut self,
372        middleware: Arc<dyn EmbeddingModelMiddleware>,
373    ) -> Self {
374        self.embedding_model_middleware.push(middleware);
375        self
376    }
377
378    /// Applies `middleware` to every resolved image model.
379    #[must_use]
380    pub fn image_model_middleware(mut self, middleware: Arc<dyn ImageModelMiddleware>) -> Self {
381        self.image_model_middleware.push(middleware);
382        self
383    }
384
385    /// Builds the registry.
386    #[must_use]
387    pub fn build(self) -> ProviderRegistry {
388        ProviderRegistry {
389            providers: self.providers,
390            separator: self.separator,
391            language_model_middleware: self.language_model_middleware,
392            embedding_model_middleware: self.embedding_model_middleware,
393            image_model_middleware: self.image_model_middleware,
394            id: ProviderId::new("registry"),
395        }
396    }
397}