mod types;
pub use types::*;
pub mod library;
pub use library::*;
use std::sync::Arc;
use crate::error::{Result, TinyAgentsError};
use crate::harness::context::RunContext;
use crate::harness::events::AgentEvent;
use crate::harness::model::{ModelDelta, ModelRequest, ModelResponse};
use crate::harness::tool::{ToolCall, ToolDelta, ToolResult};
macro_rules! run_stack_hook {
($self:ident, $ctx:ident, $iter:expr, |$mw:ident| $call:expr) => {{
for $mw in $iter {
let name = $mw.name().to_string();
$ctx.emit(AgentEvent::MiddlewareStarted { name: name.clone() });
let result = $call.await;
$ctx.emit(AgentEvent::MiddlewareCompleted { name });
if let Err(e) = result {
$self.fan_out_on_error($ctx, &e).await;
return Err(e);
}
}
Ok(())
}};
}
impl AgentRun {
pub fn new() -> Self {
Self::default()
}
pub fn text(&self) -> Option<String> {
self.final_response.as_ref().map(|r| r.text())
}
}
impl<State: Send + Sync, Ctx: Send + Sync> Default for MiddlewareStack<State, Ctx> {
fn default() -> Self {
Self::new()
}
}
impl<State: Send + Sync, Ctx: Send + Sync> MiddlewareStack<State, Ctx> {
pub fn new() -> Self {
Self {
middlewares: Vec::new(),
model_middlewares: Vec::new(),
tool_middlewares: Vec::new(),
}
}
pub fn push(&mut self, middleware: Arc<dyn Middleware<State, Ctx>>) {
self.middlewares.push(middleware);
}
pub fn push_model_middleware(&mut self, middleware: Arc<dyn ModelMiddleware<State, Ctx>>) {
self.model_middlewares.push(middleware);
}
pub fn push_tool_middleware(&mut self, middleware: Arc<dyn ToolMiddleware<State, Ctx>>) {
self.tool_middlewares.push(middleware);
}
pub fn model_middleware_len(&self) -> usize {
self.model_middlewares.len()
}
pub fn tool_middleware_len(&self) -> usize {
self.tool_middlewares.len()
}
pub fn len(&self) -> usize {
self.middlewares.len()
}
pub fn is_empty(&self) -> bool {
self.middlewares.is_empty()
}
async fn fan_out_on_error(&self, ctx: &mut RunContext<Ctx>, error: &TinyAgentsError) {
for mw in self.middlewares.iter() {
let _ = mw.on_error(ctx, error).await;
}
}
pub async fn run_before_agent(&self, ctx: &mut RunContext<Ctx>, state: &State) -> Result<()> {
run_stack_hook!(self, ctx, self.middlewares.iter(), |mw| mw
.before_agent(ctx, state))
}
pub async fn run_after_agent(
&self,
ctx: &mut RunContext<Ctx>,
state: &State,
run: &mut AgentRun,
) -> Result<()> {
run_stack_hook!(self, ctx, self.middlewares.iter().rev(), |mw| mw
.after_agent(ctx, state, run))
}
pub async fn run_before_model(
&self,
ctx: &mut RunContext<Ctx>,
state: &State,
request: &mut ModelRequest,
) -> Result<()> {
run_stack_hook!(self, ctx, self.middlewares.iter(), |mw| mw
.before_model(ctx, state, request))
}
pub async fn run_on_model_delta(
&self,
ctx: &mut RunContext<Ctx>,
state: &State,
delta: &mut ModelDelta,
) -> Result<()> {
for mw in self.middlewares.iter() {
if let Err(e) = mw.on_model_delta(ctx, state, delta).await {
self.fan_out_on_error(ctx, &e).await;
return Err(e);
}
}
Ok(())
}
pub async fn run_after_model(
&self,
ctx: &mut RunContext<Ctx>,
state: &State,
response: &mut ModelResponse,
) -> Result<()> {
run_stack_hook!(self, ctx, self.middlewares.iter().rev(), |mw| mw
.after_model(ctx, state, response))
}
pub async fn run_before_tool(
&self,
ctx: &mut RunContext<Ctx>,
state: &State,
call: &mut ToolCall,
) -> Result<()> {
run_stack_hook!(self, ctx, self.middlewares.iter(), |mw| mw
.before_tool(ctx, state, call))
}
pub async fn run_on_tool_delta(
&self,
ctx: &mut RunContext<Ctx>,
state: &State,
delta: &mut ToolDelta,
) -> Result<()> {
run_stack_hook!(self, ctx, self.middlewares.iter(), |mw| mw
.on_tool_delta(ctx, state, delta))
}
pub async fn run_after_tool(
&self,
ctx: &mut RunContext<Ctx>,
state: &State,
result: &mut ToolResult,
) -> Result<()> {
run_stack_hook!(self, ctx, self.middlewares.iter().rev(), |mw| mw
.after_tool(ctx, state, result))
}
pub async fn run_on_error(
&self,
ctx: &mut RunContext<Ctx>,
error: &TinyAgentsError,
) -> Result<()> {
for mw in self.middlewares.iter() {
ctx.emit(AgentEvent::MiddlewareStarted {
name: mw.name().to_string(),
});
let _ = mw.on_error(ctx, error).await;
ctx.emit(AgentEvent::MiddlewareCompleted {
name: mw.name().to_string(),
});
}
Ok(())
}
pub async fn run_wrapped_model(
&self,
ctx: &mut RunContext<Ctx>,
state: &State,
request: ModelRequest,
base: &dyn ModelBaseCall<State, Ctx>,
) -> Result<MiddlewareModelOutcome> {
let handler = ModelHandler {
remaining: &self.model_middlewares,
base,
};
handler.run(ctx, state, request).await
}
pub async fn run_wrapped_tool(
&self,
ctx: &mut RunContext<Ctx>,
state: &State,
call: ToolCall,
base: &dyn ToolBaseCall<State, Ctx>,
) -> Result<MiddlewareToolOutcome> {
let handler = ToolHandler {
remaining: &self.tool_middlewares,
base,
};
handler.run(ctx, state, call).await
}
}
impl<State: Send + Sync, Ctx: Send + Sync> ModelHandler<'_, State, Ctx> {
pub async fn run(
&self,
ctx: &mut RunContext<Ctx>,
state: &State,
request: ModelRequest,
) -> Result<MiddlewareModelOutcome> {
match self.remaining.split_first() {
Some((head, tail)) => {
let next = ModelHandler {
remaining: tail,
base: self.base,
};
let name = head.name().to_string();
ctx.emit(AgentEvent::MiddlewareStarted { name: name.clone() });
let outcome = head.wrap_model(ctx, state, request, next).await;
ctx.emit(AgentEvent::MiddlewareCompleted { name });
outcome
}
None => Ok(MiddlewareModelOutcome::Response(
self.base.call(ctx, state, request).await?,
)),
}
}
}
impl<State: Send + Sync, Ctx: Send + Sync> ToolHandler<'_, State, Ctx> {
pub async fn run(
&self,
ctx: &mut RunContext<Ctx>,
state: &State,
call: ToolCall,
) -> Result<MiddlewareToolOutcome> {
match self.remaining.split_first() {
Some((head, tail)) => {
let next = ToolHandler {
remaining: tail,
base: self.base,
};
let name = head.name().to_string();
ctx.emit(AgentEvent::MiddlewareStarted { name: name.clone() });
let outcome = head.wrap_tool(ctx, state, call, next).await;
ctx.emit(AgentEvent::MiddlewareCompleted { name });
outcome
}
None => Ok(MiddlewareToolOutcome::Result(
self.base.call(ctx, state, call).await?,
)),
}
}
}
#[cfg(test)]
mod test;