use std::marker::PhantomData;
use std::sync::Arc;
use serde_json::Value as JsonValue;
use tokio::sync::mpsc;
use crate::chat::message::{MsgVariant, Role};
pub type InspectorError = Box<dyn std::error::Error + Send + Sync + 'static>;
pub trait InspectorMiddleware<Evt>: Send + Sync {
fn name(&self) -> &str;
fn inspect(&self, event: &Evt) -> Result<(), InspectorError>;
}
pub enum MutatorAction<Evt> {
Pass,
Modify(Evt),
Block { reason: String },
}
pub trait MutatorMiddleware<Evt>: Send + Sync {
fn name(&self) -> &str;
fn process(&self, event: Evt) -> MutatorAction<Evt>;
}
pub trait IntoInspectors<Evt> {
fn into_inspectors(self) -> Vec<Arc<dyn InspectorMiddleware<Evt>>>;
}
pub trait IntoMutators<Evt> {
fn into_mutators(self) -> Vec<Arc<dyn MutatorMiddleware<Evt>>>;
}
macro_rules! impl_into_inspectors_tuple {
($($T:ident),+) => {
impl<Evt, $($T: InspectorMiddleware<Evt> + 'static),+> IntoInspectors<Evt>
for ($($T,)+)
{
#[allow(non_snake_case)]
fn into_inspectors(self) -> Vec<Arc<dyn InspectorMiddleware<Evt>>> {
let ($($T,)+) = self;
vec![$(Arc::new($T) as Arc<dyn InspectorMiddleware<Evt>>),+]
}
}
};
}
macro_rules! impl_into_mutators_tuple {
($($T:ident),+) => {
impl<Evt, $($T: MutatorMiddleware<Evt> + 'static),+> IntoMutators<Evt>
for ($($T,)+)
{
#[allow(non_snake_case)]
fn into_mutators(self) -> Vec<Arc<dyn MutatorMiddleware<Evt>>> {
let ($($T,)+) = self;
vec![$(Arc::new($T) as Arc<dyn MutatorMiddleware<Evt>>),+]
}
}
};
}
impl_into_inspectors_tuple!(A);
impl_into_inspectors_tuple!(A, B);
impl_into_inspectors_tuple!(A, B, C);
impl_into_inspectors_tuple!(A, B, C, D);
impl_into_inspectors_tuple!(A, B, C, D, F);
impl_into_inspectors_tuple!(A, B, C, D, F, G);
impl_into_inspectors_tuple!(A, B, C, D, F, G, H);
impl_into_inspectors_tuple!(A, B, C, D, F, G, H, I);
impl_into_inspectors_tuple!(A, B, C, D, F, G, H, I, J);
impl_into_inspectors_tuple!(A, B, C, D, F, G, H, I, J, K);
impl_into_inspectors_tuple!(A, B, C, D, F, G, H, I, J, K, L);
impl_into_inspectors_tuple!(A, B, C, D, F, G, H, I, J, K, L, M);
impl_into_mutators_tuple!(A);
impl_into_mutators_tuple!(A, B);
impl_into_mutators_tuple!(A, B, C);
impl_into_mutators_tuple!(A, B, C, D);
impl_into_mutators_tuple!(A, B, C, D, F);
impl_into_mutators_tuple!(A, B, C, D, F, G);
impl_into_mutators_tuple!(A, B, C, D, F, G, H);
impl_into_mutators_tuple!(A, B, C, D, F, G, H, I);
impl_into_mutators_tuple!(A, B, C, D, F, G, H, I, J);
impl_into_mutators_tuple!(A, B, C, D, F, G, H, I, J, K);
impl_into_mutators_tuple!(A, B, C, D, F, G, H, I, J, K, L);
impl_into_mutators_tuple!(A, B, C, D, F, G, H, I, J, K, L, M);
pub enum MiddlewareLayer<Evt> {
Inspector(Vec<Arc<dyn InspectorMiddleware<Evt>>>),
Mutator(Vec<Arc<dyn MutatorMiddleware<Evt>>>),
}
pub struct ErrorsDisabled;
pub struct ErrorsEnabled;
pub struct MiddlewareChain<Evt, ErrState = ErrorsDisabled> {
layers: Vec<MiddlewareLayer<Evt>>,
error_tx: Option<mpsc::UnboundedSender<(String, InspectorError)>>,
_err: PhantomData<ErrState>,
}
impl<Evt: Clone + Send + 'static> MiddlewareChain<Evt, ErrorsDisabled> {
pub fn new() -> Self {
Self {
layers: Vec::new(),
error_tx: None,
_err: PhantomData,
}
}
pub fn activate_error_channel(
self,
) -> (
MiddlewareChain<Evt, ErrorsEnabled>,
mpsc::UnboundedReceiver<(String, InspectorError)>,
) {
let (tx, rx) = mpsc::unbounded_channel();
let chain = MiddlewareChain::<Evt, ErrorsEnabled> {
layers: self.layers,
error_tx: Some(tx),
_err: PhantomData,
};
(chain, rx)
}
}
impl<Evt: Clone + Send + 'static> Default for MiddlewareChain<Evt, ErrorsDisabled> {
fn default() -> Self {
Self::new()
}
}
impl<Evt: Clone + Send + 'static, S> MiddlewareChain<Evt, S> {
pub fn with_inspector(mut self, i: impl InspectorMiddleware<Evt> + 'static) -> Self {
self.layers
.push(MiddlewareLayer::Inspector(vec![Arc::new(i)]));
self
}
pub fn with_inspectors(mut self, i: impl IntoInspectors<Evt>) -> Self {
let v = i.into_inspectors();
if !v.is_empty() {
self.layers.push(MiddlewareLayer::Inspector(v));
}
self
}
pub fn with_inspectors_from_iter(
mut self,
iter: impl IntoIterator<Item = Arc<dyn InspectorMiddleware<Evt>>>,
) -> Self {
let v: Vec<_> = iter.into_iter().collect();
if !v.is_empty() {
self.layers.push(MiddlewareLayer::Inspector(v));
}
self
}
pub fn with_mutator(mut self, m: impl MutatorMiddleware<Evt> + 'static) -> Self {
self.layers
.push(MiddlewareLayer::Mutator(vec![Arc::new(m)]));
self
}
pub fn with_mutators(mut self, m: impl IntoMutators<Evt>) -> Self {
let v = m.into_mutators();
if !v.is_empty() {
self.layers.push(MiddlewareLayer::Mutator(v));
}
self
}
pub fn with_mutators_from_iter(
mut self,
iter: impl IntoIterator<Item = Arc<dyn MutatorMiddleware<Evt>>>,
) -> Self {
let v: Vec<_> = iter.into_iter().collect();
if !v.is_empty() {
self.layers.push(MiddlewareLayer::Mutator(v));
}
self
}
pub fn is_empty(&self) -> bool {
self.layers.is_empty()
}
pub fn len(&self) -> usize {
self.layers.len()
}
pub fn process(&self, event: Evt) -> Result<Evt, MiddlewareBlocked> {
let mut current = event;
for layer in &self.layers {
match layer {
MiddlewareLayer::Inspector(inspectors) => {
current = self.run_inspectors(inspectors, current);
}
MiddlewareLayer::Mutator(mutators) => {
current = self.run_mutators(mutators, current)?;
}
}
}
Ok(current)
}
fn run_inspectors(&self, inspectors: &[Arc<dyn InspectorMiddleware<Evt>>], event: Evt) -> Evt {
if let Ok(handle) = tokio::runtime::Handle::try_current() {
for insp in inspectors {
let name = insp.name().to_string();
let evt = event.clone();
let tx = self.error_tx.clone();
let insp = Arc::clone(insp);
handle.spawn(async move {
if let Err(e) = insp.inspect(&evt)
&& let Some(tx) = tx
{
let _ = tx.send((name, e));
}
});
}
} else {
for insp in inspectors {
if let (Some(tx), Err(e)) = (&self.error_tx, insp.inspect(&event)) {
let _ = tx.send((insp.name().to_string(), e));
}
}
}
event
}
fn run_mutators(
&self,
mutators: &[Arc<dyn MutatorMiddleware<Evt>>],
event: Evt,
) -> Result<Evt, MiddlewareBlocked> {
let mut current = event;
for m in mutators {
match m.process(current.clone()) {
MutatorAction::Pass => {}
MutatorAction::Modify(e) => current = e,
MutatorAction::Block { reason } => {
return Err(MiddlewareBlocked {
middleware_name: m.name().to_string(),
reason,
});
}
}
}
Ok(current)
}
}
impl<Evt: Clone + Send + 'static> MiddlewareChain<Evt, ErrorsEnabled> {
pub fn error_sender(&self) -> Option<mpsc::UnboundedSender<(String, InspectorError)>> {
self.error_tx.clone()
}
}
#[derive(Debug, Clone)]
pub struct MiddlewareBlocked {
pub middleware_name: String,
pub reason: String,
}
impl std::fmt::Display for MiddlewareBlocked {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"[middleware:{}] event blocked: {}",
self.middleware_name, self.reason
)
}
}
pub type EventSenderFn<E> = Box<dyn Fn(E) + Send + Sync>;
pub trait MiddlewareEvent: Clone + Send + 'static {
type Error: std::fmt::Display + Send + Sync + 'static + From<String>;
fn assistant_text(content: String, reasoning: Option<String>) -> Self;
fn tool_call_request(call_id: Arc<str>, name: String, args: JsonValue) -> Self;
fn tool_response(call_id: Arc<str>, name: String, result: Result<String, Self::Error>) -> Self;
fn turn_start() -> Self;
fn turn_end(finish_reason: Option<String>) -> Self;
fn done() -> Self;
fn into_session_message(self) -> Option<(Role, MsgVariant)>;
}
#[cfg(test)]
mod tests {
use super::*;
struct NoopInspector;
impl InspectorMiddleware<String> for NoopInspector {
fn name(&self) -> &str {
"noop"
}
fn inspect(&self, _event: &String) -> Result<(), InspectorError> {
Ok(())
}
}
struct UpperMutator;
impl MutatorMiddleware<String> for UpperMutator {
fn name(&self) -> &str {
"upper"
}
fn process(&self, event: String) -> MutatorAction<String> {
MutatorAction::Modify(event.to_uppercase())
}
}
struct BlockMutator;
impl MutatorMiddleware<String> for BlockMutator {
fn name(&self) -> &str {
"blocker"
}
fn process(&self, _event: String) -> MutatorAction<String> {
MutatorAction::Block {
reason: "blocked".into(),
}
}
}
struct PassMutator;
impl MutatorMiddleware<String> for PassMutator {
fn name(&self) -> &str {
"pass"
}
fn process(&self, _event: String) -> MutatorAction<String> {
MutatorAction::Pass
}
}
#[test]
fn new_chain_is_empty() {
let chain = MiddlewareChain::<String>::new();
assert!(chain.is_empty());
}
#[test]
fn single_mutator_modify() {
let chain = MiddlewareChain::<String>::new().with_mutator(UpperMutator);
let result = chain.process("hello".into()).unwrap();
assert_eq!(result, "HELLO");
}
#[test]
fn single_mutator_pass() {
let chain = MiddlewareChain::<String>::new().with_mutator(PassMutator);
let result = chain.process("hello".into()).unwrap();
assert_eq!(result, "hello");
}
#[test]
fn single_mutator_block() {
let chain = MiddlewareChain::<String>::new().with_mutator(BlockMutator);
let err = chain.process("hello".into()).unwrap_err();
assert_eq!(err.middleware_name, "blocker");
assert_eq!(err.reason, "blocked");
}
#[test]
fn pass_then_modify() {
let chain = MiddlewareChain::<String>::new().with_mutators((PassMutator, UpperMutator));
let result = chain.process("hello".into()).unwrap();
assert_eq!(result, "HELLO");
}
#[test]
fn modify_then_block() {
let chain = MiddlewareChain::<String>::new().with_mutators((UpperMutator, BlockMutator));
let err = chain.process("hello".into()).unwrap_err();
assert_eq!(err.middleware_name, "blocker");
}
#[test]
fn tuple_arity_3() {
let chain = MiddlewareChain::<String>::new().with_mutators((
UpperMutator,
PassMutator,
PassMutator,
));
let result = chain.process("hello".into()).unwrap();
assert_eq!(result, "HELLO");
}
#[test]
fn inspector_without_tokio_is_noop() {
let chain = MiddlewareChain::<String>::new()
.with_inspector(NoopInspector)
.with_mutator(UpperMutator);
let result = chain.process("hello".into()).unwrap();
assert_eq!(result, "HELLO");
}
#[test]
fn activate_error_channel_transitions_state() {
let chain = MiddlewareChain::<String>::new();
let (_enabled, _rx) = chain.activate_error_channel();
}
#[test]
fn with_inspectors_tuple() {
let chain =
MiddlewareChain::<String>::new().with_inspectors((NoopInspector, NoopInspector));
assert_eq!(chain.len(), 1);
}
#[test]
fn single_inspector_tuple() {
let chain = MiddlewareChain::<String>::new().with_inspectors((NoopInspector,));
assert_eq!(chain.len(), 1);
}
#[test]
fn chain_with_tokio_spawns_inspectors() {
let rt = tokio::runtime::Builder::new_current_thread()
.build()
.unwrap();
rt.block_on(async {
let chain = MiddlewareChain::<String>::new().with_inspector(NoopInspector);
let result = chain.process("hi".into()).unwrap();
assert_eq!(result, "hi");
});
}
}