Skip to main content

ferrin_core/middleware/
provider.rs

1//! Application of middleware to every model a provider resolves.
2
3use std::fmt;
4use std::sync::Arc;
5
6use ferrin_spec::BatchRef;
7use ferrin_spec::EmbeddingModelRef;
8use ferrin_spec::FilesRef;
9use ferrin_spec::ImageModelRef;
10use ferrin_spec::LanguageModelRef;
11use ferrin_spec::ModelKind;
12use ferrin_spec::ModelRef;
13use ferrin_spec::NoSuchModelError;
14use ferrin_spec::Provider;
15use ferrin_spec::ProviderId;
16use ferrin_spec::ProviderRef;
17use ferrin_spec::RealtimeFactoryRef;
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 super::EmbeddingModelMiddleware;
26use super::ImageModelMiddleware;
27use super::LanguageModelMiddleware;
28use super::wrap_embedding_model;
29use super::wrap_image_model;
30use super::wrap_language_model;
31
32/// Middleware applied by [`wrap_provider`] to the models a provider resolves.
33///
34/// Each list is applied in order (first outermost) to every model of its
35/// kind; models of other kinds and provider services pass through.
36#[derive(Clone, Default)]
37pub struct ProviderMiddleware {
38    /// Applied to every language model.
39    pub language_model: Vec<Arc<dyn LanguageModelMiddleware>>,
40    /// Applied to every embedding model.
41    pub embedding_model: Vec<Arc<dyn EmbeddingModelMiddleware>>,
42    /// Applied to every image model.
43    pub image_model: Vec<Arc<dyn ImageModelMiddleware>>,
44}
45
46impl fmt::Debug for ProviderMiddleware {
47    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
48        f.debug_struct("ProviderMiddleware")
49            .field("language_model", &self.language_model.len())
50            .field("embedding_model", &self.embedding_model.len())
51            .field("image_model", &self.image_model.len())
52            .finish()
53    }
54}
55
56impl ProviderMiddleware {
57    /// Creates an empty set.
58    #[must_use]
59    pub fn new() -> Self {
60        Self::default()
61    }
62
63    /// Appends a language model middleware.
64    #[must_use]
65    pub fn language_model(mut self, middleware: Arc<dyn LanguageModelMiddleware>) -> Self {
66        self.language_model.push(middleware);
67        self
68    }
69
70    /// Appends an embedding model middleware.
71    #[must_use]
72    pub fn embedding_model(mut self, middleware: Arc<dyn EmbeddingModelMiddleware>) -> Self {
73        self.embedding_model.push(middleware);
74        self
75    }
76
77    /// Appends an image model middleware.
78    #[must_use]
79    pub fn image_model(mut self, middleware: Arc<dyn ImageModelMiddleware>) -> Self {
80        self.image_model.push(middleware);
81        self
82    }
83
84    /// Returns `true` when no middleware is configured.
85    #[must_use]
86    pub fn is_empty(&self) -> bool {
87        self.language_model.is_empty()
88            && self.embedding_model.is_empty()
89            && self.image_model.is_empty()
90    }
91}
92
93/// Wraps every language, embedding and image model resolved through
94/// `provider` with the matching `middleware`. Other model kinds, the provider
95/// id and the services are delegated unchanged. An empty set returns
96/// `provider` itself.
97#[must_use]
98pub fn wrap_provider(provider: ProviderRef, middleware: ProviderMiddleware) -> ProviderRef {
99    if middleware.is_empty() {
100        return provider;
101    }
102    Arc::new(WrappedProvider {
103        inner: provider,
104        middleware,
105    })
106}
107
108struct WrappedProvider {
109    inner: ProviderRef,
110    middleware: ProviderMiddleware,
111}
112
113impl fmt::Debug for WrappedProvider {
114    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
115        f.debug_struct("WrappedProvider")
116            .field("provider", self.inner.provider_id())
117            .field("middleware", &self.middleware)
118            .finish()
119    }
120}
121
122impl WrappedProvider {
123    /// Unwraps a resolved reference; providers hand out instances, so an id
124    /// form cannot be wrapped and is reported as `NoSuchModel`.
125    fn resolved<D: ?Sized>(
126        &self,
127        model: ModelRef<D>,
128        model_id: &str,
129        kind: ModelKind,
130    ) -> Result<Arc<D>, NoSuchModelError> {
131        model.into_model().map_err(|id| {
132            NoSuchModelError::new(model_id, kind)
133                .with_provider(self.inner.provider_id())
134                .with_message(format!(
135                    "provider returned the unresolved model id `{id}`, which middleware cannot wrap"
136                ))
137        })
138    }
139}
140
141impl Provider for WrappedProvider {
142    fn provider_id(&self) -> &ProviderId {
143        self.inner.provider_id()
144    }
145
146    fn language_model(&self, model_id: &str) -> Result<LanguageModelRef, NoSuchModelError> {
147        let model = self.inner.language_model(model_id)?;
148        if self.middleware.language_model.is_empty() {
149            return Ok(model);
150        }
151        let inner = self.resolved(model, model_id, ModelKind::Language)?;
152        Ok(LanguageModelRef::from_arc(wrap_language_model(
153            inner,
154            self.middleware.language_model.iter().cloned(),
155        )))
156    }
157
158    fn embedding_model(&self, model_id: &str) -> Result<EmbeddingModelRef, NoSuchModelError> {
159        let model = self.inner.embedding_model(model_id)?;
160        if self.middleware.embedding_model.is_empty() {
161            return Ok(model);
162        }
163        let inner = self.resolved(model, model_id, ModelKind::Embedding)?;
164        Ok(EmbeddingModelRef::from_arc(wrap_embedding_model(
165            inner,
166            self.middleware.embedding_model.iter().cloned(),
167        )))
168    }
169
170    fn image_model(&self, model_id: &str) -> Result<ImageModelRef, NoSuchModelError> {
171        let model = self.inner.image_model(model_id)?;
172        if self.middleware.image_model.is_empty() {
173            return Ok(model);
174        }
175        let inner = self.resolved(model, model_id, ModelKind::Image)?;
176        Ok(ImageModelRef::from_arc(wrap_image_model(
177            inner,
178            self.middleware.image_model.iter().cloned(),
179        )))
180    }
181
182    fn transcription_model(
183        &self,
184        model_id: &str,
185    ) -> Result<TranscriptionModelRef, NoSuchModelError> {
186        self.inner.transcription_model(model_id)
187    }
188
189    fn speech_model(&self, model_id: &str) -> Result<SpeechModelRef, NoSuchModelError> {
190        self.inner.speech_model(model_id)
191    }
192
193    fn reranking_model(&self, model_id: &str) -> Result<RerankingModelRef, NoSuchModelError> {
194        self.inner.reranking_model(model_id)
195    }
196
197    fn video_model(&self, model_id: &str) -> Result<VideoModelRef, NoSuchModelError> {
198        self.inner.video_model(model_id)
199    }
200
201    fn speech_translation_model(
202        &self,
203        model_id: &str,
204    ) -> Result<SpeechTranslationModelRef, NoSuchModelError> {
205        self.inner.speech_translation_model(model_id)
206    }
207
208    fn realtime(&self) -> Option<RealtimeFactoryRef> {
209        self.inner.realtime()
210    }
211
212    fn files(&self) -> Option<FilesRef> {
213        self.inner.files()
214    }
215
216    fn skills(&self) -> Option<SkillsRef> {
217        self.inner.skills()
218    }
219
220    fn batch(&self) -> Option<BatchRef> {
221        self.inner.batch()
222    }
223}