From da69c8c6d580b3acf683921b3165e964a074a21a Mon Sep 17 00:00:00 2001 From: Jacob Hull Date: Tue, 1 Sep 2026 15:59:06 -0700 Subject: [PATCH 1/3] feat(agents): give StreamingAgent one entry point returning a run handle Starting an agent went through two methods whose implementations disagreed about cancellation. One remains, returning a handle that owns a run's events, cancellation and usage. RunOptions carries the bound and, for a caller that must cancel before the run exists, a token the run takes a child of. An unbounded run has no bound rather than the Duration::MAX the orchestrator passed for one. Dropping the handle, or the stream it hands out, cancels the run. An orchestration run spawns before its caller sees it, so a consumer that goes away would otherwise leave it spending provider turns. A2A registers a task's cancel entry before the agent build awaits, so a cancelTask arriving during it has something to cancel. make test-loom models that map's interleavings. Ref: #625 Signed-off-by: Jacob Hull --- .makefiles/rust.mk | 4 + Cargo.lock | 35 ++ crates/aura-test-utils/src/mock_agent.rs | 137 +++-- crates/aura-web-server/Cargo.toml | 10 + .../aura-web-server/src/a2a/agent_executor.rs | 502 ++++++++++++++++-- crates/aura-web-server/src/handlers.rs | 14 +- .../aura-web-server/src/streaming/handlers.rs | 46 +- .../tests/agent_event_differential.rs | 15 +- crates/aura/src/builder.rs | 116 ++-- crates/aura/src/orchestration/factory.rs | 89 ++-- crates/aura/src/orchestration/mod.rs | 3 +- crates/aura/src/orchestration/orchestrator.rs | 178 ++----- crates/aura/src/orchestration/test_rig.rs | 19 +- crates/aura/src/provider_agent.rs | 84 ++- crates/aura/src/streaming.rs | 313 +++++++++-- crates/aura/src/streaming_request_hook.rs | 99 ++-- 16 files changed, 1088 insertions(+), 576 deletions(-) diff --git a/.makefiles/rust.mk b/.makefiles/rust.mk index 39b69a87c..f4979f77c 100644 --- a/.makefiles/rust.mk +++ b/.makefiles/rust.mk @@ -57,6 +57,10 @@ coverage: $(DOCKER_ENV) $(REPORT_DIR) $(GRCOV_BIN) ## Run the local test suite w nextest: $(DOCKER_ENV) $(NEXTEST_BIN) $(REPORT_DIR) $(RUN) cargo nextest run --workspace --all-targets --features integration $(if $(IS_CI),-P ci,) +.PHONY:test-loom +test-loom: $(DOCKER_ENV) ## Run the loom interleaving models (not part of `test`; they explore, so they are slow) + $(RUN) env RUSTFLAGS="--cfg aura_loom" cargo test -p aura-web-server --lib loom_ + .PHONY:lint-rust lint-rust: | $(DOCKER_ENV) $(REPORT_DIR) ## lint rust code via clippy $(RUN) cargo clippy $(if $(IS_CI),-q,) --all-targets --all-features $(if $(IS_CI),--message-format=json,) -- -D warnings $(if $(IS_CI),> $(REPORT_DIR)/clippy.json,) diff --git a/Cargo.lock b/Cargo.lock index 415a02d43..f55545425 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -394,6 +394,7 @@ dependencies = [ "futures", "futures-util", "http-body-util", + "loom", "opentelemetry", "redis", "reqwest", @@ -2026,6 +2027,21 @@ dependencies = [ "sha2 0.10.9", ] +[[package]] +name = "generator" +version = "0.8.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "54ade96dc9003043bce7c035c85a9df5a858bfb2039c5a2e6fdf00f324f6c551" +dependencies = [ + "cc", + "cfg-if", + "libc", + "log", + "rustversion", + "windows-link", + "windows-result", +] + [[package]] name = "generic-array" version = "0.14.7" @@ -2729,6 +2745,19 @@ version = "0.4.29" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897" +[[package]] +name = "loom" +version = "0.7.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "419e0dc8046cb947daa77eb95ae174acfbddb7673b4151f56d1eed8e93fbfaca" +dependencies = [ + "cfg-if", + "generator", + "scoped-tls", + "tracing", + "tracing-subscriber", +] + [[package]] name = "lru-slab" version = "0.1.2" @@ -4132,6 +4161,12 @@ dependencies = [ "syn 2.0.117", ] +[[package]] +name = "scoped-tls" +version = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e1cf6437eb19a8f4a6cc0f7dca544973b0b78843adbfeb3683d1a94a0024a294" + [[package]] name = "scopeguard" version = "1.2.0" diff --git a/crates/aura-test-utils/src/mock_agent.rs b/crates/aura-test-utils/src/mock_agent.rs index e429385e3..8ff8eb0dd 100644 --- a/crates/aura-test-utils/src/mock_agent.rs +++ b/crates/aura-test-utils/src/mock_agent.rs @@ -6,11 +6,10 @@ use std::sync::{Arc, Mutex}; use std::time::Duration; use async_trait::async_trait; +use aura::streaming::AgentRun; use aura::{Message, StreamError, StreamItem, StreamingAgent, UsageState}; use aura::{StreamedAssistantContent, StreamedUserContent, ToolCall, ToolResult}; use futures::stream::{self, BoxStream, StreamExt}; -use tokio::sync::watch; -use tokio_util::sync::CancellationToken; type StartHook = Arc Pin + Send>> + Send + Sync>; type EffectHook = Arc Pin + Send>> + Send + Sync>; @@ -107,7 +106,7 @@ impl MockAgent { } /// Runs the given steps in order, then ends the stream. The script is - /// consumed by the first `stream`/`stream_with_timeout` call; a second call + /// consumed by the first `stream` call; a second call /// on the same agent yields an empty stream. pub fn scripted(steps: Vec) -> Self { Self { @@ -116,7 +115,7 @@ impl MockAgent { } } - /// Awaited before either entry point produces its stream, and passed that + /// Awaited before `stream` produces its stream, and passed that /// call's `request_id`. pub fn on_stream_start(mut self, hook: F) -> Self where @@ -175,26 +174,18 @@ impl StreamingAgent for MockAgent { &self, _query: &str, _chat_history: Vec, - _cancel_token: CancellationToken, + options: aura::streaming::RunOptions, request_id: &str, - ) -> Result>, StreamError> { - Ok(self.start(request_id).await) - } - - async fn stream_with_timeout( - &self, - _query: &str, - _chat_history: Vec, - _timeout: Duration, - request_id: &str, - ) -> ( - BoxStream<'static, Result>, - watch::Sender, - UsageState, - ) { + ) -> AgentRun { let stream = self.start(request_id).await; - let (cancel_tx, _cancel_rx) = watch::channel(false); - (stream, cancel_tx, UsageState::new()) + // Carries a caller-supplied token so `cancel_token()` returns the one the + // caller named. The scripts do not race it, so cancelling does not end a + // mock stream. + AgentRun::new( + stream, + options.into_parts().1.unwrap_or_default(), + UsageState::new(), + ) } async fn cancel_and_close_mcp(&self, _request_id: &str, _reason: &str) -> usize { @@ -211,9 +202,9 @@ mod tests { async fn a_pending_agent_never_yields() { let agent = MockAgent::pending(); let mut stream = agent - .stream("q", vec![], CancellationToken::new(), "req_1") + .stream("q", vec![], aura::streaming::RunOptions::default(), "req_1") .await - .expect("mock stream should start"); + .into_events(); assert!( tokio::time::timeout(Duration::from_millis(50), stream.next()) .await @@ -221,58 +212,45 @@ mod tests { ); } + /// `into_events` wraps the stream, so the order is asserted rather than the + /// count — a wrapper that buffered or reordered would keep the count. #[tokio::test] - async fn the_start_hook_runs_before_both_entry_points_stream() { - for entry_point in ["stream", "stream_with_timeout"] { - let ran = Arc::new(AtomicBool::new(false)); - let flag = Arc::clone(&ran); - let agent = MockAgent::pending().on_stream_start(move |request_id| { - let flag = Arc::clone(&flag); - async move { - assert_eq!(request_id, "req_1"); - flag.store(true, Ordering::SeqCst); - } - }); - - if entry_point == "stream" { - let _ = agent - .stream("q", vec![], CancellationToken::new(), "req_1") - .await - .expect("mock stream should start"); - } else { - let _ = agent - .stream_with_timeout("q", vec![], Duration::from_secs(1), "req_1") - .await; + async fn a_yielding_agent_produces_its_items_then_ends() { + let agent = MockAgent::yielding(vec![items::text("hello "), items::text("world")]); + let mut stream = agent + .stream("q", vec![], aura::streaming::RunOptions::default(), "req_1") + .await + .into_events(); + + let mut texts = Vec::new(); + while let Some(item) = stream.next().await { + if let Ok(aura::StreamItem::StreamAssistantItem(StreamedAssistantContent::Text(t))) = + item + { + texts.push(t); } - - assert!( - ran.load(Ordering::SeqCst), - "hook should run for {entry_point}" - ); } + assert_eq!(texts, ["hello ", "world"]); } - #[tokio::test(start_paused = true)] - async fn a_yielding_agent_produces_its_items_then_ends() { - let agent = MockAgent::yielding([items::text("hello "), items::text("world")]); - let stream = agent - .stream("q", vec![], CancellationToken::new(), "req_1") + #[tokio::test] + async fn the_start_hook_runs_before_the_stream() { + let ran = Arc::new(AtomicBool::new(false)); + let flag = Arc::clone(&ran); + let agent = MockAgent::pending().on_stream_start(move |request_id| { + let flag = Arc::clone(&flag); + async move { + assert_eq!(request_id, "req_1"); + flag.store(true, Ordering::SeqCst); + } + }); + + let _ = agent + .stream("q", vec![], aura::streaming::RunOptions::default(), "req_1") .await - .expect("mock stream should start"); - - let texts: Vec = stream - .filter_map(|item| async move { - match item { - Ok(StreamItem::StreamAssistantItem(StreamedAssistantContent::Text(t))) => { - Some(t) - } - _ => None, - } - }) - .collect() - .await; + .into_events(); - assert_eq!(texts, vec!["hello ", "world"]); + assert!(ran.load(Ordering::SeqCst), "hook should run"); } #[tokio::test(start_paused = true)] @@ -290,9 +268,14 @@ mod tests { ]); let stream = agent - .stream("q", vec![], CancellationToken::new(), "req_42") + .stream( + "q", + vec![], + aura::streaming::RunOptions::default(), + "req_42", + ) .await - .expect("mock stream should start"); + .into_events(); let items: Vec<_> = stream.collect().await; assert_eq!(order.lock().expect("order lock").as_slice(), ["req_42"]); @@ -304,15 +287,15 @@ mod tests { let agent = MockAgent::yielding([items::text("once")]); let first: Vec<_> = agent - .stream("q", vec![], CancellationToken::new(), "req_1") + .stream("q", vec![], aura::streaming::RunOptions::default(), "req_1") .await - .expect("mock stream should start") + .into_events() .collect() .await; let second: Vec<_> = agent - .stream("q", vec![], CancellationToken::new(), "req_1") + .stream("q", vec![], aura::streaming::RunOptions::default(), "req_1") .await - .expect("mock stream should start") + .into_events() .collect() .await; @@ -337,9 +320,9 @@ mod tests { ]); let mut stream = agent - .stream("q", vec![], CancellationToken::new(), "req_1") + .stream("q", vec![], aura::streaming::RunOptions::default(), "req_1") .await - .expect("mock stream should start"); + .into_events(); while stream.next().await.is_some() { seen.fetch_add(1, Ordering::SeqCst); } diff --git a/crates/aura-web-server/Cargo.toml b/crates/aura-web-server/Cargo.toml index af24abaa5..cab76a00e 100644 --- a/crates/aura-web-server/Cargo.toml +++ b/crates/aura-web-server/Cargo.toml @@ -98,3 +98,13 @@ aura-events = { path = "../aura-events" } tower = "0.5" http-body-util = "0.1" tempfile = "3" + +# Exhaustively explores interleavings of the task-cancel map, which a stress +# test cannot reach: the window is one mutex release and reacquire wide. +[lints.rust] +unexpected_cfgs = { level = "warn", check-cfg = ['cfg(aura_loom)'] } + +# Not a dev-dependency: the lib itself swaps in loom's mutex under this cfg, +# which is never set by a normal build. +[target.'cfg(aura_loom)'.dependencies] +loom = "0.7" diff --git a/crates/aura-web-server/src/a2a/agent_executor.rs b/crates/aura-web-server/src/a2a/agent_executor.rs index 1119ee493..404ae5610 100644 --- a/crates/aura-web-server/src/a2a/agent_executor.rs +++ b/crates/aura-web-server/src/a2a/agent_executor.rs @@ -1,5 +1,12 @@ use std::collections::HashMap; -use std::sync::{Arc, Mutex, MutexGuard}; +use std::sync::Arc; +// Under `--cfg aura_loom` the cancel map's mutex is loom's, so its model can +// explore every interleaving of the code that locks it. The cfg is ours rather +// than loom's usual `loom`, which tokio also reacts to. +#[cfg(aura_loom)] +use loom::sync::{Mutex, MutexGuard}; +#[cfg(not(aura_loom))] +use std::sync::{Mutex, MutexGuard}; use a2a::{ A2AError, AgentCapabilities, AgentCard, AgentInterface, AgentSkill, Artifact, ListTasksRequest, @@ -35,13 +42,64 @@ pub struct AuraAgentExecutor { task_cancel_state: Arc, } +/// A live execution's cancellation handle and what a cancel needs to clean up. struct TaskCancelEntry { token: CancellationToken, - agent: Arc, + agent: Option>, request_id: String, } -/// The live executions' cancel handles, keyed by task id. +/// Whether `cancel()` got to a task before the code asking. +#[derive(PartialEq, Eq)] +enum CancelRaced { + Yes, + No, +} + +/// Hands the built agent to the task's entry, reporting whether a cancel beat it. +/// +/// One lock acquisition, because a cancel landing between the write and the +/// answer would take the entry with the agent set, close its MCP calls, and +/// leave the caller to close them again. +fn claim_agent( + state: &TaskCancelState, + task_id: &str, + agent: &Arc, +) -> CancelRaced { + match lock_cancel_state(state).get_mut(task_id) { + Some(entry) => { + entry.agent = Some(Arc::clone(agent)); + CancelRaced::No + } + // `cancel()` takes the entry before firing the token, so a missing + // entry means it has already run. + None => CancelRaced::Yes, + } +} + +/// The terminal status this run emits once its loop ends, if any. +/// +/// Taking the entry claims the terminal status. `cancel()` takes it before +/// answering `Canceled`, so a run that finds it gone answers nothing, however +/// its loop ended. A run that still holds it answers `Canceled` when shutdown +/// stopped it, and `Completed` when it succeeded. +fn post_loop_status( + success: bool, + stopped_by_cancel: bool, + shutdown: &CancellationToken, + state: &TaskCancelState, + task_id: &str, +) -> Option { + // Taking rather than reading keeps a cancel from landing between the two + // and leaving both sides to answer. + lock_cancel_state(state).remove(task_id)?; + if stopped_by_cancel { + return shutdown.is_cancelled().then_some(TaskState::Canceled); + } + success.then_some(TaskState::Completed) +} + +/// The live executions, keyed by task id. type TaskCancelState = Mutex>; /// Lock the cancel map, taking a poisoned lock's contents rather than @@ -70,6 +128,15 @@ impl Drop for TaskCancelGuard { /// this composes with the explicit removals. fn drop(&mut self) { lock_cancel_state(&self.state).remove(&self.task_id); + // The streaming hook keys a tool-call FIFO under this request id, and + // a run cancelled between a tool call and its result leaves an entry + // behind. The broker is async, so this runs as its own task — which a + // runtime already shutting down may never poll, leaving the entry for + // the process to reclaim. + if let Ok(handle) = tokio::runtime::Handle::try_current() { + let request_id = self.request_id.clone(); + handle.spawn(async move { aura::tool_event_unsubscribe(&request_id).await }); + } RequestCancellation::unregister(&self.request_id); } } @@ -226,6 +293,25 @@ impl AgentExecutor for AuraAgentExecutor { })); let request_id = format!("a2a_{}", task_id); + + // Registered before the agent build and history fetch, both of which + // await, so a cancelTask during those has a token to cancel. Its + // guard comes first, because those paths can return early. + let cancel_token = stream_shutdown_token.child_token(); + let _cancel_guard = TaskCancelGuard { + state: task_cancel_state.clone(), + task_id: task_id.clone(), + request_id: request_id.clone(), + }; + lock_cancel_state(&task_cancel_state).insert( + task_id.clone(), + TaskCancelEntry { + token: cancel_token.clone(), + agent: None, + request_id: request_id.clone(), + }, + ); + let session_id = Some(context_id.clone()); let builder = RigBuilder::new(config, pending_approvals).with_hitl_hmac(hitl_hmac); let agent = match builder @@ -247,28 +333,30 @@ impl AgentExecutor for AuraAgentExecutor { // build any history for this context that can be used in further aura reasoning let history = get_history_for_context(task_store.clone(), &request_id, &context_id, &task_id).await?; - let cancel_token = stream_shutdown_token.child_token(); // Register with the global cancellation registry for parity with the OpenAI handler // and to let any future code address this request by id. RequestCancellation::register(request_id.clone()); - let _cancel_guard = TaskCancelGuard { - state: task_cancel_state.clone(), - task_id: task_id.clone(), - request_id: request_id.clone(), - }; - lock_cancel_state(&task_cancel_state).insert(task_id.clone(), TaskCancelEntry { - token: cancel_token.clone(), - agent: agent.clone(), - request_id: request_id.clone(), - }); + // The agent exists now, so a cancel from here can close its MCP calls. + // A missing entry means `cancel()` already ran while the build was in + // flight, with no agent to close — so the run does it instead. + if claim_agent(&task_cancel_state, &task_id, &agent) == CancelRaced::Yes { + agent + .cancel_and_close_mcp(&request_id, "A2A cancelTask during build") + .await; + } - let mut stream = match agent.stream(&text, history, cancel_token.clone(), &request_id).await { - Ok(s) => s, - Err(e) => { - yield Ok(fail_status(&task_id, &context_id, &e.to_string())); - return; - } - }; + // A2A tasks have their own lifetime, so the run is unbounded here and + // ends on cancelTask or shutdown. + let run = agent + .stream( + &text, + history, + aura::streaming::RunOptions::default() + .cancelled_by(&cancel_token), + &request_id, + ) + .await; + let mut stream = run.into_events(); // RAII guard: drop on any generator exit (loop break, early return, panic, // consumer drop) produces exactly one decrement. Replaces the manual @@ -280,10 +368,17 @@ impl AgentExecutor for AuraAgentExecutor { let mut success = true; // assume everything is successful let mut reasoning_num = 0; + // Read after the loop, because a token cancelled once the stream has + // already ended says nothing about how this run finished. + let mut stopped_by_cancel = false; loop { let next = tokio::select! { biased; - _ = cancel_token.cancelled() => break, + // A child of the shutdown token, so this covers both. + () = cancel_token.cancelled() => { + stopped_by_cancel = true; + break; + } next = stream.next() => next, }; let Some(item) = next else { break }; @@ -459,43 +554,31 @@ impl AgentExecutor for AuraAgentExecutor { } } - // If cancel_token fired but our entry is still in the map, the cancel came - // from the parent stream_shutdown_token (server shutdown), not from our - // cancel() hook — cancel() removes its entry before firing the token. - // In that case the executor has to drive MCP cleanup itself and emit a - // terminal Canceled status (the OpenAI handler does the equivalent in its - // Shutdown post-loop arm). - let entry_still_present = lock_cancel_state(&task_cancel_state).remove(&task_id).is_some(); - let shutdown_initiated_cancel = cancel_token.is_cancelled() && entry_still_present; + let status = post_loop_status( + success, + stopped_by_cancel, + &stream_shutdown_token, + &task_cancel_state, + &task_id, + ); RequestCancellation::unregister(&request_id); - if shutdown_initiated_cancel { + // Shutdown is the one cancel the executor cleans up after itself; + // `cancel()` drives its own. + if status == Some(TaskState::Canceled) { agent .cancel_and_close_mcp(&request_id, "server shutdown") .await; - - yield Ok(StreamResponse::StatusUpdate(TaskStatusUpdateEvent { - task_id: task_id.clone(), - context_id: context_id.clone(), - status: TaskStatus { - state: TaskState::Canceled, - message: None, - timestamp: Some(chrono::Utc::now()), - }, - metadata: None, - })); } // _request_guard drops at end of generator scope → exactly one decrement. - // Skip Completed if cancel() or the shutdown path already emitted Canceled — - // yielding here would clobber it. - if success && !cancel_token.is_cancelled() { + if let Some(state) = status { yield Ok(StreamResponse::StatusUpdate(TaskStatusUpdateEvent { task_id, context_id, status: TaskStatus { - state: TaskState::Completed, + state, message: None, timestamp: Some(chrono::Utc::now()), }, @@ -520,10 +603,12 @@ impl AgentExecutor for AuraAgentExecutor { if let Some(entry) = entry { // Send notifications/cancelled to in-flight MCP tool calls. No-op in // orchestration mode (workers manage their own MCP cancellation). - entry - .agent - .cancel_and_close_mcp(&entry.request_id, "A2A cancelTask") - .await; + // A run cancelled before its agent was built has no MCP calls yet. + if let Some(agent) = &entry.agent { + agent + .cancel_and_close_mcp(&entry.request_id, "A2A cancelTask") + .await; + } entry.token.cancel(); RequestCancellation::unregister(&entry.request_id); } @@ -865,7 +950,7 @@ mod tests { task_id.clone(), TaskCancelEntry { token: CancellationToken::new(), - agent: Arc::new(MockAgent::pending()), + agent: Some(Arc::new(MockAgent::pending())), request_id: request_id.clone(), }, ); @@ -877,6 +962,152 @@ mod tests { assert!(RequestCancellation::token_for_id(&request_id).is_none()); } + /// Every way a run's loop can end, and the status each one answers with. + #[test] + fn the_post_loop_status_follows_how_the_run_ended() { + // Deciding takes the entry, so each case starts from its own. + let registered = || { + let task_id = format!("t_{}", uuid::Uuid::new_v4()); + let state: Arc = Arc::new(Mutex::new(HashMap::new())); + lock_cancel_state(&state).insert( + task_id.clone(), + TaskCancelEntry { + token: CancellationToken::new(), + agent: None, + request_id: format!("a2a_{task_id}"), + }, + ); + (state, task_id) + }; + let down = || { + let token = CancellationToken::new(); + token.cancel(); + token + }; + + let (state, task_id) = registered(); + assert_eq!( + post_loop_status(true, false, &CancellationToken::new(), &state, &task_id), + Some(TaskState::Completed), + "a run that reached the end of its stream completed" + ); + + let (state, task_id) = registered(); + assert_eq!( + post_loop_status(true, false, &down(), &state, &task_id), + Some(TaskState::Completed), + "shutdown firing after a run finished does not make it cancelled" + ); + + let (state, task_id) = registered(); + assert_eq!( + post_loop_status(false, false, &CancellationToken::new(), &state, &task_id), + None, + "a run that failed has already answered" + ); + + let (state, task_id) = registered(); + assert_eq!( + post_loop_status(true, true, &down(), &state, &task_id), + Some(TaskState::Canceled), + "a run that stopped on its token with its entry intact was stopped \ + by shutdown" + ); + + // `cancel()` takes the entry, then fires the token. + let (state, task_id) = registered(); + lock_cancel_state(&state).remove(&task_id); + assert_eq!( + post_loop_status(true, true, &down(), &state, &task_id), + None, + "cancel() answers for itself" + ); + } + + /// `cancel()` can take the entry while the run is parked yielding its final + /// artifact, after which the loop ends on its own. The cancel has already + /// answered `Canceled`, so the run must not also answer `Completed`. + #[test] + fn a_run_does_not_complete_after_a_cancel_took_its_entry() { + let state: Arc = Arc::new(Mutex::new(HashMap::new())); + lock_cancel_state(&state).insert( + "t".to_string(), + TaskCancelEntry { + token: CancellationToken::new(), + agent: None, + request_id: "a2a_t".to_string(), + }, + ); + lock_cancel_state(&state).remove("t"); + + assert_eq!( + post_loop_status(true, false, &CancellationToken::new(), &state, "t"), + None + ); + } + + /// Covers the contract in both directions. `loom_tests` covers the + /// interleaving, which this cannot reach. + #[test] + fn claiming_the_agent_reports_a_cancel_that_already_ran() { + let task_id = format!("t_{}", uuid::Uuid::new_v4()); + let request_id = format!("a2a_{task_id}"); + let state: Arc = Arc::new(Mutex::new(HashMap::new())); + let agent: Arc = Arc::new(MockAgent::pending()); + + lock_cancel_state(&state).insert( + task_id.clone(), + TaskCancelEntry { + token: CancellationToken::new(), + agent: None, + request_id, + }, + ); + + assert!(claim_agent(&state, &task_id, &agent) == CancelRaced::No); + assert!( + lock_cancel_state(&state) + .get(&task_id) + .is_some_and(|entry| entry.agent.is_some()), + "the entry carries the agent a later cancel closes" + ); + + // What `cancel()` does: take the entry, leaving nothing to claim. + lock_cancel_state(&state).remove(&task_id); + assert!(claim_agent(&state, &task_id, &agent) == CancelRaced::Yes); + } + + /// The guard is created before the entry, because the agent build and the + /// history fetch can both return early and would otherwise leave it behind. + #[test] + fn an_early_return_before_the_run_releases_the_entry() { + let task_id = format!("t_{}", uuid::Uuid::new_v4()); + let request_id = format!("a2a_{task_id}"); + let state: Arc = Arc::new(Mutex::new(HashMap::new())); + + { + let _guard = TaskCancelGuard { + state: state.clone(), + task_id: task_id.clone(), + request_id: request_id.clone(), + }; + lock_cancel_state(&state).insert( + task_id.clone(), + TaskCancelEntry { + token: CancellationToken::new(), + agent: None, + request_id: request_id.clone(), + }, + ); + // The build fails here and the generator returns. + } + + assert!( + !lock_cancel_state(&state).contains_key(&task_id), + "a run that never started leaves no entry" + ); + } + /// `cancel()` takes the entry before the generator unwinds, so the guard /// has to tolerate finding both already released. #[test] @@ -1098,3 +1329,172 @@ mod tests { ); } } + +/// Exhaustive interleaving checks for the task-cancel map. +/// +/// Run with `RUSTFLAGS="--cfg aura_loom" cargo test -p aura-web-server loom_`. +/// A stress test cannot reach these windows: the one that matters is a single +/// mutex release and reacquire wide. +#[cfg(all(test, aura_loom))] +mod loom_tests { + use super::*; + use aura_test_utils::mock_agent::MockAgent; + + /// A cancel and the post-loop check race over a run shutdown stopped. + /// Exactly one of them takes the entry, and with it the job of closing the + /// agent. + #[test] + fn loom_a_cancel_racing_the_shutdown_check_yields_one_closer() { + loom::model(|| { + let shutdown = CancellationToken::new(); + shutdown.cancel(); + + let state: loom::sync::Arc = + loom::sync::Arc::new(Mutex::new(HashMap::new())); + lock_cancel_state(&state).insert( + "t1".to_string(), + TaskCancelEntry { + token: CancellationToken::new(), + agent: None, + request_id: "a2a_t1".to_string(), + }, + ); + + let closes = loom::sync::Arc::new(loom::sync::atomic::AtomicUsize::new(0)); + + let post_loop = { + let (state, closes, shutdown) = (state.clone(), closes.clone(), shutdown.clone()); + loom::thread::spawn(move || { + if post_loop_status(true, true, &shutdown, &state, "t1") + == Some(TaskState::Canceled) + { + closes.fetch_add(1, loom::sync::atomic::Ordering::SeqCst); + } + }) + }; + + let canceller = { + let (state, closes) = (state.clone(), closes.clone()); + loom::thread::spawn(move || { + if lock_cancel_state(&state).remove("t1").is_some() { + closes.fetch_add(1, loom::sync::atomic::Ordering::SeqCst); + } + }) + }; + + post_loop.join().unwrap(); + canceller.join().unwrap(); + + assert_eq!( + closes.load(loom::sync::atomic::Ordering::SeqCst), + 1, + "exactly one side takes responsibility for the terminal status" + ); + }); + } + + /// A run that reached the end of its stream races a `cancel()` for its + /// entry. Exactly one of them takes it, and the run answers `Completed` + /// only when it did. This counts claims rather than statuses, since + /// `cancel()` answers `Canceled` even when it finds no entry. + #[test] + fn loom_a_cancel_racing_a_finished_run_yields_one_claimant() { + loom::model(|| { + let shutdown = CancellationToken::new(); + + let state: loom::sync::Arc = + loom::sync::Arc::new(Mutex::new(HashMap::new())); + lock_cancel_state(&state).insert( + "t1".to_string(), + TaskCancelEntry { + token: CancellationToken::new(), + agent: None, + request_id: "a2a_t1".to_string(), + }, + ); + + let claims = loom::sync::Arc::new(loom::sync::atomic::AtomicUsize::new(0)); + + let post_loop = { + let (state, claims, shutdown) = (state.clone(), claims.clone(), shutdown.clone()); + loom::thread::spawn(move || { + if post_loop_status(true, false, &shutdown, &state, "t1") + == Some(TaskState::Completed) + { + claims.fetch_add(1, loom::sync::atomic::Ordering::SeqCst); + } + }) + }; + + let canceller = { + let (state, claims) = (state.clone(), claims.clone()); + loom::thread::spawn(move || { + if lock_cancel_state(&state).remove("t1").is_some() { + claims.fetch_add(1, loom::sync::atomic::Ordering::SeqCst); + } + }) + }; + + post_loop.join().unwrap(); + canceller.join().unwrap(); + + assert_eq!( + claims.load(loom::sync::atomic::Ordering::SeqCst), + 1, + "exactly one side claims the terminal status" + ); + }); + } + + /// A cancel and a claim race. Exactly one of them takes responsibility for + /// closing the agent — the claim when it finds the entry already gone, the + /// cancel when it finds the entry carrying an agent. The close itself is + /// async and outside the model; this is the branch that decides who runs it. + #[test] + fn loom_a_cancel_racing_a_claim_yields_one_closer() { + loom::model(|| { + let state: loom::sync::Arc = + loom::sync::Arc::new(Mutex::new(HashMap::new())); + lock_cancel_state(&state).insert( + "t1".to_string(), + TaskCancelEntry { + token: CancellationToken::new(), + agent: None, + request_id: "a2a_t1".to_string(), + }, + ); + + let closes = loom::sync::Arc::new(loom::sync::atomic::AtomicUsize::new(0)); + + let claimer = { + let (state, closes) = (state.clone(), closes.clone()); + loom::thread::spawn(move || { + let agent: Arc = Arc::new(MockAgent::pending()); + if claim_agent(&state, "t1", &agent) == CancelRaced::Yes { + closes.fetch_add(1, loom::sync::atomic::Ordering::SeqCst); + } + }) + }; + + let canceller = { + let (state, closes) = (state.clone(), closes.clone()); + loom::thread::spawn(move || { + // What `cancel()` does: take the entry, close what it holds. + let entry = lock_cancel_state(&state).remove("t1"); + if entry.is_some_and(|entry| entry.agent.is_some()) { + closes.fetch_add(1, loom::sync::atomic::Ordering::SeqCst); + } + }) + }; + + claimer.join().unwrap(); + canceller.join().unwrap(); + + assert_eq!( + closes.load(loom::sync::atomic::Ordering::SeqCst), + 1, + "exactly one side takes responsibility for closing the agent" + ); + }); + } +} diff --git a/crates/aura-web-server/src/handlers.rs b/crates/aura-web-server/src/handlers.rs index 0d1d0b37c..1d322af1d 100644 --- a/crates/aura-web-server/src/handlers.rs +++ b/crates/aura-web-server/src/handlers.rs @@ -544,7 +544,7 @@ pub async fn execute_completion( tools_json, } = setup; - // Orchestration spawns inside `stream_with_timeout`, so SSE side-channel + // Orchestration spawns inside `stream`, so SSE side-channel // receivers must be subscribed before stream startup. let delivery_channels = match delivery { DeliveryMode::Collect { result_tx } => DeliveryChannels::Collect { result_tx }, @@ -562,14 +562,18 @@ pub async fn execute_completion( }; // Create stream with timeout — single path for both Agent and Orchestrator - let (stream, cancel_tx, usage_state) = streaming_agent - .stream_with_timeout( + let run = streaming_agent + .stream( &query, chat_history, - config.timeout_duration, + aura::streaming::RunOptions::bounded(Some(config.timeout_duration)) + .cancelled_by(&config.stream_shutdown_token), &config.request_id, ) .await; + let cancel_tx = run.cancel_token(); + let usage_state = run.usage().clone(); + let stream = run.into_events(); let response_content = config.response_content.clone(); let otel_ctx = StreamOtelContext { @@ -632,7 +636,7 @@ pub async fn execute_completion( stream, chunk_tx, cancel_tx, - config.timeout_duration, + Some(config.timeout_duration), heartbeat_interval, config.first_chunk_timeout, config.inactivity_timeout, diff --git a/crates/aura-web-server/src/streaming/handlers.rs b/crates/aura-web-server/src/streaming/handlers.rs index dd7294cb4..f75310b6f 100644 --- a/crates/aura-web-server/src/streaming/handlers.rs +++ b/crates/aura-web-server/src/streaming/handlers.rs @@ -20,6 +20,7 @@ use crate::streaming::types::openai::UsageInfo; use aura_events::agent::{AgentEvent, AgentEventPayload}; +use tokio_util::sync::CancellationToken; use super::types::{ CHUNK_OBJECT, ChatCompletionChunk, ChatCompletionChunkChoice, ChatCompletionChunkDelta, @@ -38,7 +39,7 @@ use bytes::Bytes; use futures_util::{Stream, StreamExt}; use std::sync::Arc; use std::time::Duration; -use tokio::sync::{mpsc, watch}; +use tokio::sync::mpsc; /// Context for cancellation and cleanup callbacks. pub struct StreamingCallbacks { @@ -103,8 +104,8 @@ pub async fn process_sse_stream_full( ctx: &TurnContext, mut stream: S, tx: mpsc::Sender>, - cancel_tx: watch::Sender, - timeout_duration: Duration, + cancel_tx: CancellationToken, + timeout_duration: Option, heartbeat_interval: Duration, first_chunk_timeout: Option, inactivity_timeout: Option, @@ -158,8 +159,14 @@ where } } - // Safety net timeout - let timeout = tokio::time::sleep(timeout_duration); + // Safety net timeout. An unbounded run waits on a future that never + // resolves rather than a zero sleep that fires on the first poll. + let timeout = async move { + match timeout_duration { + Some(duration) => tokio::time::sleep(duration).await, + None => std::future::pending().await, + } + }; tokio::pin!(timeout); // Heartbeat for proactive disconnect detection during silent tool execution @@ -360,7 +367,7 @@ where _ = &mut timeout => { tracing::warn!( "Streaming safety net timeout ({:?}) - signaling cancellation", - timeout_duration + timeout_duration.unwrap_or_default() ); break StreamTermination::Timeout; } @@ -411,13 +418,13 @@ where } StreamTermination::Disconnected => { - let _ = cancel_tx.send(true); + cancel_tx.cancel(); RequestCancellation::cancel(&callbacks.request_id, "client disconnected"); cancel_mcp(&callbacks, "client disconnected").await; } StreamTermination::Timeout => { - let _ = cancel_tx.send(true); + cancel_tx.cancel(); RequestCancellation::cancel(&callbacks.request_id, "timeout"); cancel_mcp(&callbacks, "timeout").await; send_final_events(emit_custom_events, &mut callbacks, ctx, &state, &tx).await; @@ -425,7 +432,7 @@ where StreamTermination::Shutdown => { // [DONE] before MCP cleanup so client gets clean termination regardless of MCP latency - let _ = cancel_tx.send(true); + cancel_tx.cancel(); RequestCancellation::cancel(&callbacks.request_id, "server shutdown"); send_final_events(emit_custom_events, &mut callbacks, ctx, &state, &tx).await; cancel_mcp(&callbacks, "server shutdown").await; @@ -2262,7 +2269,7 @@ mod tests { let (chunk_tx, mut chunk_rx) = mpsc::channel(8); // Drain so sends (including heartbeats) never block the loop. tokio::spawn(async move { while chunk_rx.recv().await.is_some() {} }); - let (cancel_tx, _cancel_rx) = watch::channel(false); + let cancel_tx = CancellationToken::new(); let (cb, _senders) = callbacks(); let start = tokio::time::Instant::now(); let termination = process_sse_stream_full( @@ -2271,7 +2278,7 @@ mod tests { stream, chunk_tx, cancel_tx, - Duration::from_secs(900), + Some(Duration::from_secs(900)), heartbeat, first_chunk, inactivity, @@ -2465,7 +2472,7 @@ mod tests { ); let (chunk_tx, mut chunk_rx) = mpsc::channel(64); tokio::spawn(async move { while chunk_rx.recv().await.is_some() {} }); - let (cancel_tx, _cancel_rx) = watch::channel(false); + let cancel_tx = CancellationToken::new(); let (cb, senders) = callbacks(); // The driver gets a clone; the originals stay alive past the loop. tokio::spawn(drive(senders.clone())); @@ -2476,7 +2483,7 @@ mod tests { stream, chunk_tx, cancel_tx, - Duration::from_secs(900), + Some(Duration::from_secs(900)), HB_QUIET, None, inactivity, @@ -2655,9 +2662,14 @@ mod tests { ); let stream = MockAgent::scripted(steps) - .stream("q", vec![], CancellationToken::new(), "req_tool_events") + .stream( + "q", + vec![], + aura::streaming::RunOptions::default(), + "req_tool_events", + ) .await - .expect("mock stream should start"); + .into_events(); let (chunk_tx, mut chunk_rx) = mpsc::channel::>(64); let collector = tokio::spawn(async move { @@ -2667,7 +2679,7 @@ mod tests { } body }); - let (cancel_tx, _cancel_rx) = watch::channel(false); + let cancel_tx = CancellationToken::new(); let termination = process_sse_stream_full( &config, @@ -2675,7 +2687,7 @@ mod tests { stream, chunk_tx, cancel_tx, - Duration::from_secs(900), + Some(Duration::from_secs(900)), // Far enough out that heartbeats never interleave with the script. Duration::from_secs(86_400), None, diff --git a/crates/aura-web-server/tests/agent_event_differential.rs b/crates/aura-web-server/tests/agent_event_differential.rs index bf392095b..19889dc98 100644 --- a/crates/aura-web-server/tests/agent_event_differential.rs +++ b/crates/aura-web-server/tests/agent_event_differential.rs @@ -33,7 +33,7 @@ use aura_web_server::streaming::{ process_sse_stream_full, }; use bytes::Bytes; -use tokio::sync::{mpsc, watch}; +use tokio::sync::mpsc; use tokio_util::sync::CancellationToken; const TOOL_ID: &str = "call_abc123"; @@ -106,9 +106,14 @@ async fn run_as(request_id: &str, steps: Vec, orchestration: bool) -> Vec< }; let stream = MockAgent::scripted(steps) - .stream("q", vec![], CancellationToken::new(), request_id) + .stream( + "q", + vec![], + aura::streaming::RunOptions::default(), + request_id, + ) .await - .expect("mock stream should start"); + .into_events(); let (chunk_tx, mut chunk_rx) = mpsc::channel::>(64); let collector = tokio::spawn(async move { @@ -118,7 +123,7 @@ async fn run_as(request_id: &str, steps: Vec, orchestration: bool) -> Vec< } body }); - let (cancel_tx, _cancel_rx) = watch::channel(false); + let cancel_tx = CancellationToken::new(); let termination = process_sse_stream_full( &config, @@ -126,7 +131,7 @@ async fn run_as(request_id: &str, steps: Vec, orchestration: bool) -> Vec< stream, chunk_tx, cancel_tx, - Duration::from_secs(900), + Some(Duration::from_secs(900)), // Far enough out that heartbeats never interleave with the script. Duration::from_secs(86_400), None, diff --git a/crates/aura/src/builder.rs b/crates/aura/src/builder.rs index 9bfbb21f0..707e23683 100644 --- a/crates/aura/src/builder.rs +++ b/crates/aura/src/builder.rs @@ -20,8 +20,6 @@ use rig::completion::Usage; use std::collections::HashSet; use std::pin::Pin; use std::sync::Arc; -use std::time::Duration; -use tokio::sync::watch; /// A client-side tool definition supplied with a request. /// @@ -1427,17 +1425,13 @@ impl Agent { /// Stream a query with timeout and cancellation support. /// - /// Returns the stream and a sender to trigger external cancellation. + /// Returns the run: its stream, the token that cancels it, and its usage. /// /// # Arguments /// * `query` - The user query - /// * `timeout` - Timeout duration for the request + /// * `options` - How the run is bounded and cancelled /// * `request_id` - Unique request ID for MCP tool cancellation context /// - /// # Returns - /// * Stream of multi-turn items - /// * Sender to trigger external cancellation (e.g., on client disconnect) - /// /// # Cancellation /// The StreamingRequestHook checks for cancellation at key points during streaming: /// - Before each LLM completion call @@ -1446,39 +1440,29 @@ impl Agent { /// - After each tool result (clears active request context, adds to pending_tool_ids) /// - After each streaming completion (captures usage, emits aura.tool_usage) /// - /// To cancel externally (e.g., on client disconnect), call `cancel_tx.send(true)`. - /// - /// # Returns - /// * Stream of multi-turn items - /// * Sender to trigger external cancellation - /// * UsageState for reading final usage at stream end (shared with hook) + /// To cancel externally (e.g., on client disconnect), cancel the run's token. pub async fn stream_prompt_with_timeout( &self, query: &str, - timeout: Duration, + options: crate::streaming::RunOptions, request_id: &str, - ) -> ( - Pin> + Send>>, - watch::Sender, - crate::streaming_request_hook::UsageState, - ) { + ) -> crate::streaming::AgentRun { self.seed_scratchpad_request_input(query, &[]); - let (stream, cancel_tx, usage_state) = self - .inner + self.inner .stream_prompt_with_timeout( query, self.max_depth, - timeout, + options, request_id, self.scratchpad_budget.clone(), self.client_tool_names.clone(), ) - .await; - ( - self.append_scratchpad_usage(self.maybe_wrap_with_fallback(self.count_turns(stream))), - cancel_tx, - usage_state, - ) + .await + .map_stream(|stream| { + self.append_scratchpad_usage( + self.maybe_wrap_with_fallback(self.count_turns(stream)), + ) + }) } /// Stream a chat query with timeout and cancellation support. @@ -1486,13 +1470,9 @@ impl Agent { /// # Arguments /// * `query` - The user query /// * `chat_history` - Previous conversation messages - /// * `timeout` - Timeout duration for the request + /// * `options` - How the run is bounded and cancelled /// * `request_id` - Unique request ID for MCP tool cancellation context /// - /// # Returns - /// * Stream of multi-turn items - /// * Sender to trigger external cancellation (e.g., on client disconnect) - /// * UsageState for reading final usage at stream end (shared with hook) /// /// # Cancellation /// See `stream_prompt_with_timeout` for cancellation details. @@ -1500,31 +1480,26 @@ impl Agent { &self, query: &str, chat_history: Vec, - timeout: Duration, + options: crate::streaming::RunOptions, request_id: &str, - ) -> ( - Pin> + Send>>, - watch::Sender, - crate::streaming_request_hook::UsageState, - ) { + ) -> crate::streaming::AgentRun { self.seed_scratchpad_request_input(query, &chat_history); - let (stream, cancel_tx, usage_state) = self - .inner + self.inner .stream_chat_with_timeout( query, chat_history, self.max_depth, - timeout, + options, request_id, self.scratchpad_budget.clone(), self.client_tool_names.clone(), ) - .await; - ( - self.append_scratchpad_usage(self.maybe_wrap_with_fallback(self.count_turns(stream))), - cancel_tx, - usage_state, - ) + .await + .map_stream(|stream| { + self.append_scratchpad_usage( + self.maybe_wrap_with_fallback(self.count_turns(stream)), + ) + }) } /// Seed the scratchpad budget's running estimate with the user query + @@ -1697,8 +1672,6 @@ fn record_completion_result( // Implement StreamingAgent trait for Agent use crate::streaming::StreamingAgent; use async_trait::async_trait; -use futures::stream::BoxStream; -use tokio_util::sync::CancellationToken; #[async_trait] impl StreamingAgent for Agent { @@ -1710,51 +1683,22 @@ impl StreamingAgent for Agent { &self, query: &str, chat_history: Vec, - _cancel_token: CancellationToken, - request_id: &str, - ) -> Result>, StreamError> { - if let Some(mcp_manager) = &self.mcp_manager { - mcp_manager - .set_current_call(request_id, aura_events::AgentContext::single_agent()) - .await; - } - - let stream = if chat_history.is_empty() { - self.stream_prompt(query).await - } else { - self.stream_chat(query, chat_history).await - }; - - Ok(Box::pin(stream)) - } - - async fn stream_with_timeout( - &self, - query: &str, - chat_history: Vec, - timeout: Duration, + options: crate::streaming::RunOptions, request_id: &str, - ) -> ( - BoxStream<'static, Result>, - watch::Sender, - crate::UsageState, - ) { - // Production entry point — set MCP request ID before delegating + ) -> crate::streaming::AgentRun { if let Some(mcp_manager) = &self.mcp_manager { mcp_manager .set_current_call(request_id, aura_events::AgentContext::single_agent()) .await; } - let (stream, cancel_tx, usage_state) = if chat_history.is_empty() { - self.stream_prompt_with_timeout(query, timeout, request_id) + if chat_history.is_empty() { + self.stream_prompt_with_timeout(query, options, request_id) .await } else { - self.stream_chat_with_timeout(query, chat_history, timeout, request_id) + self.stream_chat_with_timeout(query, chat_history, options, request_id) .await - }; - - (Box::pin(stream), cancel_tx, usage_state) + } } async fn cancel_and_close_mcp(&self, request_id: &str, reason: &str) -> usize { diff --git a/crates/aura/src/orchestration/factory.rs b/crates/aura/src/orchestration/factory.rs index 0566c87d4..32288e979 100644 --- a/crates/aura/src/orchestration/factory.rs +++ b/crates/aura/src/orchestration/factory.rs @@ -9,7 +9,6 @@ use std::time::Duration; use async_trait::async_trait; use futures::stream::{self, BoxStream}; -use tokio::sync::watch; use tokio_util::sync::CancellationToken; use crate::config::AgentRuntimeConfig; @@ -17,7 +16,7 @@ use crate::provider_agent::{StreamError, StreamItem}; use crate::streaming::StreamingAgent; use super::orchestrator::{ - Orchestrator, STREAM_CHUNK_SIZE, spawn_cancellation_watcher, spawn_tool_event_forwarder, + Orchestrator, STREAM_CHUNK_SIZE, spawn_timeout_watcher, spawn_tool_event_forwarder, }; /// Zero-state wrapper that implements `StreamingAgent` for orchestration mode. @@ -28,6 +27,12 @@ pub struct OrchestratorFactory { agent_config: AgentRuntimeConfig, } +/// A run's cancellation, and the signal that its task ended. +struct RunTokens { + cancel: CancellationToken, + finished: Option, +} + impl OrchestratorFactory { pub fn new(agent_config: AgentRuntimeConfig) -> Self { Self { agent_config } @@ -35,18 +40,15 @@ impl OrchestratorFactory { /// Spawn the background orchestration task and return its event stream. /// - /// Shared by [`stream`](Self::stream) and - /// [`stream_with_timeout`](Self::stream_with_timeout). The `usage_state` - /// handle is assigned to the inner `Orchestrator` so planning, worker, - /// synthesis, and evaluation turns can accumulate into it; the caller - /// (`stream_with_timeout`) retains a clone and hands it to the streaming - /// handler for the final `aura.usage` event. `stream()` passes a detached - /// state since its trait-visible callers don't observe usage. + /// The `usage_state` handle is assigned to the inner `Orchestrator` so + /// planning, worker, synthesis, and evaluation turns can accumulate into it. + /// [`stream`](Self::stream) keeps a clone on the run it returns, so the + /// streaming handler can read the totals for the final `aura.usage` event. fn spawn_orchestration_stream( &self, query: String, chat_history: Vec, - cancel_token: CancellationToken, + tokens: RunTokens, request_id: String, usage_state: crate::UsageState, outer_budget: Option, @@ -57,11 +59,17 @@ impl OrchestratorFactory { let (event_tx, event_rx) = tokio::sync::mpsc::channel::>(100); - let cancel_token_clone = cancel_token.clone(); + let RunTokens { cancel, finished } = tokens; + let cancel_token_clone = cancel.clone(); + // Marks the run finished on every exit path, which is what lets the + // timeout watcher stop rather than sleeping out its full duration. It is + // not the run's cancel token: a run that finished was not cancelled. + let done_guard = finished.map(CancellationToken::drop_guard); // Capture parent span so child spans nest correctly in tracing. let parent_span = tracing::Span::current(); tokio::spawn(tracing::Instrument::instrument( async move { + let _done_guard = done_guard; let mut orchestrator = match Orchestrator::new(agent_config).await { Ok(o) => o, Err(e) => { @@ -160,40 +168,24 @@ impl StreamingAgent for OrchestratorFactory { &self, query: &str, chat_history: Vec, - cancel_token: CancellationToken, + options: crate::streaming::RunOptions, request_id: &str, - ) -> Result>, StreamError> { - // Raw-stream callers don't observe usage; hand the spawn a detached - // UsageState so the field is populated but nobody reads it. - Ok(self.spawn_orchestration_stream( - query.to_string(), - chat_history, - cancel_token, - request_id.to_string(), - crate::UsageState::new(), - None, - )) - } - - async fn stream_with_timeout( - &self, - query: &str, - chat_history: Vec, - timeout: Duration, - request_id: &str, - ) -> ( - BoxStream<'static, Result>, - watch::Sender, - crate::UsageState, - ) { - let (cancel_tx, cancel_rx) = watch::channel(false); - let cancel_token = CancellationToken::new(); - let watcher_cancel_token = cancel_token.clone(); - let request_id_owned = request_id.to_string(); - - // Fire-and-forget: task self-terminates when cancel_tx is dropped or timeout fires. - let _watcher_handle = - spawn_cancellation_watcher(cancel_rx, timeout, watcher_cancel_token, request_id_owned); + ) -> crate::streaming::AgentRun { + let (timeout, cancel) = options.into_parts(); + let cancel_token = cancel.unwrap_or_default(); + + // Only a watcher observes this, so an unbounded run needs none. + let finished = timeout.map(|timeout| { + let finished = CancellationToken::new(); + // Fire-and-forget: self-terminates when the run ends or the timeout fires. + let _watcher_handle = spawn_timeout_watcher( + timeout, + cancel_token.clone(), + finished.clone(), + request_id.to_string(), + ); + finished + }); // Share one UsageState between the inner orchestrator (writer) and the // streaming handler (reader) so aura.usage reflects the aggregate of @@ -202,13 +194,16 @@ impl StreamingAgent for OrchestratorFactory { let stream = self.spawn_orchestration_stream( query.to_string(), chat_history, - cancel_token, + RunTokens { + cancel: cancel_token.clone(), + finished, + }, request_id.to_string(), usage_state.clone(), - (!timeout.is_zero()).then_some(timeout), + timeout, ); - (stream, cancel_tx, usage_state) + crate::streaming::AgentRun::new(stream, cancel_token, usage_state) } async fn cancel_and_close_mcp(&self, _request_id: &str, _reason: &str) -> usize { diff --git a/crates/aura/src/orchestration/mod.rs b/crates/aura/src/orchestration/mod.rs index c66e1fc85..53ae4158e 100644 --- a/crates/aura/src/orchestration/mod.rs +++ b/crates/aura/src/orchestration/mod.rs @@ -37,7 +37,8 @@ //! .build_streaming_agent_with_headers(None, None, None) //! .await?; //! -//! let stream = agent.stream(query, history, cancel_token, "req_123").await?; +//! let run = agent.stream(query, history, RunOptions::default(), "req_123").await; +//! let stream = run.into_events(); //! ``` mod config; diff --git a/crates/aura/src/orchestration/orchestrator.rs b/crates/aura/src/orchestration/orchestrator.rs index 6dadfcd79..9034eb1a0 100644 --- a/crates/aura/src/orchestration/orchestrator.rs +++ b/crates/aura/src/orchestration/orchestrator.rs @@ -49,7 +49,7 @@ use std::time::{Duration, Instant}; use aura_config::GlobPattern; use rig::client::CompletionClient; -use tokio::sync::{Mutex, watch}; +use tokio::sync::Mutex; use tokio_util::sync::CancellationToken; use crate::Agent; @@ -190,46 +190,28 @@ fn apply_worker_skills_override( } } -/// Spawns a task that monitors for external cancellation or timeout, -/// cancelling the provided token when either occurs. +/// Spawns a task that cancels `cancel_token` once `timeout` passes. +/// +/// It also stops early on either of two signals. `finished` resolves when the +/// run's task ends, so the watcher does not sleep out its full duration; a +/// finished run has not been cancelled, which is why that is a separate token. +/// An already-cancelled `cancel_token` needs nothing further and only logs. /// /// Returns a `JoinHandle` for the watcher task. The handle is intentionally /// fire-and-forget in production (the task self-terminates via `select!`), /// but callers in tests should `.await` it to assert post-conditions. -/// -/// Cleanup: when the caller drops the sender side of `cancel_rx` after an -/// explicit signal, `rx.changed()` returns `Err`, the `select!` resolves, -/// and the sleep future is dropped (cancelling the timer via tokio's -/// standard drop semantics). A drop with no explicit signal is treated as -/// an unexplained abort — see the loop body below. #[must_use = "task runs independently; bind with `let _handle =` to document fire-and-forget intent"] -pub(super) fn spawn_cancellation_watcher( - cancel_rx: watch::Receiver, +pub(super) fn spawn_timeout_watcher( timeout: Duration, cancel_token: CancellationToken, + finished: CancellationToken, request_id: String, ) -> tokio::task::JoinHandle<()> { tokio::spawn(async move { tokio::select! { - was_cancelled = async { - let mut rx = cancel_rx; - loop { - // A closed channel means the outer stream can no longer - // signal cancellation. Fail safe by cancelling the - // orchestration token (#305) — cancelling an - // already-finished inner task is a no-op. - if rx.changed().await.is_err() { - return true; - } - if *rx.borrow_and_update() { - return true; // External cancellation requested - } - } - } => { - if was_cancelled { - tracing::info!("External cancellation triggered for {}", request_id); - cancel_token.cancel(); - } + () = finished.cancelled() => {} + () = cancel_token.cancelled() => { + tracing::info!("Run cancelled for {}", request_id); } _ = tokio::time::sleep(timeout) => { tracing::warn!("Timeout reached, cancelling orchestration"); @@ -466,8 +448,8 @@ pub struct Orchestrator { /// Accumulated token usage across all LLM calls in this orchestration run /// (planning, workers, continuation routing). /// - /// Cloned from a handle owned by `OrchestratorFactory::stream_with_timeout` - /// so the streaming handler can read the final totals and emit `aura.usage`. + /// Cloned from the handle `OrchestratorFactory::stream` puts on the run, so + /// the streaming handler can read the final totals and emit `aura.usage`. /// In orchestration mode we aggregate additively via /// [`crate::UsageState::accumulate_usage`] so the reported prompt/completion /// totals reflect *billed* tokens across every internal LLM turn, not just @@ -1445,12 +1427,15 @@ impl Orchestrator { let timeout_secs = self.config.per_call_timeout_secs(); let stream_future = async { let stream = match park_key { - Some(key) => { - agent - .stream_chat_with_timeout(prompt, history, Duration::MAX, key) - .await - .0 - } + Some(key) => agent + .stream_chat_with_timeout( + prompt, + history, + crate::streaming::RunOptions::default(), + key, + ) + .await + .into_events(), None => agent.stream_chat(prompt, history).await, }; Self::drive_forward_loop( @@ -4015,18 +4000,19 @@ Assign tasks to the worker whose tools best match the required operations."#, let srd = submit_result_decision.clone(); let park_registration = crate::streaming_request_hook::ParkCellRegistration::new(&park.key, park.cell.clone()); - let (stream, _cancel_tx, _usage_state) = worker + let stream = worker .inner .stream_chat_message_with_timeout( current_prompt, continuation.history.clone(), worker.max_depth, - Duration::MAX, + crate::streaming::RunOptions::default(), &park.key, worker.scratchpad_budget.clone(), worker.client_tool_names.clone(), ) - .await; + .await + .into_events(); let stream_result = Self::drive_forward_loop( stream, &self.usage_state, @@ -6515,120 +6501,44 @@ mod tests { // Cancellation watcher tests // ======================================================================== + /// A finished run is not a cancelled one, so the watcher has to stop on a + /// signal that leaves the run's own token untouched. #[tokio::test(start_paused = true)] - async fn test_watcher_unexplained_drop_triggers_failsafe_cancel() { - let (cancel_tx, cancel_rx) = watch::channel(false); + async fn watcher_stops_when_the_run_ends() { let cancel_token = CancellationToken::new(); - let handle = spawn_cancellation_watcher( - cancel_rx, + let finished = CancellationToken::new(); + let handle = spawn_timeout_watcher( Duration::from_secs(300), cancel_token.clone(), + finished.clone(), "test-normal".to_string(), ); - // A bare drop with no explicit `false` first means the outer task - // ended without going through its normal completion path (e.g. - // aborted during shutdown) — the watcher must fail safe and cancel. - drop(cancel_tx); - tokio::task::yield_now().await; - handle.await.unwrap(); - assert!(cancel_token.is_cancelled()); - } - - #[tokio::test(start_paused = true)] - async fn test_watcher_external_cancel_triggers_token() { - let (cancel_tx, cancel_rx) = watch::channel(false); - let cancel_token = CancellationToken::new(); - let handle = spawn_cancellation_watcher( - cancel_rx, - Duration::from_secs(300), - cancel_token.clone(), - "test-cancel".to_string(), - ); - - cancel_tx.send(true).unwrap(); - tokio::task::yield_now().await; - handle.await.unwrap(); - assert!(cancel_token.is_cancelled()); - } - - #[tokio::test(start_paused = true)] - async fn test_watcher_timeout_triggers_cancellation() { - // Keep sender alive so only the timeout path can fire - let (_cancel_tx, cancel_rx) = watch::channel(false); - let cancel_token = CancellationToken::new(); - let handle = spawn_cancellation_watcher( - cancel_rx, - Duration::from_secs(60), - cancel_token.clone(), - "test-timeout".to_string(), - ); - - tokio::time::advance(Duration::from_secs(61)).await; - tokio::task::yield_now().await; - handle.await.unwrap(); - assert!(cancel_token.is_cancelled()); - } - - #[tokio::test(start_paused = true)] - async fn test_watcher_unexplained_drop_before_timeout_cancels_promptly() { - let (cancel_tx, cancel_rx) = watch::channel(false); - let cancel_token = CancellationToken::new(); - let handle = spawn_cancellation_watcher( - cancel_rx, - Duration::from_secs(60), - cancel_token.clone(), - "test-abort-mid-stream".to_string(), - ); - - // Advance to T=30s, then drop the sender with no prior explicit - // signal — the production scenario of an outer task aborted mid- - // stream (e.g. during shutdown). The watcher must fail safe and - // cancel promptly rather than assume normal completion. - tokio::time::advance(Duration::from_secs(30)).await; - tokio::task::yield_now().await; - drop(cancel_tx); - tokio::task::yield_now().await; + finished.cancel(); let start = tokio::time::Instant::now(); handle.await.unwrap(); - let elapsed = start.elapsed(); - assert!( - cancel_token.is_cancelled(), - "an unexplained sender drop must fail safe and cancel" + start.elapsed() < Duration::from_secs(1), + "watcher should exit when the run ends, not wait out its timeout" ); - // Task should exit promptly on sender drop, not wait for remaining 30s timeout assert!( - elapsed < Duration::from_secs(1), - "task should exit promptly after sender drop, not wait for timeout; elapsed: {:?}", - elapsed + !cancel_token.is_cancelled(), + "a run that finished was never cancelled" ); } #[tokio::test(start_paused = true)] - async fn test_watcher_false_signal_does_not_cancel() { - let (cancel_tx, cancel_rx) = watch::channel(false); + async fn watcher_cancels_on_timeout() { let cancel_token = CancellationToken::new(); - let handle = spawn_cancellation_watcher( - cancel_rx, - Duration::from_secs(300), + let handle = spawn_timeout_watcher( + Duration::from_secs(60), cancel_token.clone(), - "test-false-signal".to_string(), - ); - - // Send false — triggers rx.changed() but borrow_and_update() sees false, - // so the loop continues waiting - cancel_tx.send(false).unwrap(); - tokio::task::yield_now().await; - assert!( - !cancel_token.is_cancelled(), - "false signal should not cancel" + CancellationToken::new(), + "test-timeout".to_string(), ); - // Bare drop with no final explicit signal — fail safe and cancel. - drop(cancel_tx); - tokio::task::yield_now().await; + tokio::time::advance(Duration::from_secs(61)).await; handle.await.unwrap(); assert!(cancel_token.is_cancelled()); } diff --git a/crates/aura/src/orchestration/test_rig.rs b/crates/aura/src/orchestration/test_rig.rs index 828b49859..4b0953a0d 100644 --- a/crates/aura/src/orchestration/test_rig.rs +++ b/crates/aura/src/orchestration/test_rig.rs @@ -49,7 +49,8 @@ use rig::streaming::{ StreamingCompletionResponse, }; use serde::{Deserialize, Serialize}; -use tokio::sync::{Notify, watch}; +use tokio::sync::Notify; +use tokio_util::sync::CancellationToken; use crate::streaming_request_hook::{StreamingRequestHook, UsageState}; @@ -592,10 +593,9 @@ pub(crate) struct StreamRun { /// The loop's aggregated usage (`FinalResponse.usage`). pub(crate) usage: Usage, pub(crate) tool_results: Vec, - /// External cancellation handle (the hook's watch channel), for the - /// commit-3 race tests. - #[allow(dead_code)] // commit 3: race tests cancel mid-stream - pub(crate) cancel_tx: watch::Sender, + /// External cancellation handle, for the race tests. + #[allow(dead_code)] // race tests cancel mid-stream + pub(crate) cancel: CancellationToken, /// The park-aware hook's usage state, asserting hook compatibility. pub(crate) usage_state: UsageState, } @@ -621,8 +621,11 @@ pub(crate) async fn drive_worker( max_depth: usize, ) -> Result> { let request_id = format!("rig_{}", uuid::Uuid::new_v4().simple()); - let (hook, cancel_tx, usage_state) = - StreamingRequestHook::with_scratchpad_budget(Duration::from_secs(60), request_id, None); + let (hook, cancel, usage_state) = StreamingRequestHook::with_scratchpad_budget( + crate::streaming::RunOptions::bounded(Some(Duration::from_secs(60))), + request_id, + None, + ); let mut stream = rig .agent @@ -635,7 +638,7 @@ pub(crate) async fn drive_worker( final_text: None, usage: Usage::new(), tool_results: Vec::new(), - cancel_tx, + cancel, usage_state, }; diff --git a/crates/aura/src/provider_agent.rs b/crates/aura/src/provider_agent.rs index 6466f205c..5fa7cec9d 100644 --- a/crates/aura/src/provider_agent.rs +++ b/crates/aura/src/provider_agent.rs @@ -16,8 +16,6 @@ use rig::message::ToolResultContent; use rig::streaming::{StreamingChat, StreamingPrompt}; use std::collections::HashSet; use std::pin::Pin; -use std::time::Duration; -use tokio::sync::watch; use crate::scratchpad::ContextBudget; use crate::streaming_request_hook::StreamingRequestHook; @@ -109,17 +107,13 @@ impl ProviderAgent { prompt: rig::completion::Message, chat_history: Vec, max_depth: usize, - timeout: Duration, + options: crate::streaming::RunOptions, request_id: &str, scratchpad_budget: Option, client_tool_names: HashSet, - ) -> ( - Pin> + Send>>, - watch::Sender, - crate::streaming_request_hook::UsageState, - ) { + ) -> crate::streaming::AgentRun { let (hook, cancel_tx, usage_state) = - StreamingRequestHook::with_scratchpad_budget(timeout, request_id, scratchpad_budget); + StreamingRequestHook::with_scratchpad_budget(options, request_id, scratchpad_budget); let hook = hook.with_client_tool_names(client_tool_names); match self { @@ -129,7 +123,7 @@ impl ProviderAgent { .with_hook(hook) .multi_turn(max_depth) .await; - ( + crate::streaming::AgentRun::new( Box::pin(stream.map::, _>(map_stream_item)), cancel_tx, usage_state, @@ -141,7 +135,7 @@ impl ProviderAgent { .with_hook(hook) .multi_turn(max_depth) .await; - ( + crate::streaming::AgentRun::new( Box::pin(stream.map::, _>(map_stream_item)), cancel_tx, usage_state, @@ -153,7 +147,7 @@ impl ProviderAgent { .with_hook(hook) .multi_turn(max_depth) .await; - ( + crate::streaming::AgentRun::new( Box::pin(stream.map::, _>(map_stream_item)), cancel_tx, usage_state, @@ -165,7 +159,7 @@ impl ProviderAgent { .with_hook(hook) .multi_turn(max_depth) .await; - ( + crate::streaming::AgentRun::new( Box::pin(stream.map::, _>(map_stream_item)), cancel_tx, usage_state, @@ -177,7 +171,7 @@ impl ProviderAgent { .with_hook(hook) .multi_turn(max_depth) .await; - ( + crate::streaming::AgentRun::new( Box::pin(stream.map::, _>(map_stream_item)), cancel_tx, usage_state, @@ -189,7 +183,7 @@ impl ProviderAgent { .with_hook(hook) .multi_turn(max_depth) .await; - ( + crate::streaming::AgentRun::new( Box::pin(stream.map::, _>(map_stream_item)), cancel_tx, usage_state, @@ -202,7 +196,7 @@ impl ProviderAgent { .with_hook(hook) .multi_turn(max_depth) .await; - ( + crate::streaming::AgentRun::new( Box::pin(stream.map::, _>(map_stream_item)), cancel_tx, usage_state, @@ -316,25 +310,18 @@ impl ProviderAgent { /// Stream a prompt with timeout and cancellation support. /// - /// Returns (stream, cancel_sender, usage_state): - /// - stream: The actual stream of completion items - /// - cancel_sender: Send `true` to cancel the stream - /// - usage_state: Shared state for reading final usage at stream end + /// Returns the run: its stream, the token that cancels it, and its usage. pub async fn stream_prompt_with_timeout( &self, query: &str, max_depth: usize, - timeout: Duration, + options: crate::streaming::RunOptions, request_id: &str, scratchpad_budget: Option, client_tool_names: HashSet, - ) -> ( - Pin> + Send>>, - watch::Sender, - crate::streaming_request_hook::UsageState, - ) { + ) -> crate::streaming::AgentRun { let (hook, cancel_tx, usage_state) = - StreamingRequestHook::with_scratchpad_budget(timeout, request_id, scratchpad_budget); + StreamingRequestHook::with_scratchpad_budget(options, request_id, scratchpad_budget); let hook = hook.with_client_tool_names(client_tool_names); match self { @@ -344,7 +331,7 @@ impl ProviderAgent { .with_hook(hook) .multi_turn(max_depth) .await; - ( + crate::streaming::AgentRun::new( Box::pin(stream.map::, _>(map_stream_item)), cancel_tx, usage_state, @@ -356,7 +343,7 @@ impl ProviderAgent { .with_hook(hook) .multi_turn(max_depth) .await; - ( + crate::streaming::AgentRun::new( Box::pin(stream.map::, _>(map_stream_item)), cancel_tx, usage_state, @@ -368,7 +355,7 @@ impl ProviderAgent { .with_hook(hook) .multi_turn(max_depth) .await; - ( + crate::streaming::AgentRun::new( Box::pin(stream.map::, _>(map_stream_item)), cancel_tx, usage_state, @@ -380,7 +367,7 @@ impl ProviderAgent { .with_hook(hook) .multi_turn(max_depth) .await; - ( + crate::streaming::AgentRun::new( Box::pin(stream.map::, _>(map_stream_item)), cancel_tx, usage_state, @@ -392,7 +379,7 @@ impl ProviderAgent { .with_hook(hook) .multi_turn(max_depth) .await; - ( + crate::streaming::AgentRun::new( Box::pin(stream.map::, _>(map_stream_item)), cancel_tx, usage_state, @@ -404,7 +391,7 @@ impl ProviderAgent { .with_hook(hook) .multi_turn(max_depth) .await; - ( + crate::streaming::AgentRun::new( Box::pin(stream.map::, _>(map_stream_item)), cancel_tx, usage_state, @@ -417,7 +404,7 @@ impl ProviderAgent { .with_hook(hook) .multi_turn(max_depth) .await; - ( + crate::streaming::AgentRun::new( Box::pin(stream.map::, _>(map_stream_item)), cancel_tx, usage_state, @@ -428,27 +415,20 @@ impl ProviderAgent { /// Stream a chat with timeout and cancellation support. /// - /// Returns (stream, cancel_sender, usage_state): - /// - stream: The actual stream of completion items - /// - cancel_sender: Send `true` to cancel the stream - /// - usage_state: Shared state for reading final usage at stream end + /// Returns the run: its stream, the token that cancels it, and its usage. #[allow(clippy::too_many_arguments)] pub async fn stream_chat_with_timeout( &self, query: &str, chat_history: Vec, max_depth: usize, - timeout: Duration, + options: crate::streaming::RunOptions, request_id: &str, scratchpad_budget: Option, client_tool_names: HashSet, - ) -> ( - Pin> + Send>>, - watch::Sender, - crate::streaming_request_hook::UsageState, - ) { + ) -> crate::streaming::AgentRun { let (hook, cancel_tx, usage_state) = - StreamingRequestHook::with_scratchpad_budget(timeout, request_id, scratchpad_budget); + StreamingRequestHook::with_scratchpad_budget(options, request_id, scratchpad_budget); let hook = hook.with_client_tool_names(client_tool_names); match self { @@ -458,7 +438,7 @@ impl ProviderAgent { .with_hook(hook) .multi_turn(max_depth) .await; - ( + crate::streaming::AgentRun::new( Box::pin(stream.map::, _>(map_stream_item)), cancel_tx, usage_state, @@ -470,7 +450,7 @@ impl ProviderAgent { .with_hook(hook) .multi_turn(max_depth) .await; - ( + crate::streaming::AgentRun::new( Box::pin(stream.map::, _>(map_stream_item)), cancel_tx, usage_state, @@ -482,7 +462,7 @@ impl ProviderAgent { .with_hook(hook) .multi_turn(max_depth) .await; - ( + crate::streaming::AgentRun::new( Box::pin(stream.map::, _>(map_stream_item)), cancel_tx, usage_state, @@ -494,7 +474,7 @@ impl ProviderAgent { .with_hook(hook) .multi_turn(max_depth) .await; - ( + crate::streaming::AgentRun::new( Box::pin(stream.map::, _>(map_stream_item)), cancel_tx, usage_state, @@ -506,7 +486,7 @@ impl ProviderAgent { .with_hook(hook) .multi_turn(max_depth) .await; - ( + crate::streaming::AgentRun::new( Box::pin(stream.map::, _>(map_stream_item)), cancel_tx, usage_state, @@ -518,7 +498,7 @@ impl ProviderAgent { .with_hook(hook) .multi_turn(max_depth) .await; - ( + crate::streaming::AgentRun::new( Box::pin(stream.map::, _>(map_stream_item)), cancel_tx, usage_state, @@ -531,7 +511,7 @@ impl ProviderAgent { .with_hook(hook) .multi_turn(max_depth) .await; - ( + crate::streaming::AgentRun::new( Box::pin(stream.map::, _>(map_stream_item)), cancel_tx, usage_state, diff --git a/crates/aura/src/streaming.rs b/crates/aura/src/streaming.rs index d4d36c963..c6c100d5f 100644 --- a/crates/aura/src/streaming.rs +++ b/crates/aura/src/streaming.rs @@ -13,15 +13,17 @@ //! # Usage //! //! ```ignore -//! use aura::{StreamingAgent, StreamItem, StreamError}; -//! use tokio_util::sync::CancellationToken; +//! use aura::streaming::{RunOptions, StreamingAgent}; +//! use aura::{StreamError, StreamItem}; +//! use futures::StreamExt; //! //! async fn handle_request(agent: impl StreamingAgent, query: &str) { -//! let cancel_token = CancellationToken::new(); -//! let stream = agent.stream(query, vec![], cancel_token, "req_123").await?; +//! // The default leaves the run unbounded and lets it mint its own token. +//! let run = agent.stream(query, vec![], RunOptions::default(), "req_123").await; +//! let mut items = run.into_events(); //! //! // Process stream items (convert to SSE, etc.) -//! while let Some(item) = stream.next().await { +//! while let Some(item) = items.next().await { //! match item { //! Ok(StreamItem::StreamAssistantItem(content)) => { /* ... */ } //! Ok(StreamItem::StreamUserItem(content)) => { /* ... */ } @@ -37,7 +39,124 @@ use async_trait::async_trait; use futures::stream::BoxStream; use rig::completion::Message; use std::time::Duration; -use tokio_util::sync::CancellationToken; +use tokio_util::sync::{CancellationToken, DropGuard}; + +/// How a run is bounded and cancelled. +#[derive(Default)] +pub struct RunOptions { + pub timeout: Option, + cancel: Option, +} + +impl RunOptions { + /// The bound and the token, for an implementation building a run. + #[must_use] + pub fn into_parts(self) -> (Option, Option) { + (self.timeout, self.cancel) + } + + #[must_use] + pub fn bounded(timeout: Option) -> Self { + Self { + timeout, + cancel: None, + } + } + + /// The run stops when `parent` does, and stopping the run leaves `parent` + /// alone — a caller's token is often shared, so the run takes a child of it. + #[must_use] + pub fn cancelled_by(mut self, parent: &CancellationToken) -> Self { + self.cancel = Some(parent.child_token()); + self + } +} + +/// A started run: the events it produces, the token that cancels it, and the +/// usage it accumulates. +pub struct AgentRun { + events: BoxStream<'static, Result>, + cancel: CancellationToken, + usage: UsageState, + guard: DropGuard, +} + +impl AgentRun { + pub fn new( + events: BoxStream<'static, Result>, + cancel: CancellationToken, + usage: UsageState, + ) -> Self { + Self { + // On the run's own token, because that is what its work watches. + // `RunOptions::cancelled_by` is what keeps a caller's shared token + // from being that token. + guard: cancel.clone().drop_guard(), + events, + cancel, + usage, + } + } + + /// Orchestration races this token, so cancelling it stops a run at once. + /// A single agent reads it from the streaming hook's callbacks, so a run + /// stalled with no provider output needs its MCP calls cancelled too. + pub fn cancel_token(&self) -> CancellationToken { + self.cancel.clone() + } + + pub fn usage(&self) -> &UsageState { + &self.usage + } + + /// Wraps the run's stream, keeping the cancellation and usage that belong + /// with it. A layer that decorates the stream has no reason to take the + /// handle apart and rebuild it. + #[must_use] + pub fn map_stream(self, f: F) -> Self + where + F: FnOnce( + BoxStream<'static, Result>, + ) -> BoxStream<'static, Result>, + { + Self { + events: f(self.events), + ..self + } + } + + /// Dropping the returned stream cancels the run. + /// + /// A consumer that goes away without draining — an aborted task, a dropped + /// stream — would otherwise leave the run spending provider turns nobody + /// reads, and a run started with no timeout has nothing else to stop it. + /// The guard moves with the stream, so the run outlives the handle for as + /// long as something is reading it. + pub fn into_events(self) -> BoxStream<'static, Result> { + Box::pin(CancelOnDrop { + _guard: self.guard, + inner: self.events, + }) + } +} + +/// Cancels its run when dropped, by holding the guard for as long as the stream +/// it wraps. +struct CancelOnDrop { + inner: S, + _guard: DropGuard, +} + +impl futures::Stream for CancelOnDrop { + type Item = S::Item; + + fn poll_next( + self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + std::pin::Pin::new(&mut self.get_mut().inner).poll_next(cx) + } +} /// Trait for agents that produce streaming completions. /// @@ -64,59 +183,20 @@ pub trait StreamingAgent: Send + Sync { /// needs to know the concrete agent type. fn get_provider_info(&self) -> (&str, &str); - /// Stream a completion response. - /// - /// Returns a stream of `StreamItem`s. The caller is responsible for: - /// - Converting items to SSE bytes (via handlers) - /// - Sending to the client - /// - Handling cancellation on disconnect + /// Start a run. /// - /// # Arguments + /// `options` bounds the run and may hand it a token the caller already + /// holds. The returned handle owns the events, the token that cancels them, + /// and the usage they accumulate. /// - /// * `query` - The user's query/message - /// * `chat_history` - Previous messages in the conversation - /// * `cancel_token` - Token for cancellation (e.g., on client disconnect) - /// * `request_id` - HTTP request ID for MCP progress routing and tool correlation - /// - /// # Returns - /// - /// A boxed stream of `StreamItem` results, or an error if streaming cannot start. + /// `request_id` correlates MCP progress and tool events for this run. async fn stream( &self, query: &str, chat_history: Vec, - cancel_token: CancellationToken, + options: RunOptions, request_id: &str, - ) -> Result>, StreamError>; - - /// Stream with timeout support. - /// - /// This is the primary entry point for production use. It wraps the stream - /// with timeout handling and integrates with the cancellation hook. - /// - /// # Arguments - /// - /// * `query` - The user's query/message - /// * `chat_history` - Previous messages in the conversation - /// * `timeout` - Maximum duration for the entire stream - /// * `request_id` - Request ID for MCP cancellation correlation - /// - /// # Returns - /// - /// A tuple of (stream, cancel_sender, usage_state) where cancel_sender can - /// be used to signal cancellation to the underlying provider and usage_state - /// tracks token consumption via Rig hooks. - async fn stream_with_timeout( - &self, - query: &str, - chat_history: Vec, - timeout: Duration, - request_id: &str, - ) -> ( - BoxStream<'static, Result>, - tokio::sync::watch::Sender, - UsageState, - ); + ) -> AgentRun; /// Cancel in-flight MCP requests and close connections. /// @@ -150,3 +230,132 @@ pub trait StreamingAgent: Send + Sync { None } } + +#[cfg(test)] +mod tests { + use super::*; + use futures::StreamExt; + use std::sync::Arc; + + fn empty_run() -> (AgentRun, CancellationToken) { + let cancel = CancellationToken::new(); + let run = AgentRun::new( + Box::pin(futures::stream::empty()), + cancel.clone(), + UsageState::new(), + ); + (run, cancel) + } + + /// A caller that takes a run and drops it without consuming has still + /// started it — for orchestration the work is already spawned, and with no + /// timeout nothing else would stop it. + #[test] + fn dropping_the_handle_cancels_the_run() { + let (run, cancel) = empty_run(); + assert!(!cancel.is_cancelled()); + + drop(run); + assert!(cancel.is_cancelled()); + } + + /// The guard moves to the stream, so the run outlives the handle for as + /// long as someone is reading it. + #[tokio::test] + async fn the_run_survives_the_handle_while_its_stream_is_held() { + let (run, cancel) = empty_run(); + let mut events = run.into_events(); + + assert!(!cancel.is_cancelled(), "the stream still holds the run"); + assert!(events.next().await.is_none()); + assert!(!cancel.is_cancelled()); + + drop(events); + assert!(cancel.is_cancelled()); + } + + /// An orchestration run is a spawned task feeding a channel. Dropping the + /// stream stops that task, so a consumer that goes away does not leave it + /// spending provider turns nobody reads. + #[tokio::test] + async fn dropping_a_spawned_run_stops_its_task() { + let cancel = CancellationToken::new(); + let (tx, rx) = tokio::sync::mpsc::channel::>(4); + + let token = cancel.clone(); + let worked = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + let counter = Arc::clone(&worked); + let task = tokio::spawn(async move { + loop { + tokio::select! { + () = token.cancelled() => break, + _ = tokio::time::sleep(std::time::Duration::from_millis(1)) => { + counter.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + if tx.send(Ok(StreamItem::FinalMarker)).await.is_err() { + break; + } + } + } + } + }); + + let run = AgentRun::new( + Box::pin(futures::stream::unfold(rx, |mut rx| async move { + rx.recv().await.map(|item| (item, rx)) + })), + cancel.clone(), + UsageState::new(), + ); + + drop(run.into_events()); + // Asserted before awaiting, so this isolates the guard: the channel + // closing would stop the task either way. + assert!(cancel.is_cancelled(), "dropping the stream cancels the run"); + // Bounded, so a task that keeps running fails here rather than hanging. + tokio::time::timeout(std::time::Duration::from_secs(5), task) + .await + .expect("the task ends rather than running on") + .expect("the task does not panic"); + } + + /// A caller's token is often shared, so one run's stream going away must + /// stop that run and nothing else. + #[test] + fn dropping_a_run_leaves_the_token_it_inherited_alone() { + let shared = CancellationToken::new(); + let options = RunOptions::default().cancelled_by(&shared); + let run = AgentRun::new( + Box::pin(futures::stream::empty()), + options.into_parts().1.expect("a child of the shared token"), + UsageState::new(), + ); + + drop(run); + assert!( + !shared.is_cancelled(), + "the caller's token outlives one run" + ); + } + + /// Cancelling the caller's token still stops the run. + #[test] + fn cancelling_the_inherited_token_stops_the_run() { + let shared = CancellationToken::new(); + let options = RunOptions::default().cancelled_by(&shared); + let run_token = options.into_parts().1.expect("a child of the shared token"); + + shared.cancel(); + assert!(run_token.is_cancelled()); + } + + /// Decorating the stream must not drop the guard along the way. + #[test] + fn mapping_the_stream_keeps_the_run_alive() { + let (run, cancel) = empty_run(); + let mapped = run.map_stream(|stream| Box::pin(stream)); + + assert!(!cancel.is_cancelled()); + drop(mapped); + assert!(cancel.is_cancelled()); + } +} diff --git a/crates/aura/src/streaming_request_hook.rs b/crates/aura/src/streaming_request_hook.rs index 42b488ec4..341bbb9cc 100644 --- a/crates/aura/src/streaming_request_hook.rs +++ b/crates/aura/src/streaming_request_hook.rs @@ -25,13 +25,14 @@ //! # Usage //! //! ```ignore -//! let (hook, cancel_sender, usage_state) = StreamingRequestHook::new(Duration::from_secs(60), "req_123"); +//! let options = RunOptions::bounded(Some(Duration::from_secs(60))); +//! let (hook, cancel, usage_state) = StreamingRequestHook::new(options, "req_123"); //! //! // Pass hook to streaming request //! agent.stream_prompt(query).with_hook(hook).multi_turn(depth).await; //! //! // To cancel externally (e.g., on client disconnect): -//! let _ = cancel_sender.send(true); +//! cancel.cancel(); //! //! // At stream end, read final usage from usage_state //! let (prompt, completion, total) = usage_state.get_final_usage(); @@ -47,7 +48,7 @@ use std::time::{Duration, Instant}; use aura_events::agent::{AgentEvent, AgentEventPayload}; use rig::agent::{CancelSignal, StreamingPromptHook}; use rig::completion::{CompletionModel, GetTokenUsage, Message}; -use tokio::sync::watch; +use tokio_util::sync::CancellationToken; use crate::orchestration::BlockedCell; use crate::scratchpad::{self, ContextBudget}; @@ -366,9 +367,10 @@ impl ResponseContent { #[derive(Clone)] pub struct StreamingRequestHook { start_time: Instant, - timeout: Duration, + /// Absent for a run the caller left unbounded. + timeout: Option, /// External cancellation signal (e.g., from client disconnect) - cancelled: watch::Receiver, + cancelled: CancellationToken, /// Request ID for event correlation request_id: String, /// Shared usage state (returned separately for handler access) @@ -387,17 +389,12 @@ pub struct StreamingRequestHook { } impl StreamingRequestHook { - /// Create a new streaming request hook with the given timeout duration and request ID. - /// - /// Returns a tuple of (hook, cancel_sender, usage_state). - /// - `hook`: The hook to pass to stream_prompt().with_hook() - /// - `cancel_sender`: Send `true` to trigger cancellation - /// - `usage_state`: Shared state - handler keeps clone to read final usage at stream end + /// Create a new streaming request hook for one request. pub fn new( - timeout: Duration, + options: crate::streaming::RunOptions, request_id: impl Into, - ) -> (Self, watch::Sender, UsageState) { - Self::with_scratchpad_budget(timeout, request_id, None) + ) -> (Self, CancellationToken, UsageState) { + Self::with_scratchpad_budget(options, request_id, None) } /// Like `new`, but additionally wires a scratchpad `ContextBudget` so the @@ -405,23 +402,25 @@ impl StreamingRequestHook { /// after each completion turn (mirrors what orchestration workers do via /// `StreamItem::TurnUsage`). pub fn with_scratchpad_budget( - timeout: Duration, + options: crate::streaming::RunOptions, request_id: impl Into, scratchpad_budget: Option, - ) -> (Self, watch::Sender, UsageState) { - let (tx, rx) = watch::channel(false); + ) -> (Self, CancellationToken, UsageState) { + let (timeout, cancel) = options.into_parts(); + // A caller that supplied one can cancel before the run exists. + let cancel = cancel.unwrap_or_default(); let usage_state = UsageState::new(); let hook = Self { start_time: Instant::now(), timeout, - cancelled: rx, + cancelled: cancel.clone(), request_id: request_id.into(), usage_state: usage_state.clone(), scratchpad_budget, client_tool_names: HashSet::new(), client_tool_called: Arc::new(AtomicBool::new(false)), }; - (hook, tx, usage_state) + (hook, cancel, usage_state) } /// Register the names of client-side (passthrough) tools for this request. @@ -435,10 +434,13 @@ impl StreamingRequestHook { /// Check if the request should be cancelled (timeout or external signal). fn should_cancel(&self) -> bool { - if *self.cancelled.borrow() { - return true; - } - self.start_time.elapsed() > self.timeout + self.cancelled.is_cancelled() || self.deadline_passed() + } + + /// An unbounded run has no deadline to pass. + fn deadline_passed(&self) -> bool { + self.timeout + .is_some_and(|timeout| self.start_time.elapsed() > timeout) } /// Should the SSE event surface (`aura.tool_requested` / `aura.tool_complete` @@ -465,10 +467,10 @@ impl StreamingRequestHook { /// Check and cancel if needed, logging the reason. fn check_and_cancel(&self, cancel_sig: CancelSignal, context: &str) { - if *self.cancelled.borrow() { + if self.cancelled.is_cancelled() { tracing::info!("Request cancelled externally during {}", context); cancel_sig.cancel(); - } else if self.start_time.elapsed() > self.timeout { + } else if self.deadline_passed() { tracing::warn!( "Request timeout ({:?}) exceeded during {} - cancelling", self.timeout, @@ -774,8 +776,10 @@ mod tests { #[test] fn test_streaming_request_hook_creation() { - let (hook, _tx, _usage_state) = - StreamingRequestHook::new(Duration::from_secs(60), "test_req_1"); + let (hook, _tx, _usage_state) = StreamingRequestHook::new( + crate::streaming::RunOptions::bounded(Some(Duration::from_secs(60))), + "test_req_1", + ); assert!(!hook.should_cancel()); assert_eq!(hook.request_id, "test_req_1"); } @@ -816,20 +820,23 @@ mod tests { #[test] fn test_external_cancellation() { - let (hook, tx, _usage_state) = - StreamingRequestHook::new(Duration::from_secs(60), "test_req_2"); + let (hook, cancel, _usage_state) = StreamingRequestHook::new( + crate::streaming::RunOptions::bounded(Some(Duration::from_secs(60))), + "test_req_2", + ); assert!(!hook.should_cancel()); - // Signal cancellation - tx.send(true).unwrap(); + cancel.cancel(); assert!(hook.should_cancel()); } #[test] fn test_timeout_detection() { // Create hook with very short timeout - let (hook, _tx, _usage_state) = - StreamingRequestHook::new(Duration::from_millis(1), "test_req_3"); + let (hook, _tx, _usage_state) = StreamingRequestHook::new( + crate::streaming::RunOptions::bounded(Some(Duration::from_millis(1))), + "test_req_3", + ); // Wait for timeout std::thread::sleep(Duration::from_millis(5)); @@ -838,8 +845,10 @@ mod tests { #[test] fn test_usage_state_creation() { - let (_hook, _tx, usage_state) = - StreamingRequestHook::new(Duration::from_secs(60), "test_req_4"); + let (_hook, _tx, usage_state) = StreamingRequestHook::new( + crate::streaming::RunOptions::bounded(Some(Duration::from_secs(60))), + "test_req_4", + ); // Initially all zeros let (prompt, completion, total) = usage_state.get_final_usage(); @@ -965,8 +974,10 @@ mod tests { #[test] fn test_usage_state_shared_between_clones() { - let (_hook, _tx, usage_state) = - StreamingRequestHook::new(Duration::from_secs(60), "test_req_5"); + let (_hook, _tx, usage_state) = StreamingRequestHook::new( + crate::streaming::RunOptions::bounded(Some(Duration::from_secs(60))), + "test_req_5", + ); let usage_state_clone = usage_state.clone(); // Modify through original (tool turn) @@ -1026,9 +1037,15 @@ mod tests { #[test] fn test_with_scratchpad_budget_none_matches_new() { - let (hook_a, _, _) = StreamingRequestHook::new(Duration::from_secs(60), "req_a"); - let (hook_b, _, _) = - StreamingRequestHook::with_scratchpad_budget(Duration::from_secs(60), "req_b", None); + let (hook_a, _, _) = StreamingRequestHook::new( + crate::streaming::RunOptions::bounded(Some(Duration::from_secs(60))), + "req_a", + ); + let (hook_b, _, _) = StreamingRequestHook::with_scratchpad_budget( + crate::streaming::RunOptions::bounded(Some(Duration::from_secs(60))), + "req_b", + None, + ); // Both hooks should report no scratchpad budget. assert!(hook_a.scratchpad_budget.is_none()); assert!(hook_b.scratchpad_budget.is_none()); @@ -1040,7 +1057,7 @@ mod tests { let counter = Arc::new(TiktokenCounter::default_counter()); let budget = ContextBudget::new(128_000, 0.20, 0, counter); let (hook, _, _) = StreamingRequestHook::with_scratchpad_budget( - Duration::from_secs(60), + crate::streaming::RunOptions::bounded(Some(Duration::from_secs(60))), "req_with_budget", Some(budget.clone()), ); From c8c596fa87a720a1a9d6b9962a63aeb2b66b899e Mon Sep 17 00:00:00 2001 From: Jacob Hull Date: Tue, 1 Sep 2026 16:18:15 -0700 Subject: [PATCH 2/3] feat(agents): carry run identity in task-local scope Run identity lived on MCP clients as mutable state, set before every stream. It now lives in a scope the agent establishes, which anything running inside the run reads directly. Two things sit outside that scope. Rig executes tools on its own server task, so a client binds the call it serves and a tool call finds it there. Progress notifications arrive on the transport task, correlated by their call's progress token. Both hold the request id with the agent that call acts for, so events keep their attribution. rmcp mints a progress token inside the send, so a server can answer before the call has claimed it. The call the client is bound to covers that window, and a notification with no live call at all stays unowned. Warning waits for a count of them, because firing on the first named a cause it could not know. Ref: #625 Signed-off-by: Jacob Hull --- crates/aura/src/builder.rs | 10 +- crates/aura/src/lib.rs | 1 + crates/aura/src/mcp/client.rs | 265 +++++++++++++++++++---- crates/aura/src/mcp/manager.rs | 10 +- crates/aura/src/mcp/progress.rs | 244 +++++++++++++-------- crates/aura/src/orchestration/factory.rs | 11 +- crates/aura/src/run_context.rs | 140 ++++++++++++ 7 files changed, 545 insertions(+), 136 deletions(-) create mode 100644 crates/aura/src/run_context.rs diff --git a/crates/aura/src/builder.rs b/crates/aura/src/builder.rs index 707e23683..5accda705 100644 --- a/crates/aura/src/builder.rs +++ b/crates/aura/src/builder.rs @@ -1686,9 +1686,11 @@ impl StreamingAgent for Agent { options: crate::streaming::RunOptions, request_id: &str, ) -> crate::streaming::AgentRun { + // Rig runs tools on its own server task, so bind the run where a tool + // call can still find it. if let Some(mcp_manager) = &self.mcp_manager { mcp_manager - .set_current_call(request_id, aura_events::AgentContext::single_agent()) + .bind_call(request_id, aura_events::AgentContext::single_agent()) .await; } @@ -1699,6 +1701,12 @@ impl StreamingAgent for Agent { self.stream_chat_with_timeout(query, chat_history, options, request_id) .await } + .map_stream(|stream| { + Box::pin(crate::run_context::scope_stream( + request_id.to_string(), + stream, + )) + }) } async fn cancel_and_close_mcp(&self, request_id: &str, reason: &str) -> usize { diff --git a/crates/aura/src/lib.rs b/crates/aura/src/lib.rs index c8d6f1eba..c67950378 100644 --- a/crates/aura/src/lib.rs +++ b/crates/aura/src/lib.rs @@ -31,6 +31,7 @@ pub mod rag_tools; pub mod request_cancellation; pub mod request_progress; pub mod rig_builder; +pub mod run_context; mod schema_sanitize; // Private - MCP schema sanitization for OpenAI compatibility pub mod scratchpad; pub mod session_store; diff --git a/crates/aura/src/mcp/client.rs b/crates/aura/src/mcp/client.rs index 4235cec57..b5925fc93 100644 --- a/crates/aura/src/mcp/client.rs +++ b/crates/aura/src/mcp/client.rs @@ -8,7 +8,7 @@ use rmcp::{ RoleClient, model::{ CallToolRequestParam, CancelledNotificationParam, ClientRequest, ProgressNotificationParam, - Request, RequestId, Tool, + ProgressToken, Request, RequestId, Tool, }, serve_client, service::{PeerRequestOptions, RunningService}, @@ -373,7 +373,41 @@ pub struct McpClient { namespace: ToolNamespace, /// Tracks in-flight MCP requests for cancellation support in_flight: Arc, - current_call: Arc>>, + /// Which call each in-flight progress token belongs to, shared with the + /// progress handler. + token_owners: Arc>>, + /// The call this client serves. + bound_call: Arc>>, +} + +/// A call's claim on its progress token. +struct ProgressTokenGuard { + owners: Arc>>, + token: Option, +} + +impl Drop for ProgressTokenGuard { + /// Releases the token when the call ends, including a call whose future is + /// dropped mid-await. Keeping entries for finished calls would grow the map + /// for the client's lifetime. + /// + /// A notification trailing past its call's result finds the token gone and + /// resolves through the client's bound call, so it still reaches its run. + fn drop(&mut self) { + if let Some(token) = self.token.take() { + owners_of(&self.owners).remove(&token); + } + } +} + +/// A poisoned map only means a holder panicked mid-update; the entries are +/// still sound, so recover rather than propagate. +fn owners_of( + owners: &Mutex>, +) -> std::sync::MutexGuard<'_, HashMap> { + owners + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) } /// The request an MCP client is currently serving, and the agent on whose @@ -392,7 +426,8 @@ impl Clone for McpClient { server_url: self.server_url.clone(), namespace: self.namespace.clone(), in_flight: self.in_flight.clone(), - current_call: self.current_call.clone(), + token_owners: self.token_owners.clone(), + bound_call: self.bound_call.clone(), } } } @@ -412,8 +447,13 @@ impl McpClient { T: rmcp::transport::Transport + Send + 'static, T::Error: std::error::Error + Send + Sync + 'static, { - let current_call = Arc::new(RwLock::new(None)); - let handler = ProgressEnabledHandler::new(Arc::clone(¤t_call), user_agent); + let token_owners = Arc::new(Mutex::new(HashMap::new())); + let bound_call = Arc::new(RwLock::new(None)); + let handler = ProgressEnabledHandler::new( + Arc::clone(&token_owners), + Arc::clone(&bound_call), + user_agent, + ); let client = serve_client(handler, transport) .await @@ -424,7 +464,8 @@ impl McpClient { server_url, namespace, in_flight: Arc::new(InFlightRequests::new()), - current_call, + token_owners, + bound_call, }) } @@ -509,33 +550,50 @@ impl McpClient { Ok(client) } - /// Request id and agent are stored together so a reader sees one - /// request's id paired with that same request's agent. - pub async fn set_current_call(&self, request_id: &str, agent: AgentContext) { - *self.current_call.write().await = Some(CallContext { + /// Binds this client to the call it serves for that call's duration. + /// + /// Request id and agent are stored together so a reader sees one request's + /// id paired with that same request's agent. + pub async fn bind_call(&self, request_id: &str, agent: AgentContext) { + *self.bound_call.write().await = Some(CallContext { request_id: request_id.to_string(), agent, }); - debug!("Set current call for MCP client: {}", request_id); + debug!("Bound MCP client to call: {}", request_id); } pub async fn clear_current_call(&self) { - let mut guard = self.current_call.write().await; + let mut guard = self.bound_call.write().await; if let Some(call) = guard.take() { debug!("Cleared current call: {}", call.request_id); } } - pub async fn get_current_request(&self) -> Option { - let guard = self.current_call.read().await; - guard.as_ref().map(|call| call.request_id.clone()) + /// The run in task-local scope, or the bound one otherwise. + /// + /// Rig executes tools on a long-lived server task, so a tool call reaches + /// here with no scope around it and finds its run through the binding. One + /// binding is enough because a client serves one run, an agent being built + /// per request and owning its manager. Sharing a manager across runs — warm + /// MCP reuse, #578 — makes this answer last-writer-wins, and wants per-run + /// tool instances or a rig-side change rather than another field here. + pub async fn run_id(&self) -> Option> { + match crate::run_context::current_run_id() { + Some(id) => Some(id), + None => self + .bound_call + .read() + .await + .as_ref() + .map(|call| Arc::from(call.request_id.as_str())), + } } /// The agent named for `request_id`. A client serving a different request, /// or none, yields the single-agent context — the same value the SSE /// handler stamps on a frame that arrives without an agent. async fn agent_for(&self, request_id: &str) -> AgentContext { - let guard = self.current_call.read().await; + let guard = self.bound_call.read().await; guard .as_ref() .filter(|call| call.request_id == request_id) @@ -543,6 +601,16 @@ impl McpClient { .unwrap_or_else(AgentContext::single_agent) } + /// Ties a progress token to the call that minted it, so notifications + /// arriving on the transport task can be routed back. + async fn own_progress_token(&self, token: ProgressToken, request_id: &str) { + let call = CallContext { + request_id: request_id.to_string(), + agent: self.agent_for(request_id).await, + }; + owners_of(&self.token_owners).insert(token, call); + } + pub async fn discover_tools(&self) -> Result> { debug!( "🔍 Starting tool discovery from MCP server: {}", @@ -572,14 +640,14 @@ impl McpClient { Ok(tools_response.tools) } - /// Execute a tool. Auto-tracks for cancellation if `set_current_request` was called. + /// Execute a tool, tracking it for cancellation when called inside a run. pub async fn call_tool( &self, tool_name: &str, arguments: HashMap, approver_overrides: Option, ) -> Result { - if let Some(http_request_id) = self.get_current_request().await { + if let Some(http_request_id) = self.run_id().await { info!( "Tool '{}' executing WITH automatic tracking (http_request_id={})", tool_name, http_request_id @@ -655,6 +723,18 @@ impl McpClient { .context("Failed to send tool call request")?; let progress_token = handle.progress_token.clone(); + // The transport task cannot read the run's task-local, so tie the token + // to the run here, while still inside it. + if let Some(run_id) = self.run_id().await { + self.own_progress_token(progress_token.clone(), &run_id) + .await; + } + // Held from here so a run cancelled mid-await, which drops this future + // before it returns, still releases the entry. + let _token_guard = ProgressTokenGuard { + owners: Arc::clone(&self.token_owners), + token: Some(progress_token.clone()), + }; info!( "Tool '{}' started with progress token: {:?}", tool_name, progress_token @@ -687,10 +767,8 @@ impl McpClient { debug!("Progress stream ended for '{}'", tool_name_for_task); }); - let response = handle - .await_response() - .await - .context(format!("Tool '{}' execution failed", tool_name))?; + let response = handle.await_response().await; + let response = response.context(format!("Tool '{}' execution failed", tool_name))?; match response { rmcp::model::ServerResult::CallToolResult(result) => { @@ -732,6 +810,15 @@ impl McpClient { .await .context("Failed to send tool call request")?; + if let Some(run_id) = self.run_id().await { + self.own_progress_token(handle.progress_token.clone(), &run_id) + .await; + } + let _token_guard = ProgressTokenGuard { + owners: Arc::clone(&self.token_owners), + token: Some(handle.progress_token.clone()), + }; + // Extract what we need for cancellation before moving handle let request_id = handle.id.clone(); let peer = handle.peer.clone(); @@ -810,6 +897,22 @@ impl McpClient { .await .context("Failed to send tool call request")?; + // First thing after the send, because rmcp mints the token inside that + // call and the request is already on the wire when it returns. A + // notification answering before this lands routes through the client's + // bound call instead, which is this one. + // + // Unconditional, unlike the paths that discover their run: this one is + // handed the request id by its caller. + self.own_progress_token(handle.progress_token.clone(), http_request_id) + .await; + // Held from here so a run cancelled mid-await, which drops this future + // before it returns, still releases the entry. + let _token_guard = ProgressTokenGuard { + owners: Arc::clone(&self.token_owners), + token: Some(handle.progress_token.clone()), + }; + // Track this request for potential cancellation let mcp_request_id = handle.id.clone(); self.in_flight @@ -860,7 +963,6 @@ impl McpClient { // Await the tool result let result = handle.await_response().await; - // Remove from tracking (completed or failed) self.in_flight .remove(http_request_id, &mcp_request_id) .await; @@ -951,7 +1053,9 @@ impl McpClient { pub async fn cancel_and_close(&self, http_request_id: &str, reason: &str) -> usize { let count = self.cancel_all_for_request(http_request_id, reason).await; - // Also clear the call to stop routing any straggler progress notifications + // Drop this call's token ownership so straggler notifications stop + // routing, then clear the binding. + owners_of(&self.token_owners).retain(|_, call| call.request_id != http_request_id); self.clear_current_call().await; // Forcefully close connection - server is ignoring cancellation anyway @@ -976,6 +1080,58 @@ pub(crate) mod tests { use super::*; use crate::approver_headers::tests::captured_overrides; + fn owner(request_id: &str) -> CallContext { + CallContext { + request_id: request_id.to_string(), + agent: AgentContext::single_agent(), + } + } + + /// A cancelled run drops the tool future mid-await, so a release placed + /// after the await never runs. The guard is what covers that path, and an + /// entry left behind outlives its call for the client's lifetime. + #[test] + fn a_dropped_call_releases_its_progress_token() { + let owners: Arc>> = + Arc::new(Mutex::new(HashMap::new())); + let token = ProgressToken(rmcp::model::NumberOrString::Number(7)); + + owners_of(&owners).insert(token.clone(), owner("req_1")); + { + let _guard = ProgressTokenGuard { + owners: Arc::clone(&owners), + token: Some(token.clone()), + }; + assert!(owners_of(&owners).contains_key(&token)); + // The future is dropped here rather than returning. + } + + assert!( + !owners_of(&owners).contains_key(&token), + "the entry goes with the call that owned it" + ); + } + + /// One call's guard must not take another call's entry. + #[test] + fn a_guard_releases_only_its_own_token() { + let owners: Arc>> = + Arc::new(Mutex::new(HashMap::new())); + let mine = ProgressToken(rmcp::model::NumberOrString::Number(1)); + let theirs = ProgressToken(rmcp::model::NumberOrString::Number(2)); + + owners_of(&owners).insert(mine.clone(), owner("req_1")); + owners_of(&owners).insert(theirs.clone(), owner("req_2")); + + drop(ProgressTokenGuard { + owners: Arc::clone(&owners), + token: Some(mine.clone()), + }); + + assert!(!owners_of(&owners).contains_key(&mine)); + assert!(owners_of(&owners).contains_key(&theirs)); + } + #[tokio::test] async fn test_in_flight_requests_tracking() { let tracker = InFlightRequests::new(); @@ -1471,7 +1627,7 @@ pub(crate) mod tests { "an unnamed client falls back to the single-agent context" ); - client.set_current_call("req-1", worker.clone()).await; + client.bind_call("req-1", worker.clone()).await; assert_eq!(client.agent_for("req-1").await, worker); assert_eq!( client.agent_for("req-2").await, @@ -1487,7 +1643,42 @@ pub(crate) mod tests { ); } - /// `set_current_request` selects the tracked branch, so this is the same entry point a gated call takes in the server and the branch choice must not decide whether identity is delivered. + /// Rig executes tools on a long-lived server task, so a real tool call runs + /// with no run in task-local scope. Binding is the only thing that lets it + /// find its run, and without it the call silently takes the untracked + /// branch and stops emitting `aura.tool_start`. + #[tokio::test] + async fn a_bound_run_is_found_where_no_scope_reaches() { + let (_server, client) = client_and_server(&requester_headers()).await; + + assert_eq!(client.run_id().await, None, "no scope, nothing bound"); + + client + .bind_call("req_bound", AgentContext::single_agent()) + .await; + assert_eq!(client.run_id().await.as_deref(), Some("req_bound")); + } + + /// A scope still wins, so a call made inside one is attributed to that run + /// rather than whatever the client was last bound to. + #[tokio::test] + async fn a_scope_takes_precedence_over_the_binding() { + let (_server, client) = client_and_server(&requester_headers()).await; + client + .bind_call("req_bound", AgentContext::single_agent()) + .await; + + let seen = crate::run_context::with_run_id("req_scoped".to_string(), async { + client.run_id().await + }) + .await; + + assert_eq!(seen.as_deref(), Some("req_scoped")); + } + + /// Being inside a run selects the tracked branch, so this is the same entry + /// point a gated call takes in the server and the branch choice must not + /// decide whether identity is delivered. #[tokio::test] async fn call_tool_delivers_the_override_on_either_branch() { let (server, client) = client_and_server(&requester_headers()).await; @@ -1501,17 +1692,17 @@ pub(crate) mod tests { .await .expect("the untracked call succeeds"); - client - .set_current_call("http-req-1", AgentContext::single_agent()) - .await; - client - .call_tool( - "tracked", - no_args(), - Some(captured_overrides("x-forwarded-user", "bob")), - ) - .await - .expect("the tracked call succeeds"); + crate::run_context::with_run_id("http-req-1".to_string(), async { + client + .call_tool( + "tracked", + no_args(), + Some(captured_overrides("x-forwarded-user", "bob")), + ) + .await + .expect("the tracked call succeeds"); + }) + .await; let calls = server.tool_calls(); assert_eq!(calls[0].header_values("x-forwarded-user"), vec!["alice"]); diff --git a/crates/aura/src/mcp/manager.rs b/crates/aura/src/mcp/manager.rs index b4acaec69..f523b0c65 100644 --- a/crates/aura/src/mcp/manager.rs +++ b/crates/aura/src/mcp/manager.rs @@ -827,17 +827,15 @@ impl McpManager { .chain(self.stdio_clients.values()) } - /// Name the request these clients are serving and the agent it belongs to. - pub async fn set_current_call(&self, http_request_id: &str, agent: aura_events::AgentContext) { + /// Bind these clients to the call they serve and the agent it belongs to. + pub async fn bind_call(&self, http_request_id: &str, agent: aura_events::AgentContext) { let mut total_clients = 0; for client in self.clients() { - client - .set_current_call(http_request_id, agent.clone()) - .await; + client.bind_call(http_request_id, agent.clone()).await; total_clients += 1; } debug!( - "Set current call on {} MCP client(s): {}", + "Bound {} MCP client(s) to call: {}", total_clients, http_request_id ); } diff --git a/crates/aura/src/mcp/progress.rs b/crates/aura/src/mcp/progress.rs index 1b9071c93..348829093 100644 --- a/crates/aura/src/mcp/progress.rs +++ b/crates/aura/src/mcp/progress.rs @@ -21,48 +21,36 @@ use rmcp::{ ClientHandler, handler::client::progress::ProgressDispatcher, - model::{ClientInfo, Implementation, ProgressNotificationParam}, + model::{ClientInfo, Implementation, ProgressNotificationParam, ProgressToken}, service::{NotificationContext, RoleClient}, }; +use std::collections::HashMap; use std::sync::Arc; -use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; -use tokio::sync::RwLock; -use tracing::{debug, info, warn}; +use std::sync::atomic::{AtomicU64, Ordering}; +use tracing::{debug, warn}; use aura_events::Progress; use aura_events::agent::{AgentEvent, AgentEventPayload}; use crate::mcp::client::CallContext; +/// How many unowned progress notifications mean a server is ignoring +/// cancellation rather than trailing one past its call. +const ORPHANED_PROGRESS_ALARM: u64 = 16; + /// A custom ClientHandler that routes progress notifications to request-scoped channels. /// /// This handler is used instead of `()` when creating MCP clients to enable /// progress notification support. Progress notifications received from the -/// server are routed to the specific HTTP request that initiated the tool call, -/// ensuring no cross-request or cross-customer data leakage. -/// -/// # Example -/// ```ignore -/// // Create handler with a shared reference to the in-flight call -/// let current_call = Arc::new(RwLock::new(None)); -/// let handler = ProgressEnabledHandler::new(current_call.clone(), "aura/0.1.0"); -/// let client = serve_client(handler.clone(), transport).await?; -/// -/// // Name the call before tool execution -/// *current_call.write().await = Some(CallContext { -/// request_id: "req_123".to_string(), -/// agent: AgentContext::single_agent(), -/// }); -/// -/// // Progress notifications will now be routed to req_123's channel -/// ``` +/// server are routed to the call that initiated them, by the progress token +/// that call minted, so concurrent runs cannot see each other's progress. #[derive(Clone)] pub struct ProgressEnabledHandler { progress_dispatcher: ProgressDispatcher, - /// The call this client is serving, shared with the [`McpClient`] that owns it. - current_call: Arc>>, - /// Flag to log orphaned progress only once (prevents log flood from servers ignoring cancellation) - logged_orphaned_warning: Arc, + /// Which call each in-flight progress token belongs to. + token_owners: Arc>>, + /// The call this handler's client serves, shared with it. + bound_call: Arc>>, /// Counter for orphaned progress notifications (for diagnostics) orphaned_count: Arc, /// This client's MCP `clientInfo`. @@ -73,21 +61,45 @@ impl std::fmt::Debug for ProgressEnabledHandler { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("ProgressEnabledHandler") .field("progress_dispatcher", &self.progress_dispatcher) - .field("current_call", &"Arc>>") + .field("token_owners", &"Arc>>") .finish() } } impl ProgressEnabledHandler { + /// The call a progress token belongs to, or `None` once that call has ended. + /// Notifications arrive on the transport's task, which cannot read the + /// run's task-local, so the token is the only thing tying one back. + /// + /// rmcp mints a token inside the send, so a server can answer before the + /// call has claimed it. The client serves one call, which is whose that + /// notification is. + pub async fn owner_of(&self, token: &ProgressToken) -> Option { + let owned = self + .token_owners + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + .get(token) + .cloned(); + match owned { + Some(call) => Some(call), + None => self.bound_call.read().await.clone(), + } + } + /// `user_agent` is the same `product/version` token sent as the HTTP /// `User-Agent` header; the handshake announces it split into /// `clientInfo.name` and `clientInfo.version`, the shape MCP servers /// expect there. - pub fn new(current_call: Arc>>, user_agent: &str) -> Self { + pub fn new( + token_owners: Arc>>, + bound_call: Arc>>, + user_agent: &str, + ) -> Self { Self { progress_dispatcher: ProgressDispatcher::new(), - current_call, - logged_orphaned_warning: Arc::new(AtomicBool::new(false)), + token_owners, + bound_call, orphaned_count: Arc::new(AtomicU64::new(0)), client_info: ClientInfo { client_info: implementation_from_user_agent(user_agent), @@ -96,9 +108,8 @@ impl ProgressEnabledHandler { } } - /// Reset the orphaned warning flag and counter (call when setting a new request ID) + /// Reset the orphaned counter (call when setting a new request ID) pub fn reset_orphaned_tracking(&self) { - self.logged_orphaned_warning.store(false, Ordering::SeqCst); self.orphaned_count.store(0, Ordering::SeqCst); } @@ -156,11 +167,12 @@ impl ClientHandler for ProgressEnabledHandler { _context: NotificationContext, ) -> impl std::future::Future + Send + '_ { async move { - // One read, so the id and the agent describe the same call. - let current_call = self.current_call.read().await.clone(); + // One lookup, so the id and the agent describe the same call. + let call = self.owner_of(¶ms.progress_token).await; - if let Some(CallContext { request_id, agent }) = current_call { + if let Some(CallContext { request_id, agent }) = call { let req_id = &request_id; + let routed = crate::agent_events::emit( req_id, AgentEvent::new( @@ -183,20 +195,23 @@ impl ClientHandler for ProgressEnabledHandler { ); } } else { - // No request context - could be CLI mode, test, or cancelled request - // Increment counter and log at INFO so we can see the flow + // Neither the token nor the client names a call, so the + // notification has no run to reach. Cancelling a call clears + // both, and a client that never bound one never had either. let count = self.orphaned_count.fetch_add(1, Ordering::SeqCst) + 1; - // First orphaned notification gets a warning - if !self.logged_orphaned_warning.swap(true, Ordering::SeqCst) { + // A stream of these means the server kept going after being + // told to stop. + if count == ORPHANED_PROGRESS_ALARM { warn!( - "MCP server ignoring cancellation - orphaned progress notifications arriving" + "{} MCP progress notifications with no live call — the server \ + may be ignoring notifications/cancelled", + count ); } - // Log every orphaned notification at INFO for visibility - info!( - "Orphaned MCP progress #{}: progress={}, message={:?}", + debug!( + "Unowned MCP progress #{}: progress={}, message={:?}", count, params.progress, params.message ); } @@ -211,9 +226,10 @@ impl ClientHandler for ProgressEnabledHandler { #[cfg(test)] mod tests { use super::*; + use rmcp::model::NumberOrString; - fn create_test_handler() -> ProgressEnabledHandler { - ProgressEnabledHandler::new(Arc::new(RwLock::new(None)), "test/0") + fn token(n: i64) -> ProgressToken { + ProgressToken(NumberOrString::Number(n)) } fn call(request_id: &str) -> CallContext { @@ -223,10 +239,30 @@ mod tests { } } + fn handler_owning(pairs: &[(i64, &str)]) -> ProgressEnabledHandler { + let owners = pairs + .iter() + .map(|(t, run)| (token(*t), call(run))) + .collect::>(); + ProgressEnabledHandler::new( + Arc::new(std::sync::Mutex::new(owners)), + Arc::new(tokio::sync::RwLock::new(None)), + "test/0", + ) + } + + fn create_test_handler() -> ProgressEnabledHandler { + handler_owning(&[]) + } + /// A `product/version` token lands as separate name and version fields. #[test] fn handshake_info_splits_the_user_agent_into_name_and_version() { - let handler = ProgressEnabledHandler::new(Arc::new(RwLock::new(None)), "aura/1.2.3"); + let handler = ProgressEnabledHandler::new( + Arc::new(std::sync::Mutex::new(HashMap::new())), + Arc::new(tokio::sync::RwLock::new(None)), + "aura/1.2.3", + ); let info = handler.get_info().client_info; assert_eq!(info.name, "aura"); assert_eq!(info.version, "1.2.3"); @@ -237,7 +273,11 @@ mod tests { #[test] fn handshake_info_falls_back_to_the_crate_version() { for token in ["mezmo-aura", "mezmo-aura/", " mezmo-aura / "] { - let handler = ProgressEnabledHandler::new(Arc::new(RwLock::new(None)), token); + let handler = ProgressEnabledHandler::new( + Arc::new(std::sync::Mutex::new(HashMap::new())), + Arc::new(tokio::sync::RwLock::new(None)), + token, + ); let info = handler.get_info().client_info; assert_eq!(info.name, "mezmo-aura", "token {token:?}"); assert_eq!(info.version, env!("CARGO_PKG_VERSION"), "token {token:?}"); @@ -246,57 +286,94 @@ mod tests { #[test] fn test_handler_creation() { - let handler = create_test_handler(); - // Just verify it can be created and progress_dispatcher is accessible + let handler = handler_owning(&[]); let _ = handler.progress_dispatcher(); } #[test] fn test_handler_clone() { - let handler = create_test_handler(); + let handler = handler_owning(&[]); let cloned = handler.clone(); - // Both should have accessible progress dispatchers let _ = cloned.progress_dispatcher(); } #[tokio::test] - async fn test_handler_with_request_id() { - let current_call = Arc::new(RwLock::new(Some(call("req_test_123")))); - let handler = ProgressEnabledHandler::new(current_call.clone(), "test/0"); + async fn an_unowned_token_has_no_run() { + assert!(handler_owning(&[]).owner_of(&token(1)).await.is_none()); + } + + /// rmcp mints a progress token inside the send, so a server can answer + /// before the call has claimed it. The client serves one call, so that + /// notification routes to it rather than being dropped. + #[tokio::test] + async fn a_token_claimed_after_its_first_notification_still_routes() { + let owners = Arc::new(std::sync::Mutex::new(HashMap::new())); + let bound = Arc::new(tokio::sync::RwLock::new(Some(call("run_a")))); + let handler = ProgressEnabledHandler::new(owners.clone(), bound, "test/0"); - // Verify request ID is accessible - let guard = handler.current_call.read().await; + // Nothing owns the token yet: the send has returned but the claim has not + // landed. assert_eq!( - guard.as_ref().map(|c| c.request_id.as_str()), - Some("req_test_123") + handler + .owner_of(&token(1)) + .await + .map(|call| call.request_id), + Some("run_a".to_string()), + "an unclaimed token belongs to the call this client serves" + ); + + // Once claimed, the token answers for itself. + owners.lock().unwrap().insert(token(1), call("run_a_tool")); + assert_eq!( + handler + .owner_of(&token(1)) + .await + .map(|call| call.request_id), + Some("run_a_tool".to_string()), + "a claim is more precise than the binding" ); } + /// Each token keeps its own call, so concurrent runs cannot pick up each + /// other's progress. #[tokio::test] - async fn test_handler_request_id_changes() { - let current_call = Arc::new(RwLock::new(None)); - let handler = ProgressEnabledHandler::new(current_call.clone(), "test/0"); - - // Initially no request ID - { - let guard = handler.current_call.read().await; - assert!(guard.is_none()); - } + async fn concurrent_runs_route_by_their_own_token() { + let handler = handler_owning(&[(1, "run_a"), (2, "run_b")]); - // Set request ID - { - let mut guard = current_call.write().await; - *guard = Some(call("req_456")); - } + assert_eq!( + handler.owner_of(&token(1)).await.map(|c| c.request_id), + Some("run_a".to_string()) + ); + assert_eq!( + handler.owner_of(&token(2)).await.map(|c| c.request_id), + Some("run_b".to_string()) + ); + } - // Handler should see the new value - { - let guard = handler.current_call.read().await; - assert_eq!( - guard.as_ref().map(|c| c.request_id.as_str()), - Some("req_456") - ); - } + /// A call releases its token when it ends, so a notification trailing past + /// the result falls back to the call the client serves and still reaches + /// that run. Routing stops only once nothing names a call at all. + #[tokio::test] + async fn a_released_token_routes_through_the_binding() { + let owners = Arc::new(std::sync::Mutex::new(HashMap::new())); + let bound = Arc::new(tokio::sync::RwLock::new(Some(call("run_a")))); + let handler = ProgressEnabledHandler::new(owners.clone(), Arc::clone(&bound), "test/0"); + + owners.lock().unwrap().insert(token(7), call("run_a_tool")); + assert_eq!( + handler.owner_of(&token(7)).await.map(|c| c.request_id), + Some("run_a_tool".to_string()) + ); + + owners.lock().unwrap().remove(&token(7)); + assert_eq!( + handler.owner_of(&token(7)).await.map(|c| c.request_id), + Some("run_a".to_string()), + "the binding outlives the tokens of the calls it serves" + ); + + *bound.write().await = None; + assert!(handler.owner_of(&token(7)).await.is_none()); } #[test] @@ -304,7 +381,6 @@ mod tests { let handler = create_test_handler(); // Initially false and zero - assert!(!handler.logged_orphaned_warning.load(Ordering::SeqCst)); assert_eq!(handler.orphaned_count(), 0); // Simulate orphaned notifications @@ -312,14 +388,8 @@ mod tests { handler.orphaned_count.fetch_add(1, Ordering::SeqCst); assert_eq!(handler.orphaned_count(), 2); - // Warning flag - let was_logged = handler.logged_orphaned_warning.swap(true, Ordering::SeqCst); - assert!(!was_logged); - assert!(handler.logged_orphaned_warning.load(Ordering::SeqCst)); - // Reset works for both handler.reset_orphaned_tracking(); - assert!(!handler.logged_orphaned_warning.load(Ordering::SeqCst)); assert_eq!(handler.orphaned_count(), 0); } } diff --git a/crates/aura/src/orchestration/factory.rs b/crates/aura/src/orchestration/factory.rs index 32288e979..0c68c8496 100644 --- a/crates/aura/src/orchestration/factory.rs +++ b/crates/aura/src/orchestration/factory.rs @@ -68,7 +68,7 @@ impl OrchestratorFactory { // Capture parent span so child spans nest correctly in tracing. let parent_span = tracing::Span::current(); tokio::spawn(tracing::Instrument::instrument( - async move { + crate::run_context::with_run_id(request_id.clone(), async move { let _done_guard = done_guard; let mut orchestrator = match Orchestrator::new(agent_config).await { Ok(o) => o, @@ -82,13 +82,14 @@ impl OrchestratorFactory { orchestrator.usage_state = usage_state; orchestrator.outer_budget = outer_budget; - // Set MCP request ID for progress notification routing, and - // surface per-server connection status so degraded/unavailable + // Surface per-server connection status so degraded/unavailable // MCP servers are visible in orchestration mode too (workers // share this one manager). if let Some(ref mcp_manager) = orchestrator.mcp_manager { + // Workers share this manager, and their tool calls run on + // rig's server task where the run's scope does not reach. mcp_manager - .set_current_call(&request_id, aura_events::AgentContext::coordinator()) + .bind_call(&request_id, aura_events::AgentContext::coordinator()) .await; let snapshot = mcp_manager.server_status_snapshot(); if !snapshot.is_empty() { @@ -139,7 +140,7 @@ impl OrchestratorFactory { } } } - }, + }), parent_span, )); diff --git a/crates/aura/src/run_context.rs b/crates/aura/src/run_context.rs new file mode 100644 index 000000000..013ea8b74 --- /dev/null +++ b/crates/aura/src/run_context.rs @@ -0,0 +1,140 @@ +//! Ambient run identity, scoped to the task doing the work. +//! +//! Rig invokes tools as `Tool::call(args)` with no call-time context, so a tool +//! cannot be handed the run it belongs to. A task-local supplies it without any +//! component storing it: concurrent runs each see their own, and nothing has to +//! be reset between runs. +//! +//! This reaches only code running inside the run's task. MCP progress +//! notifications arrive on the transport's own task and cannot read it — they +//! are correlated by progress token instead, registered at call time from here. + +use std::pin::Pin; +use std::sync::Arc; +use std::task::{Context, Poll}; + +use futures::Stream; + +tokio::task_local! { + static RUN_ID: Arc; +} + +pub fn current_run_id() -> Option> { + RUN_ID.try_with(Arc::clone).ok() +} + +/// Runs `f` with `run_id` in scope. Task-locals do not cross `tokio::spawn`, so +/// a spawned worker needs its own call rather than inheriting its parent's. +pub async fn with_run_id(run_id: String, f: F) -> F::Output { + RUN_ID.scope(Arc::from(run_id), f).await +} + +/// Enters the scope on every poll, so a stream polled by a consumer outside the +/// run still executes its tool calls with the run's identity in scope. +pub fn scope_stream(run_id: String, inner: S) -> ScopedStream { + ScopedStream { + run_id: Arc::from(run_id), + inner, + } +} + +pub struct ScopedStream { + run_id: Arc, + inner: S, +} + +impl Stream for ScopedStream { + type Item = S::Item; + + /// Entering the scope clones the id, so it is held behind an `Arc` rather + /// than reallocated on every poll. + fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + let this = self.get_mut(); + let run_id = Arc::clone(&this.run_id); + RUN_ID.sync_scope(run_id, || Pin::new(&mut this.inner).poll_next(cx)) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use futures::StreamExt; + + #[tokio::test] + async fn there_is_no_run_id_outside_a_run() { + assert_eq!(current_run_id(), None); + } + + #[tokio::test] + async fn a_scope_supplies_the_run_id() { + let seen = with_run_id("run_1".to_string(), async { current_run_id() }).await; + assert_eq!(seen.as_deref(), Some("run_1")); + } + + #[tokio::test] + async fn concurrent_runs_do_not_see_each_other() { + let a = tokio::spawn(with_run_id("run_a".to_string(), async { + tokio::task::yield_now().await; + current_run_id() + })); + let b = tokio::spawn(with_run_id("run_b".to_string(), async { + tokio::task::yield_now().await; + current_run_id() + })); + + assert_eq!(a.await.unwrap().as_deref(), Some("run_a")); + assert_eq!(b.await.unwrap().as_deref(), Some("run_b")); + } + + #[tokio::test] + async fn a_scoped_stream_carries_identity_into_each_poll() { + let inner = futures::stream::iter(0..3).map(|_| current_run_id()); + let seen: Vec<_> = scope_stream("run_s".to_string(), inner).collect().await; + + assert_eq!( + seen.iter().map(|id| id.as_deref()).collect::>(), + vec![Some("run_s"); 3] + ); + } + + /// Orchestration drives workers with `FuturesUnordered` inside the run's + /// task rather than spawning them, which is what lets their tool calls — + /// and so their progress — resolve to the run. + #[tokio::test] + async fn workers_driven_as_futures_keep_the_run_id() { + use futures::stream::{FuturesUnordered, StreamExt}; + + let seen = with_run_id("run_w".to_string(), async { + let mut workers: FuturesUnordered<_> = (0..3) + .map(|_| async { + tokio::task::yield_now().await; + current_run_id() + }) + .collect(); + + let mut ids = Vec::new(); + while let Some(id) = workers.next().await { + ids.push(id); + } + ids + }) + .await; + + assert_eq!( + seen.iter().map(|id| id.as_deref()).collect::>(), + vec![Some("run_w"); 3] + ); + } + + /// A spawned task does not inherit its parent's scope, which is why every + /// orchestration worker establishes its own. + #[tokio::test] + async fn a_spawned_task_does_not_inherit_the_scope() { + let seen = with_run_id("run_p".to_string(), async { + tokio::spawn(async { current_run_id() }).await.unwrap() + }) + .await; + + assert_eq!(seen, None); + } +} From 473360de0f982adf76cfb04b7a0a9bac779e542b Mon Sep 17 00:00:00 2001 From: Jacob Hull Date: Tue, 1 Sep 2026 17:18:05 -0700 Subject: [PATCH 3/3] feat(agents): make hooks an extension point Everything watching a run had to live inside the single hook rig allows per streaming request. An AgentHook trait now carries those concerns independently, and the rig-facing hook fans out to them. Cancellation and client-tool passthrough move behind it. A hook that ends a run says why, so the log can tell a blown deadline from a disconnect without the hook holding the deadline itself. Tool events and usage stay together, sharing the pending tool ids they both need. with_hook registers another, so adding a concern means writing an implementation instead of extending one function. Ref: #625 Signed-off-by: Jacob Hull --- crates/aura-test-utils/src/mock_agent.rs | 2 +- crates/aura/src/hooks.rs | 267 ++++++++++++++++++++++ crates/aura/src/lib.rs | 1 + crates/aura/src/orchestration/test_rig.rs | 2 + crates/aura/src/provider_agent.rs | 27 ++- crates/aura/src/streaming_request_hook.rs | 171 +++++++------- 6 files changed, 372 insertions(+), 98 deletions(-) create mode 100644 crates/aura/src/hooks.rs diff --git a/crates/aura-test-utils/src/mock_agent.rs b/crates/aura-test-utils/src/mock_agent.rs index 8ff8eb0dd..7cc84a7f1 100644 --- a/crates/aura-test-utils/src/mock_agent.rs +++ b/crates/aura-test-utils/src/mock_agent.rs @@ -214,7 +214,7 @@ mod tests { /// `into_events` wraps the stream, so the order is asserted rather than the /// count — a wrapper that buffered or reordered would keep the count. - #[tokio::test] + #[tokio::test(start_paused = true)] async fn a_yielding_agent_produces_its_items_then_ends() { let agent = MockAgent::yielding(vec![items::text("hello "), items::text("world")]); let mut stream = agent diff --git a/crates/aura/src/hooks.rs b/crates/aura/src/hooks.rs new file mode 100644 index 000000000..2e2e13100 --- /dev/null +++ b/crates/aura/src/hooks.rs @@ -0,0 +1,267 @@ +//! Things that observe a run's turns and tool calls. +//! +//! Rig exposes one hook per streaming request, so anything watching a run has +//! to be folded into that single implementation or bolted on elsewhere. This +//! trait is the seam that lets them be independent; the rig-facing hook fans +//! out to all of them. + +use std::sync::Arc; + +use async_trait::async_trait; + +pub struct ToolCall<'a> { + pub name: &'a str, +} + +/// Why a hook ended a run. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum RunCancelReason { + /// The run outlived the bound it was started with. + Deadline { after: std::time::Duration }, + /// Something outside the run cancelled it. + External, + /// The model called a tool the caller executes, so the run yields to it. + ClientTool, +} + +/// Something watching one run. +/// +/// A hook belongs to the run whose [`Hooks`] holds it and is registered when +/// that run's rig-facing hook is built. State it keeps is that run's, so sharing +/// one instance across runs reports a state one of them did not reach. +#[async_trait] +pub trait AgentHook: Send + Sync { + /// `false` ends the run rather than taking another LLM turn. + /// `Some` ends the run before the next turn, carrying why. + async fn before_turn(&self) -> Option { + None + } + + async fn on_tool_call(&self, _tool_name: &str) {} + + /// Ends the run at the next point the agent checks. + fn should_cancel(&self) -> Option { + None + } +} + +#[derive(Clone, Default)] +pub struct Hooks(Vec>); + +impl Hooks { + pub fn new() -> Self { + Self(Vec::new()) + } + + #[must_use] + pub fn with(mut self, hook: Arc) -> Self { + self.0.push(hook); + self + } +} + +#[async_trait] +impl AgentHook for Hooks { + /// One hook declining ends the run; the rest are still asked, so none + /// silently skips the turn it was told about. + async fn before_turn(&self) -> Option { + let mut stop = None; + for hook in &self.0 { + stop = hook.before_turn().await.or(stop); + } + stop + } + + async fn on_tool_call(&self, tool_name: &str) { + for hook in &self.0 { + hook.on_tool_call(tool_name).await; + } + } + + fn should_cancel(&self) -> Option { + self.0.iter().find_map(|hook| hook.should_cancel()) + } +} + +/// Ends a run once its deadline passes or its token is cancelled. +pub struct Deadline { + start: std::time::Instant, + timeout: Option, + cancelled: tokio_util::sync::CancellationToken, +} + +impl Deadline { + pub fn new( + timeout: Option, + cancelled: tokio_util::sync::CancellationToken, + ) -> Self { + Self { + start: std::time::Instant::now(), + timeout, + cancelled, + } + } +} + +#[async_trait] +impl AgentHook for Deadline { + fn should_cancel(&self) -> Option { + if self.cancelled.is_cancelled() { + Some(RunCancelReason::External) + } else { + self.timeout + .filter(|timeout| self.start.elapsed() > *timeout) + .map(|after| RunCancelReason::Deadline { after }) + } + } +} + +/// Ends a run after the model calls a tool the client executes, so the caller +/// can run it and resume in a follow-up request. +pub struct ClientTools { + names: std::collections::HashSet, + called: std::sync::atomic::AtomicBool, +} + +impl ClientTools { + pub fn new(names: std::collections::HashSet) -> Self { + Self { + names, + called: std::sync::atomic::AtomicBool::new(false), + } + } + + pub fn was_called(&self) -> bool { + self.called.load(std::sync::atomic::Ordering::Acquire) + } +} + +#[async_trait] +impl AgentHook for ClientTools { + async fn before_turn(&self) -> Option { + self.was_called().then_some(RunCancelReason::ClientTool) + } + + async fn on_tool_call(&self, tool_name: &str) { + if self.names.contains(tool_name) { + tracing::info!( + "Client tool '{}' called — the run yields after it", + tool_name + ); + self.called + .store(true, std::sync::atomic::Ordering::Release); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::atomic::{AtomicUsize, Ordering}; + + #[derive(Default)] + struct Counting { + calls: AtomicUsize, + proceed: bool, + cancel: bool, + } + + #[async_trait] + impl AgentHook for Counting { + async fn before_turn(&self) -> Option { + self.calls.fetch_add(1, Ordering::SeqCst); + (!self.proceed).then_some(RunCancelReason::External) + } + + fn should_cancel(&self) -> Option { + self.cancel.then_some(RunCancelReason::External) + } + } + + fn hook(proceed: bool, cancel: bool) -> Arc { + Arc::new(Counting { + calls: AtomicUsize::new(0), + proceed, + cancel, + }) + } + + #[tokio::test] + async fn an_empty_set_proceeds_and_does_not_cancel() { + let hooks = Hooks::new(); + assert!(hooks.before_turn().await.is_none()); + assert!(hooks.should_cancel().is_none()); + } + + #[tokio::test] + async fn one_hook_declining_ends_the_turn() { + let yes = hook(true, false); + let no = hook(false, false); + let hooks = Hooks::new().with(yes.clone()).with(no.clone()); + + assert!(hooks.before_turn().await.is_some()); + } + + /// A hook that stops being asked stops being able to observe the run, so + /// the decision is collected from all of them rather than short-circuited. + #[tokio::test] + async fn every_hook_is_asked_even_once_one_declines() { + let no = hook(false, false); + let later = hook(true, false); + let hooks = Hooks::new().with(no.clone()).with(later.clone()); + + assert!(hooks.before_turn().await.is_some()); + assert_eq!(later.calls.load(Ordering::SeqCst), 1); + } + + #[tokio::test] + async fn any_hook_can_cancel() { + let hooks = Hooks::new().with(hook(true, false)).with(hook(true, true)); + assert!(hooks.should_cancel().is_some()); + } + + fn client_tools(names: &[&str]) -> ClientTools { + ClientTools::new(names.iter().map(|n| (*n).to_string()).collect()) + } + + #[tokio::test] + async fn a_run_with_no_client_tools_keeps_going() { + let tools = client_tools(&[]); + assert!(tools.before_turn().await.is_none()); + } + + #[tokio::test] + async fn calling_a_client_tool_ends_the_run_before_the_next_turn() { + let tools = client_tools(&["Read"]); + assert!(tools.before_turn().await.is_none()); + + tools.on_tool_call("Read").await; + + assert!(tools.was_called()); + assert_eq!( + tools.before_turn().await, + Some(RunCancelReason::ClientTool), + "the caller must run the tool, and the reason says so" + ); + } + + #[tokio::test] + async fn a_server_side_tool_does_not_end_the_run() { + let tools = client_tools(&["Read"]); + + tools.on_tool_call("list_files").await; + + assert!(!tools.was_called()); + assert!(tools.before_turn().await.is_none()); + } + + #[tokio::test] + async fn a_deadline_cancels_once_its_token_does() { + let token = tokio_util::sync::CancellationToken::new(); + let deadline = Deadline::new(Some(std::time::Duration::from_secs(300)), token.clone()); + assert!(deadline.should_cancel().is_none()); + + token.cancel(); + assert!(deadline.should_cancel().is_some()); + } +} diff --git a/crates/aura/src/lib.rs b/crates/aura/src/lib.rs index c67950378..5ef09f021 100644 --- a/crates/aura/src/lib.rs +++ b/crates/aura/src/lib.rs @@ -17,6 +17,7 @@ pub mod fallback_tool_parser; pub mod fallback_tool_stream; pub mod governance; pub mod hitl; +pub mod hooks; pub mod inactivity; pub mod instance_id; pub mod logging; diff --git a/crates/aura/src/orchestration/test_rig.rs b/crates/aura/src/orchestration/test_rig.rs index 4b0953a0d..c95b2128c 100644 --- a/crates/aura/src/orchestration/test_rig.rs +++ b/crates/aura/src/orchestration/test_rig.rs @@ -49,6 +49,7 @@ use rig::streaming::{ StreamingCompletionResponse, }; use serde::{Deserialize, Serialize}; +use std::collections::HashSet; use tokio::sync::Notify; use tokio_util::sync::CancellationToken; @@ -625,6 +626,7 @@ pub(crate) async fn drive_worker( crate::streaming::RunOptions::bounded(Some(Duration::from_secs(60))), request_id, None, + HashSet::new(), ); let mut stream = rig diff --git a/crates/aura/src/provider_agent.rs b/crates/aura/src/provider_agent.rs index 5fa7cec9d..8cbabf127 100644 --- a/crates/aura/src/provider_agent.rs +++ b/crates/aura/src/provider_agent.rs @@ -112,9 +112,12 @@ impl ProviderAgent { scratchpad_budget: Option, client_tool_names: HashSet, ) -> crate::streaming::AgentRun { - let (hook, cancel_tx, usage_state) = - StreamingRequestHook::with_scratchpad_budget(options, request_id, scratchpad_budget); - let hook = hook.with_client_tool_names(client_tool_names); + let (hook, cancel_tx, usage_state) = StreamingRequestHook::with_scratchpad_budget( + options, + request_id, + scratchpad_budget, + client_tool_names, + ); match self { Self::OpenAI(agent) => { @@ -320,9 +323,12 @@ impl ProviderAgent { scratchpad_budget: Option, client_tool_names: HashSet, ) -> crate::streaming::AgentRun { - let (hook, cancel_tx, usage_state) = - StreamingRequestHook::with_scratchpad_budget(options, request_id, scratchpad_budget); - let hook = hook.with_client_tool_names(client_tool_names); + let (hook, cancel_tx, usage_state) = StreamingRequestHook::with_scratchpad_budget( + options, + request_id, + scratchpad_budget, + client_tool_names, + ); match self { Self::OpenAI(agent) => { @@ -427,9 +433,12 @@ impl ProviderAgent { scratchpad_budget: Option, client_tool_names: HashSet, ) -> crate::streaming::AgentRun { - let (hook, cancel_tx, usage_state) = - StreamingRequestHook::with_scratchpad_budget(options, request_id, scratchpad_budget); - let hook = hook.with_client_tool_names(client_tool_names); + let (hook, cancel_tx, usage_state) = StreamingRequestHook::with_scratchpad_budget( + options, + request_id, + scratchpad_budget, + client_tool_names, + ); match self { Self::OpenAI(agent) => { diff --git a/crates/aura/src/streaming_request_hook.rs b/crates/aura/src/streaming_request_hook.rs index 341bbb9cc..842ed54e3 100644 --- a/crates/aura/src/streaming_request_hook.rs +++ b/crates/aura/src/streaming_request_hook.rs @@ -38,16 +38,15 @@ //! let (prompt, completion, total) = usage_state.get_final_usage(); //! ``` +use crate::hooks::{AgentHook, ClientTools, Deadline, Hooks, RunCancelReason}; +use aura_events::agent::{AgentEvent, AgentEventPayload}; +use rig::agent::{CancelSignal, StreamingPromptHook}; +use rig::completion::{CompletionModel, GetTokenUsage, Message}; use std::collections::{HashMap, HashSet}; use std::future::Future; use std::sync::OnceLock; use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; use std::sync::{Arc, Mutex}; -use std::time::{Duration, Instant}; - -use aura_events::agent::{AgentEvent, AgentEventPayload}; -use rig::agent::{CancelSignal, StreamingPromptHook}; -use rig::completion::{CompletionModel, GetTokenUsage, Message}; use tokio_util::sync::CancellationToken; use crate::orchestration::BlockedCell; @@ -366,11 +365,8 @@ impl ResponseContent { /// MCP cancellation is handled separately via client-level tracking (Arc-based). #[derive(Clone)] pub struct StreamingRequestHook { - start_time: Instant, - /// Absent for a run the caller left unbounded. - timeout: Option, - /// External cancellation signal (e.g., from client disconnect) - cancelled: CancellationToken, + /// Everything watching this run. + hooks: Hooks, /// Request ID for event correlation request_id: String, /// Shared usage state (returned separately for handler access) @@ -379,13 +375,6 @@ pub struct StreamingRequestHook { /// LLM-reported per-turn input/output tokens into the budget as ground /// truth so `remaining()` reflects actual context pressure. scratchpad_budget: Option, - /// Names of client-side (passthrough) tools registered for this request. - /// When the LLM calls one, the stream is terminated before the next LLM - /// turn so the client can execute the tool locally and submit results back. - client_tool_names: HashSet, - /// Set when a client-side tool has been called in this request. Read by - /// `on_completion_call` to bail out before the next LLM turn. - client_tool_called: Arc, } impl StreamingRequestHook { @@ -394,7 +383,7 @@ impl StreamingRequestHook { options: crate::streaming::RunOptions, request_id: impl Into, ) -> (Self, CancellationToken, UsageState) { - Self::with_scratchpad_budget(options, request_id, None) + Self::with_scratchpad_budget(options, request_id, None, HashSet::new()) } /// Like `new`, but additionally wires a scratchpad `ContextBudget` so the @@ -405,42 +394,33 @@ impl StreamingRequestHook { options: crate::streaming::RunOptions, request_id: impl Into, scratchpad_budget: Option, + client_tool_names: HashSet, ) -> (Self, CancellationToken, UsageState) { let (timeout, cancel) = options.into_parts(); // A caller that supplied one can cancel before the run exists. let cancel = cancel.unwrap_or_default(); let usage_state = UsageState::new(); let hook = Self { - start_time: Instant::now(), - timeout, - cancelled: cancel.clone(), + hooks: Hooks::new() + .with(Arc::new(Deadline::new(timeout, cancel.clone()))) + .with(Arc::new(ClientTools::new(client_tool_names))), request_id: request_id.into(), usage_state: usage_state.clone(), scratchpad_budget, - client_tool_names: HashSet::new(), - client_tool_called: Arc::new(AtomicBool::new(false)), }; (hook, cancel, usage_state) } - /// Register the names of client-side (passthrough) tools for this request. - /// - /// When any of these tools are called, the hook ends the stream before the - /// next LLM turn so the client can execute the tool and submit results. - pub fn with_client_tool_names(mut self, names: HashSet) -> Self { - self.client_tool_names = names; + /// Registers another observer for this run, alongside the deadline and + /// client-tool concerns the hook starts with. + #[must_use] + pub fn with_hook(mut self, hook: Arc) -> Self { + self.hooks = self.hooks.with(hook); self } - /// Check if the request should be cancelled (timeout or external signal). - fn should_cancel(&self) -> bool { - self.cancelled.is_cancelled() || self.deadline_passed() - } - - /// An unbounded run has no deadline to pass. - fn deadline_passed(&self) -> bool { - self.timeout - .is_some_and(|timeout| self.start_time.elapsed() > timeout) + fn should_cancel(&self) -> Option { + AgentHook::should_cancel(&self.hooks) } /// Should the SSE event surface (`aura.tool_requested` / `aura.tool_complete` @@ -467,17 +447,26 @@ impl StreamingRequestHook { /// Check and cancel if needed, logging the reason. fn check_and_cancel(&self, cancel_sig: CancelSignal, context: &str) { - if self.cancelled.is_cancelled() { - tracing::info!("Request cancelled externally during {}", context); - cancel_sig.cancel(); - } else if self.deadline_passed() { - tracing::warn!( + let Some(reason) = self.should_cancel() else { + return; + }; + // A blown deadline is an operator's problem; the rest are routine. + match reason { + RunCancelReason::Deadline { after } => tracing::warn!( "Request timeout ({:?}) exceeded during {} - cancelling", - self.timeout, + after, context - ); - cancel_sig.cancel(); + ), + RunCancelReason::External => { + tracing::info!("Request cancelled externally during {}", context) + } + // Raised by `before_turn`, which logs it and ends the run itself, so + // reaching the cancel path means a hook also reports it here. + RunCancelReason::ClientTool => { + tracing::info!("Run yielding to a client tool during {}", context) + } } + cancel_sig.cancel(); } } @@ -493,11 +482,9 @@ where history: &[Message], cancel_sig: CancelSignal, ) -> impl Future + Send { - let has_client_tools = !self.client_tool_names.is_empty(); - let client_tool_called = self.client_tool_called.clone(); async move { - // Checked before the client-tool, cancel, and timeout checks: a - // waiting parked call must get its snapshot whatever else is true. + // Checked before the hooks: a waiting parked call must get its + // snapshot whatever else is true. if let Some(cell) = park_cell_for(&self.request_id) && cell.snapshot_if_pending(history, prompt) { @@ -509,15 +496,14 @@ where cancel_sig.cancel_with_reason(PARK_CANCEL_REASON); return; } - // If a passthrough tool was called this turn, do not initiate - // another LLM completion. Cancel here so the stream terminates - // and the streaming layer can emit `finish_reason: "tool_calls"` - // — the client will execute the tool and resume in a follow-up - // request. - if has_client_tools && client_tool_called.load(Ordering::Acquire) { - tracing::info!( - "Client tool was called — cancelling before next LLM completion call" - ); + + // A hook ending the run says why, so the log names it. Only a + // passthrough tool call leaves the marker the streaming layer reads + // for `finish_reason: "tool_calls"`; another reason ends the run + // without one. + if let Some(reason) = self.hooks.before_turn().await { + tracing::info!("Run ended before the next LLM completion call ({reason:?})"); + cancel_sig.cancel(); return; } @@ -532,10 +518,7 @@ where cancel_sig: CancelSignal, ) -> impl Future + Send { async move { - // Only check periodically for text deltas (they're frequent) - if self.should_cancel() { - self.check_and_cancel(cancel_sig, "text streaming"); - } + self.check_and_cancel(cancel_sig, "text streaming"); } } @@ -547,9 +530,7 @@ where cancel_sig: CancelSignal, ) -> impl Future + Send { async move { - if self.should_cancel() { - self.check_and_cancel(cancel_sig, "tool call delta"); - } + self.check_and_cancel(cancel_sig, "tool call delta"); } } @@ -571,22 +552,13 @@ where // makes the matching skip so push/pop stay symmetric. Cancellation // is still checked. let publish_event = Self::should_publish_tool_event(&tool_name); - let is_client_tool = self.client_tool_names.contains(&tool_name); - let client_tool_called = self.client_tool_called.clone(); async move { // Stash the call id so the park arm can record it on the cell entry. if let Some(cell) = park_cell_for(&request_id) { cell.set_current_call_id(tool_call_id.clone()); } - if is_client_tool { - tracing::info!( - "Client tool '{}' called for request '{}' — marking for passthrough", - tool_name, - request_id - ); - client_tool_called.store(true, Ordering::Release); - } + self.hooks.on_tool_call(&tool_name).await; if publish_event { // Parse args as JSON (fallback to empty object if invalid) @@ -622,10 +594,7 @@ where tool_call_id ); - if self.should_cancel() { - tracing::info!("Cancelling before tool '{}' execution", tool_name); - self.check_and_cancel(cancel_sig, &format!("tool call ({})", tool_name)); - } + self.check_and_cancel(cancel_sig, &format!("tool call ({})", tool_name)); } } @@ -681,10 +650,7 @@ where ); } - if self.should_cancel() { - tracing::info!("Cancelling after tool '{}' result", tool_name); - self.check_and_cancel(cancel_sig, &format!("tool result ({})", tool_name)); - } + self.check_and_cancel(cancel_sig, &format!("tool result ({})", tool_name)); } } @@ -773,6 +739,33 @@ where #[cfg(test)] mod tests { use super::*; + use std::time::Duration; + + /// A hook a caller registers has to be asked alongside the ones the hook + /// starts with, or the registration is decorative. + #[test] + fn a_registered_hook_is_asked() { + struct AlwaysCancel; + + #[async_trait::async_trait] + impl AgentHook for AlwaysCancel { + fn should_cancel(&self) -> Option { + Some(RunCancelReason::External) + } + } + + let (hook, _cancel, _usage) = StreamingRequestHook::new( + crate::streaming::RunOptions::bounded(Some(Duration::from_secs(300))), + "req_registered", + ); + assert!( + hook.should_cancel().is_none(), + "nothing has asked for cancellation yet" + ); + + let hook = hook.with_hook(Arc::new(AlwaysCancel)); + assert!(hook.should_cancel().is_some()); + } #[test] fn test_streaming_request_hook_creation() { @@ -780,7 +773,7 @@ mod tests { crate::streaming::RunOptions::bounded(Some(Duration::from_secs(60))), "test_req_1", ); - assert!(!hook.should_cancel()); + assert!(hook.should_cancel().is_none()); assert_eq!(hook.request_id, "test_req_1"); } @@ -824,10 +817,10 @@ mod tests { crate::streaming::RunOptions::bounded(Some(Duration::from_secs(60))), "test_req_2", ); - assert!(!hook.should_cancel()); + assert!(hook.should_cancel().is_none()); cancel.cancel(); - assert!(hook.should_cancel()); + assert!(hook.should_cancel().is_some()); } #[test] @@ -840,7 +833,7 @@ mod tests { // Wait for timeout std::thread::sleep(Duration::from_millis(5)); - assert!(hook.should_cancel()); + assert!(hook.should_cancel().is_some()); } #[test] @@ -1045,6 +1038,7 @@ mod tests { crate::streaming::RunOptions::bounded(Some(Duration::from_secs(60))), "req_b", None, + HashSet::new(), ); // Both hooks should report no scratchpad budget. assert!(hook_a.scratchpad_budget.is_none()); @@ -1060,6 +1054,7 @@ mod tests { crate::streaming::RunOptions::bounded(Some(Duration::from_secs(60))), "req_with_budget", Some(budget.clone()), + HashSet::new(), ); let stored = hook .scratchpad_budget