diff --git a/CLAUDE.md b/CLAUDE.md index 821bc9bca..bf8754eb4 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -157,10 +157,12 @@ export AWS_REGION="your-region" # For Knowledge Base - **Rig Fork**: `mezmo/rig` branch `mshearer/LOG-23351-openai-reasoning` ### Key Modules +- `builder.rs` - `PreparedAgent` (built once from config: provider client, tools, MCP connections) and `Agent` (one run of it, via `PreparedAgent::begin_run`) +- `forwarded_headers.rs` - `ForwardedHeaders`, the `headers_from_request` values a prepared agent's MCP connections and HITL route were opened with; `begin_run` refuses a request that forwards different values, so a prepared agent never runs under another request's credentials - `provider_agent.rs` - Type-erased streaming across providers - `stream_events.rs` - Custom aura SSE events - `request_cancellation.rs` - The signal that stops a run, as awaiting work sees it -- `run_context.rs` - The run a task is working on: its event channel, and the FIFO queue for tool_call_id correlation (see critical assumption below) +- `run_context.rs` - `RunContext`, the run a task is working on: its event channel, the FIFO queue for tool_call_id correlation (see critical assumption below), and the scratchpad budget, turn-nudge counters and skill-invocation recorder an agent keeps for it; `BoundRun` is the slot through which prepare-time tools and wrappers (turn nudge, scratchpad, skills, HITL) reach it, and `RunLease` is what `begin_run` counts to serve one run at a time - `orchestration/` - Multi-agent coordinator, workers, DAG execution, orchestration SSE events ### Critical Assumption: Rig Sequential Tool Execution diff --git a/DEVELOPMENT.md b/DEVELOPMENT.md index ee1002c86..cfadbb8eb 100644 --- a/DEVELOPMENT.md +++ b/DEVELOPMENT.md @@ -214,6 +214,7 @@ Prompt routing and execution model: - Direct Mode (`orchestration.enabled = false`): single `Agent` handles the turn. - Orchestration Mode (`orchestration.enabled = true`): `Orchestrator` coordinates worker execution. - Both `Agent` and `Orchestrator` implement `StreamingAgent`, so they are interchangeable at the API boundary. +- An agent is two halves: `PreparedAgent` is built once from config (provider client, discovered tools, MCP connections) and `Agent` is one run of it, begun with `PreparedAgent::begin_run`. Per-run state (the `RunContext`: request id, event channel, tool-call queue, scratchpad budget, turn-nudge counters, skill-invocation recorder) reaches tools through the prepared agent's `BoundRun` slot, and a prepared agent serves one run at a time; see `crates/aura/src/run_context.rs`. A prepared agent also forwards the `headers_from_request` values of the request that prepared it, and `begin_run` refuses a request that forwards different ones (`crates/aura/src/forwarded_headers.rs`). Today every request prepares a fresh agent; the split is what lets a session reuse one. Ingress paths into the web server (all end in the same `StreamingAgent::stream` call): diff --git a/crates/aura/src/builder.rs b/crates/aura/src/builder.rs index 83ca931e8..92868e045 100644 --- a/crates/aura/src/builder.rs +++ b/crates/aura/src/builder.rs @@ -1,25 +1,31 @@ use crate::{ config::{AgentRuntimeConfig, LlmConfig, McpServerConfig, VectorStoreType}, error::{BuilderError, BuilderResult}, + forwarded_headers::ForwardedHeaders, mcp::McpManager, passthrough_tool::PassthroughTool, provider_agent::{ BuilderState, CompletionResponse, ProviderAgent, StreamError, StreamItem, StreamedAssistantContent, }, - scratchpad, - skill_tool::{SkillToolset, render_skill_catalog}, + run_context::{BoundRun, RunContext, RunInProgress, RunLease}, + scratchpad::{self, ContextBudget}, + skill_tool::{SkillInvocationRecorder, SkillToolset, render_skill_catalog}, tool_wrapper::{ToolCallContext, WrappedTool}, tools::{FilesystemTool, ListDirTool, ReadFileTool, WriteFileTool}, + turn_nudge::TurnNudgeState, vector_dynamic::DynamicVectorSearchTool, vector_store::VectorStoreManager, }; +use aura_events::agent::AgentEvent; use futures::StreamExt; use rig::client::CompletionClient; use rig::completion::Usage; -use std::collections::HashSet; +use std::collections::{HashMap, HashSet}; use std::pin::Pin; -use std::sync::Arc; +use std::sync::{Arc, Mutex, PoisonError, Weak}; +use tokio::sync::mpsc; +use tokio_util::sync::CancellationToken; /// A client-side tool definition supplied with a request. /// @@ -106,7 +112,10 @@ pub(crate) fn scratchpad_usage_event( }) } -/// Rig-native agent wrapper using provider-specific agents. +/// An agent built once from configuration: the provider client, the tools +/// discovered at build time, and the open MCP connections. Its runs are +/// [`Agent`]s; see [`begin_run`](Self::begin_run) and the +/// [`run_context`](crate::run_context) module. /// /// # Tool Execution Paths /// @@ -120,34 +129,25 @@ pub(crate) fn scratchpad_usage_event( /// is wrapped with `FallbackToolExecutor`. It buffers text, detects tool call /// patterns (JSON/XML), and executes via `McpManager::execute_fallback_tool()`. /// This bypasses Rig's tool infrastructure entirely. -pub struct Agent { +pub struct PreparedAgent { pub(crate) inner: ProviderAgent, pub(crate) model: String, pub(crate) max_depth: usize, pub(crate) mcp_manager: Option>, - /// Ollama text-to-tool parsing: when enabled, intercepts text output containing - /// tool calls (JSON/XML) and executes them. Only applies to Ollama provider. - /// See `maybe_wrap_with_fallback()` for the wrapping logic. + /// Whether Ollama text-to-tool parsing is on. pub(crate) fallback_tool_parsing: bool, - /// Cached tool names for fallback parsing (avoids recomputing on each stream). - /// Only populated when `fallback_tool_parsing` is enabled. + /// Tool names for fallback parsing. pub(crate) fallback_tool_names: Vec, /// The `mcp_filter` effective for `fallback_tool_names`. pub(crate) fallback_mcp_filter: Option>, - /// Configured context window size in tokens (from LLM TOML config). - /// Used for usage percentage reporting in streaming events. + /// Configured context window size in tokens (`[agent.llm].context_window`). pub(crate) context_window: Option, - /// Per-agent scratchpad budget for context tracking. - /// Set by orchestration workers (from resolved worker LLM + scratchpad config); - /// `None` for coordinator agents and non-scratchpad use. - pub(crate) scratchpad_budget: Option, - /// Names of client-side (passthrough) tools registered for this agent. - /// When the LLM calls one of these, the streaming layer terminates the - /// stream with `finish_reason: "tool_calls"` so the caller can execute - /// the tool and resume in a follow-up request. + /// Seed context budget: the limits, with no usage recorded. + pub(crate) scratchpad_budget: Option, + /// Names of the client-side (passthrough) tools registered on the agent. pub(crate) client_tool_names: HashSet, - /// Turn-limit nudge state shared with this agent's `TurnNudgeWrapper`. - pub(crate) turn_nudge: Option>, + /// Seed turn-limit tracking: the limits, with no turns completed. + pub(crate) turn_nudge: Option>, /// The HITL config gate. pub(crate) hitl_gate: Option>, /// The `request_approval` tool. @@ -159,25 +159,79 @@ pub struct Agent { pub(crate) invocation_parameters: Option, /// Skills discovered for this agent. pub(crate) skills: Vec, + /// The request headers forwarded when this agent was prepared. + pub(crate) forwarded_headers: ForwardedHeaders, + /// The run slot the tools were built with. + pub(crate) run: Arc, + /// The run in progress. + pub(crate) active: Mutex>, } -impl Agent { +/// Why a prepared agent refused to begin a run. +#[derive(Debug, thiserror::Error)] +pub enum BeginRunError { + #[error(transparent)] + RunInProgress(#[from] RunInProgress), + #[error( + "prepared agent forwards request header `{header}` and this request carries a \ + different value for it; prepare an agent for this request" + )] + ForwardedHeaderDiffers { header: String }, +} + +/// One run of a [`PreparedAgent`]. +pub struct Agent { + prepared: Arc, + lease: Arc, + /// The run's events, for its observer. + events: Mutex>>, +} + +/// A run reads its prepared agent's fields and calls its methods directly, so +/// the orchestrator and the `StreamingAgent` impl need no forwarders. +impl std::ops::Deref for Agent { + type Target = PreparedAgent; + + fn deref(&self) -> &PreparedAgent { + &self.prepared + } +} + +impl std::fmt::Debug for PreparedAgent { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("PreparedAgent") + .field("provider", &self.inner.provider_name()) + .field("model", &self.model) + .field("run", &self.run) + .finish_non_exhaustive() + } +} + +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("prepared", &self.prepared) + .finish() + } +} + +impl PreparedAgent { /// Wire up scratchpad for a single-agent config. Skipped in orchestration /// mode (workers build their own per-worker budget in `create_worker()`) /// and when no accessible MCP tool matches a scratchpad threshold. /// - /// Per-request lifecycle: this runs inside `Agent::new`, which is called - /// fresh per chat request from - /// `aura-web-server::handlers::build_agent_for_request`. The `Agent` (and - /// this `ContextBudget`) is dropped when the request stream ends — there - /// is no cross-request budget state to manage. Request-specific data - /// (user query + chat history) is seeded into the budget at stream-start - /// in `Agent::stream_*_with_timeout`, not here, so this constructor - /// stays free of request-shape parameters. + /// The wrapper and tools composed here reach a run's budget through + /// `run`; the budget returned is what each run starts from (`begin_run` + /// hands every run a fresh copy). Request-specific data (user query + + /// chat history) is seeded into the run's budget at stream-start in + /// `Agent::stream_*_with_timeout`, not here, so this constructor stays + /// free of request-shape parameters. async fn setup_single_agent_scratchpad( config: &mut AgentRuntimeConfig, mcp_manager: Option<&Arc>, - ) -> Result, Box> { + run: &Arc, + ) -> Result, Box> { if config.orchestration_enabled() { return Ok(None); } @@ -253,6 +307,7 @@ impl Agent { context_window, initial_used, token_counter, + run: Arc::clone(run), }) .await?; @@ -287,7 +342,7 @@ impl Agent { Ok(Some(build.budget)) } - /// Create a new agent from configuration with optional additional tools. + /// Build an agent from configuration with optional additional tools. /// /// `additional_tools` registers extra rig tools the agent will execute itself /// (e.g. tools other applications using Aura as a library want to expose). @@ -296,11 +351,17 @@ impl Agent { /// `client_tools` registers passthrough tools — the LLM sees them as callable, /// but the streaming layer terminates the stream when one is invoked so the /// client can execute the tool locally. Pass `None` to disable. - pub async fn new( + /// + /// Nothing here depends on a request: the HITL gate, turn nudge, and + /// scratchpad wrappers composed below reach the run they serve through + /// the agent's [`BoundRun`], filled in by [`begin_run`](Self::begin_run). + pub async fn prepare( config: &AgentRuntimeConfig, additional_tools: Vec>, client_tools: Option>, ) -> Result> { + let run = Arc::new(BoundRun::default()); + // Initialize MCP manager first (shared across all providers) let mcp_manager = if let Some(mcp_config) = &config.mcp { tracing::info!("Initializing MCP tools using dynamic adaptors"); @@ -316,7 +377,8 @@ impl Agent { // hitl_request_approval_tool). let mut config_owned = config.clone(); let agent_scratchpad_budget = - Self::setup_single_agent_scratchpad(&mut config_owned, mcp_manager.as_ref()).await?; + Self::setup_single_agent_scratchpad(&mut config_owned, mcp_manager.as_ref(), &run) + .await?; // HITL gate for single-agent mode. Orchestration workers wire their // own gate in create_worker with per-task AgentScope::Worker; this @@ -332,12 +394,10 @@ impl Agent { .clone() .map(crate::config::SessionId::new), }; - let request_id = config_owned.request_id.clone().unwrap_or_default(); let wrapper = Arc::new(crate::hitl::HitlApprovalWrapper::new( hitl.patterns.clone(), hitl.route.clone(), scope.clone(), - request_id.clone(), config_owned.agent.name.clone(), config_owned.instance_id.clone(), )); @@ -352,7 +412,6 @@ impl Agent { let approval_tool = crate::hitl::RequestApprovalTool::new( hitl.route.clone(), scope, - request_id, config_owned.agent.name.clone(), config_owned.instance_id.clone(), ); @@ -372,22 +431,23 @@ impl Agent { let max_depth = base_depth + scratchpad_bonus; // Turn-limit nudging for single-agent mode; orchestration workers - // wire their own state in `create_worker`. + // wire their own in `create_worker`. The wrapper reads the run's + // counters through the slot; `turn_nudge` here is what each run's + // counters start from. let turn_nudge = if config_owned.orchestration_enabled() { None } else { - crate::turn_nudge::TurnNudgeState::new( + TurnNudgeState::new( config_owned.agent.nudge_last_turn, config_owned.agent.nudge_turns_remaining, max_depth, ) }; - config_owned.turn_nudge = turn_nudge.clone(); - if let Some(ref state) = turn_nudge { + if turn_nudge.is_some() { // First in the vec → transform_output runs last, on the text the // LLM actually sees (after any scratchpad pointer rewrite). let nudge: Arc = - Arc::new(crate::turn_nudge::TurnNudgeWrapper::new(state.clone())); + Arc::new(crate::turn_nudge::TurnNudgeWrapper::new(Arc::clone(&run))); config_owned.tool_wrapper = Some(match config_owned.tool_wrapper.take() { Some(existing) => Arc::new(crate::tool_wrapper::ComposedWrapper::new(vec![ nudge, existing, @@ -541,9 +601,14 @@ impl Agent { } else { builder_state }; - let builder_state = - Self::add_all_tools(builder_state, config, &mcp_manager, additional_tools) - .await?; + let builder_state = Self::add_all_tools( + builder_state, + config, + &mcp_manager, + &run, + additional_tools, + ) + .await?; let agent = builder_state.build(); ProviderAgent::OpenAI(agent) @@ -602,9 +667,14 @@ impl Agent { } else { builder_state }; - let builder_state = - Self::add_all_tools(builder_state, config, &mcp_manager, additional_tools) - .await?; + let builder_state = Self::add_all_tools( + builder_state, + config, + &mcp_manager, + &run, + additional_tools, + ) + .await?; let agent = builder_state.build(); ProviderAgent::Anthropic(agent) @@ -679,9 +749,14 @@ impl Agent { } else { builder_state }; - let builder_state = - Self::add_all_tools(builder_state, config, &mcp_manager, additional_tools) - .await?; + let builder_state = Self::add_all_tools( + builder_state, + config, + &mcp_manager, + &run, + additional_tools, + ) + .await?; let agent = builder_state.build(); ProviderAgent::Bedrock(agent) @@ -730,9 +805,14 @@ impl Agent { } else { builder_state }; - let builder_state = - Self::add_all_tools(builder_state, config, &mcp_manager, additional_tools) - .await?; + let builder_state = Self::add_all_tools( + builder_state, + config, + &mcp_manager, + &run, + additional_tools, + ) + .await?; let agent = builder_state.build(); ProviderAgent::Gemini(agent) @@ -780,9 +860,14 @@ impl Agent { } else { builder_state }; - let builder_state = - Self::add_all_tools(builder_state, config, &mcp_manager, additional_tools) - .await?; + let builder_state = Self::add_all_tools( + builder_state, + config, + &mcp_manager, + &run, + additional_tools, + ) + .await?; let agent = builder_state.build(); ProviderAgent::Ollama(agent) @@ -835,9 +920,14 @@ impl Agent { } else { builder_state }; - let builder_state = - Self::add_all_tools(builder_state, config, &mcp_manager, additional_tools) - .await?; + let builder_state = Self::add_all_tools( + builder_state, + config, + &mcp_manager, + &run, + additional_tools, + ) + .await?; let agent = builder_state.build(); ProviderAgent::OpenRouter(agent) @@ -849,7 +939,7 @@ impl Agent { .map(|tools| tools.iter().map(|t| t.name.clone()).collect()) .unwrap_or_default(); - Ok(Agent { + Ok(PreparedAgent { inner: provider_agent, model: model_name, max_depth, @@ -866,6 +956,107 @@ impl Agent { system_prompt, invocation_parameters: crate::logging::llm_invocation_parameters(&config.llm), skills: config.agent.skills.clone(), + forwarded_headers: config.forwarded_headers.clone(), + run, + active: Mutex::new(Weak::new()), + }) + } + + /// The request headers forwarded when this agent was prepared. + pub fn forwarded_headers(&self) -> &ForwardedHeaders { + &self.forwarded_headers + } + + /// Begin a run of this agent for the request `request_id`, 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 + /// `prepare` computed, and `skill_recorder` recording its skill-tool + /// invocations; `start_run` binds it where the tools read. + /// + /// Refused while another run is alive (see `start_run`), and refused when + /// the request forwards a different value for a header this agent was + /// prepared with: the MCP connections were opened with the preparing + /// request's credentials, so serving this request would run its tool + /// calls as someone else. Prepare an agent for it instead. + pub fn begin_run( + self: &Arc, + request_id: impl Into, + req_headers: Option<&HashMap>, + skill_recorder: Option>, + ) -> Result { + if let Some(header) = self.forwarded_headers.first_difference(req_headers) { + return Err(BeginRunError::ForwardedHeaderDiffers { + header: header.to_owned(), + }); + } + let (run, events) = RunContext::channel_for_agent( + request_id.into(), + CancellationToken::new(), + self.scratchpad_budget.as_ref().map(ContextBudget::fresh), + self.turn_nudge.as_ref().map(|seed| seed.fresh()), + skill_recorder, + ); + Ok(self.start_run(run, Some(events))?) + } + + /// Begin a run of this agent within `parent`, for an orchestration + /// worker or coordinator: the same id, observer and cancellation as the + /// run that owns the orchestration, with fresh tool state of its own and + /// `skill_recorder` recording its skill-tool invocations. The parent's + /// request already carries the forwarded headers this agent was prepared + /// with, so there is nothing to check. + pub fn begin_run_within( + self: &Arc, + parent: &Arc, + skill_recorder: Option>, + ) -> Result { + let run = RunContext::child( + parent, + self.scratchpad_budget.as_ref().map(ContextBudget::fresh), + self.turn_nudge.as_ref().map(|seed| seed.fresh()), + skill_recorder, + ); + self.start_run(run, None) + } + + /// Bind `run` where the tools read it and lease the agent to it. + /// + /// Rig spawns the agent's tool server once, so every tool call arrives on + /// one long-lived task with no way to tell two runs apart; the slot holds + /// one run and this refuses a second while the first is alive. Alive means + /// the lease has a holder: the `Agent`, or a stream it produced that + /// `Agent::hold_run` gave a handle. The slots keep the last run bound + /// after that, which is harmless — nothing calls a tool of a run that + /// has ended — and the next run replaces it. + fn start_run( + self: &Arc, + run: Arc, + events: Option>, + ) -> Result { + let lease = { + 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(), + }); + } + let lease = Arc::new(RunLease::new(Arc::clone(&run))); + *active = Arc::downgrade(&lease); + lease + }; + self.run.bind(Arc::clone(&run)); + if let Some(gate) = &self.hitl_gate { + gate.bind_run(Arc::clone(&run)); + } + if let Some(tool) = &self.hitl_approval_tool { + tool.bind_run(run); + } + Ok(Agent { + prepared: Arc::clone(self), + lease, + events: Mutex::new(events), }) } @@ -904,6 +1095,7 @@ impl Agent { mut builder_state: BuilderState, config: &AgentRuntimeConfig, mcp_manager: &Option>, + run: &Arc, additional_tools: Vec>, ) -> Result, Box> where @@ -1076,40 +1268,39 @@ impl Agent { "Adding scratchpad tools (head, slice, grep, schema, item_schema, get_in, iterate_over, read)" ); let s = &scratchpad.storage; - let b = &scratchpad.budget; - let n = &config.turn_nudge; + let r = &scratchpad.run; builder_state = builder_state .add_tool(NudgedTool::new( - HeadTool::new(s.clone(), b.clone()), - n.clone(), + HeadTool::new(s.clone(), r.clone()), + r.clone(), )) .add_tool(NudgedTool::new( - SliceTool::new(s.clone(), b.clone()), - n.clone(), + SliceTool::new(s.clone(), r.clone()), + r.clone(), )) .add_tool(NudgedTool::new( - GrepTool::new(s.clone(), b.clone()), - n.clone(), + GrepTool::new(s.clone(), r.clone()), + r.clone(), )) .add_tool(NudgedTool::new( - SchemaTool::new(s.clone(), b.clone()), - n.clone(), + SchemaTool::new(s.clone(), r.clone()), + r.clone(), )) .add_tool(NudgedTool::new( - ItemSchemaTool::new(s.clone(), b.clone()), - n.clone(), + ItemSchemaTool::new(s.clone(), r.clone()), + r.clone(), )) .add_tool(NudgedTool::new( - GetInTool::new(s.clone(), b.clone()), - n.clone(), + GetInTool::new(s.clone(), r.clone()), + r.clone(), )) .add_tool(NudgedTool::new( - IterateOverTool::new(s.clone(), b.clone()), - n.clone(), + IterateOverTool::new(s.clone(), r.clone()), + r.clone(), )) .add_tool(NudgedTool::new( - ReadTool::new(s.clone(), b.clone()), - n.clone(), + ReadTool::new(s.clone(), r.clone()), + r.clone(), )); } @@ -1133,7 +1324,7 @@ impl Agent { ); } read_artifact = read_artifact - .with_scratchpad(scratchpad.budget.clone(), scratchpad.storage.clone()); + .with_scratchpad(scratchpad.run.clone(), scratchpad.storage.clone()); } builder_state = builder_state.add_tool(read_artifact); } @@ -1157,9 +1348,7 @@ impl Agent { builder_state = builder_state.add_tools_dyn(additional_tools); } - if let Some(toolset) = - SkillToolset::new(&config.agent.skills, config.skill_recorder.clone()) - { + if let Some(toolset) = SkillToolset::new(&config.agent.skills, Some(Arc::clone(run))) { tracing::info!( "Adding skill tools (load_skill, read_skill_file) with {} skills", config.agent.skills.len(), @@ -1216,10 +1405,177 @@ impl Agent { } } + /// Conditionally wrap stream for Ollama text-to-tool parsing. + /// + /// When `fallback_tool_parsing` is enabled (Ollama config), this wraps the stream + /// with a `FallbackToolExecutor` that: + /// 1. Buffers streamed text content + /// 2. On stream end, parses text for tool call patterns (JSON, XML, etc.) + /// 3. Executes detected tools via MCP and injects results into the stream + /// + /// This handles Ollama models (e.g., qwen3-coder) that output tool calls as text + /// instead of using native tool_call structures. + fn maybe_wrap_with_fallback( + &self, + stream: Pin< + Box> + Send>, + >, + ) -> Pin> + Send>> { + // Early return if fallback not enabled or no tools available + if !self.fallback_tool_parsing || self.fallback_tool_names.is_empty() { + return stream; + } + + // Wrap stream with fallback executor (requires MCP manager for tool execution) + if let Some(mcp_manager) = self.mcp_manager.clone() { + let executor = crate::fallback_tool_stream::FallbackToolExecutor::new( + mcp_manager, + self.fallback_tool_names.clone(), + self.fallback_mcp_filter.clone(), + ); + return executor.wrap_stream(Box::pin(stream)); + } + + stream + } + + /// Get provider information + pub fn get_provider_info(&self) -> (&str, &str) { + (self.inner.provider_name(), &self.model) + } + + /// Cancel all in-flight MCP tool requests for an HTTP request. + /// + /// This sends `notifications/cancelled` to all MCP servers that have + /// in-flight requests, allowing them to abort long-running operations. + /// Call this when a client disconnects or request times out. + /// + /// # Arguments + /// * `http_request_id` - The HTTP request ID whose MCP calls should be cancelled + /// * `reason` - Reason for cancellation (e.g., "client disconnected", "timeout") + /// + /// # Returns + /// Total number of cancellation notifications sent + pub async fn cancel_mcp_requests(&self, http_request_id: &str, reason: &str) -> usize { + if let Some(mcp_manager) = &self.mcp_manager { + mcp_manager + .cancel_all_for_request(http_request_id, reason) + .await + } else { + 0 + } + } + + /// Cancel all in-flight MCP requests and forcefully close connections. + /// + /// This sends `notifications/cancelled` to all MCP servers and then + /// terminates the connections. Use this when the server ignores + /// cancellation requests and keeps sending progress notifications. + /// + /// # Warning + /// After calling this, all MCP clients become unusable. The next request + /// will need to reinitialize them. + /// + /// # Arguments + /// * `http_request_id` - The HTTP request ID whose MCP calls should be cancelled + /// * `reason` - Reason for cancellation + /// + /// # Returns + /// Total number of cancellation notifications sent + pub async fn cancel_and_close_mcp(&self, http_request_id: &str, reason: &str) -> usize { + if let Some(mcp_manager) = &self.mcp_manager { + mcp_manager + .cancel_and_close_all(http_request_id, reason) + .await + } else { + 0 + } + } + + /// Get all available tool names from MCP servers. + /// + /// Returns a list of tool names that can be used for fallback tool execution. + pub fn get_available_tool_names(&self) -> Vec { + self.mcp_manager + .as_ref() + .map(|m| m.get_available_tool_names()) + .unwrap_or_default() + } + + /// Get reference to the MCP manager (if configured). + pub fn mcp_manager(&self) -> Option<&McpManager> { + self.mcp_manager.as_deref() + } + + /// The configured context window size in tokens. + pub fn context_window(&self) -> Option { + self.context_window + } + + /// The assembled system prompt sent to the provider. + pub fn system_prompt(&self) -> &str { + &self.system_prompt + } + + /// MCP tool schemas serialized for the `llm.tools.{i}.tool.json_schema` + /// span attributes. Empty when no MCP manager is configured. + pub fn otel_llm_tools(&self) -> Vec { + self.mcp_manager + .as_ref() + .map(|m| m.tool_schemas_json(None)) + .unwrap_or_default() + } + + /// `llm.invocation_parameters` JSON for OTel spans. + pub fn otel_invocation_parameters(&self) -> Option<&str> { + self.invocation_parameters.as_deref() + } +} + +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.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 + /// prepared half across turns call that and [`PreparedAgent::begin_run`] + /// themselves. + pub async fn new( + config: &AgentRuntimeConfig, + additional_tools: Vec>, + client_tools: Option>, + ) -> Result> { + let prepared = + 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(), + Some(&req_headers), + config.skill_recorder.clone(), + )?) + } + + /// The run. + pub fn run(&self) -> &Arc { + 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() + } + /// Process a query with the agent (no chat history). /// /// Uses the streaming pipeline internally and collects the result. - #[tracing::instrument(name = "agent.prompt", skip(self), fields(model = %self.model))] + #[tracing::instrument(name = "agent.prompt", skip(self), fields(model = %self.prepared.model))] pub async fn prompt( &self, query: &str, @@ -1228,10 +1584,10 @@ impl Agent { let span = tracing::Span::current(); record_input_attributes( &span, - self.inner.provider_name(), - &self.model, + self.prepared.inner.provider_name(), + &self.prepared.model, query, - &self.system_prompt, + &self.prepared.system_prompt, ); self.record_llm_call_attributes(&span); @@ -1244,7 +1600,7 @@ impl Agent { /// Process a chat query with conversation history. /// /// Uses the streaming pipeline internally and collects the result. - #[tracing::instrument(name = "agent.chat", skip(self, chat_history), fields(model = %self.model, history_len = chat_history.len()))] + #[tracing::instrument(name = "agent.chat", skip(self, chat_history), fields(model = %self.prepared.model, history_len = chat_history.len()))] pub async fn chat( &self, query: &str, @@ -1254,10 +1610,10 @@ impl Agent { let span = tracing::Span::current(); record_input_attributes( &span, - self.inner.provider_name(), - &self.model, + self.prepared.inner.provider_name(), + &self.prepared.model, query, - &self.system_prompt, + &self.prepared.system_prompt, ); self.record_llm_call_attributes(&span); @@ -1269,10 +1625,10 @@ impl Agent { /// Record invocation parameters and the advertised tool schemas on a span. fn record_llm_call_attributes(&self, span: &tracing::Span) { - if let Some(params) = &self.invocation_parameters { + if let Some(params) = &self.prepared.invocation_parameters { crate::logging::set_llm_invocation_parameters(span, params); } - let tools = self.otel_llm_tools(); + let tools = self.prepared.otel_llm_tools(); if !tools.is_empty() { crate::logging::set_llm_tools(span, &tools); } @@ -1328,8 +1684,15 @@ impl Agent { &self, query: &str, ) -> Pin> + Send>> { - let stream = self.inner.stream_prompt(query, self.max_depth).await; - self.maybe_wrap_with_fallback(self.count_turns(stream)) + let stream = self + .prepared + .inner + .stream_prompt(query, self.prepared.max_depth) + .await; + self.hold_run( + self.prepared + .maybe_wrap_with_fallback(self.count_turns(stream)), + ) } /// Stream a chat query with conversation history - returns true streaming response with multi-turn tool support @@ -1339,32 +1702,66 @@ impl Agent { chat_history: Vec, ) -> Pin> + Send>> { let stream = self + .prepared .inner - .stream_chat(query, chat_history, self.max_depth) + .stream_chat(query, chat_history, self.prepared.max_depth) .await; - self.maybe_wrap_with_fallback(self.count_turns(stream)) + self.hold_run( + self.prepared + .maybe_wrap_with_fallback(self.count_turns(stream)), + ) } /// Stream a chat with explicit max_depth override. /// - /// Unlike `stream_chat()` which uses `self.max_depth`, this allows callers - /// to specify depth. Used by orchestration phases that need tighter bounds. + /// Unlike `stream_chat()` which uses the prepared agent's `max_depth`, + /// this allows callers to specify depth. Used by orchestration phases + /// that need tighter bounds. #[tracing::instrument(name = "agent.stream_chat", skip(self, chat_history), - fields(model = %self.model, history_len = chat_history.len(), max_depth))] + fields(model = %self.prepared.model, history_len = chat_history.len(), max_depth))] pub async fn stream_chat_with_depth( &self, query: &str, chat_history: Vec, max_depth: usize, ) -> Pin> + Send>> { - let stream = self.inner.stream_chat(query, chat_history, max_depth).await; - self.maybe_wrap_with_fallback(self.count_turns(stream)) + let stream = self + .prepared + .inner + .stream_chat(query, chat_history, max_depth) + .await; + self.hold_run( + self.prepared + .maybe_wrap_with_fallback(self.count_turns(stream)), + ) + } + + /// Keep this run leased for as long as `stream` is alive, so a stream + /// that outlives the `Agent` it came from keeps the prepared agent from + /// beginning another run until the stream itself is dropped. + fn hold_run( + &self, + stream: Pin< + Box> + Send>, + >, + ) -> Pin> + Send>> { + let lease = Arc::clone(&self.lease); + Box::pin(stream.map(move |item| { + let _leased = &lease; + item + })) } /// Count completed turns into the turn-nudge state. Rig yields exactly /// one `StreamItem::TurnUsage` per turn, and a turn's tool calls execute /// before its `TurnUsage` arrives, so mid-turn the current turn is /// `turns_completed + 1`. No-op when nudging is disabled. + /// + /// The count restarts at every stream start: rig's depth limit is per + /// stream, so a run that streams more than once (`prompt` then `chat`) + /// is nudged against each stream's limit, not the run's total. The + /// counters are still the run's own; `begin_run` is what keeps one run's + /// count from ever reaching another. fn count_turns( &self, stream: Pin< @@ -1373,7 +1770,7 @@ impl Agent { ) -> Pin> + Send>> { use futures::StreamExt; - let Some(state) = self.turn_nudge.clone() else { + let Some(state) = self.lease.run().turn_nudge().cloned() else { return stream; }; state.reset(); @@ -1395,7 +1792,7 @@ impl Agent { Box> + Send>, >, ) -> Pin> + Send>> { - let Some(budget) = self.scratchpad_budget.clone() else { + let Some(budget) = self.scratchpad_budget().cloned() else { return stream; }; let tail = futures::stream::once(async move { @@ -1405,40 +1802,6 @@ impl Agent { Box::pin(stream.chain(tail)) } - /// Conditionally wrap stream for Ollama text-to-tool parsing. - /// - /// When `fallback_tool_parsing` is enabled (Ollama config), this wraps the stream - /// with a `FallbackToolExecutor` that: - /// 1. Buffers streamed text content - /// 2. On stream end, parses text for tool call patterns (JSON, XML, etc.) - /// 3. Executes detected tools via MCP and injects results into the stream - /// - /// This handles Ollama models (e.g., qwen3-coder) that output tool calls as text - /// instead of using native tool_call structures. - fn maybe_wrap_with_fallback( - &self, - stream: Pin< - Box> + Send>, - >, - ) -> Pin> + Send>> { - // Early return if fallback not enabled or no tools available - if !self.fallback_tool_parsing || self.fallback_tool_names.is_empty() { - return stream; - } - - // Wrap stream with fallback executor (requires MCP manager for tool execution) - if let Some(mcp_manager) = self.mcp_manager.clone() { - let executor = crate::fallback_tool_stream::FallbackToolExecutor::new( - mcp_manager, - self.fallback_tool_names.clone(), - self.fallback_mcp_filter.clone(), - ); - return executor.wrap_stream(Box::pin(stream)); - } - - stream - } - /// Stream a query with timeout and cancellation support. /// /// Returns the run: its stream, the token that cancels it, and its usage. @@ -1464,19 +1827,23 @@ impl Agent { request_id: &str, ) -> crate::streaming::AgentRun { self.seed_scratchpad_request_input(query, &[]); - self.inner + self.prepared + .inner .stream_prompt_with_timeout( query, - self.max_depth, + self.prepared.max_depth, options, request_id, - self.scratchpad_budget.clone(), - self.client_tool_names.clone(), + self.scratchpad_budget().cloned(), + self.prepared.client_tool_names.clone(), ) .await .map_stream(|stream| { - self.append_scratchpad_usage( - self.maybe_wrap_with_fallback(self.count_turns(stream)), + self.hold_run( + self.append_scratchpad_usage( + self.prepared + .maybe_wrap_with_fallback(self.count_turns(stream)), + ), ) }) } @@ -1500,20 +1867,24 @@ impl Agent { request_id: &str, ) -> crate::streaming::AgentRun { self.seed_scratchpad_request_input(query, &chat_history); - self.inner + self.prepared + .inner .stream_chat_with_timeout( query, chat_history, - self.max_depth, + self.prepared.max_depth, options, request_id, - self.scratchpad_budget.clone(), - self.client_tool_names.clone(), + self.scratchpad_budget().cloned(), + self.prepared.client_tool_names.clone(), ) .await .map_stream(|stream| { - self.append_scratchpad_usage( - self.maybe_wrap_with_fallback(self.count_turns(stream)), + self.hold_run( + self.append_scratchpad_usage( + self.prepared + .maybe_wrap_with_fallback(self.count_turns(stream)), + ), ) }) } @@ -1530,7 +1901,7 @@ impl Agent { query: &str, chat_history: &[rig::completion::Message], ) { - let Some(budget) = &self.scratchpad_budget else { + let Some(budget) = self.scratchpad_budget() else { return; }; let query_tokens = budget.count_tokens(query); @@ -1540,88 +1911,6 @@ impl Agent { .sum(); budget.record_usage(query_tokens + history_tokens); } - - /// Get provider information - pub fn get_provider_info(&self) -> (&str, &str) { - (self.inner.provider_name(), &self.model) - } - - /// Cancel all in-flight MCP tool requests for an HTTP request. - /// - /// This sends `notifications/cancelled` to all MCP servers that have - /// in-flight requests, allowing them to abort long-running operations. - /// Call this when a client disconnects or request times out. - /// - /// # Arguments - /// * `http_request_id` - The HTTP request ID whose MCP calls should be cancelled - /// * `reason` - Reason for cancellation (e.g., "client disconnected", "timeout") - /// - /// # Returns - /// Total number of cancellation notifications sent - pub async fn cancel_mcp_requests(&self, http_request_id: &str, reason: &str) -> usize { - if let Some(mcp_manager) = &self.mcp_manager { - mcp_manager - .cancel_all_for_request(http_request_id, reason) - .await - } else { - 0 - } - } - - /// Cancel all in-flight MCP requests and forcefully close connections. - /// - /// This sends `notifications/cancelled` to all MCP servers and then - /// terminates the connections. Use this when the server ignores - /// cancellation requests and keeps sending progress notifications. - /// - /// # Warning - /// After calling this, all MCP clients become unusable. The next request - /// will need to reinitialize them. - /// - /// # Arguments - /// * `http_request_id` - The HTTP request ID whose MCP calls should be cancelled - /// * `reason` - Reason for cancellation - /// - /// # Returns - /// Total number of cancellation notifications sent - pub async fn cancel_and_close_mcp(&self, http_request_id: &str, reason: &str) -> usize { - if let Some(mcp_manager) = &self.mcp_manager { - mcp_manager - .cancel_and_close_all(http_request_id, reason) - .await - } else { - 0 - } - } - - /// Get all available tool names from MCP servers. - /// - /// Returns a list of tool names that can be used for fallback tool execution. - pub fn get_available_tool_names(&self) -> Vec { - self.mcp_manager - .as_ref() - .map(|m| m.get_available_tool_names()) - .unwrap_or_default() - } - - /// Get reference to the MCP manager (if configured). - pub fn mcp_manager(&self) -> Option<&McpManager> { - self.mcp_manager.as_deref() - } - - /// MCP tool schemas serialized for the `llm.tools.{i}.tool.json_schema` - /// span attributes. Empty when no MCP manager is configured. - pub fn otel_llm_tools(&self) -> Vec { - self.mcp_manager - .as_ref() - .map(|m| m.tool_schemas_json(None)) - .unwrap_or_default() - } - - /// `llm.invocation_parameters` JSON for OTel spans. - pub fn otel_invocation_parameters(&self) -> Option<&str> { - self.invocation_parameters.as_deref() - } } // --------------------------------------------------------------------------- @@ -1692,7 +1981,7 @@ use async_trait::async_trait; #[async_trait] impl StreamingAgent for Agent { fn get_provider_info(&self) -> (&str, &str) { - Agent::get_provider_info(self) + self.prepared.get_provider_info() } async fn stream( @@ -1702,75 +1991,95 @@ impl StreamingAgent for Agent { options: crate::streaming::RunOptions, request_id: &str, ) -> crate::streaming::AgentRun { - // The run is built on the token it will stop on, so the work that awaits - // cancellation finds it there from the start. - let (timeout, cancel) = options.into_parts(); - let cancel = cancel.unwrap_or_default(); - let (run, run_events) = - crate::run_context::RunContext::channel_on(request_id, cancel.clone()); - let options = crate::streaming::RunOptions::on_token(timeout, cancel); - - // The gate and the approval tool are built with the agent, before any - // run exists, and rig runs tools on its own server task. So bind the run - // where a tool call can still find it. - if let Some(gate) = &self.hitl_gate { - gate.bind_run(std::sync::Arc::clone(&run)); + // The run began at `begin_run`, so its context is what this streams + // under; the id is the run's. + 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 Some(tool) = &self.hitl_approval_tool { - tool.bind_run(std::sync::Arc::clone(&run)); + let run_id = run.id().to_string(); + + // The run's token is what its work watches and what dropping the + // handle cancels. A caller that handed in a token of its own gets the + // run stopped when that token is, and its token left alone: the run's + // predates the options, so the link is watched rather than built in. + let (timeout, cancel) = options.into_parts(); + if let Some(parent) = cancel { + let child = run.cancel_token().clone(); + tokio::spawn(async move { + tokio::select! { + () = parent.cancelled() => child.cancel(), + () = child.cancelled() => {} + } + }); } - if let Some(mcp_manager) = &self.mcp_manager { + let options = crate::streaming::RunOptions::on_token(timeout, run.cancel_token().clone()); + + // 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.prepared.mcp_manager { mcp_manager - .bind_call( - std::sync::Arc::clone(&run), - aura_events::AgentContext::single_agent(), - ) + .bind_call(Arc::clone(&run), aura_events::AgentContext::single_agent()) .await; } let started = if chat_history.is_empty() { - self.stream_prompt_with_timeout(query, options, request_id) + self.stream_prompt_with_timeout(query, options, &run_id) .await } else { - self.stream_chat_with_timeout(query, chat_history, options, request_id) + self.stream_chat_with_timeout(query, chat_history, options, &run_id) .await }; - started - .map_stream(move |stream| { - Box::pin(crate::run_context::scope_stream( - std::sync::Arc::clone(&run), - Box::pin(crate::streaming::tee_content( - run, - aura_events::AgentContext::single_agent(), - stream, - )), - )) - }) - .observed_by(run_events) + let started = started.map_stream(move |stream| { + Box::pin(crate::run_context::scope_stream( + Arc::clone(&run), + Box::pin(crate::streaming::tee_content( + run, + aura_events::AgentContext::single_agent(), + stream, + )), + )) + }); + // One observer per run: the first stream hands the receiver over, + // a later stream of the same run has no second one to give. + match self + .events + .lock() + .unwrap_or_else(PoisonError::into_inner) + .take() + { + Some(events) => started.observed_by(events), + None => started, + } } async fn cancel_and_close_mcp(&self, request_id: &str, reason: &str) -> usize { - Agent::cancel_and_close_mcp(self, request_id, reason).await + self.prepared.cancel_and_close_mcp(request_id, reason).await } fn context_window(&self) -> Option { - self.context_window + self.prepared.context_window } fn mcp_server_status(&self) -> Vec { - self.mcp_manager + self.prepared + .mcp_manager .as_ref() .map(|m| m.server_status_snapshot()) .unwrap_or_default() } fn skills(&self) -> &[aura_config::SkillConfig] { - &self.skills + &self.prepared.skills } fn system_prompt(&self) -> Option<&str> { - Some(&self.system_prompt) + Some(&self.prepared.system_prompt) } } @@ -2151,7 +2460,7 @@ mod tests { }; let manager = Some(Arc::new(manager_serving_all_transports(server).await)); let state = BuilderState::Initial(rig::agent::AgentBuilder::new(UnpromptedModel)); - Agent::add_all_tools(state, &config, &manager, Vec::new()) + PreparedAgent::add_all_tools(state, &config, &manager, &Arc::default(), Vec::new()) .await .expect("composition succeeds") .build() @@ -2227,7 +2536,7 @@ mod tests { /// A manager offering one HTTP-streamable tool, keyed by `namespace`. /// HTTP rather than stdio so the transport's fail-closed check cannot /// be mistaken for the gate's decision. - async fn manager_serving( + pub(super) async fn manager_serving( server: &RecordingMcpServer, namespace: &str, tool: &str, @@ -2263,7 +2572,6 @@ mod tests { /// because rig calls the tool on its server task and no scope reaches /// there. fn gated_config( - request_id: &str, pattern: &str, run: Arc, ) -> AgentRuntimeConfig { @@ -2274,7 +2582,6 @@ mod tests { timeout: Duration::from_millis(50), }), AgentScope::Single { session_id: None }, - request_id.to_owned(), "test-agent".to_owned(), "test-instance-id".to_owned(), ); @@ -2287,16 +2594,15 @@ mod tests { async fn compose_gated_agent( server: &RecordingMcpServer, - request_id: &str, pattern: &str, namespace: &str, tool: &str, run: Arc, ) -> rig::agent::Agent { - let config = gated_config(request_id, pattern, run); + let config = gated_config(pattern, run); let manager = Some(Arc::new(manager_serving(server, namespace, tool).await)); let state = BuilderState::Initial(rig::agent::AgentBuilder::new(UnpromptedModel)); - Agent::add_all_tools(state, &config, &manager, Vec::new()) + PreparedAgent::add_all_tools(state, &config, &manager, &Arc::default(), Vec::new()) .await .expect("composition succeeds") .build() @@ -2310,9 +2616,7 @@ mod tests { let request_id = "req_ns_gating_match"; let server = RecordingMcpServer::start().await; let (run, mut rx) = crate::run_context::RunContext::channel(request_id); - let agent = - compose_gated_agent(&server, request_id, "github:*", "github", "list_repos", run) - .await; + let agent = compose_gated_agent(&server, "github:*", "github", "list_repos", run).await; // The parked approval expires unanswered; the call's own outcome is // not what this test is about. @@ -2344,9 +2648,7 @@ mod tests { let request_id = "req_ns_gating_miss"; let server = RecordingMcpServer::start().await; let (run, mut rx) = crate::run_context::RunContext::channel(request_id); - let agent = - compose_gated_agent(&server, request_id, "github:*", "k8s", "list_repos", run) - .await; + let agent = compose_gated_agent(&server, "github:*", "k8s", "list_repos", run).await; agent .tool_server_handle @@ -2382,10 +2684,16 @@ mod tests { ..Default::default() }; let state = BuilderState::Initial(rig::agent::AgentBuilder::new(UnpromptedModel)); - let agent = Agent::add_all_tools(state, &config, &Some(Arc::new(manager)), Vec::new()) - .await - .expect("composition succeeds") - .build(); + let agent = PreparedAgent::add_all_tools( + state, + &config, + &Some(Arc::new(manager)), + &Arc::default(), + Vec::new(), + ) + .await + .expect("composition succeeds") + .build(); agent .tool_server_handle @@ -2434,10 +2742,16 @@ mod tests { async fn compose_over(manager: McpManager) -> rig::agent::Agent { let config = AgentRuntimeConfig::default(); let state = BuilderState::Initial(rig::agent::AgentBuilder::new(UnpromptedModel)); - Agent::add_all_tools(state, &config, &Some(Arc::new(manager)), Vec::new()) - .await - .expect("composition succeeds") - .build() + PreparedAgent::add_all_tools( + state, + &config, + &Some(Arc::new(manager)), + &Arc::default(), + Vec::new(), + ) + .await + .expect("composition succeeds") + .build() } #[tokio::test] @@ -2655,8 +2969,11 @@ mod tests { let sp_budget = scratchpad::ContextBudget::new(128_000, 0.20, 0, Arc::new(counter)); let scratchpad_tools = HashMap::from([("big_tool".to_string(), 10_usize)]); - let scratchpad: Arc = - Arc::new(ScratchpadWrapper::new(scratchpad_tools, storage, sp_budget)); + let scratchpad: Arc = Arc::new(ScratchpadWrapper::new( + scratchpad_tools, + storage, + Arc::new(BoundRun::pinned_budget(sp_budget)), + )); let recording = Arc::new(RecordingWrapper::default()); let recording_dyn: Arc = recording.clone(); @@ -2709,4 +3026,295 @@ mod tests { &result.output[..result.output.len().min(120)] ); } + + /// The split the run slot exists for: a prepared agent is built once and + /// serves runs in turn, each owning state of its own, and refuses to + /// serve two at once. + mod prepared_runs { + use super::*; + use crate::orchestration::{ScriptedCompletionModel, ScriptedTurn}; + + /// 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 + /// has state to own. + fn prepared() -> Arc { + prepared_forwarding(ForwardedHeaders::default()) + } + + /// Like [`prepared`], bound to the request headers `forwarded`. + fn prepared_forwarding(forwarded_headers: ForwardedHeaders) -> Arc { + let model = ScriptedCompletionModel::new(vec![ScriptedTurn::text("done")]); + prepared_over( + ProviderAgent::Scripted(rig::agent::AgentBuilder::new(model).build()), + forwarded_headers, + None, + ) + } + + /// Like [`prepared`], over `inner` with its HITL gate `hitl_gate`. + fn prepared_over( + inner: ProviderAgent, + forwarded_headers: ForwardedHeaders, + hitl_gate: Option>, + ) -> Arc { + Arc::new(PreparedAgent { + inner, + model: "scripted".to_owned(), + max_depth: 1, + mcp_manager: None, + fallback_tool_parsing: false, + fallback_tool_names: Vec::new(), + fallback_mcp_filter: None, + context_window: None, + scratchpad_budget: Some(budget()), + client_tool_names: HashSet::new(), + turn_nudge: TurnNudgeState::new(true, None, 1), + system_prompt: String::new(), + invocation_parameters: None, + skills: Vec::new(), + forwarded_headers, + hitl_gate, + hitl_approval_tool: None, + run: Arc::new(BoundRun::default()), + active: Mutex::new(Weak::new()), + }) + } + + /// A web request's approvals are released by its request id when the + /// request ends (the server's `RequestResourceGuard`). The gate stamps + /// the id of the run `begin_run` bound, so each run of one prepared + /// agent parks under its own request id, and ending one request + /// releases only that request's approvals. + #[tokio::test] + async fn each_run_parks_its_approvals_under_its_own_request_id() { + use std::time::Duration; + + use crate::hitl::{AgentScope, DecisionRoute, HitlApprovalWrapper, PendingApprovals}; + use crate::mcp::client::tests::RecordingMcpServer; + use aura_events::agent::AgentEventPayload; + + let server = RecordingMcpServer::start().await; + let registry = PendingApprovals::new(); + let gate = Arc::new(HitlApprovalWrapper::new( + Arc::from(["github:*".into()]), + Arc::new(DecisionRoute::Conversational { + registry: registry.clone(), + timeout: Duration::from_secs(60), + }), + AgentScope::Single { session_id: None }, + "test-agent".to_owned(), + "test-instance-id".to_owned(), + )); + let config = AgentRuntimeConfig { + tool_wrapper: Some(Arc::clone(&gate) as Arc), + ..Default::default() + }; + let manager = Some(Arc::new( + super::namespace_gating::manager_serving(&server, "github", "list_repos").await, + )); + let state = BuilderState::Initial(rig::agent::AgentBuilder::new( + ScriptedCompletionModel::new(Vec::new()), + )); + let inner = + PreparedAgent::add_all_tools(state, &config, &manager, &Arc::default(), Vec::new()) + .await + .expect("composition succeeds") + .build(); + let prepared = prepared_over( + ProviderAgent::Scripted(inner), + ForwardedHeaders::default(), + Some(gate), + ); + + // Begin a run, make the gated call, and return once the call has + // parked (the conversational route registers before it announces). + async fn park( + prepared: &Arc, + request_id: &str, + ) -> (Agent, tokio::task::JoinHandle<()>) { + let agent = prepared.begin_run(request_id, None, None).unwrap(); + let mut events = agent.events.lock().unwrap().take().unwrap(); + let call = tokio::spawn({ + let prepared = Arc::clone(prepared); + async move { + let _ = prepared.inner.call_tool("list_repos", "{}").await; + } + }); + let raised = tokio::time::timeout(Duration::from_secs(5), events.recv()) + .await + .expect("the gated call raises an approval") + .expect("the run's channel is open"); + assert!( + matches!(raised.payload, AgentEventPayload::ApprovalRequested(_)), + "got: {:?}", + raised.payload, + ); + (agent, call) + } + let released = |call: tokio::task::JoinHandle<()>| async move { + tokio::time::timeout(Duration::from_secs(5), call) + .await + .is_ok_and(|joined| joined.is_ok()) + }; + + let (first, call) = park(&prepared, "req_a").await; + registry.cancel_request_local("req_a"); + assert!( + released(call).await, + "ending req_a releases the approval it parked" + ); + drop(first); + + let (_second, call) = park(&prepared, "req_b").await; + registry.cancel_request_local("req_a"); + 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"); + assert!( + released(call).await, + "ending req_b releases the approval it parked" + ); + assert!( + server.tool_calls().is_empty(), + "a cancelled approval never runs the tool", + ); + } + + #[tokio::test] + async fn a_prepared_agent_serves_one_run_at_a_time() { + let prepared = prepared(); + let first = prepared + .begin_run("req_a", None, None) + .expect("a fresh agent has no run"); + + let refused = prepared + .begin_run("req_b", None, None) + .expect_err("the slot is taken"); + assert!( + matches!(refused, BeginRunError::RunInProgress(RunInProgress { ref active }) if active == "req_a"), + "got: {refused:?}", + ); + + drop(first); + let second = prepared + .begin_run("req_b", None, None) + .expect("the slot is free again"); + assert_eq!(second.request_id(), "req_b"); + } + + /// Each run starts from the prepared seeds, and the slot the tools + /// hold resolves to that run and nothing else. + #[tokio::test] + async fn each_run_owns_fresh_state_that_the_tools_reach_through_the_slot() { + let prepared = prepared(); + let seed = prepared.scratchpad_budget.as_ref().unwrap(); + + let first = prepared.begin_run("req_a", None, None).unwrap(); + assert_eq!( + prepared + .run + .get() + .map(|run| run.id().to_string()) + .as_deref(), + Some("req_a") + ); + prepared + .run + .scratchpad_budget() + .expect("the run's budget resolves") + .record_intercepted(9); + assert_eq!( + first.scratchpad_budget().unwrap().scratchpad_usage().0, + 9, + "the slot hands out the run's own budget", + ); + assert_eq!(seed.scratchpad_usage().0, 0, "the seed stays unused"); + let nudge = prepared.run.turn_nudge().expect("the run's nudge resolves"); + nudge.record_turn_completed(); + assert!( + nudge.nudge_message().is_some(), + "one completed turn puts the first run on its final turn", + ); + + drop(first); + + let second = prepared.begin_run("req_b", None, None).unwrap(); + assert_eq!( + prepared + .run + .get() + .map(|run| run.id().to_string()) + .as_deref(), + Some("req_b") + ); + assert_eq!( + second.scratchpad_budget().unwrap().scratchpad_usage().0, + 0, + "a new run's budget starts over", + ); + assert!( + prepared.run.turn_nudge().unwrap().nudge_message().is_none(), + "a new run's turn count starts over", + ); + } + + /// The MCP connections were opened with the preparing request's + /// forwarded headers, so a request forwarding different values gets + /// a fresh agent rather than someone else's credentials. + #[tokio::test] + async fn a_request_forwarding_different_headers_is_refused() { + let token = + |value: &str| HashMap::from([("x-user-token".to_owned(), value.to_owned())]); + let prepared = prepared_forwarding(ForwardedHeaders::of( + ["x-user-token".to_owned()], + Some(&token("alice")), + )); + + let refused = prepared + .begin_run("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(), + "a request carrying no token is not alice's either", + ); + + let served = prepared + .begin_run("req_alice", Some(&token("alice")), None) + .expect("the same credentials are served"); + assert_eq!(served.request_id(), "req_alice"); + } + + /// A stream still driving tools after its `Agent` is dropped keeps + /// the run bound; the slot frees only once the stream is gone too. + #[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 stream = agent.stream_prompt("hello").await; + drop(agent); + + assert_eq!( + prepared + .run + .get() + .map(|run| run.id().to_string()) + .as_deref(), + Some("req_a"), + "the stream holds the run", + ); + assert!(prepared.begin_run("req_b", None, None).is_err()); + + drop(stream); + prepared + .begin_run("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 ed155833d..3d9550074 100644 --- a/crates/aura/src/config.rs +++ b/crates/aura/src/config.rs @@ -6,6 +6,7 @@ //! persistence handles, shared chat history, scratchpad runtime state, and the //! session id. +use crate::forwarded_headers::ForwardedHeaders; use crate::hitl::HitlRuntime; use crate::scratchpad::ScratchpadToolsConfig; use crate::tool_wrapper::{ToolCallContext, ToolWrapper}; @@ -123,9 +124,6 @@ pub struct AgentRuntimeConfig { /// `Some` when scratchpad is wired up for this agent or worker. pub scratchpad_tools_config: Option, - /// Shared turn-limit nudge state for this agent's tool calls. - pub turn_nudge: Option>, - /// Shared decision state for worker `submit_result` tool. /// When set, workers get the `submit_result` tool for structured output. pub orchestration_submit_result: Option, @@ -135,25 +133,23 @@ pub struct AgentRuntimeConfig { /// `None` disables approval gating. pub hitl: Option, - /// Request id (`req_…`) for this build, used to stamp HITL approval requests - /// and route their SSE events. Threaded from the web server so the - /// single-agent and orchestration paths share one value. + /// Request id (`req_…`) of the request this build serves. pub request_id: Option, + /// The request headers this build forwards; see [`ForwardedHeaders`]. + pub forwarded_headers: ForwardedHeaders, + /// Computed instance UUID for this agent, derived from agent config and /// host identity. Threaded into HITL approval requests so webhook /// receivers can identify which instance raised each approval. pub instance_id: String, - /// The `request_approval` tool, pre-built with the appropriate - /// [`AgentScope`]. Orchestration workers set this in `create_worker` with - /// `AgentScope::Worker`; single-agent mode sets it in `Agent::new` with - /// `AgentScope::Single`. `None` when `[hitl]` is not configured. + /// The `request_approval` tool, pre-built with its [`AgentScope`]. /// /// [`AgentScope`]: crate::hitl::AgentScope pub hitl_request_approval_tool: Option, - /// Recorder for this session's skill-tool invocations. + /// Recorder for this request's skill-tool invocations. pub skill_recorder: Option>, } @@ -177,10 +173,10 @@ impl Clone for AgentRuntimeConfig { orchestration_persistence: self.orchestration_persistence.clone(), session_id: self.session_id.clone(), scratchpad_tools_config: self.scratchpad_tools_config.clone(), - turn_nudge: self.turn_nudge.clone(), orchestration_submit_result: self.orchestration_submit_result.clone(), hitl: self.hitl.clone(), request_id: self.request_id.clone(), + forwarded_headers: self.forwarded_headers.clone(), instance_id: self.instance_id.clone(), hitl_request_approval_tool: self.hitl_request_approval_tool.clone(), skill_recorder: self.skill_recorder.clone(), @@ -216,7 +212,6 @@ impl std::fmt::Debug for AgentRuntimeConfig { .map(|_| ""), ) .field("session_id", &self.session_id) - .field("turn_nudge", &self.turn_nudge.as_ref().map(|_| "")) .field( "orchestration_submit_result", &self @@ -226,6 +221,7 @@ impl std::fmt::Debug for AgentRuntimeConfig { ) .field("hitl", &self.hitl.as_ref().map(|_| "")) .field("request_id", &self.request_id) + .field("forwarded_headers", &self.forwarded_headers) .field("instance_id", &self.instance_id) .field( "hitl_request_approval_tool", diff --git a/crates/aura/src/forwarded_headers.rs b/crates/aura/src/forwarded_headers.rs new file mode 100644 index 000000000..d32e29d26 --- /dev/null +++ b/crates/aura/src/forwarded_headers.rs @@ -0,0 +1,235 @@ +//! The request headers a prepared agent forwards. +//! +//! `headers_from_request` mappings copy inbound request headers onto the MCP +//! servers and the HITL webhook route when an agent is prepared, and the MCP +//! connections are opened with them. A prepared agent is therefore bound to +//! the request that prepared it as far as those headers go, and can only serve +//! another request that forwards the same values. + +use std::collections::{BTreeMap, BTreeSet, HashMap}; + +use aura_config::{Config, DecisionRouteConfig, McpServerConfig}; + +/// The inbound request headers `headers_from_request` mappings read, with the +/// values the preparing request carried. +#[derive(Clone, Default, PartialEq, Eq)] +pub struct ForwardedHeaders { + /// Lowercased header name → the request's value for it. + values: BTreeMap>, +} + +impl ForwardedHeaders { + /// The headers `req_headers` forwards under `config`: every inbound name a + /// `headers_from_request` mapping reads, on an MCP server or the HITL + /// webhook route, looked up case-insensitively. + pub fn resolve(config: &Config, req_headers: Option<&HashMap>) -> Self { + let mut names = BTreeSet::new(); + if let Some(mcp) = &config.mcp { + for server in mcp.servers.values() { + match server { + McpServerConfig::HttpStreamable { + headers_from_request, + .. + } + | McpServerConfig::Sse { + headers_from_request, + .. + } => names.extend(headers_from_request.values().map(|n| n.to_lowercase())), + McpServerConfig::Stdio { .. } => {} + } + } + } + if let Some(DecisionRouteConfig::Webhook { + headers_from_request, + .. + }) = config.hitl.as_ref().map(|hitl| &hitl.route) + { + names.extend(headers_from_request.values().map(|n| n.to_lowercase())); + } + Self::of(names, req_headers) + } + + /// `req_headers` seen through the names this forwards. + pub fn project(&self, req_headers: Option<&HashMap>) -> Self { + Self::of(self.values.keys().cloned(), req_headers) + } + + /// `req_headers` seen through `names`, which are lowercase: each name with + /// the value the request carries under it in any letter case, or `None` + /// when the request does not carry it. + pub(crate) fn of( + names: impl IntoIterator, + req_headers: Option<&HashMap>, + ) -> Self { + let values = names + .into_iter() + .map(|name| { + let value = req_headers.and_then(|headers| { + headers + .iter() + .find(|(k, _)| k.to_lowercase() == name) + .map(|(_, v)| v.clone()) + }); + (name, value) + }) + .collect(); + Self { values } + } + + /// The forwarded headers as a request carries them, omitting the absent + /// ones. + pub fn as_request(&self) -> HashMap { + self.values + .iter() + .filter_map(|(name, value)| value.clone().map(|v| (name.clone(), v))) + .collect() + } + + /// The first header, by name, whose value differs between `self` and + /// `req_headers` seen through the same names. + pub fn first_difference(&self, req_headers: Option<&HashMap>) -> Option<&str> { + let theirs = self.project(req_headers); + self.values + .iter() + .find(|(name, value)| theirs.values.get(*name) != Some(value)) + .map(|(name, _)| name.as_str()) + } +} + +/// Names only: the values are credentials. +impl std::fmt::Debug for ForwardedHeaders { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + let mut map = f.debug_map(); + for (name, value) in &self.values { + map.entry( + name, + &value.as_ref().map(|_| "").unwrap_or(""), + ); + } + map.finish() + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn config(mcp_mappings: &[(&str, &str)], hitl_mappings: &[(&str, &str)]) -> Config { + let mut toml = String::from( + r#" +[agent] +name = "t" +system_prompt = "t" + +[agent.llm] +provider = "openai" +model = "gpt-4o" +api_key = "k" +"#, + ); + if !mcp_mappings.is_empty() { + toml.push_str("\n[mcp.servers.s]\ntransport = \"http_streamable\"\nurl = \"https://example.com/mcp\"\n[mcp.servers.s.headers_from_request]\n"); + for (out, inbound) in mcp_mappings { + toml.push_str(&format!("{out} = \"{inbound}\"\n")); + } + } + if !hitl_mappings.is_empty() { + toml.push_str( + "\n[hitl]\nrequire_approval = [\"x_*\"]\n[hitl.route]\nmode = \"webhook\"\nurl = \"https://example.com/hook\"\n[hitl.route.headers_from_request]\n", + ); + for (out, inbound) in hitl_mappings { + toml.push_str(&format!("{out} = \"{inbound}\"\n")); + } + } + aura_config::load_config_from_str(&toml).expect("config parses") + } + + fn headers(pairs: &[(&str, &str)]) -> HashMap { + pairs + .iter() + .map(|(k, v)| (k.to_string(), v.to_string())) + .collect() + } + + #[test] + fn nothing_is_forwarded_without_mappings() { + let forwarded = + ForwardedHeaders::resolve(&config(&[], &[]), Some(&headers(&[("x-user", "a")]))); + assert_eq!(forwarded, ForwardedHeaders::default()); + assert_eq!( + forwarded.first_difference(Some(&headers(&[("x-user", "b")]))), + None + ); + } + + #[test] + fn every_mapping_source_contributes_its_inbound_names() { + let config = config( + &[("authorization", "X-User-Token")], + &[("x-approver", "x-actor")], + ); + let forwarded = + ForwardedHeaders::resolve(&config, Some(&headers(&[("x-user-token", "t1")]))); + + assert_eq!( + forwarded.as_request(), + headers(&[("x-user-token", "t1")]), + "the MCP mapping's inbound name is read case-insensitively; the absent HITL one is omitted", + ); + assert_eq!( + format!("{forwarded:?}"), + r#"{"x-actor": "", "x-user-token": ""}"#, + "debug output names the headers and never their values", + ); + } + + #[test] + fn a_request_forwarding_the_same_values_matches() { + let config = config(&[("authorization", "x-user-token")], &[]); + let forwarded = + ForwardedHeaders::resolve(&config, Some(&headers(&[("X-User-Token", "t1")]))); + + assert_eq!( + forwarded.first_difference(Some(&headers(&[("x-user-token", "t1")]))), + None + ); + assert_eq!( + forwarded.first_difference(Some(&headers(&[ + ("x-user-token", "t1"), + ("x-other", "ignored") + ]))), + None, + "headers no mapping reads do not matter", + ); + } + + #[test] + fn a_request_forwarding_a_different_value_differs() { + let config = config(&[("authorization", "x-user-token")], &[]); + let forwarded = + ForwardedHeaders::resolve(&config, Some(&headers(&[("x-user-token", "t1")]))); + + assert_eq!( + forwarded.first_difference(Some(&headers(&[("x-user-token", "t2")]))), + Some("x-user-token"), + ); + assert_eq!( + forwarded.first_difference(None), + Some("x-user-token"), + "a request without the header the agent forwarded differs too", + ); + } + + #[test] + fn an_agent_prepared_without_the_header_differs_from_a_request_carrying_it() { + let config = config(&[("authorization", "x-user-token")], &[]); + let forwarded = ForwardedHeaders::resolve(&config, None); + + assert_eq!(forwarded.first_difference(None), None); + assert_eq!( + forwarded.first_difference(Some(&headers(&[("x-user-token", "t1")]))), + Some("x-user-token"), + "the static fallback the agent connected with is not this request's token", + ); + } +} diff --git a/crates/aura/src/hitl/gate.rs b/crates/aura/src/hitl/gate.rs index 1090a714d..c85493417 100644 --- a/crates/aura/src/hitl/gate.rs +++ b/crates/aura/src/hitl/gate.rs @@ -46,8 +46,6 @@ pub struct HitlApprovalWrapper { route: Arc, /// Who this wrapper speaks for, stamped onto every request it raises. scope: AgentScope, - /// Global request id, for SSE event routing. - request_id: String, /// The run whose observer sees this gate's approvals. run: crate::run_context::BoundRun, /// `[agent].name` of the config that built this agent. @@ -69,7 +67,6 @@ impl HitlApprovalWrapper { patterns: Arc<[GlobPattern]>, route: Arc, scope: AgentScope, - request_id: String, agent_name: String, instance_id: String, ) -> Self { @@ -77,9 +74,8 @@ impl HitlApprovalWrapper { patterns, route, scope, - request_id, // A worker's gate is built inside its run; a single agent's is - // built before one exists and is bound by `stream`. + // built before one exists and is bound by `begin_run`. run: crate::run_context::BoundRun::captured(), agent_name, instance_id, @@ -104,10 +100,7 @@ impl HitlApprovalWrapper { Some(run) => { run.emit(event).await; } - None => tracing::warn!( - request_id = %self.request_id, - "no run bound to this gate; its approval reaches no observer" - ), + None => tracing::warn!("no run bound to this gate; its approval reaches no observer"), } } @@ -319,11 +312,12 @@ impl ToolWrapper for HitlApprovalWrapper { if let Some(park) = &self.park { return self.park_pre_call(park, matched, args, ctx).await; } + let run = self.run(); let request = ApprovalRequest { version: PROTOCOL_VERSION, instance_id: self.instance_id.clone(), decision_id: DecisionId::generate(), - request_id: self.request_id.clone(), + request_id: self.run.id_or_empty(), scope: self.scope.clone(), origin: ApprovalOrigin::ConfigGate { matched_pattern: matched.to_string(), @@ -336,15 +330,15 @@ impl ToolWrapper for HitlApprovalWrapper { tool_call_intent: ctx.tool_call_intent.clone(), }], }; - let cancel = self - .run() + let cancel = run + .as_ref() .map(|run| { crate::request_cancellation::RequestCancelToken::from(run.cancel_token().clone()) }) .unwrap_or_else(crate::request_cancellation::RequestCancelToken::unbound); // `DecisionRoute` emits the lifecycle itself; the scope is how those // events find the run, since rig calls this off it. - let decision = match self.run() { + let decision = match run { Some(run) => { crate::run_context::with_run(run, self.route.decide_for_gate(request, &cancel)) .await @@ -401,7 +395,6 @@ mod tests { timeout: Duration::from_secs(1), }), AgentScope::Single { session_id: None }, - "t".into(), "test-agent".to_string(), "test-instance-id".to_string(), ); @@ -425,7 +418,6 @@ mod tests { timeout: Duration::from_secs(1), }), AgentScope::Single { session_id: None }, - "t".into(), "test-agent".to_string(), "test-instance-id".to_string(), ); @@ -463,7 +455,6 @@ mod tests { timeout: Duration::from_secs(2), }), AgentScope::Single { session_id: None }, - "req-test".into(), "test-agent".to_string(), "test-instance-id".to_string(), ); @@ -559,14 +550,12 @@ mod tests { fn parked_gate( registry: &PendingApprovals, route: &Arc, - request_id: &str, cell: &Arc, ) -> HitlApprovalWrapper { HitlApprovalWrapper::new( Arc::from(["kubectl_*".into()]), route.clone(), worker_scope(), - request_id.to_string(), "test-agent".to_string(), "test-instance".to_string(), ) @@ -591,7 +580,7 @@ mod tests { ); let route = conv_route_over(registry.clone(), Duration::from_secs(60)); let cell = Arc::new(crate::orchestration::BlockedCell::default()); - let gate = parked_gate(®istry, &route, &request_id, &cell); + let gate = parked_gate(®istry, &route, &cell); let args = serde_json::json!({ "namespace": "prod" }); let ctx = ToolCallContext::new("kubectl_apply"); @@ -629,7 +618,7 @@ mod tests { let route = conv_route_over(registry.clone(), Duration::from_secs(120)); let cell = Arc::new(crate::orchestration::BlockedCell::default()); cell.set_current_call_id(Some("call_7".to_string())); - let gate = parked_gate(®istry, &route, &request_id, &cell); + let gate = parked_gate(®istry, &route, &cell); let args = serde_json::json!({ "namespace": "prod" }); let ctx = ToolCallContext::new("kubectl_apply"); @@ -697,7 +686,7 @@ mod tests { async fn two_gated_calls_append_two_cell_entries() { let (registry, route) = conv_route(Duration::from_secs(60)); let cell = Arc::new(crate::orchestration::BlockedCell::default()); - let gate = parked_gate(®istry, &route, "req-two-calls", &cell); + let gate = parked_gate(®istry, &route, &cell); let first = gate .pre_call( @@ -733,7 +722,7 @@ mod tests { async fn ungated_tool_proceeds_without_parking() { let (registry, route) = conv_route(Duration::from_secs(60)); let cell = Arc::new(crate::orchestration::BlockedCell::default()); - let gate = parked_gate(®istry, &route, "req-ungated", &cell); + let gate = parked_gate(®istry, &route, &cell); let outcome = gate .pre_call(&serde_json::json!({}), &ToolCallContext::new("ls")) @@ -761,7 +750,6 @@ mod tests { Arc::from(["kubectl_*".into()]), route, worker_scope(), - "req-guard".to_string(), "test-agent".to_string(), "test-instance".to_string(), ) @@ -829,7 +817,6 @@ mod tests { Arc::from(["kubectl_*".into()]), route, AgentScope::Single { session_id: None }, - "req-recorded".to_string(), "test-agent".to_string(), "test-instance".to_string(), ) @@ -985,7 +972,6 @@ mod tests { Arc::from(["kubectl_*".into()]), route, AgentScope::Single { session_id: None }, - "req_run_cancel".into(), "test-agent".to_string(), "test-instance-id".to_string(), )); @@ -1226,7 +1212,6 @@ mod tests { Arc::from(["kubectl_*".into()]), Arc::new(route), AgentScope::Single { session_id: None }, - request_id.to_string(), "test-agent".to_string(), "test-instance-id".to_string(), ); diff --git a/crates/aura/src/hitl/mod.rs b/crates/aura/src/hitl/mod.rs index 77ae9e38a..594f054c5 100644 --- a/crates/aura/src/hitl/mod.rs +++ b/crates/aura/src/hitl/mod.rs @@ -25,8 +25,11 @@ //! the registry; decisions arrive via `POST /v1/approvals/{decision_id}` //! (web-server) or in-process `PendingApprovals::resolve()` (CLI standalone). //! -//! Single-agent mode composes the gate and tool in [`Agent::new`](crate::builder::Agent::new); -//! orchestration workers compose them per-task in `create_worker`. +//! Single-agent mode composes the gate and tool in +//! [`PreparedAgent::prepare`](crate::builder::PreparedAgent::prepare); +//! orchestration workers compose them per-task in `create_worker`. Both +//! stamp the request id of the run they serve, reached through a +//! [`BoundRun`](crate::run_context::BoundRun). mod decision; mod events; diff --git a/crates/aura/src/hitl/tool.rs b/crates/aura/src/hitl/tool.rs index 2afff69fa..628e8346b 100644 --- a/crates/aura/src/hitl/tool.rs +++ b/crates/aura/src/hitl/tool.rs @@ -26,7 +26,6 @@ use super::route::{ApprovalError, DecisionRoute}; pub struct RequestApprovalTool { route: Arc, scope: AgentScope, - request_id: String, run: Arc, agent_name: String, /// Instance ID of the AURA process that built this tool. @@ -38,16 +37,14 @@ impl RequestApprovalTool { pub fn new( route: Arc, scope: AgentScope, - request_id: String, agent_name: String, instance_id: String, ) -> Self { Self { route, scope, - request_id, // A worker's tool is built inside its run; a single agent's is built - // before one exists and is bound by `stream`. + // before one exists and is bound by `begin_run`. run: Arc::new(crate::run_context::BoundRun::captured()), agent_name, instance_id, @@ -158,13 +155,18 @@ impl Tool for RequestApprovalTool { } async fn call(&self, args: Self::Args) -> Result { + let run = self.run(); + let request_id = run + .as_ref() + .map(|run| run.id().to_string()) + .unwrap_or_default(); // Blank reasoning collapses to absent, consistent with the config_gate path. let tool_call_intent = normalize_tool_call_intent(args.tool_call_intent.as_deref()); let request = ApprovalRequest { version: PROTOCOL_VERSION, instance_id: self.instance_id.clone(), decision_id: DecisionId::generate(), - request_id: self.request_id.clone(), + request_id: request_id.clone(), scope: self.scope.clone(), origin: ApprovalOrigin::AgentRequested { reason: args.risk_rationale.clone(), @@ -177,7 +179,6 @@ impl Tool for RequestApprovalTool { tool_call_intent, }], }; - let run = self.run(); let cancel = run .as_ref() .map(|run| { @@ -337,7 +338,6 @@ mod tests { let tool = RequestApprovalTool::new( route.clone(), AgentScope::Single { session_id: None }, - request_id.clone(), "test-agent".to_string(), "test-instance-id".to_string(), ); @@ -504,7 +504,6 @@ mod tests { let tool = RequestApprovalTool::new( route, AgentScope::Single { session_id: None }, - request_id.clone(), "test-agent".to_string(), "test-instance-id".to_string(), ); diff --git a/crates/aura/src/lib.rs b/crates/aura/src/lib.rs index e3bd4efb2..c579380be 100644 --- a/crates/aura/src/lib.rs +++ b/crates/aura/src/lib.rs @@ -13,6 +13,7 @@ pub mod env_flags; pub mod error; pub mod fallback_tool_parser; pub mod fallback_tool_stream; +pub mod forwarded_headers; pub mod governance; pub mod hitl; pub mod hooks; @@ -51,8 +52,8 @@ pub mod vector_store; pub mod webhook_utils; pub use builder::{ - Agent, AgentBuilder, FilesystemTools, RunToolFactory, build_streaming_agent, - build_streaming_agent_with_tools, no_run_tools, + Agent, AgentBuilder, BeginRunError, FilesystemTools, PreparedAgent, RunToolFactory, + build_streaming_agent, build_streaming_agent_with_tools, no_run_tools, }; pub use config::{AgentRuntimeConfig, SessionId, ToolContextFactory}; // Pure config types are owned by `aura-config` and re-exported here for @@ -63,6 +64,7 @@ pub use aura_config::{ TodoToolsConfig, ToolsConfig, VectorStoreConfig, VectorStoreType, lenient_int, }; pub use error::{BuilderError, BuilderResult}; +pub use forwarded_headers::ForwardedHeaders; pub use orchestration::tools::{ CreatePlanTool, RequestClarificationTool, RespondDirectlyTool, RoutingDecision, RoutingToolSet, }; @@ -82,6 +84,7 @@ pub use rig::message::{AssistantContent, ToolCall as RigToolCall, ToolResultCont pub use rig::one_or_many::OneOrMany; pub use rig::tool::{Tool as RigTool, ToolDyn}; pub use rig_builder::{RigBuilder, resolve_mcp_headers_in}; +pub use run_context::{BoundRun, RunContext, RunInProgress, RunLease}; pub use scratchpad::{ScratchpadConfig, ScratchpadToolEntry}; pub use streaming::StreamingAgent; diff --git a/crates/aura/src/orchestration/mod.rs b/crates/aura/src/orchestration/mod.rs index 6f7a380d8..6c147b85f 100644 --- a/crates/aura/src/orchestration/mod.rs +++ b/crates/aura/src/orchestration/mod.rs @@ -88,7 +88,7 @@ pub(crate) use park::{CallKey, ParkGuard, RecordedDecisions, run_owner_id}; // wraps the rig's scripted agent type. pub use prompt_constants::{context, fields, sections}; #[cfg(test)] -pub(crate) use test_rig::ScriptedAgent; +pub(crate) use test_rig::{ScriptedAgent, ScriptedCompletionModel, ScriptedTurn}; pub use types::{ BlockedCell, CellOutcome, ParkSnapshot, PendingCall, Plan, PlanningResponse, RunId, StepInput, StructuredTaskOutput, Task, TaskIdentity, TaskJson, TaskState, TaskStatus, diff --git a/crates/aura/src/orchestration/orchestrator.rs b/crates/aura/src/orchestration/orchestrator.rs index ee16b4f11..2de6bf707 100644 --- a/crates/aura/src/orchestration/orchestrator.rs +++ b/crates/aura/src/orchestration/orchestrator.rs @@ -52,14 +52,15 @@ use rig::client::CompletionClient; use tokio::sync::Mutex; use tokio_util::sync::CancellationToken; -use crate::Agent; use crate::config::{AgentRuntimeConfig, LlmConfig}; use crate::inactivity::{Liveness, STALL_MESSAGE, liveness_of}; use crate::mcp::McpManager; use crate::provider_agent::{BuilderState, ProviderAgent, StreamError, StreamItem}; +use crate::run_context::{BoundRun, RunContext}; use crate::scratchpad; use crate::string_utils::safe_truncate; use crate::tool_call_observer::ToolCallObserver; +use crate::{Agent, PreparedAgent}; use super::tools::RoutingToolSet; use super::tools::{InspectToolParamsTool, ListToolsTool, ReadArtifactTool, WriteArtifactTool}; @@ -777,11 +778,6 @@ impl Orchestrator { } apply_worker_skills_override(&mut worker_config, worker_name); - // Workers are per-run ephemeral and receive no chat history, so their - // skill invocations are never rehydrated into the session — recording - // them would leak task-scoped loads across turns. Coordinator-side - // invocations keep the recorder from the top-level config. - worker_config.skill_recorder = None; // Per-worker scratchpad override falls back to [agent.scratchpad]. // Each worker gets a FRESH ContextBudget scoped to its effective LLM — @@ -791,6 +787,10 @@ impl Orchestrator { .or(self.agent_config.agent.scratchpad.as_ref()) .cloned(); + // The slot every wrapper and tool built below reaches this worker's + // run through; `begin_run_within` fills it once the worker is prepared. + let run = Arc::new(BoundRun::default()); + let mut scratchpad_budget: Option = None; let mut scratchpad_tools = Vec::>::new(); if let Some(ref sp_cfg) = effective_scratchpad && sp_cfg.enabled @@ -873,10 +873,12 @@ impl Orchestrator { context_window, initial_used, token_counter, + run: Arc::clone(&run), }) .await?; scratchpad_tools.push(build.wrapper); + scratchpad_budget = Some(build.budget); worker_config.scratchpad_tools_config = Some(build.tools_config); } } @@ -925,10 +927,10 @@ impl Orchestrator { let mut wrappers: Vec> = vec![observer_wrapper, duplicate_guard]; wrappers.extend(scratchpad_tools); wrappers.push(persistence_wrapper); - if let Some(ref state) = turn_nudge { + if turn_nudge.is_some() { wrappers.insert( 0, - Arc::new(crate::turn_nudge::TurnNudgeWrapper::new(state.clone())), + Arc::new(crate::turn_nudge::TurnNudgeWrapper::new(Arc::clone(&run))), ); tracing::info!( "Worker {} turn-limit nudging enabled (last_turn={}, wrap_up_threshold={:?})", @@ -962,12 +964,10 @@ impl Orchestrator { task: super::TaskIdentity::new(task_id, worker_name.map(String::from)), session_id: session_id_owned.map(crate::config::SessionId::new), }; - let request_id = worker_config.request_id.clone().unwrap_or_default(); let mut gate = crate::hitl::HitlApprovalWrapper::new( hitl.patterns.clone(), hitl.route.clone(), scope.clone(), - request_id.clone(), worker_config.agent.name.clone(), worker_config.instance_id.clone(), ); @@ -988,7 +988,6 @@ impl Orchestrator { worker_config.hitl_request_approval_tool = Some(crate::hitl::RequestApprovalTool::new( hitl.route.clone(), scope, - request_id, worker_config.agent.name.clone(), worker_config.instance_id.clone(), )); @@ -1065,7 +1064,6 @@ impl Orchestrator { // Orchestrator owns tool wrapping decision worker_config.tool_wrapper = Some(wrapper); - worker_config.turn_nudge = turn_nudge.clone(); // Give workers access to result artifacts worker_config.orchestration_persistence = Some(self.persistence.clone()); @@ -1112,12 +1110,16 @@ impl Orchestrator { // Build worker agent using shared MCP connections. // Client-side tools are not supported in orchestration mode and are // never attached to workers (or the coordinator). - let (provider_agent, model_name) = self.build_worker_provider_agent(&worker_config).await?; + let (provider_agent, model_name) = self + .build_worker_provider_agent(&worker_config, &run) + .await?; - let agent = Agent { + // A worker is prepared for exactly one task attempt, so its single + // run begins here, within the run that owns the orchestration. + let prepared = Arc::new(PreparedAgent { // A worker's gate and approval tool captured their run when // `create_worker` built them, inside that run, so there is nothing - // for `stream` to bind. + // for `begin_run_within` to bind. hitl_gate: None, hitl_approval_tool: None, inner: provider_agent, @@ -1128,16 +1130,21 @@ impl Orchestrator { fallback_tool_names: vec![], fallback_mcp_filter: None, context_window: worker_config.llm.context_window(), - scratchpad_budget: worker_config - .scratchpad_tools_config - .as_ref() - .map(|sp| sp.budget.clone()), + scratchpad_budget, client_tool_names: Default::default(), turn_nudge, system_prompt: preamble.clone(), invocation_parameters: crate::logging::llm_invocation_parameters(&worker_config.llm), skills: worker_config.agent.skills.clone(), - }; + forwarded_headers: worker_config.forwarded_headers.clone(), + run, + active: Default::default(), + }); + // Workers are per-run ephemeral and receive no chat history, so their + // skill invocations are never rehydrated into the session — recording + // them would leak task-scoped loads across turns. Their runs record + // none; the coordinator's run records the request's. + let agent = prepared.begin_run_within(&self.orchestration_run(), None)?; Ok(AgentWithPreamble { agent, @@ -1147,6 +1154,16 @@ 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. + 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 + }) + } + /// Park-mode wiring for one worker attempt: `Some` only when this run /// parks gated calls (see [`Self::park_enabled`]). fn worker_park(&self, task_id: usize, attempt: usize) -> Option { @@ -1492,7 +1509,7 @@ impl Orchestrator { stream, &self.usage_state, self.config.stream_inactivity_timeout_secs(), - agent.scratchpad_budget.as_ref(), + agent.scratchpad_budget(), phase, event_tx, stream_context, @@ -2703,6 +2720,10 @@ Assign tasks to the worker whose tools best match the required operations."#, } } + // The slot the coordinator's skill tools reach its run through; + // `begin_run_within` fills it once the coordinator is prepared. + let run = Arc::new(BoundRun::default()); + // Bundle all coordinator tools let coordinator_tools = CoordinatorTools { list_tools: if include_recon_tools { @@ -2729,7 +2750,7 @@ Assign tasks to the worker whose tools best match the required operations."#, }, skill_tools: crate::skill_tool::SkillToolset::new( &self.agent_config.agent.skills, - self.agent_config.skill_recorder.clone(), + Some(Arc::clone(&run)), ), }; @@ -2755,27 +2776,38 @@ Assign tasks to the worker whose tools best match the required operations."#, .turn_depth .unwrap_or(crate::builder::DEFAULT_MAX_DEPTH); + // The coordinator is one run of one prepared agent, under the request + // that owns the orchestration. + let prepared = Arc::new(PreparedAgent { + inner: provider_agent, + model: model_name, + max_depth, + mcp_manager: None, // Coordinator doesn't have MCP tools + fallback_tool_parsing: false, + fallback_tool_names: vec![], + fallback_mcp_filter: None, + context_window: self.agent_config.llm.context_window(), + scratchpad_budget: None, + client_tool_names: Default::default(), + turn_nudge: None, + system_prompt: preamble.clone(), + invocation_parameters: crate::logging::llm_invocation_parameters( + &self.agent_config.llm, + ), + skills: self.agent_config.agent.skills.clone(), + forwarded_headers: self.agent_config.forwarded_headers.clone(), + hitl_gate: None, + hitl_approval_tool: None, + run, + active: Default::default(), + }); + let agent = prepared.begin_run_within( + &self.orchestration_run(), + self.agent_config.skill_recorder.clone(), + )?; + Ok(AgentWithPreamble { - agent: Agent { - hitl_gate: None, - hitl_approval_tool: None, - inner: provider_agent, - model: model_name, - max_depth, - mcp_manager: None, // Coordinator doesn't have MCP tools - fallback_tool_parsing: false, - fallback_tool_names: vec![], - fallback_mcp_filter: None, - context_window: self.agent_config.llm.context_window(), - scratchpad_budget: None, - client_tool_names: Default::default(), - turn_nudge: None, - system_prompt: preamble.clone(), - invocation_parameters: crate::logging::llm_invocation_parameters( - &self.agent_config.llm, - ), - skills: self.agent_config.agent.skills.clone(), - }, + agent, preamble, escalation_flag: Arc::new(std::sync::atomic::AtomicBool::new(false)), submit_result_decision: Arc::new(Mutex::new(None)), @@ -2988,6 +3020,7 @@ Assign tasks to the worker whose tools best match the required operations."#, async fn build_worker_provider_agent( &self, worker_config: &AgentRuntimeConfig, + run: &Arc, ) -> Result<(ProviderAgent, String), Box> { let preamble = worker_config.effective_preamble(); let temperature = worker_config.llm.temperature(); @@ -3036,8 +3069,14 @@ Assign tasks to the worker whose tools best match the required operations."#, (None, _) => state = state.add_tool(shim), } } - let state = - Agent::add_all_tools(state, worker_config, &shared_mcp, wait_for_tools()).await?; + let state = PreparedAgent::add_all_tools( + state, + worker_config, + &shared_mcp, + run, + wait_for_tools(), + ) + .await?; return Ok(( ProviderAgent::Scripted(state.build()), "scripted".to_string(), @@ -3089,9 +3128,14 @@ Assign tasks to the worker whose tools best match the required operations."#, builder = builder.max_tokens(max); } let state = BuilderState::Initial(builder); - let state = - Agent::add_all_tools(state, worker_config, &shared_mcp, wait_for_tools()) - .await?; + let state = PreparedAgent::add_all_tools( + state, + worker_config, + &shared_mcp, + run, + wait_for_tools(), + ) + .await?; Ok((ProviderAgent::OpenAI(state.build()), model.clone())) } LlmConfig::Anthropic { @@ -3128,9 +3172,14 @@ Assign tasks to the worker whose tools best match the required operations."#, builder = builder.additional_params(params.clone()); } let state = BuilderState::Initial(builder); - let state = - Agent::add_all_tools(state, worker_config, &shared_mcp, wait_for_tools()) - .await?; + let state = PreparedAgent::add_all_tools( + state, + worker_config, + &shared_mcp, + run, + wait_for_tools(), + ) + .await?; Ok((ProviderAgent::Anthropic(state.build()), model.clone())) } LlmConfig::Bedrock { @@ -3175,9 +3224,14 @@ Assign tasks to the worker whose tools best match the required operations."#, builder = builder.additional_params(params.clone()); } let state = BuilderState::Initial(builder); - let state = - Agent::add_all_tools(state, worker_config, &shared_mcp, wait_for_tools()) - .await?; + let state = PreparedAgent::add_all_tools( + state, + worker_config, + &shared_mcp, + run, + wait_for_tools(), + ) + .await?; Ok((ProviderAgent::Bedrock(state.build()), model.clone())) } LlmConfig::Gemini { @@ -3207,9 +3261,14 @@ Assign tasks to the worker whose tools best match the required operations."#, builder = builder.additional_params(params.clone()); } let state = BuilderState::Initial(builder); - let state = - Agent::add_all_tools(state, worker_config, &shared_mcp, wait_for_tools()) - .await?; + let state = PreparedAgent::add_all_tools( + state, + worker_config, + &shared_mcp, + run, + wait_for_tools(), + ) + .await?; Ok((ProviderAgent::Gemini(state.build()), model.clone())) } LlmConfig::Ollama { @@ -3238,9 +3297,14 @@ Assign tasks to the worker whose tools best match the required operations."#, } let state = BuilderState::Initial(builder); - let state = - Agent::add_all_tools(state, worker_config, &shared_mcp, wait_for_tools()) - .await?; + let state = PreparedAgent::add_all_tools( + state, + worker_config, + &shared_mcp, + run, + wait_for_tools(), + ) + .await?; Ok((ProviderAgent::Ollama(state.build()), model.clone())) } LlmConfig::OpenRouter { @@ -3273,9 +3337,14 @@ Assign tasks to the worker whose tools best match the required operations."#, builder = builder.additional_params(params.clone()); } let state = BuilderState::Initial(builder); - let state = - Agent::add_all_tools(state, worker_config, &shared_mcp, wait_for_tools()) - .await?; + let state = PreparedAgent::add_all_tools( + state, + worker_config, + &shared_mcp, + run, + wait_for_tools(), + ) + .await?; Ok((ProviderAgent::OpenRouter(state.build()), model.clone())) } } @@ -3754,7 +3823,7 @@ Assign tasks to the worker whose tools best match the required operations."#, // their hook-carrying stream (`stream_chat_with_timeout`) seeds // the budget itself, from the prompt *and* the retry history. if park.is_none() - && let Some(ref budget) = worker.scratchpad_budget + && let Some(budget) = worker.scratchpad_budget() { let task_prompt_tokens = budget.count_tokens(&prompt); budget.record_usage(task_prompt_tokens); @@ -3839,7 +3908,7 @@ Assign tasks to the worker whose tools best match the required operations."#, } // Emit per-agent ScratchpadUsage event if this worker used scratchpad. - if let (Some(budget), Some(tx)) = (worker.scratchpad_budget.as_ref(), event_tx) { + if let (Some(budget), Some(tx)) = (worker.scratchpad_budget(), event_tx) { let agent_id = worker_name .map(|n| n.to_string()) .unwrap_or_else(|| self.orchestrator_id.clone()); @@ -4082,7 +4151,7 @@ Assign tasks to the worker whose tools best match the required operations."#, worker.max_depth, crate::streaming::RunOptions::default(), &park.key, - worker.scratchpad_budget.clone(), + worker.scratchpad_budget().cloned(), worker.client_tool_names.clone(), ) .await @@ -4091,7 +4160,7 @@ Assign tasks to the worker whose tools best match the required operations."#, stream, &self.usage_state, self.config.stream_inactivity_timeout_secs(), - worker.scratchpad_budget.as_ref(), + worker.scratchpad_budget(), "Worker resume", event_tx, worker_name.map(|name| StreamContext { diff --git a/crates/aura/src/orchestration/persistence_wrapper.rs b/crates/aura/src/orchestration/persistence_wrapper.rs index 5e831631d..470250690 100644 --- a/crates/aura/src/orchestration/persistence_wrapper.rs +++ b/crates/aura/src/orchestration/persistence_wrapper.rs @@ -892,7 +892,7 @@ mod tests { let scratchpad: Arc = Arc::new(ScratchpadWrapper::new( scratchpad_tools, storage.clone(), - budget, + Arc::new(crate::run_context::BoundRun::pinned_budget(budget)), )); let persistence_inner = Arc::new(Mutex::new(ExecutionPersistence::disabled())); @@ -1455,7 +1455,6 @@ mod tests { Arc::from(["kubectl_*".into()]), route, scope, - request_id.clone(), "test-agent".to_string(), "test-instance-id".to_string(), )); diff --git a/crates/aura/src/orchestration/tools/read_artifact.rs b/crates/aura/src/orchestration/tools/read_artifact.rs index c588b5d54..2fc36d011 100644 --- a/crates/aura/src/orchestration/tools/read_artifact.rs +++ b/crates/aura/src/orchestration/tools/read_artifact.rs @@ -8,14 +8,16 @@ use std::sync::Arc; use tokio::sync::Mutex; use crate::orchestration::persistence::ExecutionPersistence; +use crate::run_context::BoundRun; +use crate::scratchpad::ScratchpadStorage; use crate::scratchpad::storage::{ContentFormat, ScratchpadPathError}; use crate::scratchpad::tools::check_and_record_budget; use crate::scratchpad::wrapper::build_file_pointer; -use crate::scratchpad::{ContextBudget, ScratchpadStorage}; #[derive(Clone)] struct ReadArtifactScratchpad { - budget: ContextBudget, + /// The prepared agent's run slot. + run: Arc, storage: Arc, } @@ -38,12 +40,8 @@ impl ReadArtifactTool { } /// With scratchpad, oversized artifacts are returned as a pointer. - pub fn with_scratchpad( - mut self, - budget: ContextBudget, - storage: Arc, - ) -> Self { - self.scratchpad = Some(ReadArtifactScratchpad { budget, storage }); + pub fn with_scratchpad(mut self, run: Arc, storage: Arc) -> Self { + self.scratchpad = Some(ReadArtifactScratchpad { run, storage }); self } @@ -59,12 +57,25 @@ impl ReadArtifactTool { let Some(sp) = &self.scratchpad else { return content; }; + // No run bound means no budget to check the artifact against; + // withhold it rather than inline an unbounded read, as the + // scratchpad read tools do with `ScratchpadToolError::NoRun`. + let Some(budget) = sp.run.scratchpad_budget() else { + tracing::warn!( + "read_artifact: artifact {} withheld — no run is bound", + filename + ); + return format!( + "[artifact '{filename}' withheld: no run is bound to count it against. \ + Retry the read_artifact call.]" + ); + }; - if check_and_record_budget(&sp.budget, &content).is_ok() { + if check_and_record_budget(&budget, &content).is_ok() { return content; } - let tokens = sp.budget.count_tokens(&content); + let tokens = budget.count_tokens(&content); let line_count = content.lines().count(); let (format, _) = ContentFormat::detect_and_parse(&content); match sp.storage.relative_ref(abs_path).await { @@ -427,7 +438,7 @@ mod tests { // Budget-aware behavior (scratchpad active) // ------------------------------------------------------------------ - use crate::scratchpad::TiktokenCounter; + use crate::scratchpad::{ContextBudget, TiktokenCounter}; /// A standard 128k-window budget with the given per-call extraction limit. fn test_budget(max_extraction_tokens: usize) -> ContextBudget { @@ -477,8 +488,10 @@ mod tests { ); let budget = test_budget(max_extraction_tokens); - let tool = ReadArtifactTool::new(Arc::new(Mutex::new(persistence))) - .with_scratchpad(budget.clone(), storage.clone()); + let tool = ReadArtifactTool::new(Arc::new(Mutex::new(persistence))).with_scratchpad( + Arc::new(BoundRun::pinned_budget(budget.clone())), + storage.clone(), + ); (tool, storage, temp_dir) } @@ -491,7 +504,9 @@ mod tests { .scratchpad .as_ref() .unwrap() - .budget + .run + .scratchpad_budget() + .unwrap() .scratchpad_usage() .1; assert_eq!(extracted_before, 0); @@ -511,7 +526,9 @@ mod tests { .scratchpad .as_ref() .unwrap() - .budget + .run + .scratchpad_budget() + .unwrap() .scratchpad_usage() .1; assert!( @@ -579,7 +596,10 @@ mod tests { // A read tool resolves the pointer's token to the artifact in place. let file_ref = file_ref_from_pointer(&result.content); - let head = HeadTool::new(storage, test_budget(10_000)); + let head = HeadTool::new( + storage, + Arc::new(BoundRun::pinned_budget(test_budget(10_000))), + ); let head_out = head .call(HeadArgs { file: file_ref, @@ -630,8 +650,10 @@ mod tests { .unwrap() .with_read_root(read_root), ); - let tool = ReadArtifactTool::new(Arc::new(Mutex::new(run_b))) - .with_scratchpad(test_budget(50), storage.clone()); + let tool = ReadArtifactTool::new(Arc::new(Mutex::new(run_b))).with_scratchpad( + Arc::new(BoundRun::pinned_budget(test_budget(50))), + storage.clone(), + ); let result = tool .call(ReadArtifactArgs { @@ -663,7 +685,10 @@ mod tests { // The token reads the sibling-run artifact in place. let file_ref = file_ref_from_pointer(&result.content); - let head = HeadTool::new(storage, test_budget(10_000)); + let head = HeadTool::new( + storage, + Arc::new(BoundRun::pinned_budget(test_budget(10_000))), + ); let head_out = head .call(HeadArgs { file: file_ref, diff --git a/crates/aura/src/rig_builder.rs b/crates/aura/src/rig_builder.rs index d98fd82e5..223ef6ee3 100644 --- a/crates/aura/src/rig_builder.rs +++ b/crates/aura/src/rig_builder.rs @@ -9,10 +9,12 @@ //! web server can inject per-request credentials into MCP calls. use crate::builder::{ - Agent, ClientTool, RunToolFactory, build_streaming_agent_with_tools, no_run_tools, + Agent, ClientTool, PreparedAgent, RunToolFactory, build_streaming_agent_with_tools, + no_run_tools, }; use crate::config::{AgentRuntimeConfig, WorkerSkills}; use crate::error::BuilderError; +use crate::forwarded_headers::ForwardedHeaders; use crate::hitl::PendingApprovals; use crate::streaming::StreamingAgent; use aura_config::{AgentSettings, Config, McpConfig, McpServerConfig}; @@ -43,8 +45,8 @@ impl RigBuilder { self } - /// Set the recorder the built agent's skill tools persist invocations - /// with (see [`crate::skill_tool::SkillInvocationRecorder`]). + /// Set the recorder the runs this builder begins persist skill-tool + /// invocations with (see [`crate::skill_tool::SkillInvocationRecorder`]). #[must_use] pub fn with_skill_recorder( mut self, @@ -107,6 +109,7 @@ impl RigBuilder { ) }), instance_id: crate::instance_id::instance_id(&self.config.agent).to_string(), + forwarded_headers: ForwardedHeaders::resolve(&self.config, req_headers), ..Default::default() } } @@ -145,26 +148,58 @@ impl RigBuilder { Ok(agent_config) } - /// Build an agent with optional request headers, additional tools, and client-side tools. + /// Prepare an agent with optional request headers, additional tools, and client-side tools. + /// + /// The result is the reusable half of an agent: it holds no run state, so + /// one prepared agent can serve a session's turns through + /// [`PreparedAgent::begin_run`]. It does forward `req_headers` wherever + /// `headers_from_request` says to, and `begin_run` refuses a request that + /// forwards different values, so a session whose credentials change + /// prepares a new agent. /// /// - `req_headers`: HTTP headers for MCP `headers_from_request` resolution. Pass `None` when not in an HTTP context. /// - `additional_tools`: Extra rig tools the agent will execute itself (e.g. CLI/library-supplied tools). Pass `vec![]` when none needed. /// - `client_tools`: Passthrough tools the LLM may call but the *client* executes. Pass `None` when client-side tools are not in use. - pub async fn build_agent( + /// - `session_id`: The chat session the agent serves; scopes its HITL approvals. + pub async fn prepare_agent( &self, req_headers: Option<&HashMap>, additional_tools: Vec>, client_tools: Option>, - request_id: Option, session_id: Option, - ) -> Result { + ) -> Result, BuilderError> { let mut agent_config = self.discovered_agent_config(req_headers)?; resolve_mcp_headers(&mut agent_config, req_headers); - agent_config.request_id = request_id; agent_config.session_id = session_id; - agent_config.skill_recorder = self.skill_recorder.clone(); - Agent::new(&agent_config, additional_tools, client_tools) + PreparedAgent::prepare(&agent_config, additional_tools, client_tools) .await + .map(Arc::new) + .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. + /// + /// See [`Self::prepare_agent`] for the parameters. Each call prepares a + /// fresh agent; callers that want to reuse one across requests call + /// `prepare_agent` once and `begin_run` per request. + pub async fn build_agent( + &self, + req_headers: Option<&HashMap>, + additional_tools: Vec>, + client_tools: Option>, + request_id: Option, + session_id: Option, + ) -> Result { + let prepared = self + .prepare_agent(req_headers, additional_tools, client_tools, session_id) + .await?; + prepared + .begin_run( + request_id.unwrap_or_default(), + req_headers, + self.skill_recorder.clone(), + ) .map_err(|e| BuilderError::AgentError(format!("Failed to build agent: {e}"))) } diff --git a/crates/aura/src/run_context.rs b/crates/aura/src/run_context.rs index acba9c87f..5a31ebaac 100644 --- a/crates/aura/src/run_context.rs +++ b/crates/aura/src/run_context.rs @@ -23,15 +23,26 @@ use tokio::sync::mpsc; use aura_events::ToolCallId; +use crate::scratchpad::ContextBudget; +use crate::skill_tool::SkillInvocationRecorder; +use crate::turn_nudge::TurnNudgeState; + /// Events a run may buffer before its observer reads them. pub const EVENT_CHANNEL_CAPACITY: usize = 1024; -/// One run — what its own work needs to correlate, and where its events go. +/// 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, tool_calls: Mutex>, events: mpsc::Sender, cancel: CancellationToken, + /// The run's context budget. + scratchpad_budget: Option, + /// The run's turn-limit tracking. + turn_nudge: Option>, + /// Where the run's skill-tool invocations are recorded. + skill_recorder: Option>, } /// Pending tool ids before warning, in case results never arrive to pop them. @@ -49,6 +60,18 @@ impl RunContext { pub fn channel_on( id: impl Into>, cancel: CancellationToken, + ) -> (Arc, mpsc::Receiver) { + Self::channel_for_agent(id, cancel, None, None, None) + } + + /// 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>, + cancel: CancellationToken, + scratchpad_budget: Option, + turn_nudge: Option>, + skill_recorder: Option>, ) -> (Arc, mpsc::Receiver) { let (events, receiver) = mpsc::channel(EVENT_CHANNEL_CAPACITY); let run = Arc::new(Self { @@ -56,10 +79,34 @@ impl RunContext { tool_calls: Mutex::new(VecDeque::new()), events, cancel, + scratchpad_budget, + turn_nudge, + skill_recorder, }); (run, receiver) } + /// A run within `parent` for one agent of it, an orchestration worker or + /// coordinator: the same id, observer and cancellation, with tool state + /// of its own. Its tool-call queue stays empty, an agent within a run + /// streaming under a key of its own rather than the run's id. + pub fn child( + parent: &Arc, + scratchpad_budget: Option, + turn_nudge: Option>, + skill_recorder: Option>, + ) -> Arc { + Arc::new(Self { + id: Arc::clone(&parent.id), + tool_calls: Mutex::new(VecDeque::new()), + events: parent.events.clone(), + cancel: parent.cancel.clone(), + scratchpad_budget, + turn_nudge, + skill_recorder, + }) + } + /// A run nobody observes, for a test that needs one to exist without /// reading what it emits. Production names a run it can reach an observer /// through, or names none. @@ -68,6 +115,38 @@ impl RunContext { Self::channel(id).0 } + /// [`detached`](Self::detached), carrying tool state. + #[cfg(test)] + pub(crate) fn detached_with( + id: impl Into>, + scratchpad_budget: Option, + turn_nudge: Option>, + ) -> Arc { + Self::channel_for_agent( + id, + CancellationToken::new(), + scratchpad_budget, + turn_nudge, + None, + ) + .0 + } + + /// The run's context budget. + pub fn scratchpad_budget(&self) -> Option<&ContextBudget> { + self.scratchpad_budget.as_ref() + } + + /// The run's turn-limit tracking. + pub fn turn_nudge(&self) -> Option<&Arc> { + self.turn_nudge.as_ref() + } + + /// Where the run's skill-tool invocations are recorded. + pub fn skill_recorder(&self) -> Option<&Arc> { + self.skill_recorder.as_ref() + } + /// Hands an event to whoever is observing the run. `false` when nothing is /// reading, which a producer that only wanted it logged can ignore. pub async fn emit(&self, event: AgentEvent) -> bool { @@ -144,8 +223,14 @@ impl BoundRun { Self(Mutex::new(current_run())) } - /// Replaces any run already bound. One slot is enough while an agent is - /// built per request and owns the tools it gates. + /// A slot holding `run`. + pub fn holding(run: Arc) -> Self { + Self(Mutex::new(Some(run))) + } + + /// Replaces any run already bound. One slot is enough because a prepared + /// agent serves one run at a time; `PreparedAgent::begin_run` is what + /// refuses a second while the first is alive. pub fn bind(&self, run: Arc) { *self .0 @@ -161,6 +246,72 @@ impl BoundRun { .clone() .or_else(current_run) } + + /// The run's id, or empty outside a run so an approval raised then still + /// carries a well-formed id even though nothing routes it. + pub fn id_or_empty(&self) -> String { + self.get() + .map(|run| run.id().to_string()) + .unwrap_or_default() + } + + /// The run's context budget. + pub fn scratchpad_budget(&self) -> Option { + self.get().and_then(|run| run.scratchpad_budget().cloned()) + } + + /// The run's turn-limit tracking. + pub fn turn_nudge(&self) -> Option> { + self.get().and_then(|run| run.turn_nudge().cloned()) + } + + /// Where the run's skill-tool invocations are recorded. + pub fn skill_recorder(&self) -> Option> { + self.get().and_then(|run| run.skill_recorder().cloned()) + } + + /// 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)) + } + + /// 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))) + } +} + +impl std::fmt::Debug for BoundRun { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("BoundRun") + .field("run_id", &self.get().map(|run| run.id().to_string())) + .finish() + } +} + +/// A run in progress on a prepared agent. +pub struct RunLease { + run: Arc, +} + +impl RunLease { + pub(crate) fn new(run: Arc) -> Self { + Self { run } + } + + pub fn run(&self) -> &Arc { + &self.run + } +} + +/// 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")] +pub struct RunInProgress { + /// Id of the run holding the agent. + pub active: String, } tokio::task_local! { @@ -363,6 +514,77 @@ mod tests { assert_eq!(seen, None); } + #[test] + fn an_unbound_slot_outside_a_run_resolves_nothing() { + let slot = BoundRun::default(); + assert!(slot.get().is_none()); + assert_eq!(slot.id_or_empty(), ""); + assert!(slot.scratchpad_budget().is_none()); + assert!(slot.turn_nudge().is_none()); + } + + /// The state a prepared agent's tools keep for a run reaches them through + /// the slot they were built with, and follows whichever run is bound. + #[test] + fn tool_state_reaches_the_slot_from_the_bound_run() { + use crate::scratchpad::TiktokenCounter; + + let budget = + ContextBudget::new(1_000, 0.0, 0, Arc::new(TiktokenCounter::default_counter())); + let nudge = TurnNudgeState::new(true, None, 2).unwrap(); + let slot = BoundRun::default(); + slot.bind(RunContext::detached_with( + "req_a", + Some(budget.clone()), + Some(Arc::clone(&nudge)), + )); + + assert_eq!(slot.id_or_empty(), "req_a"); + slot.scratchpad_budget().unwrap().record_intercepted(7); + assert_eq!( + budget.scratchpad_usage().0, + 7, + "the slot hands out the run's own budget, counters shared", + ); + assert!(Arc::ptr_eq(&slot.turn_nudge().unwrap(), &nudge)); + + slot.bind(RunContext::detached("req_b")); + assert_eq!(slot.id_or_empty(), "req_b"); + assert!(slot.scratchpad_budget().is_none()); + } + + /// A worker's run is the orchestration run as its observer and its + /// 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 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!(Arc::ptr_eq(child.turn_nudge().unwrap(), &nudge)); + assert!(parent.turn_nudge().is_none(), "the parent keeps none of it"); + + child + .emit(AgentEvent::new( + aura_events::AgentContext::single_agent(), + aura_events::agent::AgentEventPayload::TextDelta { + content: "hi".into(), + }, + )) + .await; + assert!( + events.try_recv().is_ok(), + "what the child emits reaches the parent's observer", + ); + + parent.cancel_token().cancel(); + assert!( + child.cancel_token().is_cancelled(), + "stopping the run stops the worker", + ); + } + #[test] fn a_run_with_no_calls_in_flight_has_nothing_to_report() { let run = run("run_empty"); diff --git a/crates/aura/src/scratchpad/context_budget.rs b/crates/aura/src/scratchpad/context_budget.rs index 11f638512..b3c84d069 100644 --- a/crates/aura/src/scratchpad/context_budget.rs +++ b/crates/aura/src/scratchpad/context_budget.rs @@ -34,27 +34,15 @@ impl std::fmt::Debug for dyn TokenCounter { } } -/// Token counter using tiktoken-rs BPE tokenizers. -/// -/// For OpenAI models, resolves the exact tokenizer (o200k_base for GPT-5/GPT-4o/o-series, -/// cl100k_base for GPT-4/3.5). For all other providers, falls back to o200k_base. -/// -/// The wrapped `CoreBPE` is borrowed from tiktoken-rs's process-wide -/// `lazy_static` singletons (e.g. `o200k_base_singleton()`), so constructing -/// a `TiktokenCounter` is effectively free — no vocabulary parsing, no -/// HashMap allocation. The previous implementation called `o200k_base()` / -/// `get_bpe_from_model()`, which rebuild a fresh `CoreBPE` from the embedded -/// vocab file on every call (~3 MB include_str + ~200K HashMap inserts + -/// ~20–40 MB heap, ~50–200 ms). With per-request `Agent::new()` and -/// per-worker `create_worker()` paths in the hot path, that churn caused -/// noticeable RSS bloat and host-level slowdown over long-running sessions. +/// Token counter using tiktoken-rs's process-wide BPE tokenizers. pub struct TiktokenCounter { bpe: &'static tiktoken_rs::CoreBPE, } impl TiktokenCounter { - /// Create a counter for a specific model. - /// Falls back to `o200k_base` if the model isn't recognized. + /// Create a counter for a specific model: the exact tokenizer for an + /// OpenAI model (o200k_base for GPT-5/GPT-4o/o-series, cl100k_base for + /// GPT-4/3.5), and `o200k_base` for any model it doesn't recognize. pub fn for_model(model: &str) -> Self { Self { bpe: bpe_singleton_for_model(model), @@ -72,10 +60,12 @@ impl TiktokenCounter { /// Resolve a model name to the matching tiktoken singleton `CoreBPE`. /// /// `tiktoken_rs::get_bpe_from_model` exists but builds a fresh `CoreBPE` -/// every call. The `_singleton()` variants return a `&'static CoreBPE` from -/// a `lazy_static`, so we map model → tokenizer ourselves and dispatch to -/// the right one. Unknown models fall back to `o200k_base` (matches the -/// previous default). +/// every call from the embedded vocab file: ~3 MB of `include_str`, ~200K +/// HashMap inserts, ~20–40 MB of heap, and ~50–200 ms, which a counter built +/// for every agent would pay on every request. The `_singleton()` variants +/// return a `&'static CoreBPE` from a `lazy_static`, so we map model → +/// tokenizer ourselves and dispatch to the right one. Unknown models fall +/// back to `o200k_base`. fn bpe_singleton_for_model(model: &str) -> &'static tiktoken_rs::CoreBPE { use tiktoken_rs::tokenizer::{Tokenizer, get_tokenizer}; match get_tokenizer(model) { @@ -259,6 +249,20 @@ impl ContextBudget { self.max_extraction_tokens } + /// A budget with the same limits and no usage: the counters start over + /// from `initial_used`, as they did when this budget was created. + pub fn fresh(&self) -> Self { + Self { + max_extraction_tokens: self.max_extraction_tokens, + ..Self::new( + self.context_window, + self.safety_margin, + self.initial_used, + Arc::clone(&self.token_counter), + ) + } + } + /// Usable token budget (context window minus safety margin). pub fn usable_budget(&self) -> usize { ((self.context_window as f64) * (1.0 - self.safety_margin as f64)) as usize diff --git a/crates/aura/src/scratchpad/mod.rs b/crates/aura/src/scratchpad/mod.rs index d77bd148c..b8f90516a 100644 --- a/crates/aura/src/scratchpad/mod.rs +++ b/crates/aura/src/scratchpad/mod.rs @@ -175,8 +175,8 @@ pub use aura_config::{ScratchpadConfig, ScratchpadToolEntry}; pub struct ScratchpadToolsConfig { /// Shared storage for this request's scratchpad files. pub storage: Arc, - /// Context budget tracker shared across all scratchpad tools. - pub budget: ContextBudget, + /// The prepared agent's run slot. + pub run: Arc, /// Map of bare tool name → min_tokens threshold. Glob patterns from /// `[mcp.servers..scratchpad]` are expanded by `scratchpad_tool_map` /// when the agent is constructed for a request (per-server, diff --git a/crates/aura/src/scratchpad/setup.rs b/crates/aura/src/scratchpad/setup.rs index 7a24af79e..6c39d8725 100644 --- a/crates/aura/src/scratchpad/setup.rs +++ b/crates/aura/src/scratchpad/setup.rs @@ -10,6 +10,7 @@ use super::{ ScratchpadWrapper, TokenCounter, scratchpad_tool_schema_tokens, }; use crate::mcp::AuraTool; +use crate::run_context::BoundRun; use crate::tool_wrapper::ToolWrapper; use std::collections::HashMap; use std::path::Path; @@ -28,12 +29,13 @@ pub struct ScratchpadBuildInputs<'a> { pub context_window: usize, pub initial_used: usize, pub token_counter: Arc, + /// The prepared agent's run slot. + pub run: Arc, } -/// Output of `build_scratchpad`: the budget the caller records on its `Agent` -/// struct, the wrapper it composes into its tool pipeline, the storage handle, -/// and a ready-to-assign `ScratchpadToolsConfig` for `AgentRuntimeConfig`. +/// Output of `build_scratchpad`. pub struct ScratchpadBuild { + /// Seed context budget: the limits, with no usage recorded. pub budget: ContextBudget, pub storage: Arc, pub wrapper: Arc, @@ -71,12 +73,12 @@ pub async fn build_scratchpad( let wrapper: Arc = Arc::new(ScratchpadWrapper::new( inputs.scratchpad_tool_map.clone(), storage.clone(), - budget.clone(), + inputs.run.clone(), )); let tools_config = ScratchpadToolsConfig { storage: storage.clone(), - budget: budget.clone(), + run: inputs.run, scratchpad_tools: inputs.scratchpad_tool_map, }; @@ -280,6 +282,7 @@ mod tests { }; let mut tool_map = HashMap::new(); tool_map.insert("search_*".into(), 512); + let run = Arc::new(BoundRun::default()); let build = build_scratchpad(ScratchpadBuildInputs { sp_cfg: &sp_cfg, @@ -289,6 +292,7 @@ mod tests { context_window: 128_000, initial_used: 1_000, token_counter: counter(), + run: run.clone(), }) .await .expect("build should succeed"); @@ -296,8 +300,28 @@ mod tests { assert_eq!(build.budget.max_extraction_tokens(), Some(5_000)); assert_eq!(build.tools_config.scratchpad_tools.len(), 1); assert!(tmp.path().join("scratchpad").exists()); - // Budget in the returned struct and in tools_config share the same counters. - build.budget.record_intercepted(42); - assert_eq!(build.tools_config.budget.scratchpad_usage().0, 42); + + // The tools reach a run's budget through the slot they were built + // with; the returned budget is only what a run starts from. + assert!(build.tools_config.run.scratchpad_budget().is_none()); + let run_budget = build.budget.fresh(); + run.bind(crate::run_context::RunContext::detached_with( + "req", + Some(run_budget.clone()), + None, + )); + build + .tools_config + .run + .scratchpad_budget() + .expect("the bound run's budget resolves") + .record_intercepted(42); + assert_eq!(run_budget.scratchpad_usage().0, 42); + assert_eq!(run_budget.max_extraction_tokens(), Some(5_000)); + assert_eq!( + build.budget.scratchpad_usage().0, + 0, + "a run's usage never lands on the seed budget", + ); } } diff --git a/crates/aura/src/scratchpad/tools.rs b/crates/aura/src/scratchpad/tools.rs index 4c5b3d725..a239b725c 100644 --- a/crates/aura/src/scratchpad/tools.rs +++ b/crates/aura/src/scratchpad/tools.rs @@ -9,6 +9,7 @@ use super::schema::{ format_schema, }; use super::storage::{ScratchpadPathError, ScratchpadStorage}; +use crate::run_context::BoundRun; use rig::completion::ToolDefinition; use rig::tool::Tool; use serde::{Deserialize, Serialize}; @@ -33,6 +34,8 @@ pub enum ScratchpadToolError { NotJson, #[error("Key path not found: {0}")] KeyNotFound(String), + #[error("scratchpad tools are only available while a run is bound")] + NoRun, } impl From for ScratchpadToolError { @@ -188,6 +191,12 @@ pub(crate) fn check_and_record_budget( } } +/// The bound run's budget. A read tool called with no run bound has nothing +/// to count against and reports that instead of guessing. +fn run_budget(run: &BoundRun) -> Result { + run.scratchpad_budget().ok_or(ScratchpadToolError::NoRun) +} + // ============================================================================ // head — First N lines // ============================================================================ @@ -195,12 +204,13 @@ pub(crate) fn check_and_record_budget( #[derive(Clone)] pub struct HeadTool { storage: Arc, - budget: ContextBudget, + /// The prepared agent's run slot. + run: Arc, } impl HeadTool { - pub fn new(storage: Arc, budget: ContextBudget) -> Self { - Self { storage, budget } + pub fn new(storage: Arc, run: Arc) -> Self { + Self { storage, run } } pub fn tool_definition() -> ToolDefinition { @@ -252,6 +262,7 @@ impl Tool for HeadTool { } async fn call(&self, args: Self::Args) -> Result { + let budget = run_budget(&self.run)?; tracing::debug!("scratchpad head: file={}, lines={}", args.file, args.lines); let content = read_scratchpad_file(&self.storage, &args.file).await?; let selected: String = content @@ -265,7 +276,7 @@ impl Tool for HeadTool { // line numbers + footer add ~10-15% to the raw content's token count; // checking on the raw content would silently undercount. let numbered = add_line_numbers(&selected); - let meta = build_metadata(selected.lines().count(), &selected, &self.budget); + let meta = build_metadata(selected.lines().count(), &selected, &budget); let final_output = format!( "{}\n\n--- scratchpad head: showing {}/{} lines | {} ---", numbered, @@ -274,7 +285,7 @@ impl Tool for HeadTool { meta ); - if let Err(e) = check_and_record_budget(&self.budget, &final_output) { + if let Err(e) = check_and_record_budget(&budget, &final_output) { return Ok(format_budget_error( e, "head_too_large", @@ -299,12 +310,13 @@ impl Tool for HeadTool { #[derive(Clone)] pub struct SliceTool { storage: Arc, - budget: ContextBudget, + /// The prepared agent's run slot. + run: Arc, } impl SliceTool { - pub fn new(storage: Arc, budget: ContextBudget) -> Self { - Self { storage, budget } + pub fn new(storage: Arc, run: Arc) -> Self { + Self { storage, run } } pub fn tool_definition() -> ToolDefinition { @@ -356,6 +368,7 @@ impl Tool for SliceTool { } async fn call(&self, args: Self::Args) -> Result { + let budget = run_budget(&self.run)?; if args.start == 0 || args.end < args.start { return Err(ScratchpadToolError::InvalidArg( "start must be >= 1 and end >= start".to_string(), @@ -388,13 +401,13 @@ impl Tool for SliceTool { .join("\n"); let actual_lines = selected.lines().count(); - let meta = build_metadata(actual_lines, &selected, &self.budget); + let meta = build_metadata(actual_lines, &selected, &budget); let final_output = format!( "{}\n\n--- scratchpad slice: lines {}-{} of {} | {} ---", numbered, args.start, args.end, total_lines, meta ); - if let Err(e) = check_and_record_budget(&self.budget, &final_output) { + if let Err(e) = check_and_record_budget(&budget, &final_output) { let suggested_end = args.start + (args.end - args.start) / 2; return Ok(format_budget_error( e, @@ -420,12 +433,13 @@ impl Tool for SliceTool { #[derive(Clone)] pub struct GrepTool { storage: Arc, - budget: ContextBudget, + /// The prepared agent's run slot. + run: Arc, } impl GrepTool { - pub fn new(storage: Arc, budget: ContextBudget) -> Self { - Self { storage, budget } + pub fn new(storage: Arc, run: Arc) -> Self { + Self { storage, run } } pub fn tool_definition() -> ToolDefinition { @@ -482,6 +496,7 @@ impl Tool for GrepTool { } async fn call(&self, args: Self::Args) -> Result { + let budget = run_budget(&self.run)?; tracing::debug!( "scratchpad grep: file={}, pattern={}, context={}", args.file, @@ -548,7 +563,7 @@ impl Tool for GrepTool { } // Token-count per section and accumulate (O(output_size) total) // rather than re-tokenizing the growing result each iteration. - let section_tokens = self.budget.count_tokens(§ion); + let section_tokens = budget.count_tokens(§ion); if accumulated_tokens + section_tokens > GREP_MAX_OUTPUT_TOKENS { truncated = true; break; @@ -567,7 +582,7 @@ impl Tool for GrepTool { // Build the final formatted output FIRST (matched-region body + // footer), then budget-check on it. - let meta = build_metadata(result.lines().count(), &result, &self.budget); + let meta = build_metadata(result.lines().count(), &result, &budget); let final_output = format!( "{}\n--- scratchpad grep: {} matches in {} regions of {} | {} ---", result, @@ -577,7 +592,7 @@ impl Tool for GrepTool { meta ); - if let Err(e) = check_and_record_budget(&self.budget, &final_output) { + if let Err(e) = check_and_record_budget(&budget, &final_output) { return Ok(format_budget_error( e, "grep_too_large", @@ -621,12 +636,13 @@ fn merge_ranges(ranges: &[(usize, usize)]) -> Vec<(usize, usize)> { #[derive(Clone)] pub struct SchemaTool { storage: Arc, - budget: ContextBudget, + /// The prepared agent's run slot. + run: Arc, } impl SchemaTool { - pub fn new(storage: Arc, budget: ContextBudget) -> Self { - Self { storage, budget } + pub fn new(storage: Arc, run: Arc) -> Self { + Self { storage, run } } pub fn tool_definition() -> ToolDefinition { @@ -679,6 +695,7 @@ impl Tool for SchemaTool { } async fn call(&self, args: Self::Args) -> Result { + let budget = run_budget(&self.run)?; tracing::debug!( "scratchpad schema: file={}, max_depth={}", args.file, @@ -701,13 +718,13 @@ impl Tool for SchemaTool { }; // Build the final formatted output FIRST, then budget-check on it. - let meta = build_metadata(schema.lines().count(), &schema, &self.budget); + let meta = build_metadata(schema.lines().count(), &schema, &budget); let final_output = format!( "{}\n--- scratchpad schema: {} | {} ---", schema, args.file, meta ); - if let Err(e) = check_and_record_budget(&self.budget, &final_output) { + if let Err(e) = check_and_record_budget(&budget, &final_output) { let suggestions = if is_markdown { json!([ format!( @@ -750,12 +767,13 @@ impl Tool for SchemaTool { #[derive(Clone)] pub struct GetInTool { storage: Arc, - budget: ContextBudget, + /// The prepared agent's run slot. + run: Arc, } impl GetInTool { - pub fn new(storage: Arc, budget: ContextBudget) -> Self { - Self { storage, budget } + pub fn new(storage: Arc, run: Arc) -> Self { + Self { storage, run } } pub fn tool_definition() -> ToolDefinition { @@ -802,6 +820,7 @@ impl GetInTool { /// Return a paginated slice of a string value's lines. fn get_in_paginated( &self, + budget: &ContextBudget, path: &str, lines: &[&str], total_lines: usize, @@ -826,7 +845,7 @@ impl GetInTool { // Build the final formatted output FIRST, then budget-check on it. let numbered = add_line_numbers(&chunk); - let meta = build_metadata(end - offset, &chunk, &self.budget); + let meta = build_metadata(end - offset, &chunk, budget); let final_output = format!( "{}\n\n--- scratchpad get_in: $.{} (string, lines {}-{} of {}) | {} ---", numbered, @@ -837,7 +856,7 @@ impl GetInTool { meta ); - if let Err(e) = check_and_record_budget(&self.budget, &final_output) { + if let Err(e) = check_and_record_budget(budget, &final_output) { return Ok(format_budget_error( e, "get_in_too_large", @@ -881,6 +900,7 @@ impl Tool for GetInTool { } async fn call(&self, args: Self::Args) -> Result { + let budget = run_budget(&self.run)?; tracing::debug!("scratchpad get_in: file={}, path={}", args.file, args.path); let content = read_scratchpad_file(&self.storage, &args.file).await?; let root: serde_json::Value = @@ -910,13 +930,13 @@ impl Tool for GetInTool { serde_json::to_string_pretty(current).unwrap_or_else(|_| current.to_string()); let numbered = add_line_numbers(&result); - let meta = build_metadata(result.lines().count(), &result, &self.budget); + let meta = build_metadata(result.lines().count(), &result, &budget); let final_output = format!( "{}\n\n--- scratchpad get_in: $.{} | {} ---", numbered, args.path, meta ); - if let Err(e) = check_and_record_budget(&self.budget, &final_output) { + if let Err(e) = check_and_record_budget(&budget, &final_output) { return Ok(format_budget_error( e, "get_in_too_large", @@ -940,6 +960,7 @@ impl Tool for GetInTool { if args.offset.is_some() || args.limit.is_some() { return self.get_in_paginated( + &budget, &args.path, &lines, total_lines, @@ -950,13 +971,13 @@ impl Tool for GetInTool { // No pagination — build final formatted output, then budget-check. let numbered = add_line_numbers(raw_str); - let meta = build_metadata(total_lines, raw_str, &self.budget); + let meta = build_metadata(total_lines, raw_str, &budget); let final_output = format!( "{}\n\n--- scratchpad get_in: $.{} | {} ---", numbered, args.path, meta ); - if let Err(e) = check_and_record_budget(&self.budget, &final_output) { + if let Err(e) = check_and_record_budget(&budget, &final_output) { return Ok(format_budget_error( e, "get_in_too_large", @@ -983,12 +1004,13 @@ impl Tool for GetInTool { #[derive(Clone)] pub struct IterateOverTool { storage: Arc, - budget: ContextBudget, + /// The prepared agent's run slot. + run: Arc, } impl IterateOverTool { - pub fn new(storage: Arc, budget: ContextBudget) -> Self { - Self { storage, budget } + pub fn new(storage: Arc, run: Arc) -> Self { + Self { storage, run } } pub fn tool_definition() -> ToolDefinition { @@ -1061,6 +1083,7 @@ impl Tool for IterateOverTool { } async fn call(&self, args: Self::Args) -> Result { + let budget = run_budget(&self.run)?; tracing::debug!( "scratchpad iterate_over: file={}, path={}, fields={}, offset={:?}, limit={:?}", args.file, @@ -1136,13 +1159,13 @@ impl Tool for IterateOverTool { total_items ) }; - let meta = build_metadata(result.lines().count(), &result, &self.budget); + let meta = build_metadata(result.lines().count(), &result, &budget); let final_output = format!( "{}\n\n--- scratchpad iterate_over: $.{} ({}, fields: [{}]) | {} ---", result, args.path, window_desc, args.fields, meta ); - if let Err(e) = check_and_record_budget(&self.budget, &final_output) { + if let Err(e) = check_and_record_budget(&budget, &final_output) { let mut suggestions = Vec::new(); // A retry is only proposed when it strictly narrows the window; // a single over-budget item needs fewer fields, not fewer items. @@ -1180,12 +1203,13 @@ impl Tool for IterateOverTool { #[derive(Clone)] pub struct ItemSchemaTool { storage: Arc, - budget: ContextBudget, + /// The prepared agent's run slot. + run: Arc, } impl ItemSchemaTool { - pub fn new(storage: Arc, budget: ContextBudget) -> Self { - Self { storage, budget } + pub fn new(storage: Arc, run: Arc) -> Self { + Self { storage, run } } pub fn tool_definition() -> ToolDefinition { @@ -1251,6 +1275,7 @@ impl Tool for ItemSchemaTool { } async fn call(&self, args: Self::Args) -> Result { + let budget = run_budget(&self.run)?; tracing::debug!( "scratchpad item_schema: file={}, path={}, offset={:?}, limit={:?}", args.file, @@ -1328,13 +1353,13 @@ impl Tool for ItemSchemaTool { } // Build the final formatted output FIRST, then budget-check on it. - let meta = build_metadata(result.lines().count(), &result, &self.budget); + let meta = build_metadata(result.lines().count(), &result, &budget); let final_output = format!( "{}\n--- scratchpad item_schema: $.{} | {} ---", result, args.path, meta ); - if let Err(e) = check_and_record_budget(&self.budget, &final_output) { + if let Err(e) = check_and_record_budget(&budget, &final_output) { // Suggest halving the current window (min 10) as a starting point. let suggested_limit = (window_size / 2).max(10); return Ok(format_budget_error( @@ -1605,12 +1630,13 @@ fn format_navigation_failure( #[derive(Clone)] pub struct ReadTool { storage: Arc, - budget: ContextBudget, + /// The prepared agent's run slot. + run: Arc, } impl ReadTool { - pub fn new(storage: Arc, budget: ContextBudget) -> Self { - Self { storage, budget } + pub fn new(storage: Arc, run: Arc) -> Self { + Self { storage, run } } pub fn tool_definition() -> ToolDefinition { @@ -1651,6 +1677,7 @@ impl Tool for ReadTool { } async fn call(&self, args: Self::Args) -> Result { + let budget = run_budget(&self.run)?; tracing::debug!("scratchpad read: file={}", args.file); // Pre-flight size check: skip loading a multi-MB file into memory @@ -1660,7 +1687,7 @@ impl Tool for ReadTool { // budget check happens after the file is read; this is a fast-path // rejection only, biased toward the common case of obviously oversized // files. If the estimate already exceeds the per-call limit, bail out. - if let Some(limit) = self.budget.max_extraction_tokens() + if let Some(limit) = budget.max_extraction_tokens() && let Ok(path) = self.storage.validate_path(&args.file).await && let Ok(meta) = tokio::fs::metadata(&path).await { @@ -1692,7 +1719,7 @@ impl Tool for ReadTool { // Build the final formatted output FIRST, then budget-check on it. let numbered = add_line_numbers(&content); - let meta = build_metadata(content.lines().count(), &content, &self.budget); + let meta = build_metadata(content.lines().count(), &content, &budget); let final_output = format!( "{}\n\n--- scratchpad read: {} ({} lines) | {} ---", numbered, @@ -1701,7 +1728,7 @@ impl Tool for ReadTool { meta ); - if let Err(e) = check_and_record_budget(&self.budget, &final_output) { + if let Err(e) = check_and_record_budget(&budget, &final_output) { return Ok(format_budget_error( e, "read_too_large", @@ -1817,6 +1844,7 @@ pub fn emit_scratchpad_tool_events_enabled() -> bool { #[cfg(test)] mod tests { use super::*; + use crate::run_context::BoundRun; use crate::scratchpad::context_budget::{TiktokenCounter, TokenCounter}; use tempfile::TempDir; @@ -1862,7 +1890,7 @@ mod tests { let (_tmp, storage, budget) = setup().await; storage.write_output("test", sample_json()).await.unwrap(); - let tool = HeadTool::new(storage, budget); + let tool = HeadTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(HeadArgs { file: "test.json".to_string(), @@ -1898,7 +1926,7 @@ mod tests { let tight_budget = ContextBudget::new(100_000, 0.20, 0, std::sync::Arc::new(counter)) .with_max_extraction_tokens(raw_tokens + 1); - let tool = HeadTool::new(storage, tight_budget); + let tool = HeadTool::new(storage, Arc::new(BoundRun::pinned_budget(tight_budget))); let result = tool .call(HeadArgs { file: "preview.txt".to_string(), @@ -1928,7 +1956,7 @@ mod tests { let (_tmp, storage, budget) = setup().await; storage.write_output("test", sample_json()).await.unwrap(); - let tool = SliceTool::new(storage, budget); + let tool = SliceTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(SliceArgs { file: "test.json".to_string(), @@ -1946,7 +1974,7 @@ mod tests { let (_tmp, storage, budget) = setup().await; storage.write_output("test", sample_json()).await.unwrap(); - let tool = GrepTool::new(storage, budget); + let tool = GrepTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(GrepArgs { file: "test.json".to_string(), @@ -1965,7 +1993,7 @@ mod tests { let (_tmp, storage, budget) = setup().await; storage.write_output("test", sample_json()).await.unwrap(); - let tool = GrepTool::new(storage, budget); + let tool = GrepTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let long_pattern = "a".repeat(GREP_MAX_PATTERN_LEN + 1); let err = tool .call(GrepArgs { @@ -1990,7 +2018,7 @@ mod tests { let huge_content = (0..100_000).map(|_| "aa").collect::>().join("\n"); storage.write_output("huge", &huge_content).await.unwrap(); - let tool = GrepTool::new(storage, budget); + let tool = GrepTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(GrepArgs { file: "huge.txt".to_string(), @@ -2016,7 +2044,7 @@ mod tests { let (_tmp, storage, budget) = setup().await; storage.write_output("test", sample_json()).await.unwrap(); - let tool = GrepTool::new(storage, budget); + let tool = GrepTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(GrepArgs { file: "test.json".to_string(), @@ -2041,7 +2069,7 @@ mod tests { .await .unwrap(); - let tool = GrepTool::new(storage, budget); + let tool = GrepTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(GrepArgs { file: "metrics.json".to_string(), @@ -2062,7 +2090,7 @@ mod tests { let (_tmp, storage, budget) = setup().await; storage.write_output("test", sample_json()).await.unwrap(); - let tool = SchemaTool::new(storage, budget); + let tool = SchemaTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(SchemaArgs { file: "test.json".to_string(), @@ -2079,7 +2107,7 @@ mod tests { let (_tmp, storage, budget) = setup().await; storage.write_output("test", sample_json()).await.unwrap(); - let tool = GetInTool::new(storage, budget); + let tool = GetInTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(GetInArgs { file: "test.json".to_string(), @@ -2102,7 +2130,7 @@ mod tests { let (_tmp, storage, budget) = setup().await; storage.write_output("test", sample_json()).await.unwrap(); - let tool = GetInTool::new(storage, budget); + let tool = GetInTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(GetInArgs { file: "test.json".to_string(), @@ -2139,7 +2167,7 @@ mod tests { serde_json::json!({ "kv_markdown": "### Section A\n- key: value" }).to_string(); storage.write_output("test", &json_str).await.unwrap(); - let tool = GetInTool::new(storage, budget); + let tool = GetInTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(GetInArgs { file: "test.json".to_string(), @@ -2180,7 +2208,7 @@ mod tests { "test setup: expected a companion file for kv_markdown" ); - let tool = GetInTool::new(storage, budget); + let tool = GetInTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(GetInArgs { file: "test.json".to_string(), @@ -2221,7 +2249,7 @@ mod tests { let (_tmp, storage, budget) = setup().await; storage.write_output("test", sample_json()).await.unwrap(); - let tool = GetInTool::new(storage, budget); + let tool = GetInTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(GetInArgs { file: "test.json".to_string(), @@ -2245,7 +2273,7 @@ mod tests { let json = r#"{"kv_markdown": "line1\nline2\nline3\nline4\nline5"}"#; storage.write_output("md", json).await.unwrap(); - let tool = GetInTool::new(storage, budget); + let tool = GetInTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); // Without pagination, should return the raw string content let result = tool @@ -2283,7 +2311,7 @@ mod tests { let json = r#"{"data": "a\nb\nc"}"#; storage.write_output("small", json).await.unwrap(); - let tool = GetInTool::new(storage, budget); + let tool = GetInTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(GetInArgs { file: "small.json".to_string(), @@ -2305,7 +2333,7 @@ mod tests { let json = r#"{"data": "line1\nline2\nline3\nline4\nline5"}"#; storage.write_output("ovf", json).await.unwrap(); - let tool = GetInTool::new(storage, budget); + let tool = GetInTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); // limit = usize::MAX would wrap when added to offset without // saturating_add. Should be clamped to total_lines instead. let result = tool @@ -2331,7 +2359,7 @@ mod tests { let (_tmp, storage, budget) = setup().await; storage.write_output("test", sample_json()).await.unwrap(); - let tool = IterateOverTool::new(storage, budget); + let tool = IterateOverTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(IterateOverArgs { file: "test.json".to_string(), @@ -2353,7 +2381,7 @@ mod tests { let (_tmp, storage, budget) = setup().await; storage.write_output("test", sample_json()).await.unwrap(); - let tool = IterateOverTool::new(storage, budget); + let tool = IterateOverTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(IterateOverArgs { file: "test.json".to_string(), @@ -2373,7 +2401,7 @@ mod tests { let (_tmp, storage, budget) = setup().await; storage.write_output("test", sample_json()).await.unwrap(); - let tool = IterateOverTool::new(storage, budget); + let tool = IterateOverTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(IterateOverArgs { file: "test.json".to_string(), @@ -2405,7 +2433,7 @@ mod tests { .await .unwrap(); - let tool = IterateOverTool::new(storage, budget); + let tool = IterateOverTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(IterateOverArgs { file: "trunc.json".to_string(), @@ -2438,7 +2466,7 @@ mod tests { .await .unwrap(); - let tool = IterateOverTool::new(storage, budget); + let tool = IterateOverTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(IterateOverArgs { file: "page.json".to_string(), @@ -2465,7 +2493,7 @@ mod tests { let (_tmp, storage, budget) = setup().await; storage.write_output("test", sample_json()).await.unwrap(); - let tool = IterateOverTool::new(storage, budget); + let tool = IterateOverTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(IterateOverArgs { file: "test.json".to_string(), @@ -2486,7 +2514,7 @@ mod tests { let (_tmp, storage, budget) = setup().await; storage.write_output("test", sample_json()).await.unwrap(); - let tool = IterateOverTool::new(storage, budget); + let tool = IterateOverTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(IterateOverArgs { file: "test.json".to_string(), @@ -2507,7 +2535,7 @@ mod tests { let (_tmp, storage, budget) = setup().await; storage.write_output("test", sample_json()).await.unwrap(); - let tool = IterateOverTool::new(storage, budget); + let tool = IterateOverTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); // offset + usize::MAX must not wrap; window clamps to total items. let result = tool .call(IterateOverArgs { @@ -2529,7 +2557,7 @@ mod tests { let (_tmp, storage, budget) = setup().await; storage.write_output("test", sample_json()).await.unwrap(); - let tool = IterateOverTool::new(storage, budget); + let tool = IterateOverTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); // offset=0 + limit covering the whole array reads as a full scan. let result = tool .call(IterateOverArgs { @@ -2550,7 +2578,7 @@ mod tests { let (_tmp, storage, budget) = setup().await; storage.write_output("test", sample_json()).await.unwrap(); - let tool = IterateOverTool::new(storage, budget); + let tool = IterateOverTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); // Schema advertises minimum 1, but providers don't enforce it. let result = tool .call(IterateOverArgs { @@ -2576,7 +2604,7 @@ mod tests { let json = format!(r#"{{"items":[{}]}}"#, items.join(",")); storage.write_output("big", &json).await.unwrap(); - let tool = IterateOverTool::new(storage, tiny_budget); + let tool = IterateOverTool::new(storage, Arc::new(BoundRun::pinned_budget(tiny_budget))); let result = tool .call(IterateOverArgs { file: "big.json".to_string(), @@ -2610,7 +2638,7 @@ mod tests { let json = format!(r#"{{"items":[{{"id":1,"content":"{big}"}}]}}"#); storage.write_output("one", &json).await.unwrap(); - let tool = IterateOverTool::new(storage, tiny_budget); + let tool = IterateOverTool::new(storage, Arc::new(BoundRun::pinned_budget(tiny_budget))); let result = tool .call(IterateOverArgs { file: "one.json".to_string(), @@ -2643,7 +2671,7 @@ mod tests { let (_tmp, storage, budget) = setup().await; storage.write_output("test", sample_json()).await.unwrap(); - let tool = ItemSchemaTool::new(storage, budget); + let tool = ItemSchemaTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(ItemSchemaArgs { file: "test.json".to_string(), @@ -2673,7 +2701,7 @@ mod tests { }"#; storage.write_output("hetero", json).await.unwrap(); - let tool = ItemSchemaTool::new(storage, budget); + let tool = ItemSchemaTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(ItemSchemaArgs { file: "hetero.json".to_string(), @@ -2707,7 +2735,7 @@ mod tests { }"#; storage.write_output("paged", json).await.unwrap(); - let tool = ItemSchemaTool::new(storage, budget); + let tool = ItemSchemaTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(ItemSchemaArgs { file: "paged.json".to_string(), @@ -2735,7 +2763,7 @@ mod tests { let (_tmp, storage, budget) = setup().await; storage.write_output("test", sample_json()).await.unwrap(); - let tool = ItemSchemaTool::new(storage, budget); + let tool = ItemSchemaTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(ItemSchemaArgs { file: "test.json".to_string(), @@ -2763,7 +2791,7 @@ mod tests { let json = format!(r#"{{"items":[{}]}}"#, items.join(",")); storage.write_output("big", &json).await.unwrap(); - let tool = ItemSchemaTool::new(storage, tiny_budget); + let tool = ItemSchemaTool::new(storage, Arc::new(BoundRun::pinned_budget(tiny_budget))); let result = tool .call(ItemSchemaArgs { file: "big.json".to_string(), @@ -2803,7 +2831,7 @@ mod tests { let (_tmp, storage, budget) = setup().await; storage.write_output("test", sample_json()).await.unwrap(); - let tool = ReadTool::new(storage, budget); + let tool = ReadTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(ReadArgs { file: "test.json".to_string(), @@ -2825,7 +2853,7 @@ mod tests { let large_content = "x".repeat(1000); storage.write_output("large", &large_content).await.unwrap(); - let tool = ReadTool::new(storage, tiny_budget); + let tool = ReadTool::new(storage, Arc::new(BoundRun::pinned_budget(tiny_budget))); let result = tool .call(ReadArgs { file: "large.txt".to_string(), @@ -2843,7 +2871,7 @@ mod tests { #[tokio::test] async fn test_path_traversal_rejected() { let (_tmp, storage, budget) = setup().await; - let tool = HeadTool::new(storage, budget); + let tool = HeadTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(HeadArgs { file: "../../etc/passwd".to_string(), @@ -2861,7 +2889,7 @@ mod tests { let path = storage.dir().join("test.md"); tokio::fs::write(&path, md).await.unwrap(); - let tool = SchemaTool::new(storage, budget); + let tool = SchemaTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(SchemaArgs { file: "test.md".to_string(), @@ -2890,7 +2918,7 @@ mod tests { let path = storage.dir().join("test.payload.json"); tokio::fs::write(&path, &pretty).await.unwrap(); - let tool = SchemaTool::new(storage, budget); + let tool = SchemaTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(SchemaArgs { file: "test.payload.json".to_string(), @@ -2912,7 +2940,7 @@ mod tests { .await .unwrap(); - let tool = GetInTool::new(storage, budget); + let tool = GetInTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); let result = tool .call(GetInArgs { file: "raw.json".to_string(), @@ -2947,7 +2975,7 @@ mod tests { .await .unwrap(); - let tool = GetInTool::new(storage, tiny_budget); + let tool = GetInTool::new(storage, Arc::new(BoundRun::pinned_budget(tiny_budget))); // With offset/limit, the chunk should still exceed the tiny per-call limit let result = tool .call(GetInArgs { @@ -2974,13 +3002,16 @@ mod tests { async fn test_tool_definition_matches_trait_definition() { let (_tmp, storage, budget) = setup().await; - let head = HeadTool::new(storage.clone(), budget.clone()); + let head = HeadTool::new( + storage.clone(), + Arc::new(BoundRun::pinned_budget(budget.clone())), + ); assert_eq!( HeadTool::tool_definition(), head.definition(String::new()).await ); - let read = ReadTool::new(storage, budget); + let read = ReadTool::new(storage, Arc::new(BoundRun::pinned_budget(budget))); assert_eq!( ReadTool::tool_definition(), read.definition(String::new()).await @@ -3073,4 +3104,22 @@ mod tests { assert!(!should_forward_tool_event("some_mcp_tool", false)); assert!(!should_forward_tool_event("some_mcp_tool", true)); } + + /// A read tool has nothing to count against outside a run, and says so + /// rather than reading unbudgeted. + #[tokio::test] + async fn read_tools_refuse_when_no_run_is_bound() { + let (_tmp, storage, _budget) = setup().await; + storage.write_output("test", sample_json()).await.unwrap(); + + let tool = HeadTool::new(storage, Arc::new(BoundRun::default())); + let err = tool + .call(HeadArgs { + file: "test.json".to_string(), + lines: 5, + }) + .await + .expect_err("no run is bound"); + assert!(matches!(err, ScratchpadToolError::NoRun), "got: {err}"); + } } diff --git a/crates/aura/src/scratchpad/wrapper.rs b/crates/aura/src/scratchpad/wrapper.rs index 07a2c7f45..1ecb98c3b 100644 --- a/crates/aura/src/scratchpad/wrapper.rs +++ b/crates/aura/src/scratchpad/wrapper.rs @@ -1,10 +1,10 @@ //! ScratchpadWrapper — intercepts large MCP tool outputs and writes them //! to the scratchpad directory, returning a summary pointer to the LLM. -use super::context_budget::ContextBudget; use super::storage::ScratchpadStorage; use crate::mcp::CallOutcome; use crate::orchestration::persistence_wrapper::strip_artifact_footer; +use crate::run_context::BoundRun; use crate::tool_wrapper::{ToolCallContext, ToolWrapper, TransformOutputResult}; use async_trait::async_trait; use std::collections::HashMap; @@ -41,34 +41,34 @@ pub(crate) fn build_file_pointer(headline: &str, file_ref: &str) -> String { /// ToolWrapper that intercepts large outputs from flagged tools and writes /// them to scratchpad files, replacing the output with a compact pointer. pub struct ScratchpadWrapper { - /// Map of bare tool name → `min_tokens` threshold. Resolved per-request - /// during `Agent::new` (single-agent) or `Orchestrator::create_worker` - /// (orchestration) by `scratchpad::scratchpad_tool_map` — server-aware, - /// glob patterns expanded against each server's tool list. Runtime - /// lookup is an exact `HashMap::get`. + /// Map of bare tool name → `min_tokens` threshold. scratchpad_tools: HashMap, /// Storage backend for writing scratchpad files. storage: Arc, - /// Budget tracker for recording intercepted tokens. - budget: ContextBudget, + /// The prepared agent's run slot. + run: Arc, } impl ScratchpadWrapper { pub fn new( scratchpad_tools: HashMap, storage: Arc, - budget: ContextBudget, + run: Arc, ) -> Self { Self { scratchpad_tools, storage, - budget, + run, } } } #[async_trait] impl ToolWrapper for ScratchpadWrapper { + /// Counts and records against the budget of the run bound when the call + /// happens. Exact-name lookup against `scratchpad_tools`, which + /// `scratchpad::scratchpad_tool_map` resolves from the per-server glob + /// patterns when the agent is prepared. async fn transform_output( &self, output: String, @@ -101,7 +101,24 @@ impl ToolWrapper for ScratchpadWrapper { // and all. let content = strip_artifact_footer(&output); - let output_tokens = self.budget.count_tokens(content); + // With no run bound there is no budget to count against, and a + // flagged tool's output is presumed large enough to need one. + // Withhold it rather than pass through the overflow scratchpad + // exists to prevent — the same choice the read tools make with + // `ScratchpadToolError::NoRun`. + let Some(budget) = self.run.scratchpad_budget() else { + tracing::warn!( + "Scratchpad: {} output withheld — no run is bound", + ctx.tool_name + ); + return TransformOutputResult::new(format!( + "[scratchpad: {} output withheld: no run is bound to count it against. \ + Retry the tool call.]", + ctx.tool_name + )); + }; + + let output_tokens = budget.count_tokens(content); if output_tokens < min_tokens { tracing::debug!( "Scratchpad: {} output (~{} tokens) below threshold ({}), passing through", @@ -193,7 +210,7 @@ impl ToolWrapper for ScratchpadWrapper { pointer.push_str(&tool_list); } - self.budget.record_intercepted(token_count); + budget.record_intercepted(token_count); tracing::debug!( "Scratchpad: intercepted {} output (~{} tokens) → {} ({} companions)", @@ -230,6 +247,7 @@ impl ToolWrapper for ScratchpadWrapper { #[cfg(test)] mod tests { use super::*; + use crate::scratchpad::ContextBudget; use crate::scratchpad::context_budget::{TiktokenCounter, TokenCounter}; use crate::tool_wrapper::ToolCallContext; use tempfile::TempDir; @@ -266,7 +284,11 @@ mod tests { let counter = TiktokenCounter::default_counter(); let budget = ContextBudget::new(128_000, 0.20, 0, std::sync::Arc::new(counter)); - let wrapper = ScratchpadWrapper::new(tools, storage.clone(), budget); + let wrapper = ScratchpadWrapper::new( + tools, + storage.clone(), + Arc::new(BoundRun::pinned_budget(budget)), + ); let large_output = (0..500) .map(|i| format!("entry_{} ", i)) @@ -308,7 +330,8 @@ mod tests { let tools = HashMap::from([("echo_large".to_string(), 10)]); let counter = TiktokenCounter::default_counter(); let budget = ContextBudget::new(128_000, 0.20, 0, std::sync::Arc::new(counter)); - let wrapper = ScratchpadWrapper::new(tools, storage, budget); + let wrapper = + ScratchpadWrapper::new(tools, storage, Arc::new(BoundRun::pinned_budget(budget))); let large_output = (0..200).map(|i| format!("entry_{i} ")).collect::(); let mut ctx = ToolCallContext::new("echo_large"); @@ -341,7 +364,11 @@ mod tests { let counter = TiktokenCounter::default_counter(); let budget = ContextBudget::new(128_000, 0.20, 0, std::sync::Arc::new(counter)); - let wrapper = ScratchpadWrapper::new(tools, storage.clone(), budget); + let wrapper = ScratchpadWrapper::new( + tools, + storage.clone(), + Arc::new(BoundRun::pinned_budget(budget)), + ); // Use varied content to avoid tokenizer compression of repeated chars let large_output = (0..500).map(|i| format!("item_{} ", i)).collect::(); @@ -391,7 +418,11 @@ mod tests { let tools = HashMap::from([("execute_range_query".to_string(), 10)]); let counter = TiktokenCounter::default_counter(); let budget = ContextBudget::new(128_000, 0.20, 0, std::sync::Arc::new(counter)); - let wrapper = ScratchpadWrapper::new(tools, storage.clone(), budget); + let wrapper = ScratchpadWrapper::new( + tools, + storage.clone(), + Arc::new(BoundRun::pinned_budget(budget)), + ); // Valid JSON payload, large enough to be intercepted... let items: String = (0..200) @@ -446,7 +477,11 @@ mod tests { let counter = TiktokenCounter::default_counter(); let budget = ContextBudget::new(128_000, 0.20, 0, std::sync::Arc::new(counter)); - let wrapper = ScratchpadWrapper::new(tools, storage.clone(), budget.clone()); + let wrapper = ScratchpadWrapper::new( + tools, + storage.clone(), + Arc::new(BoundRun::pinned_budget(budget.clone())), + ); let small_output = "small result".to_string(); let ctx = ToolCallContext::new("search_knowledge_base"); @@ -484,7 +519,8 @@ mod tests { let counter = TiktokenCounter::default_counter(); let budget = ContextBudget::new(128_000, 0.20, 0, std::sync::Arc::new(counter)); - let wrapper = ScratchpadWrapper::new(tools, storage, budget); + let wrapper = + ScratchpadWrapper::new(tools, storage, Arc::new(BoundRun::pinned_budget(budget))); let large_output = "x".repeat(500); let ctx = ToolCallContext::new("other_tool"); @@ -512,7 +548,11 @@ mod tests { let counter = TiktokenCounter::default_counter(); let budget = ContextBudget::new(128_000, 0.20, 0, std::sync::Arc::new(counter)); - let wrapper = ScratchpadWrapper::new(tools, storage.clone(), budget); + let wrapper = ScratchpadWrapper::new( + tools, + storage.clone(), + Arc::new(BoundRun::pinned_budget(budget)), + ); let large_output = (0..500).map(|i| format!("line_{} ", i)).collect::(); for tool_name in ["load_skill", "read_skill_file"] { @@ -550,7 +590,11 @@ mod tests { // Set threshold to exactly the token count — should be intercepted (>=) let tools = HashMap::from([("tool_at_boundary".to_string(), exact_tokens)]); let budget = ContextBudget::new(128_000, 0.20, 0, Arc::new(counter)); - let wrapper = ScratchpadWrapper::new(tools, storage.clone(), budget); + let wrapper = ScratchpadWrapper::new( + tools, + storage.clone(), + Arc::new(BoundRun::pinned_budget(budget)), + ); let mut ctx = ToolCallContext::new("tool_at_boundary"); ctx.task_id = Some(1); @@ -570,7 +614,11 @@ mod tests { let tools_above = HashMap::from([("tool_at_boundary".to_string(), exact_tokens + 1)]); let counter2 = TiktokenCounter::default_counter(); let budget2 = ContextBudget::new(128_000, 0.20, 0, Arc::new(counter2)); - let wrapper2 = ScratchpadWrapper::new(tools_above, storage, budget2); + let wrapper2 = ScratchpadWrapper::new( + tools_above, + storage, + Arc::new(BoundRun::pinned_budget(budget2)), + ); let result2 = wrapper2 .transform_output(content.clone(), &ok(), &ctx, None) @@ -598,7 +646,11 @@ mod tests { let expected_tokens = counter.count_tokens(&content); let budget = ContextBudget::new(128_000, 0.20, 0, Arc::new(counter)); - let wrapper = ScratchpadWrapper::new(tools, storage, budget.clone()); + let wrapper = ScratchpadWrapper::new( + tools, + storage, + Arc::new(BoundRun::pinned_budget(budget.clone())), + ); let mut ctx = ToolCallContext::new("counted_tool"); ctx.task_id = Some(1); @@ -634,7 +686,11 @@ mod tests { let tools = HashMap::from([("failing_tool".to_string(), 10)]); let counter = TiktokenCounter::default_counter(); let budget = ContextBudget::new(128_000, 0.20, 0, Arc::new(counter)); - let wrapper = ScratchpadWrapper::new(tools, storage, budget.clone()); + let wrapper = ScratchpadWrapper::new( + tools, + storage, + Arc::new(BoundRun::pinned_budget(budget.clone())), + ); let large_output = (0..200).map(|i| format!("item_{} ", i)).collect::(); let mut ctx = ToolCallContext::new("failing_tool"); @@ -688,7 +744,11 @@ mod tests { let tools = HashMap::from([("nested_call".to_string(), 10)]); let counter = TiktokenCounter::default_counter(); let budget = ContextBudget::new(128_000, 0.20, 0, std::sync::Arc::new(counter)); - let wrapper = ScratchpadWrapper::new(tools, storage.clone(), budget); + let wrapper = ScratchpadWrapper::new( + tools, + storage.clone(), + Arc::new(BoundRun::pinned_budget(budget)), + ); // Outer JSON whose `payload` value is itself an escaped JSON string // long enough to trigger companion extraction (>= COMPANION_MIN_LINES @@ -752,7 +812,11 @@ mod tests { let counter = TiktokenCounter::default_counter(); let budget = ContextBudget::new(128_000, 0.20, 0, std::sync::Arc::new(counter)); - let wrapper = ScratchpadWrapper::new(tools, storage.clone(), budget); + let wrapper = ScratchpadWrapper::new( + tools, + storage.clone(), + Arc::new(BoundRun::pinned_budget(budget)), + ); // JSON with a large markdown string value that will be extracted as a companion let md_lines = (0..15) diff --git a/crates/aura/src/skill_rehydration.rs b/crates/aura/src/skill_rehydration.rs index cb4e0a7c7..c6dfcda48 100644 --- a/crates/aura/src/skill_rehydration.rs +++ b/crates/aura/src/skill_rehydration.rs @@ -410,7 +410,15 @@ 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 toolset = SkillToolset::new(&skills, Some(recorder)).unwrap(); + let (run, _events) = crate::run_context::RunContext::channel_for_agent( + "req-turn-n", + tokio_util::sync::CancellationToken::new(), + None, + None, + Some(recorder), + ); + let slot = Arc::new(crate::run_context::BoundRun::holding(run)); + let toolset = SkillToolset::new(&skills, Some(slot)).unwrap(); toolset .load .call(LoadSkillArgs { diff --git a/crates/aura/src/skill_tool.rs b/crates/aura/src/skill_tool.rs index a37f55c53..b308be76a 100644 --- a/crates/aura/src/skill_tool.rs +++ b/crates/aura/src/skill_tool.rs @@ -11,6 +11,7 @@ //! system prompt small. use crate::config::SkillConfig; +use crate::run_context::BoundRun; use crate::session_store::{ SKILL_INVOCATION_RECORD_VERSION, SkillInvocation, SkillInvocationRecord, SkillInvocationStore, SkillLogKey, @@ -110,7 +111,7 @@ impl std::fmt::Debug for SkillInvocationRecorder { #[derive(Debug, Clone)] pub struct LoadSkillTool { skills: Arc<[SkillConfig]>, - recorder: Option>, + run: Option>, } #[derive(Debug, thiserror::Error)] @@ -272,7 +273,7 @@ impl SkillResourcePath { #[derive(Debug, Clone)] pub struct ReadSkillFileTool { skills: Arc<[SkillConfig]>, - recorder: Option>, + run: Option>, } #[derive(Debug, Deserialize, Serialize)] @@ -321,7 +322,7 @@ impl Tool for ReadSkillFileTool { .ok_or_else(|| SkillError::UnknownSkill(args.skill.clone()))?; let content = render_read_skill_file_output(skill, &args.path).await?; - if let Some(recorder) = &self.recorder { + if let Some(recorder) = run_recorder(self.run.as_deref()) { recorder .record(SkillInvocation::ReadSkillFile { skill: args.skill, @@ -368,11 +369,9 @@ pub struct SkillToolset { impl SkillToolset { /// Build both skill tools, or `None` when no skills are configured. - /// Both share `skills` and `recorder`. - pub fn new( - skills: &[SkillConfig], - recorder: Option>, - ) -> Option { + /// Both share `skills`, and record invocations to the recorder of + /// whichever run `run` holds when they are called. + pub fn new(skills: &[SkillConfig], run: Option>) -> Option { if skills.is_empty() { return None; } @@ -380,13 +379,19 @@ impl SkillToolset { Some(Self { load: LoadSkillTool { skills: Arc::clone(&skills), - recorder: recorder.clone(), + run: run.clone(), }, - read_file: ReadSkillFileTool { skills, recorder }, + read_file: ReadSkillFileTool { skills, run }, }) } } +/// The recorder of the run bound in `run`. A tool built without a slot, or +/// called with no run bound, records nothing. +fn run_recorder(run: Option<&BoundRun>) -> Option> { + run.and_then(BoundRun::skill_recorder) +} + impl LoadSkillTool { /// Create a new LoadSkillTool from discovered skill configs, without /// invocation recording. @@ -395,7 +400,7 @@ impl LoadSkillTool { pub fn new(skills: &[SkillConfig]) -> Self { Self { skills: skills.into(), - recorder: None, + run: None, } } @@ -442,7 +447,7 @@ impl Tool for LoadSkillTool { .ok_or_else(|| SkillError::UnknownSkill(args.name.clone()))?; let result = render_load_skill_output(skill).await?; - if let Some(recorder) = &self.recorder { + if let Some(recorder) = run_recorder(self.run.as_deref()) { recorder .record(SkillInvocation::LoadSkill { name: args.name }) .await; @@ -854,6 +859,50 @@ mod tests { assert!(SkillToolset::new(&[], None).is_none()); } + /// A toolset built once records each invocation under the turn of the + /// run bound when it is called, and records nothing with no run bound. + #[tokio::test] + async fn skill_tools_record_under_the_run_bound_at_call_time() { + use crate::run_context::RunContext; + use crate::session_store::InMemorySkillInvocationStore; + use tokio_util::sync::CancellationToken; + + let dir = TempDir::new().unwrap(); + let configs = make_skill_configs(dir.path()); + let store = Arc::new(InMemorySkillInvocationStore::new()); + let log = SkillLogKey::new(crate::config::SessionId::new("sess-reuse"), "agent"); + let turn = |id: &str, anchor| { + let recorder = Arc::new(SkillInvocationRecorder::new( + store.clone(), + log.clone(), + anchor, + )); + RunContext::channel_for_agent(id, CancellationToken::new(), None, None, Some(recorder)) + .0 + }; + let load = |name: &str| LoadSkillArgs { + name: name.to_string(), + }; + let slot = Arc::new(BoundRun::default()); + let toolset = SkillToolset::new(&configs, Some(Arc::clone(&slot))).unwrap(); + + toolset.load.call(load("another-skill")).await.unwrap(); + assert!(store.list(&log).await.unwrap().is_empty()); + + slot.bind(turn("req_a", 1)); + toolset.load.call(load("test-skill")).await.unwrap(); + slot.bind(turn("req_b", 3)); + toolset.load.call(load("another-skill")).await.unwrap(); + + let records = store.list(&log).await.unwrap(); + let positions: Vec<_> = records.iter().map(|r| (r.anchor, r.seq)).collect(); + assert_eq!( + positions, + vec![(1, 0), (3, 0)], + "each turn's invocation carries that turn's anchor, counted from zero", + ); + } + #[test] fn test_is_skill_tool() { assert!(is_skill_tool(LOAD_SKILL_TOOL_NAME)); diff --git a/crates/aura/src/turn_nudge.rs b/crates/aura/src/turn_nudge.rs index a041420f6..6ab647f96 100644 --- a/crates/aura/src/turn_nudge.rs +++ b/crates/aura/src/turn_nudge.rs @@ -5,7 +5,8 @@ //! `Agent::count_turns` tracks the turn number (one `StreamItem::TurnUsage` //! per rig turn); [`TurnNudgeWrapper`] appends a notice to MCP tool output //! and [`NudgedTool`] to scratchpad read tool output, which rig feeds back -//! as the next turn's prompt. Enabled via `[agent].nudge_last_turn` and +//! as the next turn's prompt. Both reach the counters of the run they serve +//! through a [`BoundRun`]. Enabled via `[agent].nudge_last_turn` and //! `[agent].nudge_turns_remaining`. use std::sync::Arc; @@ -15,9 +16,10 @@ use async_trait::async_trait; use serde_json::Value; use crate::mcp::CallOutcome; +use crate::run_context::BoundRun; use crate::tool_wrapper::{ToolCallContext, ToolWrapper, TransformOutputResult}; -/// Shared turn-limit tracking for one agent stream. +/// Turn-limit tracking for one run of an agent. pub struct TurnNudgeState { /// Total turns rig will execute before `MaxDepthError`. max_turns: usize, @@ -59,16 +61,41 @@ impl TurnNudgeState { if !nudge_last_turn && nudge_turns_remaining.is_none() { return None; } - Some(Arc::new(Self { - // Rig's streaming loop breaks when its pre-increment counter - // exceeds max_depth + 1, i.e. it runs max_depth + 2 turns. - max_turns: max_depth + 2, + // Rig's streaming loop breaks when its pre-increment counter + // exceeds max_depth + 1, i.e. it runs max_depth + 2 turns. + Some(Self::with_limits( + max_depth + 2, nudge_last_turn, - wrap_up_threshold: nudge_turns_remaining, + nudge_turns_remaining, + has_submit_tool, + )) + } + + /// Tracking against these limits with no turns completed. + fn with_limits( + max_turns: usize, + nudge_last_turn: bool, + wrap_up_threshold: Option, + has_submit_tool: bool, + ) -> Arc { + Arc::new(Self { + max_turns, + nudge_last_turn, + wrap_up_threshold, has_submit_tool, turns_completed: AtomicUsize::new(0), last_nudged_turn: AtomicUsize::new(0), - })) + }) + } + + /// Tracking with the same limits and no turns completed. + pub fn fresh(&self) -> Arc { + Self::with_limits( + self.max_turns, + self.nudge_last_turn, + self.wrap_up_threshold, + self.has_submit_tool, + ) } /// Reset counters at stream start. @@ -150,14 +177,28 @@ impl TurnNudgeState { } } +/// Append the bound run's nudge, if one is due, to `tool`'s output. The +/// output passes through untouched when no run is bound or the bound run has +/// nudging off. +fn append_nudge(run: &BoundRun, tool: &str, output: String) -> String { + match run.turn_nudge().and_then(|state| state.nudge_message()) { + Some(nudge) => { + tracing::debug!(tool, "appending turn-limit nudge to tool output"); + format!("{output}{nudge}") + } + None => output, + } +} + /// ToolWrapper that appends turn-limit nudges to tool output. pub struct TurnNudgeWrapper { - state: Arc, + /// The prepared agent's run slot. + run: Arc, } impl TurnNudgeWrapper { - pub fn new(state: Arc) -> Self { - Self { state } + pub fn new(run: Arc) -> Self { + Self { run } } } @@ -170,13 +211,7 @@ impl ToolWrapper for TurnNudgeWrapper { ctx: &ToolCallContext, _extracted: Option<&Value>, ) -> TransformOutputResult { - match self.state.nudge_message() { - Some(nudge) => { - tracing::debug!(tool = %ctx.tool_name, "appending turn-limit nudge to tool output"); - TransformOutputResult::new(format!("{output}{nudge}")) - } - None => TransformOutputResult::new(output), - } + TransformOutputResult::new(append_nudge(&self.run, &ctx.tool_name, output)) } } @@ -186,12 +221,13 @@ impl ToolWrapper for TurnNudgeWrapper { #[derive(Clone)] pub struct NudgedTool { inner: T, - state: Option>, + /// The prepared agent's run slot. + run: Arc, } impl NudgedTool { - pub fn new(inner: T, state: Option>) -> Self { - Self { inner, state } + pub fn new(inner: T, run: Arc) -> Self { + Self { inner, run } } } @@ -214,13 +250,7 @@ where async fn call(&self, args: Self::Args) -> Result { let output = self.inner.call(args).await?; - match self.state.as_ref().and_then(|s| s.nudge_message()) { - Some(nudge) => { - tracing::debug!(tool = %self.inner.name(), "appending turn-limit nudge to tool output"); - Ok(format!("{output}{nudge}")) - } - None => Ok(output), - } + Ok(append_nudge(&self.run, &self.inner.name(), output)) } } @@ -331,6 +361,29 @@ mod tests { assert!(state.nudge_message().is_none()); } + #[test] + fn fresh_keeps_the_limits_and_starts_the_count_over() { + let seed = TurnNudgeState::new(true, None, 1).unwrap(); // 3 turns total + advance(&seed, 1); // turn 2, remaining 1 + assert!(seed.nudge_message().is_some()); + + let run = seed.fresh(); + assert!( + run.nudge_message().is_none(), + "a fresh run is back in turn 1" + ); + advance(&run, 1); + assert!( + run.nudge_message().is_some(), + "the fresh run nudges at the same limit as the seed", + ); + assert_eq!( + seed.turns_completed.load(Ordering::Acquire), + 1, + "advancing the fresh run leaves the seed untouched", + ); + } + #[derive(Clone)] struct EchoTool; @@ -358,7 +411,7 @@ mod tests { use rig::tool::Tool; let state = TurnNudgeState::new(true, None, 1).unwrap(); // 3 turns total - let tool = NudgedTool::new(EchoTool, Some(state.clone())); + let tool = NudgedTool::new(EchoTool, Arc::new(BoundRun::pinned_nudge(state.clone()))); // Turn 1 (remaining 2): output passes through untouched. let out = tool.call("hello".to_string()).await.unwrap(); @@ -376,12 +429,46 @@ mod tests { } #[tokio::test] - async fn nudged_tool_without_state_is_passthrough() { + async fn nudged_tool_without_a_run_is_passthrough() { use rig::tool::Tool; - let tool = NudgedTool::new(EchoTool, None); + let tool = NudgedTool::new(EchoTool, Arc::new(BoundRun::default())); assert_eq!(tool.name(), "echo"); let out = tool.call("hello".to_string()).await.unwrap(); assert_eq!(out, "hello"); } + + /// The wrapper and tool hold the slot, not the counters, so the nudge + /// follows whichever run is bound when the call happens. + #[tokio::test] + async fn nudged_tool_follows_the_run_bound_at_call_time() { + use crate::run_context::RunContext; + use rig::tool::Tool; + + let slot = Arc::new(BoundRun::default()); + let tool = NudgedTool::new(EchoTool, Arc::clone(&slot)); + let seed = TurnNudgeState::new(true, None, 1).unwrap(); // 3 turns total + + let first_nudge = seed.fresh(); + slot.bind(RunContext::detached_with( + "req_a", + None, + Some(Arc::clone(&first_nudge)), + )); + first_nudge.record_turn_completed(); + assert!( + tool.call("hello".to_string()) + .await + .unwrap() + .contains("FINAL TURN"), + "the first run is on its penultimate turn", + ); + + slot.bind(RunContext::detached_with("req_b", None, Some(seed.fresh()))); + assert_eq!( + tool.call("hello".to_string()).await.unwrap(), + "hello", + "the second run starts its count over", + ); + } }