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::LanguageModelMiddleware;
25use crate::middleware::wrap_language_model;
26
27#[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
49pub 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 #[must_use]
64 pub fn builder() -> ProviderRegistryBuilder {
65 ProviderRegistryBuilder::default()
66 }
67
68 pub fn provider_ids(&self) -> impl Iterator<Item = &str> + '_ {
70 self.providers.keys().map(String::as_str)
71 }
72
73 #[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 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 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 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 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 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 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 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 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 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
273pub 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 #[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 #[must_use]
313 pub fn separator(mut self, separator: impl Into<String>) -> Self {
314 self.separator = separator.into();
315 self
316 }
317
318 #[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 #[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}