diff --git a/Cargo.lock b/Cargo.lock index 19753da9f..025d106fe 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -414,6 +414,7 @@ dependencies = [ "tower-http", "tracing", "tracing-opentelemetry", + "tracing-subscriber", "uuid", "wiremock", ] diff --git a/crates/aura-test-utils/src/mock_agent.rs b/crates/aura-test-utils/src/mock_agent.rs index 9796ced3e..c9efa2525 100644 --- a/crates/aura-test-utils/src/mock_agent.rs +++ b/crates/aura-test-utils/src/mock_agent.rs @@ -84,6 +84,8 @@ enum Script { pub struct MockAgent { on_stream_start: Option, script: Script, + run_id: aura::RunId, + stream_claim: aura::streaming::StreamClaim, } impl MockAgent { @@ -92,6 +94,8 @@ impl MockAgent { Self { on_stream_start: None, script: Script::Pending, + run_id: aura::RunId::mint(), + stream_claim: aura::streaming::StreamClaim::default(), } } @@ -103,13 +107,13 @@ impl MockAgent { Self::scripted(items.into_iter().map(Step::Item).collect()) } - /// Runs the given steps in order, then ends the stream. The script is - /// consumed by the first `stream` call; a second call - /// on the same agent yields an empty stream. + /// Runs the given steps in order, then ends the stream. pub fn scripted(steps: Vec) -> Self { Self { on_stream_start: None, script: Script::Steps(Mutex::new(Some(steps))), + run_id: aura::RunId::mint(), + stream_claim: aura::streaming::StreamClaim::default(), } } @@ -168,6 +172,10 @@ impl StreamingAgent for MockAgent { ("test", "fake") } + fn run_id(&self) -> aura::RunId { + self.run_id + } + async fn stream( &self, _query: &str, @@ -175,6 +183,11 @@ impl StreamingAgent for MockAgent { options: aura::streaming::RunOptions, request_id: &str, ) -> AgentRun { + // Held to the contract a real agent keeps, so a test that names the + // wrong run or streams twice fails the way production would. + if let Err(refused) = self.stream_claim.claim(self.run_id, request_id) { + return AgentRun::refused(refused); + } let stream = self.start(request_id).await; // 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 @@ -200,7 +213,12 @@ mod tests { async fn a_pending_agent_never_yields() { let agent = MockAgent::pending(); let mut stream = agent - .stream("q", vec![], aura::streaming::RunOptions::default(), "req_1") + .stream( + "q", + vec![], + aura::streaming::RunOptions::default(), + &agent.run_id().to_string(), + ) .await .into_events(); assert!( @@ -216,7 +234,12 @@ mod tests { 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") + .stream( + "q", + vec![], + aura::streaming::RunOptions::default(), + &agent.run_id().to_string(), + ) .await .into_events(); @@ -234,21 +257,31 @@ mod tests { #[tokio::test] async fn the_start_hook_runs_before_the_stream() { let ran = Arc::new(AtomicBool::new(false)); - let flag = Arc::clone(&ran); + let seen = Arc::new(Mutex::new(None::)); + let (flag, seen_by_hook) = (Arc::clone(&ran), Arc::clone(&seen)); let agent = MockAgent::pending().on_stream_start(move |request_id| { - let flag = Arc::clone(&flag); + let (flag, seen) = (Arc::clone(&flag), Arc::clone(&seen_by_hook)); async move { - assert_eq!(request_id, "req_1"); + *seen.lock().expect("seen lock") = Some(request_id); flag.store(true, Ordering::SeqCst); } }); let _ = agent - .stream("q", vec![], aura::streaming::RunOptions::default(), "req_1") + .stream( + "q", + vec![], + aura::streaming::RunOptions::default(), + &agent.run_id().to_string(), + ) .await .into_events(); assert!(ran.load(Ordering::SeqCst), "hook should run"); + assert_eq!( + seen.lock().expect("seen lock").as_deref(), + Some(agent.run_id().to_string().as_str()) + ); } #[tokio::test(start_paused = true)] @@ -270,35 +303,82 @@ mod tests { "q", vec![], aura::streaming::RunOptions::default(), - "req_42", + &agent.run_id().to_string(), ) .await .into_events(); let items: Vec<_> = stream.collect().await; - assert_eq!(order.lock().expect("order lock").as_slice(), ["req_42"]); + assert_eq!( + order.lock().expect("order lock").as_slice(), + [agent.run_id().to_string()] + ); assert_eq!(items.len(), 1, "effects do not yield stream items"); } + /// The mock refuses what a real agent refuses: a stream naming another + /// run, which never reaches the start hook. + #[tokio::test] + async fn a_stream_naming_another_run_is_refused_before_the_hook() { + let ran = Arc::new(AtomicBool::new(false)); + let flag = Arc::clone(&ran); + let agent = MockAgent::yielding([items::text("never")]).on_stream_start(move |_| { + let flag = Arc::clone(&flag); + async move { flag.store(true, Ordering::SeqCst) } + }); + + let items: Vec<_> = agent + .stream("q", vec![], aura::streaming::RunOptions::default(), "req_1") + .await + .into_events() + .collect() + .await; + + let [Err(error)] = items.as_slice() else { + panic!("expected one refusal, got {items:?}"); + }; + assert!(matches!( + error.downcast_ref::(), + Some(aura::streaming::StreamRefused::OtherId { .. }) + )); + assert!(!ran.load(Ordering::SeqCst), "a refused stream runs nothing"); + } + #[tokio::test(start_paused = true)] - async fn a_script_is_consumed_by_the_first_stream_call() { + async fn a_second_stream_is_refused() { let agent = MockAgent::yielding([items::text("once")]); let first: Vec<_> = agent - .stream("q", vec![], aura::streaming::RunOptions::default(), "req_1") + .stream( + "q", + vec![], + aura::streaming::RunOptions::default(), + &agent.run_id().to_string(), + ) .await .into_events() .collect() .await; let second: Vec<_> = agent - .stream("q", vec![], aura::streaming::RunOptions::default(), "req_1") + .stream( + "q", + vec![], + aura::streaming::RunOptions::default(), + &agent.run_id().to_string(), + ) .await .into_events() .collect() .await; assert_eq!(first.len(), 1); - assert!(second.is_empty()); + let [Err(error)] = second.as_slice() else { + panic!("expected one refusal, got {second:?}"); + }; + assert!(matches!( + error.downcast_ref::(), + Some(aura::streaming::StreamRefused::AlreadyStreamed { .. }) + )); } #[tokio::test(start_paused = true)] @@ -318,7 +398,12 @@ mod tests { ]); let mut stream = agent - .stream("q", vec![], aura::streaming::RunOptions::default(), "req_1") + .stream( + "q", + vec![], + aura::streaming::RunOptions::default(), + &agent.run_id().to_string(), + ) .await .into_events(); while stream.next().await.is_some() { diff --git a/crates/aura-web-server/Cargo.toml b/crates/aura-web-server/Cargo.toml index 40ecaf885..dff8a4013 100644 --- a/crates/aura-web-server/Cargo.toml +++ b/crates/aura-web-server/Cargo.toml @@ -104,6 +104,7 @@ http-body-util = "0.1" tempfile = "3" aura-config = { path = "../aura-config", features = ["test_util"] } wiremock = "0.6" +tracing-subscriber = { workspace = true } # Exhaustively explores interleavings of the task-cancel map, which a stress # test cannot reach: the window is one mutex release and reacquire wide. diff --git a/crates/aura-web-server/src/a2a/agent_executor.rs b/crates/aura-web-server/src/a2a/agent_executor.rs index 2fed3d215..55272d6a6 100644 --- a/crates/aura-web-server/src/a2a/agent_executor.rs +++ b/crates/aura-web-server/src/a2a/agent_executor.rs @@ -243,7 +243,18 @@ impl AgentExecutor for AuraAgentExecutor { let run_tools = self.app_state.run_tools(&config); let mut append_tracker: HashMap<(String, String, String), bool> = HashMap::new(); - Box::pin(async_stream::stream! { + // The run is polled inside this span, so its log lines, the agent's + // included, carry the run's id and the task it executes. + let run_id = aura::RunId::mint(); + let span = tracing::info_span!( + parent: None, + "agent.stream", + run.id = %run_id, + a2a.task_id = %ctx.task_id, + a2a.context_id = %ctx.context_id, + ); + + let execution = async_stream::stream! { let task_id = ctx.task_id.clone(); let context_id = ctx.context_id.clone(); @@ -269,7 +280,9 @@ impl AgentExecutor for AuraAgentExecutor { metadata: None, })); - let request_id = format!("a2a_{}", task_id); + // In string form, the run's id is the request id everything + // request-keyed reads: MCP cancellation, approvals, the cancel map. + let request_id = run_id.to_string(); // Registered before the agent build and history fetch, both of which // await, so a cancelTask during those has a token to cancel. Its @@ -295,7 +308,7 @@ impl AgentExecutor for AuraAgentExecutor { Some(&req_headers), session_id, None, - Some(request_id.clone()), + Some(run_id), run_tools, ) .await @@ -558,7 +571,8 @@ impl AgentExecutor for AuraAgentExecutor { metadata: None, })); } - }) + }; + in_span(span, Box::pin(execution)) } fn cancel(&self, ctx: ExecutorContext) -> BoxStream<'static, Result> { @@ -621,6 +635,18 @@ fn extract_text(parts: Vec) -> Result { Ok(strings.join("\n")) } +/// `stream`, polled inside `span`, so its log lines carry the span's fields. +fn in_span( + span: tracing::Span, + mut stream: BoxStream<'static, T>, +) -> BoxStream<'static, T> { + futures_util::stream::poll_fn(move |cx| { + let _entered = span.enter(); + stream.poll_next_unpin(cx) + }) + .boxed() +} + pub(super) fn fail_status(task_id: &str, context_id: &str, error_msg: &str) -> StreamResponse { StreamResponse::StatusUpdate(TaskStatusUpdateEvent { task_id: task_id.to_string(), @@ -1046,6 +1072,50 @@ mod tests { assert!(claim_agent(&state, &task_id, &agent) == CancelRaced::Yes); } + /// A line the run logs carries its span's fields, which is what ties a + /// run's id to the task it executes. + #[tokio::test] + async fn a_line_logged_inside_the_span_carries_its_fields() { + let written = Arc::new(std::sync::Mutex::new(Vec::::new())); + let writer = { + let written = Arc::clone(&written); + move || SharedWriter(Arc::clone(&written)) + }; + let subscriber = tracing_subscriber::fmt() + .with_writer(writer) + .with_ansi(false) + .finish(); + let _default = tracing::subscriber::set_default(subscriber); + + let span = tracing::info_span!("agent.stream", a2a.task_id = %"t_1"); + let run = futures_util::stream::once(async { + tracing::warn!("inside the run"); + 1 + }) + .boxed(); + assert_eq!(in_span(span, run).collect::>().await, vec![1]); + + let written = String::from_utf8(written.lock().unwrap().clone()).unwrap(); + let line = written + .lines() + .find(|line| line.contains("inside the run")) + .expect("the line is written"); + assert!(line.contains("agent.stream{a2a.task_id=t_1}"), "{line}"); + } + + struct SharedWriter(Arc>>); + + impl std::io::Write for SharedWriter { + fn write(&mut self, buf: &[u8]) -> std::io::Result { + self.0.lock().unwrap().extend_from_slice(buf); + Ok(buf.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } + } + /// 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] diff --git a/crates/aura-web-server/src/handlers.rs b/crates/aura-web-server/src/handlers.rs index 1eae4aef2..80c3848df 100644 --- a/crates/aura-web-server/src/handlers.rs +++ b/crates/aura-web-server/src/handlers.rs @@ -176,7 +176,7 @@ async fn build_agent_for_request( req_headers: &HashMap, additional_tools: Vec>, client_tools: Option<&[ClientToolDefinition]>, - request_id: String, + run_id: aura::RunId, session_id: String, ) -> Result, PrepareError> { let client_tool_defs = @@ -186,7 +186,7 @@ async fn build_agent_for_request( Some(req_headers), additional_tools, client_tool_defs, - Some(request_id), + Some(run_id), Some(session_id), ) .await @@ -243,11 +243,12 @@ pub async fn prepare_request( // `finish_reason: "tool_calls"` when one fires. let has_client_tools = req.tools.is_some(); - // Generate the request id up front so the agent build (single-agent or - // orchestration) shares one value with the completion stream. The HITL gate - // and approval events stamp this id; previously it was minted later in - // `build_completion_config`, after the agent was already built. - let request_id = format!("req_{}", Uuid::new_v4().simple()); + // The run's id, minted before the agent build (single-agent or + // orchestration) so the build and the completion stream share one value. + // Its string form is the request id every request-keyed registry reads: + // the HITL gate's approvals, their sweep, and MCP cancellation. + let run_id = aura::RunId::mint(); + let request_id = run_id.to_string(); // Find the matching config: single-config passthrough > explicit model > DEFAULT_AGENT // Single-config servers accept any model field value (clients like LibreChat always send one). @@ -323,7 +324,7 @@ pub async fn prepare_request( Some(req_headers_map), Some(chat_session_id.to_string()), client_tools_vec.clone(), - Some(request_id.clone()), + Some(run_id), run_tools, ) .await @@ -351,7 +352,7 @@ pub async fn prepare_request( req_headers_map, additional_tools, client_tools, - request_id.clone(), + run_id, chat_session_id.to_string(), ) .await?; @@ -768,9 +769,10 @@ async fn handle_non_streaming_completion( let (result_tx, result_rx) = oneshot::channel(); + let agent_span = tracing::info_span!(parent: None, "agent.stream", run.id = %config.request_id); let handle = tokio::spawn( execute_completion(setup, config, DeliveryMode::Collect { result_tx }) - .instrument(tracing::info_span!(parent: None, "agent.stream")), + .instrument(agent_span), ); data.active_requests.track_task(handle); @@ -813,6 +815,7 @@ async fn handle_streaming_completion( let heartbeat_interval = std::time::Duration::from_secs(15); + let agent_span = tracing::info_span!(parent: None, "agent.stream", run.id = %config.request_id); let handle = tokio::spawn( execute_completion( setup, @@ -822,7 +825,7 @@ async fn handle_streaming_completion( heartbeat_interval, }, ) - .instrument(tracing::info_span!(parent: None, "agent.stream")), + .instrument(agent_span), ); data.active_requests.track_task(handle); diff --git a/crates/aura-web-server/src/slack/runner.rs b/crates/aura-web-server/src/slack/runner.rs index abbf0d9c7..57aba8980 100644 --- a/crates/aura-web-server/src/slack/runner.rs +++ b/crates/aura-web-server/src/slack/runner.rs @@ -235,6 +235,16 @@ fn queue_answer(ingress: &Arc, inbound: Inbound, earlier: Option, inbound: Inbound, earlier: Option>) { - let request_id = format!("slack_{}_{}", inbound.channel, inbound.ts); + async fn answer( + &self, + inbound: Inbound, + prefetched: Option>, + run_id: aura::RunId, + ) { + // In string form, the run's id is the request id everything + // request-keyed reads. + let request_id = run_id.to_string(); let earlier = match prefetched { Some(earlier) => earlier, None => { @@ -291,7 +306,7 @@ impl SlackIngress { warn!(request_id, error = %e, "could not react to slack message"); } - let reply = match self.run_agent(&inbound, &earlier, &request_id).await { + let reply = match self.run_agent(&inbound, &earlier, run_id).await { Ok(text) if text.trim().is_empty() => EMPTY_REPLY.to_owned(), Ok(text) => text, // The server is going down; a reply would race the shutdown and @@ -344,8 +359,10 @@ impl SlackIngress { &self, inbound: &Inbound, earlier: &[SlackMessage], - request_id: &str, + run_id: aura::RunId, ) -> Result { + let request_id = run_id.to_string(); + let request_id = request_id.as_str(); let history = thread_history(earlier, &self.identity, &inbound.ts); let session_id = match inbound.reply_thread() { Some(thread) => format!("slack:{}:{thread}", inbound.channel), @@ -359,13 +376,7 @@ impl SlackIngress { ); let agent = RigBuilder::new(config, self.state.pending_approvals.clone()) .with_hitl_hmac(self.state.hitl_webhook_hmac.clone()) - .build_streaming_agent_with_tools( - None, - Some(session_id), - None, - Some(request_id.to_owned()), - tools, - ) + .build_streaming_agent_with_tools(None, Some(session_id), None, Some(run_id), tools) .await .map_err(|e| RunError::Build(e.to_string()))?; diff --git a/crates/aura-web-server/src/streaming/handlers.rs b/crates/aura-web-server/src/streaming/handlers.rs index 28d8c352b..027ae7502 100644 --- a/crates/aura-web-server/src/streaming/handlers.rs +++ b/crates/aura-web-server/src/streaming/handlers.rs @@ -2912,7 +2912,7 @@ mod tests { where F: FnOnce(&Senders) -> Vec, { - let (senders, callbacks) = channels(); + let (senders, mut callbacks) = channels(); let steps = build(&senders); let config = StreamConfig::new(emit_custom_events, false, ToolResultMode::Aura, 0); let ctx = TurnContext::new( @@ -2923,13 +2923,13 @@ mod tests { session_id, ); - let stream = MockAgent::scripted(steps) - .stream( - "q", - vec![], - aura::streaming::RunOptions::default(), - "req_tool_events", - ) + // The handler passes the run's id to both the stream and the + // callbacks, so the test does too. + let agent = MockAgent::scripted(steps); + let run_id = agent.run_id().to_string(); + callbacks.request_id = run_id.clone(); + let stream = agent + .stream("q", vec![], aura::streaming::RunOptions::default(), &run_id) .await .into_events(); diff --git a/crates/aura/src/builder.rs b/crates/aura/src/builder.rs index 92868e045..4fe91872c 100644 --- a/crates/aura/src/builder.rs +++ b/crates/aura/src/builder.rs @@ -17,6 +17,7 @@ use crate::{ vector_dynamic::DynamicVectorSearchTool, vector_store::VectorStoreManager, }; +use aura_events::RunId; use aura_events::agent::AgentEvent; use futures::StreamExt; use rig::client::CompletionClient; @@ -185,6 +186,8 @@ pub struct Agent { lease: Arc, /// The run's events, for its observer. events: Mutex>>, + /// The run's one stream through [`StreamingAgent::stream`]. + stream_claim: crate::streaming::StreamClaim, } /// A run reads its prepared agent's fields and calls its methods directly, so @@ -210,7 +213,7 @@ impl std::fmt::Debug for PreparedAgent { impl std::fmt::Debug for Agent { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("Agent") - .field("request_id", &self.request_id()) + .field("run_id", &self.run_id()) .field("prepared", &self.prepared) .finish() } @@ -967,8 +970,8 @@ impl PreparedAgent { &self.forwarded_headers } - /// Begin a run of this agent for the request `request_id`, whose headers - /// are `req_headers`. + /// Begin the run `run_id` of this agent, for a request whose headers are + /// `req_headers`. /// /// The run is a [`RunContext`] of its own, on a token of its own, with a /// fresh scratchpad budget and turn-limit counters starting from what @@ -982,7 +985,7 @@ impl PreparedAgent { /// calls as someone else. Prepare an agent for it instead. pub fn begin_run( self: &Arc, - request_id: impl Into, + run_id: RunId, req_headers: Option<&HashMap>, skill_recorder: Option>, ) -> Result { @@ -992,7 +995,7 @@ impl PreparedAgent { }); } let (run, events) = RunContext::channel_for_agent( - request_id.into(), + run_id, CancellationToken::new(), self.scratchpad_budget.as_ref().map(ContextBudget::fresh), self.turn_nudge.as_ref().map(|seed| seed.fresh()), @@ -1039,7 +1042,7 @@ impl PreparedAgent { let mut active = self.active.lock().unwrap_or_else(PoisonError::into_inner); if let Some(alive) = active.upgrade() { return Err(RunInProgress { - active: alive.run().id().to_string(), + active: alive.run().id(), }); } let lease = Arc::new(RunLease::new(Arc::clone(&run))); @@ -1057,6 +1060,7 @@ impl PreparedAgent { prepared: Arc::clone(self), lease, events: Mutex::new(events), + stream_claim: crate::streaming::StreamClaim::default(), }) } @@ -1536,7 +1540,7 @@ impl Agent { /// Prepare an agent from configuration and begin its single run. /// /// The run is for the request `config` was resolved for: its id is - /// `config.request_id` (empty when unset), its headers are the ones + /// `config.run_id` (a fresh one when unset), its headers are the ones /// `config.forwarded_headers` recorded, and `config.skill_recorder` /// records its skill-tool invocations. See [`PreparedAgent::prepare`] for /// `additional_tools` and `client_tools`; callers that want to reuse the @@ -1551,7 +1555,7 @@ impl Agent { Arc::new(PreparedAgent::prepare(config, additional_tools, client_tools).await?); let req_headers = config.forwarded_headers.as_request(); Ok(prepared.begin_run( - config.request_id.clone().unwrap_or_default(), + config.run_id.unwrap_or_else(RunId::mint), Some(&req_headers), config.skill_recorder.clone(), )?) @@ -1562,11 +1566,6 @@ impl Agent { self.lease.run() } - /// Request id of this run. - pub fn request_id(&self) -> &str { - self.lease.run().id() - } - /// This run's scratchpad budget, when scratchpad is wired up. pub fn scratchpad_budget(&self) -> Option<&ContextBudget> { self.lease.run().scratchpad_budget() @@ -1984,6 +1983,10 @@ impl StreamingAgent for Agent { self.prepared.get_provider_info() } + fn run_id(&self) -> RunId { + self.lease.run().id() + } + async fn stream( &self, query: &str, @@ -1992,14 +1995,12 @@ impl StreamingAgent for Agent { request_id: &str, ) -> crate::streaming::AgentRun { // The run began at `begin_run`, so its context is what this streams - // under; the id is the run's. + // under. Only this trait method claims the run's stream: the + // orchestrator re-streams a coordinator through the inherent methods on + // a transient retry. let run = Arc::clone(self.lease.run()); - if request_id != run.id().as_ref() { - tracing::debug!( - run_id = %run.id(), - request_id, - "streaming under the run's id rather than the one passed", - ); + if let Err(refused) = self.stream_claim.claim(run.id(), request_id) { + return crate::streaming::AgentRun::refused(refused); } let run_id = run.id().to_string(); @@ -2613,9 +2614,10 @@ mod tests { /// one `add_mcp_tool` stamped, carried through `pre_call`. #[tokio::test] async fn a_namespace_scoped_pattern_gates_a_tool_from_that_server() { - let request_id = "req_ns_gating_match"; let server = RecordingMcpServer::start().await; - let (run, mut rx) = crate::run_context::RunContext::channel(request_id); + let (run, mut rx) = crate::run_context::RunContext::channel( + crate::run_context::named_run_id("req_ns_gating_match"), + ); let agent = compose_gated_agent(&server, "github:*", "github", "list_repos", run).await; // The parked approval expires unanswered; the call's own outcome is @@ -2645,9 +2647,10 @@ mod tests { /// namespace were ignored and the bare name alone matched. #[tokio::test] async fn a_namespace_scoped_pattern_ignores_a_tool_from_another_server() { - let request_id = "req_ns_gating_miss"; let server = RecordingMcpServer::start().await; - let (run, mut rx) = crate::run_context::RunContext::channel(request_id); + let (run, mut rx) = crate::run_context::RunContext::channel( + crate::run_context::named_run_id("req_ns_gating_miss"), + ); let agent = compose_gated_agent(&server, "github:*", "k8s", "list_repos", run).await; agent @@ -3033,6 +3036,169 @@ mod tests { mod prepared_runs { use super::*; use crate::orchestration::{ScriptedCompletionModel, ScriptedTurn}; + use crate::streaming::StreamRefused; + + /// A prepared agent over `turns`, and the log of requests its model + /// served. + fn prepared_scripted( + turns: Vec, + ) -> ( + Arc, + Arc>>, + ) { + let model = ScriptedCompletionModel::new(turns); + let requests = model.requests(); + let prepared = prepared_over( + ProviderAgent::Scripted(rig::agent::AgentBuilder::new(model).build()), + ForwardedHeaders::default(), + None, + ); + (prepared, requests) + } + + /// Streams `agent` through the trait, as a server path does, and + /// collects what it yields. + async fn streamed_as( + agent: &Agent, + request_id: &str, + ) -> Vec> { + StreamingAgent::stream( + agent, + "q", + Vec::new(), + crate::streaming::RunOptions::default(), + request_id, + ) + .await + .into_events() + .collect() + .await + } + + fn refusal(items: &[Result]) -> Option<&StreamRefused> { + match items { + [Err(error)] => error.downcast_ref::(), + _ => None, + } + } + + #[tokio::test] + async fn a_stream_naming_another_run_is_refused_and_runs_nothing() { + let (prepared, requests) = prepared_scripted(vec![ScriptedTurn::text("done")]); + let run_id = crate::run_context::named_run_id("req_a"); + let agent = prepared.begin_run(run_id, None, None).unwrap(); + + let items = streamed_as(&agent, "req_123").await; + assert!( + matches!( + refusal(&items), + Some(StreamRefused::OtherId { run, requested }) + if *run == run_id && requested == "req_123" + ), + "got: {items:?}", + ); + assert!( + requests.lock().unwrap().is_empty(), + "a refused stream never reaches the model", + ); + + let items = streamed_as(&agent, &run_id.to_string()).await; + assert!( + refusal(&items).is_none(), + "the run still streams under its own id: {items:?}", + ); + assert_eq!(requests.lock().unwrap().len(), 1); + } + + #[tokio::test] + async fn a_second_stream_is_refused_and_leaves_the_first_running() { + let (prepared, requests) = prepared_scripted(vec![ + ScriptedTurn::text("first"), + ScriptedTurn::text("second"), + ]); + let run_id = crate::run_context::named_run_id("req_a"); + let agent = prepared.begin_run(run_id, None, None).unwrap(); + let id = run_id.to_string(); + + let first = StreamingAgent::stream( + &agent, + "q", + Vec::new(), + crate::streaming::RunOptions::default(), + &id, + ) + .await; + let second = streamed_as(&agent, &id).await; + + assert!( + matches!(refusal(&second), Some(StreamRefused::AlreadyStreamed { run }) if *run == run_id), + "got: {second:?}", + ); + assert!( + !agent.run().cancel_token().is_cancelled(), + "the refusal leaves the run's token alone", + ); + + let first: Vec<_> = first.into_events().collect().await; + assert!( + !first.is_empty() && first.iter().all(Result::is_ok), + "the first stream runs to the end: {first:?}", + ); + assert_eq!( + requests.lock().unwrap().len(), + 1, + "only the first stream reaches the model", + ); + } + + /// The orchestrator re-streams one coordinator `Agent` through + /// `stream_chat_with_depth` on a transient retry. The claim guards only + /// the trait method, so that path still streams twice. + #[tokio::test] + async fn the_coordinator_retry_path_still_streams_twice() { + let (prepared, requests) = + prepared_scripted(vec![ScriptedTurn::text("one"), ScriptedTurn::text("two")]); + let agent = prepared + .begin_run(crate::run_context::named_run_id("req_a"), None, None) + .unwrap(); + + for attempt in 0..2 { + let items: Vec<_> = agent + .stream_chat_with_depth("q", Vec::new(), agent.max_depth) + .await + .collect() + .await; + assert!( + !items.is_empty() && items.iter().all(Result::is_ok), + "attempt {attempt}: {items:?}", + ); + } + assert_eq!(requests.lock().unwrap().len(), 2); + } + + /// A caller holding only the trait object, as every builder returns, + /// reads the run's id from it and streams under that. + #[tokio::test] + async fn a_caller_holding_only_the_trait_streams_under_its_run_id() { + let (prepared, requests) = prepared_scripted(vec![ScriptedTurn::text("done")]); + let agent: Arc = + Arc::new(prepared.begin_run(RunId::mint(), None, None).unwrap()); + + let items: Vec<_> = agent + .stream( + "q", + Vec::new(), + crate::streaming::RunOptions::default(), + &agent.run_id().to_string(), + ) + .await + .into_events() + .collect() + .await; + + assert!(refusal(&items).is_none(), "got: {items:?}"); + assert_eq!(requests.lock().unwrap().len(), 1); + } /// A prepared agent over a scripted model, seeded with a scratchpad /// budget and a turn nudge that fires on the second turn, so a run @@ -3132,7 +3298,9 @@ mod tests { prepared: &Arc, request_id: &str, ) -> (Agent, tokio::task::JoinHandle<()>) { - let agent = prepared.begin_run(request_id, None, None).unwrap(); + let agent = prepared + .begin_run(crate::run_context::named_run_id(request_id), None, None) + .unwrap(); let mut events = agent.events.lock().unwrap().take().unwrap(); let call = tokio::spawn({ let prepared = Arc::clone(prepared); @@ -3158,7 +3326,7 @@ mod tests { }; let (first, call) = park(&prepared, "req_a").await; - registry.cancel_request_local("req_a"); + registry.cancel_request_local(&crate::run_context::named_run_id("req_a").to_string()); assert!( released(call).await, "ending req_a releases the approval it parked" @@ -3166,13 +3334,13 @@ mod tests { drop(first); let (_second, call) = park(&prepared, "req_b").await; - registry.cancel_request_local("req_a"); + registry.cancel_request_local(&crate::run_context::named_run_id("req_a").to_string()); tokio::time::sleep(Duration::from_millis(100)).await; assert!( !call.is_finished(), "ending req_a leaves req_b's approval parked" ); - registry.cancel_request_local("req_b"); + registry.cancel_request_local(&crate::run_context::named_run_id("req_b").to_string()); assert!( released(call).await, "ending req_b releases the approval it parked" @@ -3187,22 +3355,22 @@ mod tests { async fn a_prepared_agent_serves_one_run_at_a_time() { let prepared = prepared(); let first = prepared - .begin_run("req_a", None, None) + .begin_run(crate::run_context::named_run_id("req_a"), None, None) .expect("a fresh agent has no run"); let refused = prepared - .begin_run("req_b", None, None) + .begin_run(crate::run_context::named_run_id("req_b"), None, None) .expect_err("the slot is taken"); assert!( - matches!(refused, BeginRunError::RunInProgress(RunInProgress { ref active }) if active == "req_a"), + matches!(refused, BeginRunError::RunInProgress(RunInProgress { ref active }) if *active == crate::run_context::named_run_id("req_a")), "got: {refused:?}", ); drop(first); let second = prepared - .begin_run("req_b", None, None) + .begin_run(crate::run_context::named_run_id("req_b"), None, None) .expect("the slot is free again"); - assert_eq!(second.request_id(), "req_b"); + assert_eq!(second.run_id(), crate::run_context::named_run_id("req_b")); } /// Each run starts from the prepared seeds, and the slot the tools @@ -3212,14 +3380,12 @@ mod tests { let prepared = prepared(); let seed = prepared.scratchpad_budget.as_ref().unwrap(); - let first = prepared.begin_run("req_a", None, None).unwrap(); + let first = prepared + .begin_run(crate::run_context::named_run_id("req_a"), None, None) + .unwrap(); assert_eq!( - prepared - .run - .get() - .map(|run| run.id().to_string()) - .as_deref(), - Some("req_a") + prepared.run.get().map(|run| run.id()), + Some(crate::run_context::named_run_id("req_a")) ); prepared .run @@ -3241,14 +3407,12 @@ mod tests { drop(first); - let second = prepared.begin_run("req_b", None, None).unwrap(); + let second = prepared + .begin_run(crate::run_context::named_run_id("req_b"), None, None) + .unwrap(); assert_eq!( - prepared - .run - .get() - .map(|run| run.id().to_string()) - .as_deref(), - Some("req_b") + prepared.run.get().map(|run| run.id()), + Some(crate::run_context::named_run_id("req_b")) ); assert_eq!( second.scratchpad_budget().unwrap().scratchpad_usage().0, @@ -3274,21 +3438,34 @@ mod tests { )); let refused = prepared - .begin_run("req_bob", Some(&token("bob")), None) + .begin_run( + crate::run_context::named_run_id("req_bob"), + Some(&token("bob")), + None, + ) .expect_err("bob's token is not alice's"); assert!( matches!(refused, BeginRunError::ForwardedHeaderDiffers { ref header } if header == "x-user-token"), "got: {refused:?}", ); assert!( - prepared.begin_run("req_none", None, None).is_err(), + prepared + .begin_run(crate::run_context::named_run_id("req_none"), None, None) + .is_err(), "a request carrying no token is not alice's either", ); let served = prepared - .begin_run("req_alice", Some(&token("alice")), None) + .begin_run( + crate::run_context::named_run_id("req_alice"), + Some(&token("alice")), + None, + ) .expect("the same credentials are served"); - assert_eq!(served.request_id(), "req_alice"); + assert_eq!( + served.run_id(), + crate::run_context::named_run_id("req_alice") + ); } /// A stream still driving tools after its `Agent` is dropped keeps @@ -3296,24 +3473,26 @@ mod tests { #[tokio::test] async fn a_live_stream_keeps_the_run_bound_after_the_agent_drops() { let prepared = prepared(); - let agent = prepared.begin_run("req_a", None, None).unwrap(); + let agent = prepared + .begin_run(crate::run_context::named_run_id("req_a"), None, None) + .unwrap(); let stream = agent.stream_prompt("hello").await; drop(agent); assert_eq!( - prepared - .run - .get() - .map(|run| run.id().to_string()) - .as_deref(), - Some("req_a"), + prepared.run.get().map(|run| run.id()), + Some(crate::run_context::named_run_id("req_a")), "the stream holds the run", ); - assert!(prepared.begin_run("req_b", None, None).is_err()); + assert!( + prepared + .begin_run(crate::run_context::named_run_id("req_b"), None, None) + .is_err() + ); drop(stream); prepared - .begin_run("req_b", None, None) + .begin_run(crate::run_context::named_run_id("req_b"), None, None) .expect("the slot frees with the stream"); } } diff --git a/crates/aura/src/config.rs b/crates/aura/src/config.rs index 7e19c726d..e51d49400 100644 --- a/crates/aura/src/config.rs +++ b/crates/aura/src/config.rs @@ -39,7 +39,7 @@ pub enum WorkerSkills { Override(Vec), } -pub use aura_events::SessionId; +pub use aura_events::{RunId, SessionId}; /// Runtime build context for constructing agents. /// @@ -111,8 +111,8 @@ pub struct AgentRuntimeConfig { /// `None` disables approval gating. pub hitl: Option, - /// Request id (`req_…`) of the request this build serves. - pub request_id: Option, + /// The run this build serves. + pub run_id: Option, /// The request headers this build forwards; see [`ForwardedHeaders`]. pub forwarded_headers: ForwardedHeaders, @@ -153,7 +153,7 @@ impl Clone for AgentRuntimeConfig { scratchpad_tools_config: self.scratchpad_tools_config.clone(), orchestration_submit_result: self.orchestration_submit_result.clone(), hitl: self.hitl.clone(), - request_id: self.request_id.clone(), + run_id: self.run_id, forwarded_headers: self.forwarded_headers.clone(), instance_id: self.instance_id.clone(), hitl_request_approval_tool: self.hitl_request_approval_tool.clone(), @@ -198,7 +198,7 @@ impl std::fmt::Debug for AgentRuntimeConfig { .map(|_| ""), ) .field("hitl", &self.hitl.as_ref().map(|_| "")) - .field("request_id", &self.request_id) + .field("run_id", &self.run_id) .field("forwarded_headers", &self.forwarded_headers) .field("instance_id", &self.instance_id) .field( diff --git a/crates/aura/src/hitl/gate.rs b/crates/aura/src/hitl/gate.rs index c85493417..2dc35cc92 100644 --- a/crates/aura/src/hitl/gate.rs +++ b/crates/aura/src/hitl/gate.rs @@ -977,8 +977,10 @@ mod tests { )); let cancel = tokio_util::sync::CancellationToken::new(); - let (run, mut events) = - crate::run_context::RunContext::channel_on("req_run_cancel", cancel.clone()); + let (run, mut events) = crate::run_context::RunContext::channel_on( + crate::run_context::named_run_id("req_run_cancel"), + cancel.clone(), + ); gate.bind_run(run); let gated = Arc::clone(&gate); @@ -1215,7 +1217,8 @@ mod tests { "test-agent".to_string(), "test-instance-id".to_string(), ); - let (run, events) = crate::run_context::RunContext::channel(request_id); + let (run, events) = + crate::run_context::RunContext::channel(request_id.parse().expect("a run id")); gate.bind_run(run); ( WrappedTool::new(inner, Arc::new(gate) as Arc), @@ -1241,7 +1244,7 @@ mod tests { } fn unique_request_id() -> String { - format!("req_span_{}", uuid::Uuid::new_v4().simple()) + aura_events::RunId::mint().to_string() } /// The correlation the whole feature exists for: the id the approver diff --git a/crates/aura/src/hitl/route.rs b/crates/aura/src/hitl/route.rs index 56c7ecc7f..e30815c90 100644 --- a/crates/aura/src/hitl/route.rs +++ b/crates/aura/src/hitl/route.rs @@ -1221,8 +1221,9 @@ mod tests { #[tokio::test] async fn conversational_resolve_at_requested_event_succeeds() { - let request_id = format!("req_test_{}", uuid::Uuid::new_v4().simple()); - let (run, mut rx) = crate::run_context::RunContext::channel(request_id.as_str()); + let run_id = aura_events::RunId::mint(); + let request_id = run_id.to_string(); + let (run, mut rx) = crate::run_context::RunContext::channel(run_id); let (registry, route) = conv_route(Duration::from_secs(60)); let request = single_request( @@ -2219,8 +2220,9 @@ mod tests { #[tokio::test] async fn webhook_route_emits_requested_and_completed_on_channel_error() { - let request_id = format!("req_test_{}", uuid::Uuid::new_v4().simple()); - let (run, mut rx) = crate::run_context::RunContext::channel(request_id.as_str()); + let run_id = aura_events::RunId::mint(); + let request_id = run_id.to_string(); + let (run, mut rx) = crate::run_context::RunContext::channel(run_id); let route = super::DecisionRoute::Webhook { client: super::WebhookClient::new( super::build_webhook_client(), diff --git a/crates/aura/src/hitl/tool.rs b/crates/aura/src/hitl/tool.rs index 628e8346b..85f8d8e24 100644 --- a/crates/aura/src/hitl/tool.rs +++ b/crates/aura/src/hitl/tool.rs @@ -334,7 +334,7 @@ mod tests { route: &Arc, args: RequestApprovalArgs, ) -> ApprovalItem { - let request_id = format!("req_w2_{}", uuid::Uuid::new_v4().simple()); + let run_id = aura_events::RunId::mint(); let tool = RequestApprovalTool::new( route.clone(), AgentScope::Single { session_id: None }, @@ -343,7 +343,7 @@ mod tests { ); // The scope goes inside the spawn, because task-locals do not cross one. - let (run, mut rx) = crate::run_context::RunContext::channel(request_id.as_str()); + let (run, mut rx) = crate::run_context::RunContext::channel(run_id); let call_handle: tokio::task::JoinHandle> = tokio::spawn(crate::run_context::with_run(run, async move { tool.call(args).await @@ -499,8 +499,8 @@ mod tests { timeout: std::time::Duration::from_secs(60), }); - let request_id = format!("req_tool_span_{}", uuid::Uuid::new_v4().simple()); - let (run, mut events) = crate::run_context::RunContext::channel(request_id.as_str()); + let run_id = aura_events::RunId::mint(); + let (run, mut events) = crate::run_context::RunContext::channel(run_id); let tool = RequestApprovalTool::new( route, AgentScope::Single { session_id: None }, diff --git a/crates/aura/src/mcp/client.rs b/crates/aura/src/mcp/client.rs index 9b1c6f6d7..edad45165 100644 --- a/crates/aura/src/mcp/client.rs +++ b/crates/aura/src/mcp/client.rs @@ -29,9 +29,9 @@ use crate::approver_headers::ApproverHeaders; use crate::mcp::progress::ProgressEnabledHandler; use crate::mcp::response::extract_tool_result; use crate::mcp::types::ToolNamespace; -use aura_events::AgentContext; use aura_events::ToolName; use aura_events::agent::{AgentEvent, AgentEventPayload}; +use aura_events::{AgentContext, RunId}; /// Custom HTTP client that captures the underlying HTTP status when a request /// fails. @@ -574,7 +574,7 @@ impl McpClient { /// 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> { + pub async fn run_id(&self) -> Option { match crate::run_context::current_run_id() { Some(id) => Some(id), None => self @@ -582,7 +582,7 @@ impl McpClient { .read() .await .as_ref() - .map(|call| Arc::clone(call.run.id())), + .map(|call| call.run.id()), } } @@ -597,7 +597,7 @@ impl McpClient { .read() .await .as_ref() - .filter(|call| call.run.id().as_ref() == request_id) + .filter(|call| call.run.has_id(request_id)) .cloned() } @@ -659,7 +659,12 @@ impl McpClient { tool_name, http_request_id ); return self - .call_tool_tracked(tool_name, arguments, &http_request_id, approver_overrides) + .call_tool_tracked( + tool_name, + arguments, + &http_request_id.to_string(), + approver_overrides, + ) .await; } @@ -732,7 +737,7 @@ impl McpClient { // 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) + self.own_progress_token(progress_token.clone(), &run_id.to_string()) .await; } // Held from here so a run cancelled mid-await, which drops this future @@ -817,7 +822,7 @@ impl McpClient { .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) + self.own_progress_token(handle.progress_token.clone(), &run_id.to_string()) .await; } let _token_guard = ProgressTokenGuard { @@ -1062,7 +1067,7 @@ impl McpClient { // Drop this call's token ownership so straggler notifications stop // routing, then clear the binding. - owners_of(&self.token_owners).retain(|_, call| call.run.id().as_ref() != http_request_id); + owners_of(&self.token_owners).retain(|_, call| !call.run.has_id(http_request_id)); self.clear_current_call().await; // Forcefully close connection - server is ignoring cancellation anyway @@ -1087,9 +1092,9 @@ pub(crate) mod tests { use super::*; use crate::approver_headers::tests::captured_overrides; - fn owner(request_id: &str) -> CallContext { + fn owner(name: &str) -> CallContext { CallContext { - run: crate::run_context::RunContext::detached(request_id), + run: crate::run_context::RunContext::detached(crate::run_context::named_run_id(name)), agent: AgentContext::single_agent(), } } @@ -1627,30 +1632,32 @@ pub(crate) mod tests { async fn a_call_answers_only_for_the_request_that_named_it() { let (_server, client) = client_and_server(&requester_headers()).await; let worker = AgentContext::worker("log_worker", None, "coordinator"); + let req_1 = crate::run_context::named_run_id("req-1").to_string(); + let req_2 = crate::run_context::named_run_id("req-2").to_string(); assert!( - client.call_for("req-1").await.is_none(), + client.call_for(&req_1).await.is_none(), "an unbound client has no call to attribute work to" ); client .bind_call( - crate::run_context::RunContext::detached("req-1"), + crate::run_context::RunContext::detached(crate::run_context::named_run_id("req-1")), worker.clone(), ) .await; assert_eq!( - client.call_for("req-1").await.map(|call| call.agent), + client.call_for(&req_1).await.map(|call| call.agent), Some(worker) ); assert!( - client.call_for("req-2").await.is_none(), + client.call_for(&req_2).await.is_none(), "another request's id must not pick up this call" ); client.clear_current_call().await; assert!( - client.call_for("req-1").await.is_none(), + client.call_for(&req_1).await.is_none(), "clearing the call drops it with the request id" ); } @@ -1667,11 +1674,16 @@ pub(crate) mod tests { client .bind_call( - crate::run_context::RunContext::detached("req_bound"), + crate::run_context::RunContext::detached(crate::run_context::named_run_id( + "req_bound", + )), AgentContext::single_agent(), ) .await; - assert_eq!(client.run_id().await.as_deref(), Some("req_bound")); + assert_eq!( + client.run_id().await, + Some(crate::run_context::named_run_id("req_bound")) + ); } /// A scope still wins, so a call made inside one is attributed to that run @@ -1681,18 +1693,22 @@ pub(crate) mod tests { let (_server, client) = client_and_server(&requester_headers()).await; client .bind_call( - crate::run_context::RunContext::detached("req_bound"), + crate::run_context::RunContext::detached(crate::run_context::named_run_id( + "req_bound", + )), AgentContext::single_agent(), ) .await; let seen = crate::run_context::with_run( - crate::run_context::RunContext::detached("req_scoped"), + crate::run_context::RunContext::detached(crate::run_context::named_run_id( + "req_scoped", + )), async { client.run_id().await }, ) .await; - assert_eq!(seen.as_deref(), Some("req_scoped")); + assert_eq!(seen, Some(crate::run_context::named_run_id("req_scoped"))); } /// Being inside a run selects the tracked branch, so this is the same entry @@ -1712,7 +1728,9 @@ pub(crate) mod tests { .expect("the untracked call succeeds"); crate::run_context::with_run( - crate::run_context::RunContext::detached("http-req-1"), + crate::run_context::RunContext::detached(crate::run_context::named_run_id( + "http-req-1", + )), async { client .call_tool( diff --git a/crates/aura/src/mcp/progress.rs b/crates/aura/src/mcp/progress.rs index 888f937e5..6bc150e63 100644 --- a/crates/aura/src/mcp/progress.rs +++ b/crates/aura/src/mcp/progress.rs @@ -172,7 +172,7 @@ impl ClientHandler for ProgressEnabledHandler { let call = self.owner_of(¶ms.progress_token).await; if let Some(CallContext { run, agent }) = call { - let req_id = run.id().as_ref(); + let req_id = run.id(); let routed = run .emit(AgentEvent::new( @@ -231,9 +231,9 @@ mod tests { ProgressToken(NumberOrString::Number(n)) } - fn call(request_id: &str) -> CallContext { + fn call(name: &str) -> CallContext { CallContext { - run: crate::run_context::RunContext::detached(request_id), + run: crate::run_context::RunContext::detached(crate::run_context::named_run_id(name)), agent: aura_events::AgentContext::single_agent(), } } @@ -313,22 +313,16 @@ mod tests { // Nothing owns the token yet: the send has returned but the claim has not // landed. assert_eq!( - handler - .owner_of(&token(1)) - .await - .map(|call| call.run.id().to_string()), - Some("run_a".to_string()), + handler.owner_of(&token(1)).await.map(|call| call.run.id()), + Some(crate::run_context::named_run_id("run_a")), "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.run.id().to_string()), - Some("run_a_tool".to_string()), + handler.owner_of(&token(1)).await.map(|call| call.run.id()), + Some(crate::run_context::named_run_id("run_a_tool")), "a claim is more precise than the binding" ); } @@ -340,18 +334,12 @@ mod tests { let handler = handler_owning(&[(1, "run_a"), (2, "run_b")]); assert_eq!( - handler - .owner_of(&token(1)) - .await - .map(|c| c.run.id().to_string()), - Some("run_a".to_string()) + handler.owner_of(&token(1)).await.map(|c| c.run.id()), + Some(crate::run_context::named_run_id("run_a")) ); assert_eq!( - handler - .owner_of(&token(2)) - .await - .map(|c| c.run.id().to_string()), - Some("run_b".to_string()) + handler.owner_of(&token(2)).await.map(|c| c.run.id()), + Some(crate::run_context::named_run_id("run_b")) ); } @@ -366,20 +354,14 @@ mod tests { owners.lock().unwrap().insert(token(7), call("run_a_tool")); assert_eq!( - handler - .owner_of(&token(7)) - .await - .map(|c| c.run.id().to_string()), - Some("run_a_tool".to_string()) + handler.owner_of(&token(7)).await.map(|c| c.run.id()), + Some(crate::run_context::named_run_id("run_a_tool")) ); owners.lock().unwrap().remove(&token(7)); assert_eq!( - handler - .owner_of(&token(7)) - .await - .map(|c| c.run.id().to_string()), - Some("run_a".to_string()), + handler.owner_of(&token(7)).await.map(|c| c.run.id()), + Some(crate::run_context::named_run_id("run_a")), "the binding outlives the tokens of the calls it serves" ); diff --git a/crates/aura/src/orchestration/factory.rs b/crates/aura/src/orchestration/factory.rs index 9043b3735..9b0b614ec 100644 --- a/crates/aura/src/orchestration/factory.rs +++ b/crates/aura/src/orchestration/factory.rs @@ -11,7 +11,7 @@ use async_trait::async_trait; use futures::stream::{self, BoxStream}; use tokio_util::sync::CancellationToken; -use crate::config::AgentRuntimeConfig; +use crate::config::{AgentRuntimeConfig, RunId}; use crate::provider_agent::{StreamError, StreamItem}; use crate::streaming::StreamingAgent; @@ -26,6 +26,9 @@ use super::orchestrator::{ pub struct OrchestratorFactory { agent_config: AgentRuntimeConfig, run_tools: crate::builder::RunToolFactory, + /// The one run this factory streams. + run_id: RunId, + stream_claim: crate::streaming::StreamClaim, } /// A run's cancellation, and the signal that its task ended. @@ -64,10 +67,15 @@ fn final_response( } impl OrchestratorFactory { - pub fn new(agent_config: AgentRuntimeConfig) -> Self { + /// A factory for the run `agent_config.run_id`, or for a fresh run when + /// the config names none. The orchestration it spawns sees the same id. + pub fn new(mut agent_config: AgentRuntimeConfig) -> Self { + let run_id = *agent_config.run_id.get_or_insert_with(RunId::mint); Self { agent_config, run_tools: crate::builder::no_run_tools(), + run_id, + stream_claim: crate::streaming::StreamClaim::default(), } } @@ -100,7 +108,7 @@ impl OrchestratorFactory { let (event_tx, event_rx) = tokio::sync::mpsc::channel::>(100); - let run_id = std::sync::Arc::clone(run.id()); + let run_id = run.id().to_string(); let RunTokens { cancel, finished } = tokens; let cancel_token_clone = cancel.clone(); // Marks the run finished on every exit path, which is what lets the @@ -174,7 +182,7 @@ impl OrchestratorFactory { tracing::info!("Orchestration cancelled"); if let Some(ref mcp_manager) = orchestrator.mcp_manager { let cancelled = mcp_manager - .cancel_and_close_all(run_id.as_ref(), "Client disconnected or timeout") + .cancel_and_close_all(&run_id, "Client disconnected or timeout") .await; if cancelled > 0 { tracing::info!("Cancelled {} MCP request(s) during orchestration shutdown", cancelled); @@ -201,6 +209,10 @@ impl StreamingAgent for OrchestratorFactory { self.agent_config.llm.model_info() } + fn run_id(&self) -> RunId { + self.run_id + } + /// The coordinator's window: it holds the persistent conversation, so it /// is the context a client measures the session against. fn context_window(&self) -> Option { @@ -218,6 +230,10 @@ impl StreamingAgent for OrchestratorFactory { options: crate::streaming::RunOptions, request_id: &str, ) -> crate::streaming::AgentRun { + let run_id = self.run_id; + if let Err(refused) = self.stream_claim.claim(run_id, request_id) { + return crate::streaming::AgentRun::refused(refused); + } let (timeout, cancel) = options.into_parts(); let cancel_token = cancel.unwrap_or_default(); @@ -229,7 +245,7 @@ impl StreamingAgent for OrchestratorFactory { timeout, cancel_token.clone(), finished.clone(), - request_id.to_string(), + run_id.to_string(), ); finished }); @@ -239,7 +255,7 @@ impl StreamingAgent for OrchestratorFactory { // all orchestration LLM turns. let usage_state = crate::UsageState::new(); let (run, run_events) = - crate::run_context::RunContext::channel_on(request_id, cancel_token.clone()); + crate::run_context::RunContext::channel_on(run_id, cancel_token.clone()); let stream = self.spawn_orchestration_stream( query.to_string(), chat_history, @@ -272,6 +288,54 @@ impl StreamingAgent for OrchestratorFactory { #[cfg(test)] mod tests { use super::*; + use futures::StreamExt; + + /// A factory's run id is fixed when it is built, and the orchestration it + /// spawns carries the same one. + #[test] + fn a_factory_streams_the_run_it_was_built_for() { + let named = crate::run_context::named_run_id("req_a"); + let config = AgentRuntimeConfig { + run_id: Some(named), + ..AgentRuntimeConfig::default() + }; + let factory = OrchestratorFactory::new(config); + assert_eq!(factory.run_id(), named); + assert_eq!(factory.agent_config.run_id, Some(named)); + + let unnamed = OrchestratorFactory::new(AgentRuntimeConfig::default()); + let other = OrchestratorFactory::new(AgentRuntimeConfig::default()); + assert_eq!(unnamed.agent_config.run_id, Some(unnamed.run_id())); + assert_ne!( + unnamed.run_id(), + other.run_id(), + "each factory mints its own" + ); + } + + #[tokio::test] + async fn a_factory_refuses_a_stream_naming_another_run() { + let factory = OrchestratorFactory::new(AgentRuntimeConfig::default()); + let items: Vec<_> = factory + .stream( + "q", + Vec::new(), + crate::streaming::RunOptions::default(), + "req_123", + ) + .await + .into_events() + .collect() + .await; + + let [Err(error)] = items.as_slice() else { + panic!("expected one refusal, got {items:?}"); + }; + assert!(matches!( + error.downcast_ref::(), + Some(crate::streaming::StreamRefused::OtherId { run, .. }) if *run == factory.run_id() + )); + } /// A reader that finds usage on the response takes the cache split from /// there too, so the response carries the run's split with its totals. diff --git a/crates/aura/src/orchestration/mod.rs b/crates/aura/src/orchestration/mod.rs index 6c147b85f..91637861b 100644 --- a/crates/aura/src/orchestration/mod.rs +++ b/crates/aura/src/orchestration/mod.rs @@ -27,17 +27,22 @@ //! # Example Usage //! //! ```ignore -//! use aura::{RigBuilder, StreamingAgent}; +//! use aura::hitl::PendingApprovals; +//! use aura::{RigBuilder, RunId, StreamingAgent}; //! use aura_config::load_config_from_str; //! //! // `RigBuilder` returns an `Orchestrator` (wrapped as `StreamingAgent`) when //! // `orchestration.enabled = true`, or a standard `Agent` otherwise. //! let config = load_config_from_str(toml_str)?; -//! let agent: std::sync::Arc = RigBuilder::new(config) -//! .build_streaming_agent_with_headers(None, None, None) +//! let run_id = RunId::mint(); +//! let agent: std::sync::Arc = RigBuilder::new(config, PendingApprovals::new()) +//! .build_streaming_agent_with_headers(None, None, None, Some(run_id)) //! .await?; //! -//! let run = agent.stream(query, history, RunOptions::default(), "req_123").await; +//! // The run streams under the id it was built for. +//! let run = agent +//! .stream(query, history, RunOptions::default(), &run_id.to_string()) +//! .await; //! let stream = run.into_events(); //! ``` diff --git a/crates/aura/src/orchestration/orchestrator.rs b/crates/aura/src/orchestration/orchestrator.rs index 21c6d204c..61d46299e 100644 --- a/crates/aura/src/orchestration/orchestrator.rs +++ b/crates/aura/src/orchestration/orchestrator.rs @@ -450,6 +450,9 @@ pub struct Orchestrator { /// The underlying agent configuration (for creating workers) agent_config: AgentRuntimeConfig, + /// The run the orchestration serves. + run_id: crate::config::RunId, + /// Tool call observer for coordinator visibility into worker tool execution. /// Wired to emit AgentEventPayload for real-time SSE streaming via spawn_tool_event_forwarder. pub(super) tool_call_observer: ToolCallObserver, @@ -595,10 +598,16 @@ enum LoopStep { } impl Orchestrator { - /// Create a new orchestrator from configuration. + /// Create a new orchestrator from configuration, for the run + /// `agent_config.run_id` names, or for a fresh run when it names none. pub async fn new( - agent_config: AgentRuntimeConfig, + mut agent_config: AgentRuntimeConfig, ) -> Result> { + // Minted once, here, so every agent of the orchestration begins its + // run within the same one. + let run_id = *agent_config + .run_id + .get_or_insert_with(crate::config::RunId::mint); let orchestration_config = agent_config.orchestration.clone().unwrap_or_default(); // Initialize MCP manager (shared across coordinator and all workers via Arc) @@ -635,6 +644,9 @@ impl Orchestrator { let orchestrator_id = uuid::Uuid::new_v4().to_string(); + // Persistence names the run by an id of its own, which the park + // owner, the checkpoint and `RunParked` carry; the run's `RunContext` + // is named by its `RunId`. The two are distinct. let run_id_str = persistence.lock().await.run_id().to_string(); // One guard per park-mode run; `ParkGuard` documents arming and drop. let park_guard = agent_config @@ -662,6 +674,7 @@ impl Orchestrator { Ok(Self { orchestrator_id, + run_id, config: orchestration_config, agent_config, tool_call_observer, @@ -1156,12 +1169,10 @@ impl Orchestrator { /// The run the orchestration serves, for its workers and coordinator to /// begin theirs within. That is the run in scope; a test driving the - /// orchestrator outside one gets a run nobody observes, under the - /// request id the config carries. + /// orchestrator outside one gets a run nobody observes, under the id the + /// orchestrator was created with. fn orchestration_run(&self) -> Arc { - crate::run_context::current_run().unwrap_or_else(|| { - RunContext::channel(self.agent_config.request_id.clone().unwrap_or_default()).0 - }) + crate::run_context::current_run().unwrap_or_else(|| RunContext::channel(self.run_id).0) } /// Park-mode wiring for one worker attempt: `Some` only when this run @@ -7499,6 +7510,33 @@ mod tests { assert_eq!(other["content"], "kept"); } + /// Every agent of an orchestration begins its run within the + /// orchestration's, so outside a scope they all get the same unobserved + /// run, under the id the orchestrator was created with. + #[tokio::test] + async fn an_unscoped_orchestration_names_one_run() { + let orchestrator = Orchestrator::new(crate::config::AgentRuntimeConfig::default()) + .await + .unwrap(); + let first = orchestrator.orchestration_run(); + let second = orchestrator.orchestration_run(); + assert_eq!(first.id(), second.id()); + assert_eq!( + orchestrator.agent_config.run_id, + Some(first.id()), + "the config carries the id the orchestrator minted" + ); + + let named = crate::run_context::named_run_id("req_named"); + let orchestrator = Orchestrator::new(crate::config::AgentRuntimeConfig { + run_id: Some(named), + ..Default::default() + }) + .await + .unwrap(); + assert_eq!(orchestrator.orchestration_run().id(), named); + } + /// Attached artifacts are listed in the worker's context after any /// dependency results, and alone when the task has no dependencies. #[tokio::test] @@ -7703,8 +7741,8 @@ mod tests { }; use crate::session_store::{InMemoryApprovalStore, InMemoryEventBus}; - let request_id = format!("req_cancel_{}", uuid::Uuid::new_v4().simple()); - let (run, mut events) = crate::run_context::RunContext::channel(request_id.as_str()); + let run_id = aura_events::RunId::mint(); + let (run, mut events) = crate::run_context::RunContext::channel(run_id); let store: Arc = Arc::new(InMemoryApprovalStore::new()); @@ -7719,7 +7757,7 @@ mod tests { }), park_enabled: true, }), - request_id: Some(request_id.clone()), + run_id: Some(run_id), ..AgentRuntimeConfig::default() }; let orchestrator = Orchestrator::new(config).await.unwrap(); @@ -7835,7 +7873,7 @@ mod tests { }), memory_dir: Some(memory_dir.to_string_lossy().into_owned()), session_id: Some("park-sess".to_string()), - request_id: Some(format!("req_park_{}", uuid::Uuid::new_v4().simple())), + run_id: Some(aura_events::RunId::mint()), ..AgentRuntimeConfig::default() }; let orchestrator = Orchestrator::new(config).await.unwrap(); @@ -8083,12 +8121,11 @@ mod tests { let dir = tempfile::tempdir().unwrap(); let (orchestrator, store, registry, run_id) = park_orchestrator(dir.path()).await; let (plan, records, pending) = awaiting_plan_with_parked_calls(®istry, &run_id).await; - let request_id = orchestrator + let this_run = orchestrator .agent_config - .request_id - .clone() - .unwrap_or_default(); - let (run, mut events) = crate::run_context::RunContext::channel(request_id.as_str()); + .run_id + .expect("a park orchestrator is built for a run"); + let (run, mut events) = crate::run_context::RunContext::channel(this_run); // Arming inside the scope is what gives the guard the run its drop // sweep reports to. crate::run_context::with_run(Arc::clone(&run), arm_guard(&orchestrator, &plan)).await; @@ -8178,12 +8215,11 @@ mod tests { let dir = tempfile::tempdir().unwrap(); let (orchestrator, store, registry, run_id) = park_orchestrator(dir.path()).await; let (plan, records, pending) = awaiting_plan_with_parked_calls(®istry, &run_id).await; - let request_id = orchestrator + let this_run = orchestrator .agent_config - .request_id - .clone() - .unwrap_or_default(); - let (run, mut events) = crate::run_context::RunContext::channel(request_id.as_str()); + .run_id + .expect("a park orchestrator is built for a run"); + let (run, mut events) = crate::run_context::RunContext::channel(this_run); // The human decides the first call before the commit is attempted. let decided = pending[0].decision_id; @@ -8270,12 +8306,11 @@ mod tests { let (orchestrator, registry, run_id) = park_orchestrator_over(store.clone(), dir.path()).await; let (plan, records, pending) = awaiting_plan_with_parked_calls(®istry, &run_id).await; - let request_id = orchestrator + let this_run = orchestrator .agent_config - .request_id - .clone() - .unwrap_or_default(); - let (run, mut events) = crate::run_context::RunContext::channel(request_id.as_str()); + .run_id + .expect("a park orchestrator is built for a run"); + let (run, mut events) = crate::run_context::RunContext::channel(this_run); crate::run_context::with_run(Arc::clone(&run), arm_guard(&orchestrator, &plan)).await; let (event_tx, mut event_rx) = tokio::sync::mpsc::channel(32); @@ -8398,7 +8433,8 @@ mod tests { skills: None, }, )]); - let request_id = format!("req_orphan_{}", uuid::Uuid::new_v4().simple()); + let run_id = aura_events::RunId::mint(); + let request_id = run_id.to_string(); let config = AgentRuntimeConfig { hitl: Some(crate::hitl::HitlRuntime { patterns: Arc::from(["echo_tool".into()]), @@ -8410,7 +8446,7 @@ mod tests { }), memory_dir: Some(memory_dir.to_string_lossy().into_owned()), session_id: Some("orphan-sess".to_string()), - request_id: Some(request_id.clone()), + run_id: Some(run_id), orchestration: Some(OrchestrationConfig { enabled: true, workers, @@ -8522,7 +8558,8 @@ mod tests { let dir = tempfile::tempdir().unwrap(); let (orchestrator, store, _registry, request_id) = override_park_orchestrator(dir.path(), 1).await; - let (run, mut events) = crate::run_context::RunContext::channel(request_id.as_str()); + let (run, mut events) = + crate::run_context::RunContext::channel(request_id.parse().expect("a run id")); // Depth 1 gives the loop three turns (the rig's +1 safety net), so // the gated call must land on the third: the first two turns burn @@ -8581,7 +8618,8 @@ mod tests { let dir = tempfile::tempdir().unwrap(); let (orchestrator, store, _registry, request_id) = override_park_orchestrator(dir.path(), 4).await; - let (run, mut events) = crate::run_context::RunContext::channel(request_id.as_str()); + let (run, mut events) = + crate::run_context::RunContext::channel(request_id.parse().expect("a run id")); let (_model, gated_invocations) = gated_worker_override(vec![ScriptedTurn::tool_calls_then_stream_failure(vec![ @@ -8701,7 +8739,7 @@ mod tests { }), memory_dir: Some(memory_dir.to_string_lossy().into_owned()), session_id: Some(session_id.to_string()), - request_id: Some(format!("req_resume_{}", uuid::Uuid::new_v4().simple())), + run_id: Some(aura_events::RunId::mint()), orchestration: Some(OrchestrationConfig { enabled: true, workers, diff --git a/crates/aura/src/orchestration/park/commit.rs b/crates/aura/src/orchestration/park/commit.rs index c385962a6..fbb813efe 100644 --- a/crates/aura/src/orchestration/park/commit.rs +++ b/crates/aura/src/orchestration/park/commit.rs @@ -558,8 +558,8 @@ mod tests { let (registry, store) = conv_registry(); let run_id = "0191e8c0-ffff-7000-8000-000000000006"; let owner = run_owner_id(run_id); - let request_id = format!("req_sweep_{}", uuid::Uuid::new_v4().simple()); - let (run, mut events) = crate::run_context::RunContext::channel(request_id.as_str()); + let this_run = aura_events::RunId::mint(); + let (run, mut events) = crate::run_context::RunContext::channel(this_run); let now = chrono::Utc::now(); let decided = DecisionId::generate(); @@ -626,8 +626,8 @@ mod tests { let (registry, store) = conv_registry(); let run_id = "0191e8c0-aaaa-7000-8000-000000000007"; let owner = run_owner_id(run_id); - let request_id = format!("req_sweep_{}", uuid::Uuid::new_v4().simple()); - let (run, mut events) = crate::run_context::RunContext::channel(request_id.as_str()); + let this_run = aura_events::RunId::mint(); + let (run, mut events) = crate::run_context::RunContext::channel(this_run); let now = chrono::Utc::now(); let first = DecisionId::generate(); diff --git a/crates/aura/src/orchestration/park/guard.rs b/crates/aura/src/orchestration/park/guard.rs index 2297b3c19..ec2f2e05d 100644 --- a/crates/aura/src/orchestration/park/guard.rs +++ b/crates/aura/src/orchestration/park/guard.rs @@ -144,8 +144,8 @@ mod tests { async fn unpublished_guard_drop_cancels_run_approvals() { let (registry, store) = registry_with_store(); let run_id: RunId = "0191e8c0-2222-7000-8000-000000000042".parse().unwrap(); - let request_id = format!("req_guard_{}", uuid::Uuid::new_v4().simple()); - let (run, mut events) = crate::run_context::RunContext::channel(request_id.as_str()); + let this_run = aura_events::RunId::mint(); + let (run, mut events) = crate::run_context::RunContext::channel(this_run); let scope = worker_scope(run_id); let decision_id = DecisionId::generate(); diff --git a/crates/aura/src/orchestration/persistence_wrapper.rs b/crates/aura/src/orchestration/persistence_wrapper.rs index 470250690..2033f1c99 100644 --- a/crates/aura/src/orchestration/persistence_wrapper.rs +++ b/crates/aura/src/orchestration/persistence_wrapper.rs @@ -1449,7 +1449,7 @@ mod tests { session_id: None, }; - let request_id = format!("req_w2_{}", uuid::Uuid::new_v4().simple()); + let this_run = aura_events::RunId::mint(); let gate = Arc::new(HitlApprovalWrapper::new( Arc::from(["kubectl_*".into()]), @@ -1460,7 +1460,7 @@ mod tests { )); // `WrappedTool` runs `pre_call` in its own task, which no scope // crosses, so the gate is bound the way `stream` binds it. - let (run, mut rx) = crate::run_context::RunContext::channel(request_id.as_str()); + let (run, mut rx) = crate::run_context::RunContext::channel(this_run); gate.bind_run(run); let gate: Arc = gate; let persistence: Arc = Arc::new(test_wrapper(Arc::new(Mutex::new( diff --git a/crates/aura/src/orchestration/test_rig.rs b/crates/aura/src/orchestration/test_rig.rs index c95b2128c..1ab821de6 100644 --- a/crates/aura/src/orchestration/test_rig.rs +++ b/crates/aura/src/orchestration/test_rig.rs @@ -742,7 +742,7 @@ pub(crate) async fn park_orchestrator_in( }), memory_dir: Some(memory_dir.to_string_lossy().into_owned()), session_id: Some("park-sess".to_string()), - request_id: Some(format!("req_rig_{}", uuid::Uuid::new_v4().simple())), + run_id: Some(aura_events::RunId::mint()), orchestration: Some(super::OrchestrationConfig { enabled: true, workers, diff --git a/crates/aura/src/orchestration/types.rs b/crates/aura/src/orchestration/types.rs index d10d4a427..2eab5b543 100644 --- a/crates/aura/src/orchestration/types.rs +++ b/crates/aura/src/orchestration/types.rs @@ -4,12 +4,9 @@ //! queries into tasks, track their execution, and manage dependencies. use std::collections::HashMap; -use std::fmt; -use std::str::FromStr; use std::sync::Mutex; use serde::{Deserialize, Serialize}; -use uuid::Uuid; use crate::hitl::DecisionId; @@ -23,28 +20,7 @@ const MAX_STEP_NESTING: usize = 2; // Domain identifiers for orchestration runs and tasks, modeled as simple types // (opaque newtypes reached through canonical conversion traits). -/// Identifier for a single orchestration run. -/// -/// A run is an orchestration concept; single-agent requests have none. Run ids -/// are v4 UUIDs; parse one from its string form via `FromStr`. Serializes as -/// the bare UUID string. -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)] -#[serde(transparent)] -pub struct RunId(Uuid); - -impl fmt::Display for RunId { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - self.0.fmt(f) - } -} - -impl FromStr for RunId { - type Err = uuid::Error; - - fn from_str(s: &str) -> Result { - Uuid::from_str(s).map(Self) - } -} +pub use aura_events::RunId; /// Identity of a worker task within a run. /// diff --git a/crates/aura/src/rig_builder.rs b/crates/aura/src/rig_builder.rs index 223ef6ee3..0a032ab61 100644 --- a/crates/aura/src/rig_builder.rs +++ b/crates/aura/src/rig_builder.rs @@ -12,7 +12,7 @@ use crate::builder::{ Agent, ClientTool, PreparedAgent, RunToolFactory, build_streaming_agent_with_tools, no_run_tools, }; -use crate::config::{AgentRuntimeConfig, WorkerSkills}; +use crate::config::{AgentRuntimeConfig, RunId, WorkerSkills}; use crate::error::BuilderError; use crate::forwarded_headers::ForwardedHeaders; use crate::hitl::PendingApprovals; @@ -177,8 +177,8 @@ impl RigBuilder { .map_err(|e| BuilderError::AgentError(format!("Failed to build agent: {e}"))) } - /// Prepare an agent and begin its run for `request_id`, recording skill - /// invocations with the builder's skill recorder. + /// Prepare an agent and begin its run `run_id` (a fresh one when `None`), + /// recording skill invocations with the builder's skill recorder. /// /// See [`Self::prepare_agent`] for the parameters. Each call prepares a /// fresh agent; callers that want to reuse one across requests call @@ -188,7 +188,7 @@ impl RigBuilder { req_headers: Option<&HashMap>, additional_tools: Vec>, client_tools: Option>, - request_id: Option, + run_id: Option, session_id: Option, ) -> Result { let prepared = self @@ -196,7 +196,7 @@ impl RigBuilder { .await?; prepared .begin_run( - request_id.unwrap_or_default(), + run_id.unwrap_or_else(RunId::mint), req_headers, self.skill_recorder.clone(), ) @@ -218,13 +218,13 @@ impl RigBuilder { req_headers: Option<&HashMap>, session_id: Option, client_tools: Option>, - request_id: Option, + run_id: Option, ) -> Result, BuilderError> { self.build_streaming_agent_with_tools( req_headers, session_id, client_tools, - request_id, + run_id, no_run_tools(), ) .await @@ -239,13 +239,13 @@ impl RigBuilder { req_headers: Option<&HashMap>, session_id: Option, client_tools: Option>, - request_id: Option, + run_id: Option, run_tools: RunToolFactory, ) -> Result, BuilderError> { let mut agent_config = self.discovered_agent_config(req_headers)?; resolve_mcp_headers(&mut agent_config, req_headers); agent_config.session_id = session_id; - agent_config.request_id = request_id; + agent_config.run_id = run_id; agent_config.skill_recorder = self.skill_recorder.clone(); build_streaming_agent_with_tools(&agent_config, client_tools, run_tools) diff --git a/crates/aura/src/run_context.rs b/crates/aura/src/run_context.rs index 5a31ebaac..822536419 100644 --- a/crates/aura/src/run_context.rs +++ b/crates/aura/src/run_context.rs @@ -21,7 +21,7 @@ use futures::Stream; use aura_events::agent::AgentEvent; use tokio::sync::mpsc; -use aura_events::ToolCallId; +use aura_events::{RunId, ToolCallId}; use crate::scratchpad::ContextBudget; use crate::skill_tool::SkillInvocationRecorder; @@ -33,7 +33,7 @@ pub const EVENT_CHANNEL_CAPACITY: usize = 1024; /// One run — what its own work needs to correlate, where its events go, and /// the state an agent's tools keep for it. pub struct RunContext { - id: Arc, + id: RunId, tool_calls: Mutex>, events: mpsc::Sender, cancel: CancellationToken, @@ -50,7 +50,7 @@ const MAX_PENDING_TOOL_CALLS: usize = 256; impl RunContext { /// A run and the receiver its observer reads, on a token of its own. - pub fn channel(id: impl Into>) -> (Arc, mpsc::Receiver) { + pub fn channel(id: RunId) -> (Arc, mpsc::Receiver) { Self::channel_on(id, CancellationToken::new()) } @@ -58,7 +58,7 @@ impl RunContext { /// to stop on — a child of its own caller's, so one run ending leaves the /// others alone. pub fn channel_on( - id: impl Into>, + id: RunId, cancel: CancellationToken, ) -> (Arc, mpsc::Receiver) { Self::channel_for_agent(id, cancel, None, None, None) @@ -67,7 +67,7 @@ impl RunContext { /// A run on `cancel` carrying the state a prepared agent's tools keep for /// it, and the receiver its observer reads. pub fn channel_for_agent( - id: impl Into>, + id: RunId, cancel: CancellationToken, scratchpad_budget: Option, turn_nudge: Option>, @@ -75,7 +75,7 @@ impl RunContext { ) -> (Arc, mpsc::Receiver) { let (events, receiver) = mpsc::channel(EVENT_CHANNEL_CAPACITY); let run = Arc::new(Self { - id: id.into(), + id, tool_calls: Mutex::new(VecDeque::new()), events, cancel, @@ -97,7 +97,7 @@ impl RunContext { skill_recorder: Option>, ) -> Arc { Arc::new(Self { - id: Arc::clone(&parent.id), + id: parent.id, tool_calls: Mutex::new(VecDeque::new()), events: parent.events.clone(), cancel: parent.cancel.clone(), @@ -111,14 +111,14 @@ impl RunContext { /// reading what it emits. Production names a run it can reach an observer /// through, or names none. #[cfg(test)] - pub fn detached(id: impl Into>) -> Arc { + pub fn detached(id: RunId) -> Arc { Self::channel(id).0 } /// [`detached`](Self::detached), carrying tool state. #[cfg(test)] pub(crate) fn detached_with( - id: impl Into>, + id: RunId, scratchpad_budget: Option, turn_nudge: Option>, ) -> Arc { @@ -163,8 +163,17 @@ impl RunContext { delivered } - pub fn id(&self) -> &Arc { - &self.id + pub fn id(&self) -> RunId { + self.id + } + + /// Whether `id` is this run's id as the run spells it — the form its + /// `Display` gives, which is the key every request-keyed registry holds. + /// The same UUID spelled another way names no run, and neither does a + /// string that is not one. + pub fn has_id(&self, id: &str) -> bool { + let mut spelled = uuid::Uuid::encode_buffer(); + *self.id.as_uuid().hyphenated().encode_lower(&mut spelled) == *id } /// The token that cancels this run. @@ -273,13 +282,17 @@ impl BoundRun { /// A slot holding an unobserved run that carries only `budget`. #[cfg(test)] pub(crate) fn pinned_budget(budget: ContextBudget) -> Self { - Self::holding(RunContext::detached_with("pinned", Some(budget), None)) + Self::holding(RunContext::detached_with(RunId::mint(), Some(budget), None)) } /// A slot holding an unobserved run that carries only `turn_nudge`. #[cfg(test)] pub(crate) fn pinned_nudge(turn_nudge: Arc) -> Self { - Self::holding(RunContext::detached_with("pinned", None, Some(turn_nudge))) + Self::holding(RunContext::detached_with( + RunId::mint(), + None, + Some(turn_nudge), + )) } } @@ -308,10 +321,10 @@ impl RunLease { /// A prepared agent was asked to begin a run while it still serves another. #[derive(Debug, thiserror::Error)] -#[error("prepared agent already serves request `{active}`; it runs one request at a time")] +#[error("prepared agent already serves run `{active}`; it serves one run at a time")] pub struct RunInProgress { /// Id of the run holding the agent. - pub active: String, + pub active: RunId, } tokio::task_local! { @@ -341,8 +354,8 @@ pub fn current_run() -> Option> { RUN.try_with(Arc::clone).ok() } -pub fn current_run_id() -> Option> { - RUN.try_with(|run| Arc::clone(run.id())).ok() +pub fn current_run_id() -> Option { + RUN.try_with(|run| run.id()).ok() } /// Runs `f` with `run` in scope. Task-locals do not cross `tokio::spawn`, so @@ -377,8 +390,8 @@ impl Stream for ScopedStream { /// Runs `f` with a fresh run in scope and returns what it emitted, for a test /// that asserts on a run's events without standing up an observer. #[cfg(test)] -pub(crate) async fn observing(id: &str, f: F) -> (F::Output, Vec) { - let (run, mut events) = RunContext::channel(id); +pub(crate) async fn observing(name: &str, f: F) -> (F::Output, Vec) { + let (run, mut events) = RunContext::channel(named_run_id(name)); let out = with_run(run, f).await; let mut seen = Vec::new(); @@ -388,22 +401,33 @@ pub(crate) async fn observing(id: &str, f: F) -> (F::Output, Vec RunId { + RunId::try_from(uuid::Uuid::new_v5( + &uuid::Uuid::NAMESPACE_OID, + name.as_bytes(), + )) + .expect("a v5 UUID is never nil") +} + #[cfg(test)] mod tests { use super::*; use futures::StreamExt; - fn run(id: &str) -> Arc { - RunContext::detached(id) + fn run(name: &str) -> Arc { + RunContext::detached(named_run_id(name)) } #[tokio::test] async fn a_scope_established_inside_a_spawn_holds() { - let (run, _rx) = RunContext::channel("spawned"); + let (run, _rx) = RunContext::channel(named_run_id("spawned")); let seen = tokio::spawn(with_run(run, async { current_run_id() })) .await .unwrap(); - assert_eq!(seen.as_deref(), Some("spawned")); + assert_eq!(seen, Some(named_run_id("spawned"))); } /// A run built on the caller's token stops when the caller does. The work @@ -413,7 +437,7 @@ mod tests { #[tokio::test] async fn a_run_stops_on_the_token_it_was_built_on() { let caller = CancellationToken::new(); - let (run, _events) = RunContext::channel_on("run_on_token", caller.clone()); + let (run, _events) = RunContext::channel_on(named_run_id("run_on_token"), caller.clone()); assert!(!run.cancel_token().is_cancelled()); caller.cancel(); @@ -424,8 +448,8 @@ mod tests { /// cancels something else. #[tokio::test] async fn a_run_given_no_token_has_its_own() { - let (a, _ea) = RunContext::channel("run_a"); - let (b, _eb) = RunContext::channel("run_b"); + let (a, _ea) = RunContext::channel(named_run_id("run_a")); + let (b, _eb) = RunContext::channel(named_run_id("run_b")); a.cancel_token().cancel(); assert!(a.cancel_token().is_cancelled()); @@ -441,10 +465,26 @@ mod tests { assert_eq!(current_run_id(), None); } + /// Registries key a run by the string its id displays as, so that is the + /// one spelling a run answers to: not the same UUID in another case or + /// form, and not a string that is no run id at all. + #[test] + fn a_run_answers_only_to_its_id_as_it_spells_it() { + let run = run("run_spelled"); + let spelled = run.id().to_string(); + + assert!(run.has_id(&spelled)); + assert!(!run.has_id(&spelled.to_uppercase())); + assert!(!run.has_id(&run.id().as_uuid().simple().to_string())); + assert!(!run.has_id(&format!("urn:uuid:{spelled}"))); + assert!(!run.has_id("req_1")); + assert!(!run.has_id(&named_run_id("run_other").to_string())); + } + #[tokio::test] async fn a_scope_supplies_the_run() { let seen = with_run(run("run_1"), async { current_run_id() }).await; - assert_eq!(seen.as_deref(), Some("run_1")); + assert_eq!(seen, Some(named_run_id("run_1"))); } #[tokio::test] @@ -458,8 +498,8 @@ mod tests { current_run_id() })); - assert_eq!(a.await.unwrap().as_deref(), Some("run_a")); - assert_eq!(b.await.unwrap().as_deref(), Some("run_b")); + assert_eq!(a.await.unwrap(), Some(named_run_id("run_a"))); + assert_eq!(b.await.unwrap(), Some(named_run_id("run_b"))); } #[tokio::test] @@ -467,10 +507,7 @@ mod tests { let inner = futures::stream::iter(0..3).map(|_| current_run_id()); let seen: Vec<_> = scope_stream(run("run_s"), inner).collect().await; - assert_eq!( - seen.iter().map(|id| id.as_deref()).collect::>(), - vec![Some("run_s"); 3] - ); + assert_eq!(seen, vec![Some(named_run_id("run_s")); 3]); } /// Orchestration drives workers with `FuturesUnordered` inside the run's @@ -496,10 +533,7 @@ mod tests { }) .await; - assert_eq!( - seen.iter().map(|id| id.as_deref()).collect::>(), - vec![Some("run_w"); 3] - ); + assert_eq!(seen, vec![Some(named_run_id("run_w")); 3]); } /// A spawned task does not inherit its parent's scope, which is why every @@ -534,12 +568,12 @@ mod tests { let nudge = TurnNudgeState::new(true, None, 2).unwrap(); let slot = BoundRun::default(); slot.bind(RunContext::detached_with( - "req_a", + named_run_id("req_a"), Some(budget.clone()), Some(Arc::clone(&nudge)), )); - assert_eq!(slot.id_or_empty(), "req_a"); + assert_eq!(slot.id_or_empty(), named_run_id("req_a").to_string()); slot.scratchpad_budget().unwrap().record_intercepted(7); assert_eq!( budget.scratchpad_usage().0, @@ -548,8 +582,8 @@ mod tests { ); assert!(Arc::ptr_eq(&slot.turn_nudge().unwrap(), &nudge)); - slot.bind(RunContext::detached("req_b")); - assert_eq!(slot.id_or_empty(), "req_b"); + slot.bind(RunContext::detached(named_run_id("req_b"))); + assert_eq!(slot.id_or_empty(), named_run_id("req_b").to_string()); assert!(slot.scratchpad_budget().is_none()); } @@ -557,11 +591,11 @@ mod tests { /// cancellation see it, with the worker's own tool state. #[tokio::test] async fn a_child_shares_its_parents_identity_and_keeps_its_own_state() { - let (parent, mut events) = RunContext::channel("req_parent"); + let (parent, mut events) = RunContext::channel(named_run_id("req_parent")); let nudge = TurnNudgeState::new(true, None, 2).unwrap(); let child = RunContext::child(&parent, None, Some(Arc::clone(&nudge)), None); - assert_eq!(child.id().as_ref(), "req_parent"); + assert_eq!(child.id(), parent.id()); assert!(Arc::ptr_eq(child.turn_nudge().unwrap(), &nudge)); assert!(parent.turn_nudge().is_none(), "the parent keeps none of it"); diff --git a/crates/aura/src/scratchpad/setup.rs b/crates/aura/src/scratchpad/setup.rs index 6c39d8725..d13161547 100644 --- a/crates/aura/src/scratchpad/setup.rs +++ b/crates/aura/src/scratchpad/setup.rs @@ -306,7 +306,7 @@ mod tests { assert!(build.tools_config.run.scratchpad_budget().is_none()); let run_budget = build.budget.fresh(); run.bind(crate::run_context::RunContext::detached_with( - "req", + crate::run_context::named_run_id("req"), Some(run_budget.clone()), None, )); diff --git a/crates/aura/src/skill_rehydration.rs b/crates/aura/src/skill_rehydration.rs index c6dfcda48..465096bbe 100644 --- a/crates/aura/src/skill_rehydration.rs +++ b/crates/aura/src/skill_rehydration.rs @@ -411,7 +411,7 @@ mod tests { // Turn N: history [user], anchor = 0 + 1. The LLM calls load_skill. let recorder = Arc::new(SkillInvocationRecorder::new(store.clone(), log.clone(), 1)); let (run, _events) = crate::run_context::RunContext::channel_for_agent( - "req-turn-n", + crate::run_context::named_run_id("req-turn-n"), tokio_util::sync::CancellationToken::new(), None, None, diff --git a/crates/aura/src/skill_tool.rs b/crates/aura/src/skill_tool.rs index b308be76a..4203c45d4 100644 --- a/crates/aura/src/skill_tool.rs +++ b/crates/aura/src/skill_tool.rs @@ -877,8 +877,14 @@ mod tests { log.clone(), anchor, )); - RunContext::channel_for_agent(id, CancellationToken::new(), None, None, Some(recorder)) - .0 + RunContext::channel_for_agent( + crate::run_context::named_run_id(id), + CancellationToken::new(), + None, + None, + Some(recorder), + ) + .0 }; let load = |name: &str| LoadSkillArgs { name: name.to_string(), diff --git a/crates/aura/src/streaming.rs b/crates/aura/src/streaming.rs index 794f7dfac..c98dafbda 100644 --- a/crates/aura/src/streaming.rs +++ b/crates/aura/src/streaming.rs @@ -14,12 +14,16 @@ //! //! ```ignore //! use aura::streaming::{RunOptions, StreamingAgent}; -//! use aura::{StreamError, StreamItem}; +//! use aura::{RunId, StreamError, StreamItem}; //! use futures::StreamExt; //! -//! async fn handle_request(agent: impl StreamingAgent, query: &str) { +//! // `agent` was built for `run_id` (`AgentRuntimeConfig::run_id`), the one +//! // id it streams under. +//! async fn handle_request(agent: impl StreamingAgent, run_id: RunId, query: &str) { //! // 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 run = agent +//! .stream(query, vec![], RunOptions::default(), &run_id.to_string()) +//! .await; //! let mut items = run.into_events(); //! //! // Process stream items (convert to SSE, etc.) @@ -37,6 +41,7 @@ use crate::provider_agent::{StreamError, StreamItem, StreamedAssistantContent}; use crate::run_context::RunContext; use crate::streaming_request_hook::UsageState; use async_trait::async_trait; +use aura_events::RunId; use futures::stream::BoxStream; use rig::completion::Message; use std::time::Duration; @@ -83,6 +88,44 @@ impl RunOptions { } } +/// Why a run would not stream. +#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)] +pub enum StreamRefused { + /// The caller named a different run. + #[error("run `{run}` cannot stream under `{requested}`; pass the agent's `run_id()`")] + OtherId { run: RunId, requested: String }, + /// The run has streamed already. + #[error("run `{run}` has streamed already; begin a new run to stream again")] + AlreadyStreamed { run: RunId }, +} + +/// The claim on a run's one stream. +#[derive(Debug, Default)] +pub struct StreamClaim { + taken: std::sync::atomic::AtomicBool, +} + +impl StreamClaim { + /// Claims the stream of run `run` for a caller naming `request_id`. + /// + /// The id is checked before the claim is taken, so a call naming the wrong + /// run leaves the stream to the caller that names it. The id must be the + /// run's own string form, the key every request-keyed registry holds, not + /// merely the same UUID spelled differently. + pub fn claim(&self, run: RunId, request_id: &str) -> Result<(), StreamRefused> { + if run.to_string() != request_id { + return Err(StreamRefused::OtherId { + run, + requested: request_id.to_owned(), + }); + } + if self.taken.swap(true, std::sync::atomic::Ordering::AcqRel) { + return Err(StreamRefused::AlreadyStreamed { run }); + } + Ok(()) + } +} + /// A started run: the events it produces, the token that cancels it, and the /// usage it accumulates. pub struct AgentRun { @@ -111,6 +154,16 @@ impl AgentRun { } } + /// A run that does nothing, whose stream yields `refused` and ends. + pub fn refused(refused: StreamRefused) -> Self { + let error: StreamError = Box::new(refused); + Self::new( + Box::pin(futures::stream::once(async move { Err(error) })), + CancellationToken::new(), + UsageState::new(), + ) + } + /// Hands the run's events to its observer. A run whose producers emit /// through [`crate::run_context::RunContext::emit`] pairs the sender it /// scopes with the receiver named here. @@ -277,13 +330,18 @@ pub trait StreamingAgent: Send + Sync { /// needs to know the concrete agent type. fn get_provider_info(&self) -> (&str, &str); + /// The run this agent streams. + fn run_id(&self) -> RunId; + /// Start a run. /// /// `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. /// - /// `request_id` correlates MCP progress and tool events for this run. + /// `request_id` names the run: [`run_id`](Self::run_id), as a string. A + /// run streams once, and only under that id; any other call returns a run + /// whose stream yields one [`StreamRefused`] and does nothing. async fn stream( &self, query: &str, @@ -338,6 +396,83 @@ mod tests { use futures::StreamExt; use std::sync::Arc; + fn run(name: &str) -> RunId { + crate::run_context::named_run_id(name) + } + + #[test] + fn the_first_claim_naming_the_run_streams() { + let claim = StreamClaim::default(); + assert_eq!(claim.claim(run("a"), &run("a").to_string()), Ok(())); + } + + #[test] + fn a_second_claim_is_refused() { + let claim = StreamClaim::default(); + let id = run("a").to_string(); + claim.claim(run("a"), &id).unwrap(); + + assert_eq!( + claim.claim(run("a"), &id), + Err(StreamRefused::AlreadyStreamed { run: run("a") }) + ); + } + + /// A call naming the wrong run is refused without taking the stream, so + /// the caller that names the run still gets it. + #[test] + fn a_claim_naming_another_run_leaves_the_stream_untaken() { + let claim = StreamClaim::default(); + + assert_eq!( + claim.claim(run("a"), "req_123"), + Err(StreamRefused::OtherId { + run: run("a"), + requested: "req_123".to_owned(), + }) + ); + assert_eq!(claim.claim(run("a"), &run("a").to_string()), Ok(())); + } + + /// The same UUID spelled another way is not the key the run's registries + /// hold, so it is refused too. + #[test] + fn the_run_id_must_be_spelled_as_the_run_spells_it() { + let claim = StreamClaim::default(); + let shouted = run("a").to_string().to_uppercase(); + + assert!(matches!( + claim.claim(run("a"), &shouted), + Err(StreamRefused::OtherId { .. }) + )); + } + + /// A caller refused for naming the wrong run is told where the right one is. + #[test] + fn a_refusal_for_another_run_points_at_run_id() { + let refused = StreamRefused::OtherId { + run: run("a"), + requested: "req_123".to_owned(), + }; + assert!(refused.to_string().contains("run_id()"), "got: {refused}"); + } + + /// A refused run reports why and ends, and dropping it cancels nothing but + /// its own token. + #[tokio::test] + async fn a_refused_run_yields_its_refusal_and_ends() { + let refused = AgentRun::refused(StreamRefused::AlreadyStreamed { run: run("a") }); + let items: Vec<_> = refused.into_events().collect().await; + + let [Err(error)] = items.as_slice() else { + panic!("expected one error, got {items:?}"); + }; + assert_eq!( + error.downcast_ref::(), + Some(&StreamRefused::AlreadyStreamed { run: run("a") }) + ); + } + fn empty_run() -> (AgentRun, CancellationToken) { let cancel = CancellationToken::new(); let run = AgentRun::new( @@ -481,7 +616,8 @@ mod tests { agent: aura_events::AgentContext, items: Vec>, ) -> (Vec, usize, Vec) { - let (run, mut events) = RunContext::channel("run_tee"); + let (run, mut events) = + RunContext::channel(crate::run_context::named_run_id("run_tee")); let passed = tee_content(run, agent, futures::stream::iter(items)) .collect::>() .await @@ -564,7 +700,8 @@ mod tests { /// the items must still pass. #[tokio::test] async fn an_unobserved_run_still_streams_its_items() { - let (run, events) = RunContext::channel("run_unobserved"); + let (run, events) = + RunContext::channel(crate::run_context::named_run_id("run_unobserved")); drop(events); let passed = tee_content( diff --git a/crates/aura/src/streaming_request_hook.rs b/crates/aura/src/streaming_request_hook.rs index 44c9909a3..b9ab7f8fc 100644 --- a/crates/aura/src/streaming_request_hook.rs +++ b/crates/aura/src/streaming_request_hook.rs @@ -26,7 +26,8 @@ //! //! ```ignore //! let options = RunOptions::bounded(Some(Duration::from_secs(60))); -//! let (hook, cancel, usage_state) = StreamingRequestHook::new(options, "req_123"); +//! // Keyed by the run's id, so this stream owns the run's tool-call queue. +//! let (hook, cancel, usage_state) = StreamingRequestHook::new(options, run_id.to_string()); //! //! // Pass hook to streaming request //! agent.stream_prompt(query).with_hook(hook).multi_turn(depth).await; @@ -106,7 +107,7 @@ impl Drop for ParkCellRegistration { /// correlates nothing until #732. Its tool events stay off the run for the same /// reason, and orchestration reports the worker's calls itself. fn queue_owner(stream_id: &str) -> Option> { - current_run().filter(|run| run.id().as_ref() == stream_id) + current_run().filter(|run| run.has_id(stream_id)) } /// Sends a tool event raised by this stream to the run [`queue_owner`] gives it. @@ -778,10 +779,11 @@ mod tests { /// scope, and its tool events must not reach the run as the run's own. #[tokio::test] async fn only_the_run_s_own_stream_sends_it_tool_events() { - let (run, mut events) = RunContext::channel("req_1"); + let id = crate::run_context::named_run_id("req_1").to_string(); + let (run, mut events) = RunContext::channel(id.parse().unwrap()); with_run(run, async { - emit_from_stream("req_1:task:0:attempt:1", requested("worker")).await; - emit_from_stream("req_1", requested("own")).await; + emit_from_stream(&format!("{id}:task:0:attempt:1"), requested("worker")).await; + emit_from_stream(&id, requested("own")).await; }) .await; @@ -799,17 +801,19 @@ mod tests { #[tokio::test] async fn the_run_s_own_stream_owns_the_queue() { - let run = RunContext::detached("req_1"); - let owned = with_run(run, async { queue_owner("req_1").is_some() }).await; + let id = crate::run_context::named_run_id("req_1"); + let run = RunContext::detached(id); + let owned = with_run(run, async { queue_owner(&id.to_string()).is_some() }).await; assert!(owned); } /// An orchestration worker streams under its task attempt. #[tokio::test] async fn a_worker_s_stream_owns_no_queue() { - let run = RunContext::detached("req_1"); + let id = crate::run_context::named_run_id("req_1"); + let run = RunContext::detached(id); let owned = with_run(run, async { - queue_owner("req_1:task:0:attempt:1").is_some() + queue_owner(&format!("{id}:task:0:attempt:1")).is_some() }) .await; assert!( diff --git a/crates/aura/src/turn_nudge.rs b/crates/aura/src/turn_nudge.rs index 6ab647f96..9eead3502 100644 --- a/crates/aura/src/turn_nudge.rs +++ b/crates/aura/src/turn_nudge.rs @@ -451,7 +451,7 @@ mod tests { let first_nudge = seed.fresh(); slot.bind(RunContext::detached_with( - "req_a", + crate::run_context::named_run_id("req_a"), None, Some(Arc::clone(&first_nudge)), )); @@ -464,7 +464,11 @@ mod tests { "the first run is on its penultimate turn", ); - slot.bind(RunContext::detached_with("req_b", None, Some(seed.fresh()))); + slot.bind(RunContext::detached_with( + crate::run_context::named_run_id("req_b"), + None, + Some(seed.fresh()), + )); assert_eq!( tool.call("hello".to_string()).await.unwrap(), "hello",