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