1use 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#[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
63pub 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 #[must_use]
78 pub fn builder() -> ProviderRegistryBuilder {
79 ProviderRegistryBuilder::default()
80 }
81
82 pub fn provider_ids(&self) -> impl Iterator<Item = &str> + '_ {
84 self.providers.keys().map(String::as_str)
85 }
86
87 #[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 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 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 pub fn register_provider(&mut self, id: impl Into<String>, provider: ProviderRef) {
145 self.providers.insert(id.into(), provider);
146 }
147
148 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 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 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 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 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 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 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 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 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
343pub 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 #[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 #[must_use]
392 pub fn separator(mut self, separator: impl Into<String>) -> Self {
393 self.separator = separator.into();
394 self
395 }
396
397 #[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 #[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 #[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 #[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}