1use 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#[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
60pub 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 #[must_use]
75 pub fn builder() -> ProviderRegistryBuilder {
76 ProviderRegistryBuilder::default()
77 }
78
79 pub fn provider_ids(&self) -> impl Iterator<Item = &str> + '_ {
81 self.providers.keys().map(String::as_str)
82 }
83
84 #[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 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 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 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 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 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 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 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 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 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
304pub 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 #[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 #[must_use]
353 pub fn separator(mut self, separator: impl Into<String>) -> Self {
354 self.separator = separator.into();
355 self
356 }
357
358 #[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 #[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 #[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 #[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}