Skip to main content

ferrin_core/middleware/
mod.rs

1//! 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. [`EmbeddingModelMiddleware`] / [`wrap_embedding_model`] and
7//! [`ImageModelMiddleware`] / [`wrap_image_model`] do the same for embedding
8//! and image models, and [`wrap_provider`] applies all three kinds to every
9//! model a provider resolves.
10
11use std::fmt;
12
13use ferrin_spec::BoxFuture;
14use ferrin_spec::CallOptions;
15use ferrin_spec::DynLanguageModel;
16use ferrin_spec::GenerateResult;
17use ferrin_spec::ModelId;
18use ferrin_spec::ProviderId;
19use ferrin_spec::StreamResult;
20use ferrin_spec::SupportedUrls;
21use ferrin_spec::error::ProviderError;
22
23pub mod builtin;
24mod embedding;
25mod image;
26mod provider;
27pub(crate) mod tool_contract;
28mod wrap;
29
30pub use embedding::EmbedNext;
31pub use embedding::EmbeddingMiddlewareContext;
32pub use embedding::EmbeddingModelMiddleware;
33pub use embedding::wrap_embedding_model;
34pub use image::ImageGenerateNext;
35pub use image::ImageMiddlewareContext;
36pub use image::ImageModelMiddleware;
37pub use image::wrap_image_model;
38pub use provider::ProviderMiddleware;
39pub use provider::wrap_provider;
40pub use wrap::wrap_language_model;
41
42/// Which model method is being called.
43#[derive(Debug, Clone, Copy, PartialEq, Eq)]
44#[non_exhaustive]
45pub enum CallKind {
46    /// `do_generate`.
47    Generate,
48    /// `do_stream`.
49    Stream,
50}
51
52/// The model being wrapped and the kind of call.
53#[derive(Clone, Copy)]
54pub struct MiddlewareContext<'a> {
55    /// The wrapped (inner) model.
56    pub model: &'a dyn DynLanguageModel,
57    /// The kind of call.
58    pub kind: CallKind,
59}
60
61impl fmt::Debug for MiddlewareContext<'_> {
62    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
63        f.debug_struct("MiddlewareContext")
64            .field("provider", self.model.provider())
65            .field("model_id", self.model.model_id())
66            .field("kind", &self.kind)
67            .finish()
68    }
69}
70
71/// Continuation of a wrapped `do_generate`.
72pub type GenerateNext<'a> = Box<
73    dyn FnOnce(CallOptions) -> BoxFuture<'a, Result<GenerateResult, ProviderError>> + Send + 'a,
74>;
75
76/// Continuation of a wrapped `do_stream`.
77pub type StreamNext<'a> =
78    Box<dyn FnOnce(CallOptions) -> BoxFuture<'a, Result<StreamResult, ProviderError>> + Send + 'a>;
79
80/// Intercepts language model calls. Every method has a pass-through default.
81pub trait LanguageModelMiddleware: Send + Sync + 'static {
82    /// Rewrites the call options before the call.
83    fn transform_params<'a>(
84        &'a self,
85        options: CallOptions,
86        _ctx: MiddlewareContext<'a>,
87    ) -> BoxFuture<'a, Result<CallOptions, ProviderError>> {
88        Box::pin(async move { Ok(options) })
89    }
90
91    /// Wraps `do_generate`.
92    fn wrap_generate<'a>(
93        &'a self,
94        options: CallOptions,
95        next: GenerateNext<'a>,
96        _ctx: MiddlewareContext<'a>,
97    ) -> BoxFuture<'a, Result<GenerateResult, ProviderError>> {
98        next(options)
99    }
100
101    /// Wraps `do_stream`.
102    fn wrap_stream<'a>(
103        &'a self,
104        options: CallOptions,
105        next: StreamNext<'a>,
106        _ctx: MiddlewareContext<'a>,
107    ) -> BoxFuture<'a, Result<StreamResult, ProviderError>> {
108        next(options)
109    }
110
111    /// Overrides the provider id reported by the wrapped model.
112    fn override_provider(&self, _model: &dyn DynLanguageModel) -> Option<ProviderId> {
113        None
114    }
115
116    /// Overrides the model id reported by the wrapped model.
117    fn override_model_id(&self, _model: &dyn DynLanguageModel) -> Option<ModelId> {
118        None
119    }
120
121    /// Overrides the supported URLs of the wrapped model.
122    fn override_supported_urls<'a>(
123        &'a self,
124        _model: &'a dyn DynLanguageModel,
125    ) -> Option<BoxFuture<'a, SupportedUrls>> {
126        None
127    }
128}