ferrin_core/middleware/
wrap.rs1use 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#[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}