use std::collections::VecDeque;
use std::future::Future;
use std::sync::Arc;
use everruns_core::traits::EventEmitter;
use everruns_core::turn::TurnStopReason;
use everruns_core::typed_id::{MessageId, TurnId};
use everruns_core::{AgentLoopError, InputMessage, SessionId};
use everruns_host::{
AcceptedTurnInput, InProcessRuntime, TurnResult, TurnSteering, TurnSteeringPushError,
};
use tokio::sync::{OnceCell, mpsc, oneshot, watch};
use crate::Agent;
use crate::events::{EventStream, FacadeEventBus, RunOptions};
use crate::hooks::{
AgentStartContext, CompletionContext, HookFailure, HookRunState, TurnStartContext,
};
#[derive(Clone)]
pub struct Session {
inner: Arc<SessionInner>,
}
struct SessionInner {
agent: Agent,
session_id: SessionId,
event_bus: Arc<FacadeEventBus>,
hook_state: Arc<HookRunState>,
commands: OnceCell<mpsc::Sender<Command>>,
}
const SESSION_COMMAND_CAPACITY: usize = 64;
impl Session {
pub(crate) fn new(agent: Agent, session_id: SessionId) -> Self {
let hook_state = HookRunState::new(agent.lifecycle_hooks());
Self {
inner: Arc::new(SessionInner {
agent,
session_id,
event_bus: Arc::new(FacadeEventBus::new()),
hook_state,
commands: OnceCell::new(),
}),
}
}
pub fn id(&self) -> String {
self.inner.session_id.to_string()
}
pub fn session_id(&self) -> SessionId {
self.inner.session_id
}
pub fn history(&self) -> crate::HistoryQuery {
crate::HistoryQuery::new(self.inner.agent.clone(), self.inner.session_id)
}
pub fn work(&self, queue: &crate::work::WorkQueue) -> crate::work::SessionWork {
queue.for_session(self.id())
}
pub fn events(&self) -> EventStream {
self.inner.event_bus.subscribe()
}
pub async fn send(&self, input: impl Into<InputMessage>) -> Result<SentMessage, RunError> {
self.send_internal(input.into(), None).await
}
pub async fn send_and_wait(&self, input: impl Into<InputMessage>) -> Result<Turn, RunError> {
self.send(input).await?.wait().await
}
pub async fn run(&self, input: impl Into<InputMessage>) -> Result<Turn, RunError> {
self.send_and_wait(input).await
}
pub async fn run_with(
&self,
input: impl Into<InputMessage>,
options: RunOptions,
) -> Result<Turn, RunError> {
let Some(token) = options.cancel else {
return self.send_and_wait(input).await;
};
let sent = self
.send_internal(input.into(), Some(token.clone()))
.await?;
tokio::select! {
biased;
result = sent.wait() => result,
() = token.cancelled() => {
let _ = sent.turn.cancel().await;
sent.wait().await
},
}
}
pub async fn inspect(&self) -> Result<crate::SessionContext, RunError> {
let (response, result) = oneshot::channel();
self.command_sender()
.await
.send(Command::Inspect { response })
.await
.map_err(|_| RunError::SessionClosed)?;
result.await.map_err(|_| RunError::SessionClosed)?
}
async fn send_internal(
&self,
input: InputMessage,
cancel: Option<crate::CancellationToken>,
) -> Result<SentMessage, RunError> {
let (response, result) = oneshot::channel();
self.command_sender()
.await
.send(Command::Send {
input: Box::new(AcceptedTurnInput::new(input)),
cancel,
response,
})
.await
.map_err(|_| RunError::SessionClosed)?;
let ack = result.await.map_err(|_| RunError::SessionClosed)??;
Ok(SentMessage::new(self.clone(), ack))
}
async fn command_sender(&self) -> mpsc::Sender<Command> {
self.inner
.commands
.get_or_init(|| async {
let (sender, receiver) = mpsc::channel(SESSION_COMMAND_CAPACITY);
tokio::spawn(SessionActor::new(&self.inner).run(receiver));
sender
})
.await
.clone()
}
async fn cancel_turn(&self, turn_id: TurnId) -> Result<(), CancelError> {
let (response, result) = oneshot::channel();
self.command_sender()
.await
.send(Command::Cancel { turn_id, response })
.await
.map_err(|_| CancelError::SessionClosed)?;
match result.await.map_err(|_| CancelError::SessionClosed)? {
true => Ok(()),
false => Err(CancelError::TurnFinished),
}
}
}
enum Command {
Send {
input: Box<AcceptedTurnInput>,
cancel: Option<crate::CancellationToken>,
response: oneshot::Sender<Result<ActorSentMessage, RunError>>,
},
Inspect {
response: oneshot::Sender<Result<crate::SessionContext, RunError>>,
},
Cancel {
turn_id: TurnId,
response: oneshot::Sender<bool>,
},
}
struct ActorSentMessage {
message_id: MessageId,
turn_id: TurnId,
disposition: SendDisposition,
completion: watch::Receiver<TurnCompletion>,
}
#[derive(Clone, Debug)]
enum TurnCompletion {
Pending,
Ready(Result<Turn, RunError>),
}
struct SessionActor {
agent: Agent,
session_id: SessionId,
event_bus: Arc<FacadeEventBus>,
hook_state: Arc<HookRunState>,
runtime: Option<InProcessRuntime>,
agent_started: bool,
deferred: VecDeque<Command>,
}
impl SessionActor {
fn new(inner: &SessionInner) -> Self {
Self {
agent: inner.agent.clone(),
session_id: inner.session_id,
event_bus: inner.event_bus.clone(),
hook_state: inner.hook_state.clone(),
runtime: None,
agent_started: false,
deferred: VecDeque::new(),
}
}
async fn run(mut self, mut commands: mpsc::Receiver<Command>) {
loop {
let command = match self.deferred.pop_front() {
Some(command) => command,
None => match commands.recv().await {
Some(command) => command,
None => break,
},
};
match command {
Command::Send {
input,
cancel,
response,
} => {
if !self
.start_turn(input, cancel, response, &mut commands)
.await
{
break;
}
}
Command::Inspect { response } => {
let result = self.inspect().await;
let _ = response.send(result);
}
Command::Cancel { response, .. } => {
let _ = response.send(false);
}
}
}
}
async fn start_turn(
&mut self,
input: Box<AcceptedTurnInput>,
cancel: Option<crate::CancellationToken>,
response: oneshot::Sender<Result<ActorSentMessage, RunError>>,
commands: &mut mpsc::Receiver<Command>,
) -> bool {
let input = *input;
let turn_id = TurnId::new();
let message_id = input.message_id();
let steering = TurnSteering::new();
let (completion_tx, completion_rx) = watch::channel(TurnCompletion::Pending);
let _ = response.send(Ok(ActorSentMessage {
message_id,
turn_id,
disposition: SendDisposition::Started,
completion: completion_rx,
}));
match self
.prepare_turn(input.input().clone(), cancel.as_ref())
.await
{
HookRun::Cancelled => {
steering.close();
self.hook_state.take_failures();
let result = self.emit_cancelled(turn_id).await;
let _ = completion_tx.send(TurnCompletion::Ready(result));
return true;
}
HookRun::Completed(Err(error)) => {
steering.close();
let _ = completion_tx.send(TurnCompletion::Ready(Err(error)));
return true;
}
HookRun::Completed(Ok(())) => {}
}
self.drive_turn(input, turn_id, steering, completion_tx, commands)
.await
}
async fn prepare_turn(
&mut self,
input: InputMessage,
cancel: Option<&crate::CancellationToken>,
) -> HookRun<Result<(), RunError>> {
self.hook_state.begin_turn();
if !self.agent_started {
let context = AgentStartContext {
agent_name: self.agent.name().to_string(),
session_id: self.session_id,
};
match cancellable(cancel, self.hook_state.hooks().run_agent_start(context)).await {
HookRun::Cancelled => return HookRun::Cancelled,
HookRun::Completed(Err(failure)) => {
return HookRun::Completed(Err(RunError::Hook(failure)));
}
HookRun::Completed(Ok(())) => self.agent_started = true,
}
}
let context = TurnStartContext {
agent_name: self.agent.name().to_string(),
session_id: self.session_id,
input,
};
match cancellable(cancel, self.hook_state.hooks().run_turn_start(context)).await {
HookRun::Cancelled => return HookRun::Cancelled,
HookRun::Completed(Err(failure)) => {
return HookRun::Completed(Err(RunError::Hook(failure)));
}
HookRun::Completed(Ok(())) => {}
}
HookRun::Completed(self.ensure_runtime().await)
}
async fn drive_turn(
&mut self,
input: AcceptedTurnInput,
turn_id: TurnId,
steering: TurnSteering,
completion: watch::Sender<TurnCompletion>,
commands: &mut mpsc::Receiver<Command>,
) -> bool {
let runtime = self.runtime.as_ref().expect("runtime built above").clone();
let (outcome, cancelled) = {
let run = runtime.run_steerable_turn(self.session_id, input, turn_id, steering.clone());
tokio::pin!(run);
let mut cancelled = false;
let outcome = loop {
tokio::select! {
biased;
result = &mut run => break Some(result.map(Turn::from).map_err(RunError::from)),
command = commands.recv(), if self.deferred.len() < SESSION_COMMAND_CAPACITY => match command {
None => {
steering.close();
return false;
}
Some(Command::Send { input, cancel, response }) => {
let message_id = input.message_id();
match steering.try_push(*input) {
Ok(()) => {
let _ = response.send(Ok(ActorSentMessage {
message_id,
turn_id,
disposition: SendDisposition::Steered,
completion: completion.subscribe(),
}));
}
Err(TurnSteeringPushError::Closed(input)) => self.deferred.push_back(Command::Send {
input,
cancel,
response,
}),
Err(TurnSteeringPushError::Full(_)) => {
let _ = response.send(Err(RunError::SteeringQueueFull));
}
}
}
Some(Command::Inspect { response }) => {
self.deferred.push_back(Command::Inspect { response });
}
Some(Command::Cancel { turn_id: requested, response }) => {
if requested == turn_id {
steering.close();
let _ = response.send(true);
self.hook_state.take_failures();
cancelled = true;
break Some(Ok(Turn::cancelled(turn_id)));
}
let _ = response.send(false);
}
},
}
};
(outcome, cancelled)
};
let finalization = async {
let remaining = steering.close_and_drain();
runtime
.append_accepted_inputs(self.session_id, remaining)
.await?;
if cancelled {
let (_, request) = self
.event_bus
.cancellation_request_for_turn(self.session_id, turn_id);
runtime.host_event_emitter().emit(request).await?;
}
Ok::<_, RunError>(())
}
.await;
let result = match (finalization, outcome.expect("turn outcome set")) {
(Err(error), _) => Err(error),
(Ok(()), Ok(turn)) if turn.stop_reason == TurnStopReason::Cancelled => Ok(turn),
(Ok(()), Ok(mut turn)) => {
turn.hook_failures.extend(self.hook_state.take_failures());
let context = CompletionContext {
agent_name: self.agent.name().to_string(),
session_id: self.session_id,
turn: turn.clone(),
};
turn.hook_failures
.extend(self.hook_state.hooks().run_completion(context).await);
Ok(turn)
}
(Ok(()), Err(error)) => Err(error),
};
let _ = completion.send(TurnCompletion::Ready(result));
true
}
async fn inspect(&mut self) -> Result<crate::SessionContext, RunError> {
self.ensure_runtime().await?;
let context = self
.runtime
.as_ref()
.expect("runtime built above")
.load_context(self.session_id)
.await?;
Ok(crate::SessionContext::from_runtime(
context,
self.agent.plugin_warnings(),
))
}
async fn ensure_runtime(&mut self) -> Result<(), RunError> {
if self.runtime.is_none() {
self.runtime = Some(
self.agent
.build_runtime_with_event_sink(
self.session_id,
self.event_bus.clone(),
self.hook_state.clone(),
)
.await?,
);
}
Ok(())
}
async fn emit_cancelled(&mut self, turn_id: TurnId) -> Result<Turn, RunError> {
self.ensure_runtime().await?;
let runtime = self.runtime.as_ref().expect("runtime built above");
let (_, request) = self
.event_bus
.cancellation_request_for_turn(self.session_id, turn_id);
runtime.host_event_emitter().emit(request).await?;
Ok(Turn::cancelled(turn_id))
}
}
enum HookRun<T> {
Completed(T),
Cancelled,
}
async fn cancellable<T>(
token: Option<&crate::CancellationToken>,
future: impl Future<Output = T>,
) -> HookRun<T> {
match token {
None => HookRun::Completed(future.await),
Some(token) => {
tokio::select! {
biased;
() = token.cancelled() => HookRun::Cancelled,
output = future => HookRun::Completed(output),
}
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum SendDisposition {
Started,
Steered,
}
#[derive(Clone)]
pub struct SentMessage {
pub message_id: String,
pub turn_id: String,
pub disposition: SendDisposition,
turn: TurnHandle,
}
impl SentMessage {
fn new(session: Session, message: ActorSentMessage) -> Self {
let turn = TurnHandle {
session,
turn_id: message.turn_id,
completion: message.completion,
};
Self {
message_id: message.message_id.to_string(),
turn_id: message.turn_id.to_string(),
disposition: message.disposition,
turn,
}
}
pub async fn wait(&self) -> Result<Turn, RunError> {
self.turn.wait().await
}
pub fn turn(&self) -> TurnHandle {
self.turn.clone()
}
}
#[derive(Clone)]
pub struct TurnHandle {
session: Session,
turn_id: TurnId,
completion: watch::Receiver<TurnCompletion>,
}
impl TurnHandle {
pub fn id(&self) -> String {
self.turn_id.to_string()
}
pub async fn wait(&self) -> Result<Turn, RunError> {
let mut completion = self.completion.clone();
loop {
let state = completion.borrow().clone();
match state {
TurnCompletion::Pending => completion
.changed()
.await
.map_err(|_| RunError::SessionClosed)?,
TurnCompletion::Ready(result) => return result,
}
}
}
pub async fn cancel(&self) -> Result<(), CancelError> {
self.session.cancel_turn(self.turn_id).await
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum CancelError {
SessionClosed,
TurnFinished,
}
impl std::fmt::Display for CancelError {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str(match self {
Self::SessionClosed => "session is closed",
Self::TurnFinished => "turn already finished",
})
}
}
impl std::error::Error for CancelError {}
#[derive(Clone, Debug)]
#[non_exhaustive]
pub struct Turn {
pub response: String,
pub turn_id: String,
pub stop_reason: TurnStopReason,
pub iterations: usize,
pub tool_calls: usize,
pub success: bool,
pub error: Option<String>,
pub hook_failures: Vec<HookFailure>,
}
impl Turn {
pub(crate) fn cancelled(turn_id: TurnId) -> Self {
Self {
response: String::new(),
turn_id: turn_id.to_string(),
stop_reason: TurnStopReason::Cancelled,
iterations: 0,
tool_calls: 0,
success: false,
error: Some("turn cancelled".to_string()),
hook_failures: Vec::new(),
}
}
}
impl From<TurnResult> for Turn {
fn from(result: TurnResult) -> Self {
Self {
response: result.response,
turn_id: result.turn_id.to_string(),
stop_reason: result.stop_reason,
iterations: result.iterations,
tool_calls: result.tool_calls_count,
success: result.success,
error: result.error,
hook_failures: Vec::new(),
}
}
}
#[derive(Clone, Debug)]
#[non_exhaustive]
pub enum RunError {
Runtime(Arc<AgentLoopError>),
Hook(HookFailure),
SessionClosed,
SteeringQueueFull,
}
impl std::fmt::Display for RunError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
RunError::Runtime(err) => write!(f, "session run failed: {err}"),
RunError::Hook(err) => write!(f, "session hook failed: {err}"),
RunError::SessionClosed => f.write_str("session is closed"),
RunError::SteeringQueueFull => f.write_str("active turn steering queue is full"),
}
}
}
impl std::error::Error for RunError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
RunError::Runtime(err) => Some(err.as_ref()),
RunError::Hook(err) => Some(err),
RunError::SessionClosed => None,
RunError::SteeringQueueFull => None,
}
}
}
impl From<AgentLoopError> for RunError {
fn from(err: AgentLoopError) -> Self {
RunError::Runtime(Arc::new(err))
}
}
#[cfg(test)]
mod tests {
use std::sync::{Arc, Mutex};
use everruns_core::events::EventData;
use everruns_core::turn::TurnStopReason;
use everruns_core::{ContentPart, InputMessage, MessageRole, TurnId};
use everruns_host::{
EventHistory, EventHistoryReadLimit, EventHistoryReadRequest, EventReadLimit,
EventReadRequest, TurnResult,
};
use super::Turn;
use crate::{Agent, Model};
#[tokio::test]
async fn history_accumulates_across_turns() {
let capture = Arc::new(Mutex::new(Vec::new()));
let agent = Agent::builder()
.instructions("You are concise.")
.model(Model::simulated_capturing("ok", capture.clone()))
.build()
.expect("valid agent");
let session = agent.session();
session.run("hello").await.expect("first turn");
session.run("continue").await.expect("second turn");
let calls = capture.lock().unwrap();
assert_eq!(calls.len(), 2, "two turns => two LLM calls");
assert!(
calls[1].len() > calls[0].len(),
"the second turn's request must include the first turn's messages"
);
}
#[tokio::test]
async fn normal_session_history_is_rebuilt_from_canonical_events() {
let agent = Agent::builder()
.instructions("You are concise.")
.model(Model::simulated("ok"))
.build()
.expect("valid agent");
let session = agent.session();
let session_id = session.session_id();
session.run("hello").await.expect("turn runs");
let event_log = agent
.shared_backends()
.await
.unwrap_or_else(|_| panic!("run built the shared backends"))
.event_log
.clone();
let events = event_log
.read_page(EventReadRequest::new(session_id, EventReadLimit::default()))
.await
.expect("canonical events replay");
assert!(
events.events.iter().all(|event| event.sequence.is_some()),
"durable replay excludes sequence-less live deltas"
);
let canonical_messages: Vec<_> = events
.events
.iter()
.filter_map(|event| match &event.data {
EventData::InputMessage(data) => Some(data.message.clone()),
EventData::OutputMessageCompleted(data) => Some(data.message.clone()),
_ => None,
})
.collect();
let history = EventHistory::new(event_log);
let page = history
.read_page(EventHistoryReadRequest::new(
session_id,
EventHistoryReadLimit::new(8).expect("valid message limit"),
))
.await
.expect("event-derived history page");
assert_eq!(
serde_json::to_value(&page.messages).expect("history serializes"),
serde_json::to_value(&canonical_messages).expect("events serialize")
);
assert_eq!(page.messages.len(), 2);
assert_eq!(page.messages[0].text(), Some("hello"));
assert_eq!(page.messages[1].text(), Some("ok"));
assert!(page.next_cursor.is_none());
}
#[tokio::test]
async fn two_sessions_do_not_share_history() {
let capture = Arc::new(Mutex::new(Vec::new()));
let agent = Agent::builder()
.instructions("You are concise.")
.model(Model::simulated_capturing("ok", capture.clone()))
.build()
.expect("valid agent");
let first = agent.session();
first.run("a1").await.expect("a1");
first.run("a2").await.expect("a2");
let second = agent.session();
second.run("b1").await.expect("b1");
assert_ne!(first.id(), second.id(), "sessions have distinct ids");
let calls = capture.lock().unwrap();
assert_eq!(calls.len(), 3);
assert_eq!(
calls[2].len(),
calls[0].len(),
"a second session must not inherit the first session's history"
);
assert!(calls[1].len() > calls[2].len());
}
#[tokio::test]
async fn accepts_multimodal_input() {
let agent = Agent::builder()
.instructions("You are concise.")
.model(Model::simulated("ok"))
.build()
.expect("valid agent");
let session = agent.session();
let message = InputMessage {
role: MessageRole::User,
content: vec![
ContentPart::text("describe"),
ContentPart::text("this attachment"),
],
controls: None,
metadata: None,
tags: vec![],
};
let turn = session.run(message).await.expect("turn runs");
assert!(turn.success);
}
#[test]
fn turn_preserves_failure_and_stop_reason() {
let result = TurnResult {
response: String::new(),
iterations: 3,
tool_calls_count: 0,
success: false,
error: Some("hit the ceiling".to_string()),
stop_reason: TurnStopReason::MaxTurnRequests,
turn_id: TurnId::new(),
};
let turn = Turn::from(result);
assert!(!turn.success);
assert_eq!(turn.stop_reason, TurnStopReason::MaxTurnRequests);
assert_eq!(turn.error.as_deref(), Some("hit the ceiling"));
}
use std::time::Duration;
use everruns_core::ToolCall;
use serde_json::json;
use crate::{CancellationToken, RunOptions, SessionEvent, SessionEventKind};
async fn drain(mut stream: crate::EventStream) -> Vec<SessionEvent> {
let mut events = Vec::new();
while let Some(event) = stream.recv().await.expect("event stream stays lossless") {
events.push(event);
}
events
}
#[tokio::test]
async fn tool_events_correlate_with_parent_turn() {
let tool = crate::FunctionTool::new(
"ping",
"Respond to a ping.",
json!({ "type": "object", "properties": {} }),
|_args: serde_json::Value| async move { Ok::<_, String>(json!({ "ok": true })) },
);
let agent = Agent::builder()
.instructions("Call ping when asked.")
.model(Model::simulated_scripted(
"done",
vec![
vec![ToolCall {
id: "call_ping_1".into(),
name: "ping".into(),
arguments: json!({}),
}],
vec![],
],
))
.tool(tool)
.build()
.expect("valid agent");
let session = agent.session();
let stream = session.events();
let turn = session.run("please ping").await.expect("turn runs");
assert!(turn.success, "turn should succeed: {:?}", turn.error);
assert_eq!(turn.tool_calls, 1);
drop(session);
let events = drain(stream).await;
let tool_started = events
.iter()
.find(|e| matches!(e.kind, SessionEventKind::ToolStarted { .. }))
.expect("a tool.started event");
let tool_completed = events
.iter()
.find(|e| matches!(e.kind, SessionEventKind::ToolCompleted { .. }))
.expect("a tool.completed event");
assert_eq!(tool_started.turn_id.as_deref(), Some(turn.turn_id.as_str()));
assert_eq!(
tool_completed.turn_id.as_deref(),
Some(turn.turn_id.as_str())
);
let SessionEventKind::ToolStarted {
tool_call_id: started_id,
tool_name,
} = &tool_started.kind
else {
unreachable!("matched ToolStarted above")
};
assert_eq!(tool_name, "ping");
let SessionEventKind::ToolCompleted {
tool_call_id: completed_id,
success,
..
} = &tool_completed.kind
else {
unreachable!("matched ToolCompleted above")
};
assert_eq!(started_id, completed_id, "same tool call across the pair");
assert!(success, "the ping tool succeeded");
}
#[tokio::test]
async fn cancellation_stops_a_running_turn_with_cancelled_stop_reason() {
let agent = Agent::builder()
.instructions("You are slow.")
.model(Model::simulated_delayed(
"eventually",
Duration::from_secs(30),
))
.build()
.expect("valid agent");
let session = agent.session();
let token = CancellationToken::new();
let canceller = token.clone();
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(100)).await;
canceller.cancel();
});
let turn = session
.run_with("hi", RunOptions::new().cancel_token(token))
.await
.expect("run_with resolves");
assert!(!turn.success, "a cancelled turn is not a success");
assert_eq!(turn.stop_reason, TurnStopReason::Cancelled);
}
#[tokio::test]
async fn an_uncancelled_run_with_matches_run() {
let agent = Agent::builder()
.instructions("You are concise.")
.model(Model::simulated("ok"))
.build()
.expect("valid agent");
let session = agent.session();
let turn = session
.run_with("hi", RunOptions::new())
.await
.expect("turn runs");
assert!(turn.success);
assert_eq!(turn.response, "ok");
assert!(turn.hook_failures.is_empty());
}
#[tokio::test]
async fn lifecycle_hooks_wrap_a_tool_call_in_registration_order() {
let order = Arc::new(Mutex::new(Vec::new()));
let tool_order = order.clone();
let tool = crate::FunctionTool::new(
"ping",
"Respond to a ping.",
json!({ "type": "object", "properties": {} }),
move |_args: serde_json::Value| {
let tool_order = tool_order.clone();
async move {
tool_order.lock().unwrap().push("tool");
Ok::<_, String>(json!({ "ok": true }))
}
},
);
let start_one = order.clone();
let start_two = order.clone();
let end_one = order.clone();
let end_two = order.clone();
let completion = order.clone();
let agent = Agent::builder()
.instructions("Call ping when asked.")
.model(Model::simulated_scripted(
"done",
vec![
vec![ToolCall {
id: "call_ping_hooks".into(),
name: "ping".into(),
arguments: json!({}),
}],
vec![],
],
))
.tool(tool)
.on_tool_start(move |context| {
let start_one = start_one.clone();
async move {
assert_eq!(context.tool_name, "ping");
assert!(context.turn_id.is_some());
start_one.lock().unwrap().push("start-1");
}
})
.on_tool_start(move |_context| {
let start_two = start_two.clone();
async move { start_two.lock().unwrap().push("start-2") }
})
.on_tool_end(move |context| {
let end_one = end_one.clone();
async move {
assert!(context.success());
end_one.lock().unwrap().push("end-1");
}
})
.on_tool_end(move |_context| {
let end_two = end_two.clone();
async move { end_two.lock().unwrap().push("end-2") }
})
.on_completion(move |context| {
let completion = completion.clone();
async move {
assert!(context.turn.success);
completion.lock().unwrap().push("completion");
}
})
.build()
.expect("valid agent");
let turn = agent.session().run("please ping").await.expect("turn runs");
assert!(turn.hook_failures.is_empty());
assert_eq!(
*order.lock().unwrap(),
["start-1", "start-2", "tool", "end-1", "end-2", "completion"]
);
}
#[tokio::test]
async fn tool_start_error_blocks_call_and_skips_later_start_hooks() {
use std::sync::atomic::{AtomicBool, Ordering};
let tool_ran = Arc::new(AtomicBool::new(false));
let tool_ran_in_handler = tool_ran.clone();
let later_ran = Arc::new(AtomicBool::new(false));
let later = later_ran.clone();
let end_context = Arc::new(Mutex::new(None));
let end_context_in_hook = end_context.clone();
let tool = crate::FunctionTool::new(
"ping",
"Respond to a ping.",
json!({ "type": "object", "properties": {} }),
move |_args: serde_json::Value| {
let tool_ran_in_handler = tool_ran_in_handler.clone();
async move {
tool_ran_in_handler.store(true, Ordering::SeqCst);
Ok::<_, String>(json!({ "ok": true }))
}
},
);
let agent = Agent::builder()
.instructions("Call ping when asked.")
.model(Model::simulated_scripted(
"recovered",
vec![
vec![ToolCall {
id: "call_blocked_by_framework_hook".into(),
name: "ping".into(),
arguments: json!({}),
}],
vec![],
],
))
.tool(tool)
.on_tool_start(
|_context| async move { Err::<(), _>("policy backend diagnostic: secret") },
)
.on_tool_start(move |_context| {
let later = later.clone();
async move { later.store(true, Ordering::SeqCst) }
})
.on_tool_end(move |context| {
let end_context_in_hook = end_context_in_hook.clone();
async move {
*end_context_in_hook.lock().unwrap() = Some(context);
}
})
.build()
.expect("valid agent");
let turn = agent.session().run("ping").await.expect("turn settles");
assert!(turn.success, "model can recover from a blocked tool call");
assert!(!tool_ran.load(Ordering::SeqCst));
assert!(!later_ran.load(Ordering::SeqCst));
let end_context = end_context.lock().unwrap();
let end_context = end_context.as_ref().expect("blocked call still ends");
assert!(!end_context.success());
let model_visible_error = end_context.error.as_deref().expect("blocked call error");
assert!(model_visible_error.contains("tool call blocked by tool_start hook #0"));
assert!(!model_visible_error.contains("secret"));
assert_eq!(turn.hook_failures.len(), 1);
assert_eq!(turn.hook_failures[0].point, crate::HookPoint::ToolStart);
assert_eq!(
turn.hook_failures[0].message,
"policy backend diagnostic: secret"
);
assert_eq!(
turn.hook_failures[0].tool_call_id.as_deref(),
Some("call_blocked_by_framework_hook")
);
}
#[tokio::test]
async fn tool_end_error_is_isolated_and_later_handlers_run() {
use std::sync::atomic::{AtomicBool, Ordering};
let later_ran = Arc::new(AtomicBool::new(false));
let later = later_ran.clone();
let tool = crate::FunctionTool::new(
"ping",
"Respond to a ping.",
json!({ "type": "object", "properties": {} }),
|_args: serde_json::Value| async move { Ok::<_, String>(json!({ "ok": true })) },
);
let agent = Agent::builder()
.instructions("Call ping when asked.")
.model(Model::simulated_scripted(
"done",
vec![
vec![ToolCall {
id: "call_post_hook_error".into(),
name: "ping".into(),
arguments: json!({}),
}],
vec![],
],
))
.tool(tool)
.on_tool_end(|_context| async move { Err::<(), _>("audit sink offline") })
.on_tool_end(move |_context| {
let later = later.clone();
async move { later.store(true, Ordering::SeqCst) }
})
.build()
.expect("valid agent");
let turn = agent.session().run("ping").await.expect("turn runs");
assert!(turn.success);
assert!(later_ran.load(Ordering::SeqCst));
assert_eq!(turn.hook_failures.len(), 1);
assert_eq!(turn.hook_failures[0].point, crate::HookPoint::ToolEnd);
}
#[tokio::test]
async fn cancellation_drops_an_in_flight_hook_and_skips_remaining_hooks() {
use std::sync::atomic::{AtomicBool, Ordering};
let started = Arc::new(tokio::sync::Notify::new());
let started_in_hook = started.clone();
let later_ran = Arc::new(AtomicBool::new(false));
let later = later_ran.clone();
let completion_ran = Arc::new(AtomicBool::new(false));
let completion = completion_ran.clone();
let agent = Agent::builder()
.instructions("You are concise.")
.model(Model::simulated("unreachable"))
.on_turn_start(move |_context| {
let started_in_hook = started_in_hook.clone();
async move {
started_in_hook.notify_one();
std::future::pending::<()>().await;
}
})
.on_turn_start(move |_context| {
let later = later.clone();
async move { later.store(true, Ordering::SeqCst) }
})
.on_completion(move |_context| {
let completion = completion.clone();
async move { completion.store(true, Ordering::SeqCst) }
})
.build()
.expect("valid agent");
let token = CancellationToken::new();
let canceller = token.clone();
tokio::spawn(async move {
started.notified().await;
canceller.cancel();
});
let turn = tokio::time::timeout(
Duration::from_secs(2),
agent
.session()
.run_with("hello", RunOptions::new().cancel_token(token)),
)
.await
.expect("cancellation is prompt")
.expect("run resolves");
assert_eq!(turn.stop_reason, TurnStopReason::Cancelled);
assert!(!later_ran.load(Ordering::SeqCst));
assert!(!completion_ran.load(Ordering::SeqCst));
}
#[tokio::test]
async fn cancellation_drops_an_in_flight_tool_hook() {
use std::sync::atomic::{AtomicBool, Ordering};
let hook_started = Arc::new(tokio::sync::Notify::new());
let hook_started_inside = hook_started.clone();
let tool_ran = Arc::new(AtomicBool::new(false));
let tool_ran_inside = tool_ran.clone();
let completion_ran = Arc::new(AtomicBool::new(false));
let completion = completion_ran.clone();
let tool = crate::FunctionTool::new(
"ping",
"Respond to a ping.",
json!({ "type": "object", "properties": {} }),
move |_args: serde_json::Value| {
let tool_ran_inside = tool_ran_inside.clone();
async move {
tool_ran_inside.store(true, Ordering::SeqCst);
Ok::<_, String>(json!({ "ok": true }))
}
},
);
let agent = Agent::builder()
.instructions("Call ping when asked.")
.model(Model::simulated_scripted(
"unreachable",
vec![vec![ToolCall {
id: "call_cancelled_hook".into(),
name: "ping".into(),
arguments: json!({}),
}]],
))
.tool(tool)
.on_tool_start(move |_context| {
let hook_started_inside = hook_started_inside.clone();
async move {
hook_started_inside.notify_one();
std::future::pending::<()>().await;
}
})
.on_completion(move |_context| {
let completion = completion.clone();
async move { completion.store(true, Ordering::SeqCst) }
})
.build()
.expect("valid agent");
let token = CancellationToken::new();
let canceller = token.clone();
tokio::spawn(async move {
hook_started.notified().await;
canceller.cancel();
});
let turn = tokio::time::timeout(
Duration::from_secs(2),
agent
.session()
.run_with("ping", RunOptions::new().cancel_token(token)),
)
.await
.expect("cancellation is prompt")
.expect("run resolves");
assert_eq!(turn.stop_reason, TurnStopReason::Cancelled);
assert!(!tool_ran.load(Ordering::SeqCst));
assert!(!completion_ran.load(Ordering::SeqCst));
}
#[tokio::test]
async fn completion_finishes_after_the_runtime_commits_even_if_token_is_cancelled() {
use std::sync::atomic::{AtomicUsize, Ordering};
let token = CancellationToken::new();
let cancel_inside = token.clone();
let completions = Arc::new(AtomicUsize::new(0));
let first = completions.clone();
let second = completions.clone();
let agent = Agent::builder()
.instructions("You are concise.")
.model(Model::simulated("ok"))
.on_completion(move |_context| {
let cancel_inside = cancel_inside.clone();
let first = first.clone();
async move {
cancel_inside.cancel();
tokio::task::yield_now().await;
first.fetch_add(1, Ordering::SeqCst);
}
})
.on_completion(move |_context| {
let second = second.clone();
async move {
second.fetch_add(1, Ordering::SeqCst);
}
})
.build()
.expect("valid agent");
let turn = agent
.session()
.run_with("hello", RunOptions::new().cancel_token(token))
.await
.expect("committed turn completes its hooks");
assert!(turn.success);
assert_eq!(completions.load(Ordering::SeqCst), 2);
}
}