ferrin_core/middleware/
mod.rs1use 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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
44#[non_exhaustive]
45pub enum CallKind {
46 Generate,
48 Stream,
50}
51
52#[derive(Clone, Copy)]
54pub struct MiddlewareContext<'a> {
55 pub model: &'a dyn DynLanguageModel,
57 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
71pub type GenerateNext<'a> = Box<
73 dyn FnOnce(CallOptions) -> BoxFuture<'a, Result<GenerateResult, ProviderError>> + Send + 'a,
74>;
75
76pub type StreamNext<'a> =
78 Box<dyn FnOnce(CallOptions) -> BoxFuture<'a, Result<StreamResult, ProviderError>> + Send + 'a>;
79
80pub trait LanguageModelMiddleware: Send + Sync + 'static {
82 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 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 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 fn override_provider(&self, _model: &dyn DynLanguageModel) -> Option<ProviderId> {
113 None
114 }
115
116 fn override_model_id(&self, _model: &dyn DynLanguageModel) -> Option<ModelId> {
118 None
119 }
120
121 fn override_supported_urls<'a>(
123 &'a self,
124 _model: &'a dyn DynLanguageModel,
125 ) -> Option<BoxFuture<'a, SupportedUrls>> {
126 None
127 }
128}