Skip to main content

ferrin_core/middleware/
mod.rs

1//! Language model middleware.
2//!
3//! A middleware wraps a model: it may rewrite call options, wrap the
4//! generate/stream calls, and override identity or supported URLs. Apply
5//! with [`wrap_language_model`]; the first middleware in the list is the
6//! outermost.
7
8use std::fmt;
9
10use ferrin_spec::BoxFuture;
11use ferrin_spec::CallOptions;
12use ferrin_spec::DynLanguageModel;
13use ferrin_spec::GenerateResult;
14use ferrin_spec::ModelId;
15use ferrin_spec::ProviderId;
16use ferrin_spec::StreamResult;
17use ferrin_spec::SupportedUrls;
18use ferrin_spec::error::ProviderError;
19
20pub mod builtin;
21mod wrap;
22
23pub use wrap::wrap_language_model;
24
25/// Which model method is being called.
26#[derive(Debug, Clone, Copy, PartialEq, Eq)]
27#[non_exhaustive]
28pub enum CallKind {
29    /// `do_generate`.
30    Generate,
31    /// `do_stream`.
32    Stream,
33}
34
35/// The model being wrapped and the kind of call.
36#[derive(Clone, Copy)]
37pub struct MiddlewareContext<'a> {
38    /// The wrapped (inner) model.
39    pub model: &'a dyn DynLanguageModel,
40    /// The kind of call.
41    pub kind: CallKind,
42}
43
44impl fmt::Debug for MiddlewareContext<'_> {
45    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
46        f.debug_struct("MiddlewareContext")
47            .field("provider", self.model.provider())
48            .field("model_id", self.model.model_id())
49            .field("kind", &self.kind)
50            .finish()
51    }
52}
53
54/// Continuation of a wrapped `do_generate`.
55pub type GenerateNext<'a> = Box<
56    dyn FnOnce(CallOptions) -> BoxFuture<'a, Result<GenerateResult, ProviderError>> + Send + 'a,
57>;
58
59/// Continuation of a wrapped `do_stream`.
60pub type StreamNext<'a> =
61    Box<dyn FnOnce(CallOptions) -> BoxFuture<'a, Result<StreamResult, ProviderError>> + Send + 'a>;
62
63/// Intercepts language model calls. Every method has a pass-through default.
64pub trait LanguageModelMiddleware: Send + Sync + 'static {
65    /// Rewrites the call options before the call.
66    fn transform_params<'a>(
67        &'a self,
68        options: CallOptions,
69        _ctx: MiddlewareContext<'a>,
70    ) -> BoxFuture<'a, Result<CallOptions, ProviderError>> {
71        Box::pin(async move { Ok(options) })
72    }
73
74    /// Wraps `do_generate`.
75    fn wrap_generate<'a>(
76        &'a self,
77        options: CallOptions,
78        next: GenerateNext<'a>,
79        _ctx: MiddlewareContext<'a>,
80    ) -> BoxFuture<'a, Result<GenerateResult, ProviderError>> {
81        next(options)
82    }
83
84    /// Wraps `do_stream`.
85    fn wrap_stream<'a>(
86        &'a self,
87        options: CallOptions,
88        next: StreamNext<'a>,
89        _ctx: MiddlewareContext<'a>,
90    ) -> BoxFuture<'a, Result<StreamResult, ProviderError>> {
91        next(options)
92    }
93
94    /// Overrides the provider id reported by the wrapped model.
95    fn override_provider(&self, _model: &dyn DynLanguageModel) -> Option<ProviderId> {
96        None
97    }
98
99    /// Overrides the model id reported by the wrapped model.
100    fn override_model_id(&self, _model: &dyn DynLanguageModel) -> Option<ModelId> {
101        None
102    }
103
104    /// Overrides the supported URLs of the wrapped model.
105    fn override_supported_urls<'a>(
106        &'a self,
107        _model: &'a dyn DynLanguageModel,
108    ) -> Option<BoxFuture<'a, SupportedUrls>> {
109        None
110    }
111}