#[path = "common/mod.rs"]
mod common;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use common::{assistant_tool_calls, base_config};
use rpi_agent::{AbortHandle, AgentContext, AgentEvent, AgentToolResult};
use rpi_ai::event_stream::create_assistant_message_event_stream;
use rpi_ai::types::{AssistantMessage, AssistantMessageEvent, ErrorReason, StopReason};
use tokio_util::sync::CancellationToken;
fn saw_agent_end(events: &[AgentEvent]) -> bool {
events
.iter()
.any(|e| matches!(e, AgentEvent::AgentEnd { .. }))
}
#[tokio::test]
async fn abort_during_stream_produces_agent_end() {
let handle = AbortHandle::new();
let token = handle.token();
let stream_fn = rpi_agent::stream_fn(move |_model, _ctx, opts| {
let token = opts.signal.clone();
let (mut prod, stream) = create_assistant_message_event_stream();
tokio::spawn(async move {
let partial = AssistantMessage::empty(
rpi_ai::types::Api::Other("openai-responses".into()),
"mock",
"mock",
0,
);
prod.push(AssistantMessageEvent::Start {
partial: Arc::new(partial),
});
loop {
if token.is_cancelled() {
let aborted = AssistantMessage::terminal(
rpi_ai::types::Api::Other("mock".into()),
"mock",
"mock",
StopReason::Aborted,
"Aborted",
0,
);
prod.push(AssistantMessageEvent::Error {
reason: ErrorReason::Aborted,
error: aborted,
});
return;
}
tokio::time::sleep(std::time::Duration::from_millis(2)).await;
}
});
stream
});
let mut config = base_config();
config.signal = token;
let (collector, events_buf) = rpi_agent::CollectorEmitter::new();
let emit: Arc<dyn rpi_agent::AgentEmitter> = Arc::new(collector);
let run_handle = tokio::spawn(async move {
rpi_agent::run_agent_loop(
vec![common::user_message("hello")],
AgentContext::default(),
config,
emit,
stream_fn,
)
.await
});
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
handle.abort();
let new_messages = run_handle
.await
.expect("run task did not panic")
.expect("run resolves Ok on abort-with-Error-event");
let events = events_buf.lock().expect("events lock").clone();
assert!(
saw_agent_end(&events),
"abort during stream must emit AgentEnd, got: {events:?}"
);
let last = new_messages.last().expect("non-empty new messages");
assert_eq!(last.role().as_str(), "assistant");
}
struct BlockingTool {
schema: rpi_ai::types::Tool,
started: Arc<tokio::sync::Notify>,
}
#[async_trait::async_trait]
impl rpi_agent::AgentTool for BlockingTool {
fn schema(&self) -> &rpi_ai::types::Tool {
&self.schema
}
fn label(&self) -> &str {
"Blocking"
}
async fn execute(
&self,
_tool_call_id: &str,
_params: serde_json::Value,
signal: CancellationToken,
_on_update: Arc<dyn Fn(rpi_agent::ToolResultPartial) + Send + Sync>,
) -> Result<AgentToolResult, rpi_agent::AgentError> {
self.started.notify_one();
signal.cancelled().await;
Ok(AgentToolResult::text("unblocked"))
}
}
fn empty_schema(name: &str) -> rpi_ai::types::Tool {
rpi_ai::types::Tool {
name: name.to_string(),
description: "Blocking tool".to_string(),
parameters: rpi_ai::types::Schema::new(serde_json::json!({
"type": "object",
"properties": {},
"additionalProperties": false,
})),
constrained_sampling: None,
}
}
#[tokio::test]
async fn abort_during_tool_unblocks_and_settles() {
let started = Arc::new(tokio::sync::Notify::new());
let tool = BlockingTool {
schema: empty_schema("blocking"),
started: Arc::clone(&started),
};
let context = AgentContext {
system_prompt: String::new(),
messages: Vec::new(),
tools: vec![Arc::new(tool)],
};
let handle = AbortHandle::new();
let token = handle.token();
let mut config = base_config();
config.signal = token;
let stream_fn = common::mock_stream_fn(vec![assistant_tool_calls(
vec![("tool-1", "blocking", serde_json::json!({}))],
StopReason::ToolUse,
)]);
let (collector, events_buf) = rpi_agent::CollectorEmitter::new();
let emit: Arc<dyn rpi_agent::AgentEmitter> = Arc::new(collector);
let run_handle = tokio::spawn(async move {
rpi_agent::run_agent_loop(
vec![common::user_message("run the tool")],
context,
config,
emit,
stream_fn,
)
.await
});
started.notified().await;
handle.abort();
let _new_messages = run_handle.await.expect("run task did not panic");
let events = events_buf.lock().expect("events lock").clone();
assert!(
saw_agent_end(&events),
"abort during tool must still emit AgentEnd: {events:?}"
);
}
#[tokio::test]
async fn abort_with_no_active_run_is_a_no_op() {
let handle = AbortHandle::new();
handle.abort();
assert!(handle.is_aborted());
handle.abort();
assert!(handle.is_aborted());
let child = handle.child();
let saw = Arc::new(AtomicBool::new(false));
let saw2 = Arc::clone(&saw);
let t = tokio::spawn(async move {
child.cancelled().await;
saw2.store(true, Ordering::SeqCst);
});
t.await.expect("child task did not panic");
assert!(
saw.load(Ordering::SeqCst),
"born-cancelled child should resolve cancelled()"
);
}