Skip to main content

ferrin_core/middleware/
wrap.rs

1//! Application of middleware to a model.
2
3use std::sync::Arc;
4
5use ferrin_spec::CallOptions;
6use ferrin_spec::DynLanguageModel;
7use ferrin_spec::GenerateResult;
8use ferrin_spec::LanguageModel;
9use ferrin_spec::ModelId;
10use ferrin_spec::ProviderId;
11use ferrin_spec::StreamResult;
12use ferrin_spec::SupportedUrls;
13use ferrin_spec::error::ProviderError;
14
15use super::CallKind;
16use super::LanguageModelMiddleware;
17use super::MiddlewareContext;
18use super::tool_contract::ToolContract;
19
20/// Wraps `model` with `middleware`; the first entry becomes the outermost
21/// layer. An empty list returns `model` unchanged.
22#[must_use]
23pub fn wrap_language_model(
24    model: Arc<dyn DynLanguageModel>,
25    middleware: impl IntoIterator<
26        Item = Arc<dyn LanguageModelMiddleware>,
27        IntoIter: DoubleEndedIterator,
28    >,
29) -> Arc<dyn DynLanguageModel> {
30    middleware.into_iter().rev().fold(model, |inner, layer| {
31        let provider = layer
32            .override_provider(inner.as_ref())
33            .unwrap_or_else(|| inner.provider().clone());
34        let model_id = layer
35            .override_model_id(inner.as_ref())
36            .unwrap_or_else(|| inner.model_id().clone());
37        Arc::new(WrappedLanguageModel {
38            inner,
39            layer,
40            provider,
41            model_id,
42        })
43    })
44}
45
46struct WrappedLanguageModel {
47    inner: Arc<dyn DynLanguageModel>,
48    layer: Arc<dyn LanguageModelMiddleware>,
49    provider: ProviderId,
50    model_id: ModelId,
51}
52
53impl std::fmt::Debug for WrappedLanguageModel {
54    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
55        f.debug_struct("WrappedLanguageModel")
56            .field("provider", &self.provider)
57            .field("model_id", &self.model_id)
58            .finish_non_exhaustive()
59    }
60}
61
62impl LanguageModel for WrappedLanguageModel {
63    fn provider(&self) -> &ProviderId {
64        &self.provider
65    }
66
67    fn model_id(&self) -> &ModelId {
68        &self.model_id
69    }
70
71    async fn supported_urls(&self) -> SupportedUrls {
72        match self.layer.override_supported_urls(self.inner.as_ref()) {
73            Some(urls) => urls.await,
74            None => self.inner.supported_urls().await,
75        }
76    }
77
78    async fn do_generate(&self, options: CallOptions) -> Result<GenerateResult, ProviderError> {
79        let ctx = MiddlewareContext {
80            model: self.inner.as_ref(),
81            kind: CallKind::Generate,
82        };
83        let contract = ToolContract::current();
84        let options = self.layer.transform_params(options, ctx).await?;
85        if let Some(contract) = &contract {
86            contract.observe(&options);
87        }
88        let inner = &self.inner;
89        self.layer
90            .wrap_generate(
91                options,
92                Box::new(move |options| {
93                    if let Some(contract) = &contract {
94                        contract.observe(&options);
95                    }
96                    Box::pin(ToolContract::continue_call(
97                        contract,
98                        inner.do_generate(options),
99                    ))
100                }),
101                ctx,
102            )
103            .await
104    }
105
106    async fn do_stream(&self, options: CallOptions) -> Result<StreamResult, ProviderError> {
107        let ctx = MiddlewareContext {
108            model: self.inner.as_ref(),
109            kind: CallKind::Stream,
110        };
111        let contract = ToolContract::current();
112        let options = self.layer.transform_params(options, ctx).await?;
113        if let Some(contract) = &contract {
114            contract.observe(&options);
115        }
116        let inner = &self.inner;
117        self.layer
118            .wrap_stream(
119                options,
120                Box::new(move |options| {
121                    if let Some(contract) = &contract {
122                        contract.observe(&options);
123                    }
124                    Box::pin(ToolContract::continue_call(
125                        contract,
126                        inner.do_stream(options),
127                    ))
128                }),
129                ctx,
130            )
131            .await
132    }
133}