use std::fmt;
use ferrin_spec::BoxFuture;
use ferrin_spec::CallOptions;
use ferrin_spec::DynLanguageModel;
use ferrin_spec::GenerateResult;
use ferrin_spec::ModelId;
use ferrin_spec::ProviderId;
use ferrin_spec::StreamResult;
use ferrin_spec::SupportedUrls;
use ferrin_spec::error::ProviderError;
pub mod builtin;
mod wrap;
pub use wrap::wrap_language_model;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum CallKind {
Generate,
Stream,
}
#[derive(Clone, Copy)]
pub struct MiddlewareContext<'a> {
pub model: &'a dyn DynLanguageModel,
pub kind: CallKind,
}
impl fmt::Debug for MiddlewareContext<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("MiddlewareContext")
.field("provider", self.model.provider())
.field("model_id", self.model.model_id())
.field("kind", &self.kind)
.finish()
}
}
pub type GenerateNext<'a> = Box<
dyn FnOnce(CallOptions) -> BoxFuture<'a, Result<GenerateResult, ProviderError>> + Send + 'a,
>;
pub type StreamNext<'a> =
Box<dyn FnOnce(CallOptions) -> BoxFuture<'a, Result<StreamResult, ProviderError>> + Send + 'a>;
pub trait LanguageModelMiddleware: Send + Sync + 'static {
fn transform_params<'a>(
&'a self,
options: CallOptions,
_ctx: MiddlewareContext<'a>,
) -> BoxFuture<'a, Result<CallOptions, ProviderError>> {
Box::pin(async move { Ok(options) })
}
fn wrap_generate<'a>(
&'a self,
options: CallOptions,
next: GenerateNext<'a>,
_ctx: MiddlewareContext<'a>,
) -> BoxFuture<'a, Result<GenerateResult, ProviderError>> {
next(options)
}
fn wrap_stream<'a>(
&'a self,
options: CallOptions,
next: StreamNext<'a>,
_ctx: MiddlewareContext<'a>,
) -> BoxFuture<'a, Result<StreamResult, ProviderError>> {
next(options)
}
fn override_provider(&self, _model: &dyn DynLanguageModel) -> Option<ProviderId> {
None
}
fn override_model_id(&self, _model: &dyn DynLanguageModel) -> Option<ModelId> {
None
}
fn override_supported_urls<'a>(
&'a self,
_model: &'a dyn DynLanguageModel,
) -> Option<BoxFuture<'a, SupportedUrls>> {
None
}
}