ferrin_core/middleware/
provider.rs1use 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#[derive(Clone, Default)]
37pub struct ProviderMiddleware {
38 pub language_model: Vec<Arc<dyn LanguageModelMiddleware>>,
40 pub embedding_model: Vec<Arc<dyn EmbeddingModelMiddleware>>,
42 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 #[must_use]
59 pub fn new() -> Self {
60 Self::default()
61 }
62
63 #[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 #[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 #[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 #[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#[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 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}