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