From 41ed0a0f9814d754c17df80c14d263ae10e09b45 Mon Sep 17 00:00:00 2001 From: Illia Polosukhin Date: Wed, 25 Mar 2026 08:35:41 -0700 Subject: [PATCH 1/8] feat(agent): thread per-tool reasoning through provider, session, and all surfaces (#1513) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat(agent): thread per-tool reasoning from LLM through to REPL, HTTP, SSE, and DB Add end-to-end agent reasoning summaries so users can see *why* the agent chose specific tools, not just what it did. - Add `reasoning: Option` to `ToolCall` (all providers) - Populate from LLM response content in `Reasoning::respond_with_tools` and `select_tools`, with per-tool override when providers supply it - Extend `Turn` with `narrative` and `TurnToolCall` with `rationale` + `tool_call_id` for identity-based result matching - Persist reasoning in DB via existing tool_calls JSON (no migration) - Add `StatusUpdate::ReasoningUpdate` and `SseEvent::ReasoningUpdate` + `SseEvent::JobReasoning` for real-time streaming - Emit reasoning events in both chat dispatcher and worker job path - Add `/reasoning [N|all]` command for inspecting turn reasoning - Surface `narrative` and `rationale` in HTTP `/api/chat/history` Based on the design from #361 and #456, reconstructed cleanly with Option to minimize blast radius (vs mandatory String that broke compilation in #456). Closes #456 Co-Authored-By: panosAthDBX <47406510+panosAthDBX@users.noreply.github.com> Co-Authored-By: Claude Opus 4.6 (1M context) * fix: address PR review feedback from Gemini and Copilot - Fix `_ => Ok(None)` in agent_loop.rs to avoid accidental shutdown - Fix fallback in record_tool_result_for/record_tool_error_for to use first pending call instead of last_mut (parallel execution safety) - Include per-tool decisions in WASM channel reasoning messages - Apply truncate_at_tool_tags + clean_response to shared_reasoning in select_tools (parity with respond_with_tools) - Persist turn-level narrative to DB in tool_calls JSON wrapper - Parse both old (array) and new (object) tool_calls formats in build_turns_from_db_messages for backward compatibility - Populate reasoning from action.reasoning in execute_plan ToolCalls [skip-regression-check] Co-Authored-By: Claude Opus 4.6 (1M context) * fix: address second round of review comments + merge fixes - Add reasoning: None to new github_copilot.rs ToolCall sites (from staging merge) - Run cargo fmt on 4 files with formatting diffs - Truncate narrative to 1000 chars before DB persistence - Clone turn data and drop session lock in /reasoning command - Extract ToolDecisionDto::from_json_array shared helper (deduplicate worker/job.rs and orchestrator/api.rs) - Add unit tests for wrapped tool_calls JSON format with narrative [skip-regression-check] Co-Authored-By: Claude Opus 4.6 (1M context) * fix: address third round of review comments (Copilot + serrrfirat) - Reword ToolCall.reasoning docstring to reflect provider-supplied or fallback contract - Sanitize narrative through SafetyLayer before storage/emission - Clean per-tool reasoning via truncate_at_tool_tags + clean_response in select_tools (parity with shared reasoning) - Convert 4 approval-path recording sites in thread_ops.rs to identity-based record_tool_result_for/record_tool_error_for - Preserve tool_call_id and reasoning through restore_from_messages - Fix has_result/has_error to reject JSON null values - Truncate tool_call_id to 128 chars before DB persistence - Add 4 unit tests for record_tool_result_for/error_for edge cases Co-Authored-By: Claude Opus 4.6 (1M context) * fix: address zmanian review — sanitize JobDelegate reasoning + warn on dropped results - Sanitize narrative and per-tool rationale through SafetyLayer in JobDelegate reasoning events (parity with ChatDelegate) - Add tracing::warn when record_tool_result_for/error_for drops a result because no matching or pending tool call exists - Add 3 unit tests for reasoning normalization (thinking tags, tool tags, empty-after-cleaning) Co-Authored-By: Claude Opus 4.6 (1M context) * fix: address 4 remaining unreplied review comments - Clean per-tool reasoning in respond_with_tools via truncate_at_tool_tags + clean_response (parity with select_tools) - Handle wrapped JSON format in rebuild_chat_messages_from_db so cold hydration works after persist_tool_calls format change - Update persist_tool_calls doc comment to describe new JSON shape - Sanitize per-tool rationale through SafetyLayer in ChatDelegate before emission and storage (parity with JobDelegate) Co-Authored-By: Claude Opus 4.6 (1M context) * fix: address zmanian review round 2 - Add tracing::debug on fallback-to-pending path in record_tool_result_for and record_tool_error_for (item 1) - Add comment explaining why /reasoning is special-cased in agent_loop.rs (item 4) - Items 2 (narrative persistence), 3 (rationale sanitization), and 5 (catch-all fix) were already addressed in prior commits Co-Authored-By: Claude Opus 4.6 (1M context) --------- Co-authored-by: panosAthDBX <47406510+panosAthDBX@users.noreply.github.com> Co-authored-by: Claude Opus 4.6 (1M context) --- crates/ironclaw_common/src/event.rs | 55 ++++++++ crates/ironclaw_common/src/lib.rs | 2 +- src/agent/agent_loop.rs | 16 +++ src/agent/agentic_loop.rs | 1 + src/agent/commands.rs | 89 +++++++++++++ src/agent/dispatcher.rs | 84 +++++++++++- src/agent/session.rs | 193 +++++++++++++++++++++++++++- src/agent/submission.rs | 11 ++ src/agent/thread_ops.rs | 69 ++++++++-- src/channels/channel.rs | 16 +++ src/channels/mod.rs | 2 +- src/channels/repl.rs | 14 ++ src/channels/wasm/wrapper.rs | 14 ++ src/channels/web/handlers/chat.rs | 2 + src/channels/web/mod.rs | 14 ++ src/channels/web/openai_compat.rs | 2 + src/channels/web/server.rs | 2 + src/channels/web/types.rs | 8 +- src/channels/web/util.rs | 99 ++++++++++++-- src/llm/anthropic_oauth.rs | 2 + src/llm/bedrock.rs | 7 + src/llm/codex_chatgpt.rs | 2 + src/llm/gemini_oauth.rs | 1 + src/llm/github_copilot.rs | 2 + src/llm/nearai_chat.rs | 7 + src/llm/openai_codex_provider.rs | 5 + src/llm/provider.rs | 8 ++ src/llm/reasoning.rs | 97 ++++++++++++-- src/llm/rig_adapter.rs | 7 + src/orchestrator/api.rs | 15 +++ src/worker/job.rs | 68 +++++++++- tests/openai_compat_integration.rs | 1 + tests/support/trace_llm.rs | 1 + 33 files changed, 871 insertions(+), 45 deletions(-) diff --git a/crates/ironclaw_common/src/event.rs b/crates/ironclaw_common/src/event.rs index 83592c95..256aba3d 100644 --- a/crates/ironclaw_common/src/event.rs +++ b/crates/ironclaw_common/src/event.rs @@ -7,6 +7,32 @@ use serde::{Deserialize, Serialize}; +/// A single tool decision in a reasoning update (SSE DTO). +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ToolDecisionDto { + pub tool_name: String, + pub rationale: String, +} + +impl ToolDecisionDto { + /// Parse a list of tool decisions from a JSON array value. + pub fn from_json_array(value: &serde_json::Value) -> Vec { + value + .as_array() + .map(|arr| { + arr.iter() + .filter_map(|d| { + Some(Self { + tool_name: d.get("tool_name")?.as_str()?.to_string(), + rationale: d.get("rationale")?.as_str()?.to_string(), + }) + }) + .collect() + }) + .unwrap_or_default() + } +} + #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(tag = "type")] pub enum AppEvent { @@ -163,6 +189,23 @@ pub enum AppEvent { #[serde(skip_serializing_if = "Option::is_none")] message: Option, }, + + /// Agent reasoning update (why it chose specific tools). + #[serde(rename = "reasoning_update")] + ReasoningUpdate { + narrative: String, + decisions: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + + /// Reasoning update for a sandbox job. + #[serde(rename = "job_reasoning")] + JobReasoning { + job_id: String, + narrative: String, + decisions: Vec, + }, } impl AppEvent { @@ -191,6 +234,8 @@ impl AppEvent { Self::Suggestions { .. } => "suggestions", Self::TurnCost { .. } => "turn_cost", Self::ExtensionStatus { .. } => "extension_status", + Self::ReasoningUpdate { .. } => "reasoning_update", + Self::JobReasoning { .. } => "job_reasoning", } } } @@ -311,6 +356,16 @@ mod tests { status: String::new(), message: None, }, + AppEvent::ReasoningUpdate { + narrative: String::new(), + decisions: vec![], + thread_id: None, + }, + AppEvent::JobReasoning { + job_id: String::new(), + narrative: String::new(), + decisions: vec![], + }, ]; for variant in &variants { diff --git a/crates/ironclaw_common/src/lib.rs b/crates/ironclaw_common/src/lib.rs index 6822bad1..f52dc0aa 100644 --- a/crates/ironclaw_common/src/lib.rs +++ b/crates/ironclaw_common/src/lib.rs @@ -3,5 +3,5 @@ mod event; mod util; -pub use event::AppEvent; +pub use event::{AppEvent, ToolDecisionDto}; pub use util::truncate_preview; diff --git a/src/agent/agent_loop.rs b/src/agent/agent_loop.rs index 7e950146..f51a8db1 100644 --- a/src/agent/agent_loop.rs +++ b/src/agent/agent_loop.rs @@ -1250,6 +1250,22 @@ impl Agent { command, message.channel ); + // /reasoning is special-cased here (not in handle_system_command) + // because it needs the session + thread_id to read turn reasoning + // data, which handle_system_command's signature doesn't provide. + if command == "reasoning" { + let result = self + .handle_reasoning_command(&args, &session, thread_id) + .await; + return match result { + SubmissionResult::Response { content } => Ok(Some(content)), + SubmissionResult::Ok { message } => Ok(message), + SubmissionResult::Error { message } => { + Ok(Some(format!("Error: {}", message))) + } + _ => Ok(Some(String::new())), + }; + } // Authorization checks (including restart channel check) are enforced in handle_system_command self.handle_system_command(&command, &args, &message.channel) .await diff --git a/src/agent/agentic_loop.rs b/src/agent/agentic_loop.rs index cc6fd486..e61856dc 100644 --- a/src/agent/agentic_loop.rs +++ b/src/agent/agentic_loop.rs @@ -414,6 +414,7 @@ mod tests { id: "call_1".to_string(), name: "echo".to_string(), arguments: serde_json::json!({}), + reasoning: None, }; let delegate = MockDelegate::new(vec![ tool_calls_output(vec![tool_call]), diff --git a/src/agent/commands.rs b/src/agent/commands.rs index b6aff3c0..e02b33db 100644 --- a/src/agent/commands.rs +++ b/src/agent/commands.rs @@ -465,6 +465,94 @@ impl Agent { } } + /// Handle `/reasoning [N|all]` — show reasoning history for the active thread. + pub(super) async fn handle_reasoning_command( + &self, + args: &[String], + session: &Arc>, + thread_id: Uuid, + ) -> SubmissionResult { + // Clone the turn data we need, then drop the session lock. + let turns_snapshot: Vec<( + usize, + Option, + Vec, + )>; + { + let sess = session.lock().await; + let thread = match sess.threads.get(&thread_id) { + Some(t) => t, + None => return SubmissionResult::error("No active thread."), + }; + + if thread.turns.is_empty() { + return SubmissionResult::ok_with_message("No turns yet."); + } + + // Parse argument: default=last turn, "all"=all turns, N=specific turn (1-based). + let selected: Vec<&crate::agent::session::Turn> = match args.first().map(|s| s.as_str()) + { + Some("all") => thread.turns.iter().collect(), + Some(n) => match n.parse::() { + Ok(0) => return SubmissionResult::error("Turn numbers start at 1."), + Ok(num) if num > thread.turns.len() => { + return SubmissionResult::error(format!( + "Turn {} does not exist (max: {}).", + num, + thread.turns.len() + )); + } + Ok(num) => vec![&thread.turns[num - 1]], + Err(_) => return SubmissionResult::error("Usage: /reasoning [N|all]"), + }, + None => { + // Default: last turn that has tool calls + match thread.turns.iter().rev().find(|t| !t.tool_calls.is_empty()) { + Some(t) => vec![t], + None => { + return SubmissionResult::ok_with_message("No turns with tool calls."); + } + } + } + }; + + turns_snapshot = selected + .into_iter() + .map(|t| (t.turn_number, t.narrative.clone(), t.tool_calls.clone())) + .collect(); + } + // Session lock is now dropped — format output without holding it. + + let mut output = String::new(); + for (turn_number, narrative, tool_calls) in &turns_snapshot { + output.push_str(&format!("--- Turn {} ---\n", turn_number + 1)); + if let Some(narrative) = narrative { + output.push_str(&format!("Reasoning: {}\n", narrative)); + } + if tool_calls.is_empty() { + output.push_str(" (no tool calls)\n"); + } else { + for tc in tool_calls { + let status = if tc.error.is_some() { + "error" + } else if tc.result.is_some() { + "ok" + } else { + "pending" + }; + output.push_str(&format!(" {} [{}]", tc.name, status)); + if let Some(ref rationale) = tc.rationale { + output.push_str(&format!(" — {}", rationale)); + } + output.push('\n'); + } + } + output.push('\n'); + } + + SubmissionResult::response(output.trim_end()) + } + /// Handle system commands that bypass thread-state checks entirely. pub(super) async fn handle_system_command( &self, @@ -480,6 +568,7 @@ impl Agent { " /version Show version info\n", " /tools List available tools\n", " /debug Toggle debug mode\n", + " /reasoning [N|all] Show agent reasoning for turns\n", " /ping Connectivity check\n", "\n", "Jobs:\n", diff --git a/src/agent/dispatcher.rs b/src/agent/dispatcher.rs index a195458d..cba84c35 100644 --- a/src/agent/dispatcher.rs +++ b/src/agent/dispatcher.rs @@ -420,6 +420,19 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { content: Option, reason_ctx: &mut ReasoningContext, ) -> Result, Error> { + // Extract and sanitize the narrative before consuming `content`. + let narrative = content + .as_deref() + .filter(|c| !c.trim().is_empty()) + .map(|c| { + let sanitized = self + .agent + .safety() + .sanitize_tool_output("agent_narrative", c); + sanitized.content + }) + .filter(|c| !c.trim().is_empty()); + // Add the assistant message with tool_calls to context. // OpenAI protocol requires this before tool-result messages. reason_ctx @@ -440,6 +453,41 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { ) .await; + // Build per-tool decisions for the reasoning update. + // Sanitize each rationale through SafetyLayer (parity with JobDelegate). + let decisions: Vec = tool_calls + .iter() + .filter_map(|tc| { + tc.reasoning.as_ref().map(|r| { + let sanitized = self + .agent + .safety() + .sanitize_tool_output("tool_rationale", r) + .content; + crate::channels::ToolDecision { + tool_name: tc.name.clone(), + rationale: sanitized, + } + }) + }) + .collect(); + + // Emit reasoning update to channels. + if narrative.is_some() || !decisions.is_empty() { + let _ = self + .agent + .channels + .send_status( + &self.message.channel, + StatusUpdate::ReasoningUpdate { + narrative: narrative.clone().unwrap_or_default(), + decisions: decisions.clone(), + }, + &self.message.metadata, + ) + .await; + } + // Record tool calls in the thread with sensitive params redacted. { let mut redacted_args: Vec = Vec::with_capacity(tool_calls.len()); @@ -455,8 +503,23 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { if let Some(thread) = sess.threads.get_mut(&self.thread_id) && let Some(turn) = thread.last_turn_mut() { + // Set turn-level narrative. + if turn.narrative.is_none() { + turn.narrative = narrative; + } for (tc, safe_args) in tool_calls.iter().zip(redacted_args) { - turn.record_tool_call(&tc.name, safe_args); + let sanitized_rationale = tc.reasoning.as_ref().map(|r| { + self.agent + .safety() + .sanitize_tool_output("tool_rationale", r) + .content + }); + turn.record_tool_call_with_reasoning( + &tc.name, + safe_args, + sanitized_rationale, + Some(tc.id.clone()), + ); } } } @@ -726,7 +789,7 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { if let Some(thread) = sess.threads.get_mut(&self.thread_id) && let Some(turn) = thread.last_turn_mut() { - turn.record_tool_error(error_msg.clone()); + turn.record_tool_error_for(&tc.id, error_msg.clone()); } } reason_ctx @@ -852,16 +915,19 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { Err(e) => format!("Tool '{}' failed: {}", tc.name, e), }; - // Record sanitized result in thread + // Record sanitized result in thread (identity-based matching). { let mut sess = self.session.lock().await; if let Some(thread) = sess.threads.get_mut(&self.thread_id) && let Some(turn) = thread.last_turn_mut() { if is_tool_error { - turn.record_tool_error(result_content.clone()); + turn.record_tool_error_for(&tc.id, result_content.clone()); } else { - turn.record_tool_result(serde_json::json!(result_content)); + turn.record_tool_result_for( + &tc.id, + serde_json::json!(result_content), + ); } } } @@ -1462,11 +1528,13 @@ mod tests { id: "call_2".to_string(), name: "http".to_string(), arguments: serde_json::json!({"url": "https://example.com"}), + reasoning: None, }, ToolCall { id: "call_3".to_string(), name: "echo".to_string(), arguments: serde_json::json!({"message": "done"}), + reasoning: None, }, ], user_timezone: None, @@ -1652,6 +1720,7 @@ mod tests { id: "call_1".to_string(), name: "echo".to_string(), arguments: serde_json::json!({"message": "hi"}), + reasoning: None, }], ), ChatMessage::tool_result("call_1", "echo", "hi"), @@ -1744,11 +1813,13 @@ mod tests { id: "c1".to_string(), name: "http".to_string(), arguments: serde_json::json!({}), + reasoning: None, }, ToolCall { id: "c2".to_string(), name: "echo".to_string(), arguments: serde_json::json!({}), + reasoning: None, }, ], ), @@ -1782,6 +1853,7 @@ mod tests { id: "c1".to_string(), name: "echo".to_string(), arguments: serde_json::json!({}), + reasoning: None, }], ), ChatMessage::tool_result("c1", "echo", "done"), @@ -1912,6 +1984,7 @@ mod tests { id: crate::llm::generate_tool_call_id(0, 0), name: "echo".to_string(), arguments: serde_json::json!({"message": "looping"}), + reasoning: None, }], input_tokens: 0, output_tokens: 5, @@ -2065,6 +2138,7 @@ mod tests { id: crate::llm::generate_tool_call_id(0, 0), name: "nonexistent_tool".to_string(), arguments: serde_json::json!({}), + reasoning: None, }], input_tokens: 0, output_tokens: 5, diff --git a/src/agent/session.rs b/src/agent/session.rs index 7ec2023f..6c873e46 100644 --- a/src/agent/session.rs +++ b/src/agent/session.rs @@ -449,6 +449,7 @@ impl Thread { id: call_id.clone(), name: tc.name.clone(), arguments: tc.parameters.clone(), + reasoning: None, }) .collect(); @@ -522,7 +523,12 @@ impl Thread { && let Some(ref tcs) = assistant_msg.tool_calls { for tc in tcs { - turn.record_tool_call(&tc.name, tc.arguments.clone()); + turn.record_tool_call_with_reasoning( + &tc.name, + tc.arguments.clone(), + tc.reasoning.clone(), + Some(tc.id.clone()), + ); } } @@ -602,6 +608,10 @@ pub struct Turn { pub completed_at: Option>, /// Error message (if failed). pub error: Option, + /// Agent's reasoning narrative for this turn. + /// Cleaned via `clean_response` and sanitized through `SafetyLayer` before storage. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub narrative: Option, /// Transient image content parts for multimodal LLM input. /// Not serialized — images are only needed for the current LLM call. /// The text description in `user_input` persists for compaction/context. @@ -621,6 +631,7 @@ impl Turn { started_at: Utc::now(), completed_at: None, error: None, + narrative: None, image_content_parts: Vec::new(), } } @@ -656,6 +667,26 @@ impl Turn { parameters: params, result: None, error: None, + rationale: None, + tool_call_id: None, + }); + } + + /// Record a tool call with reasoning context. + pub fn record_tool_call_with_reasoning( + &mut self, + name: impl Into, + params: serde_json::Value, + rationale: Option, + tool_call_id: Option, + ) { + self.tool_calls.push(TurnToolCall { + name: name.into(), + parameters: params, + result: None, + error: None, + rationale, + tool_call_id, }); } @@ -672,6 +703,60 @@ impl Turn { call.error = Some(error.into()); } } + + /// Record a tool result by tool_call_id, with fallback to first pending call. + pub fn record_tool_result_for(&mut self, tool_call_id: &str, result: serde_json::Value) { + if let Some(call) = self + .tool_calls + .iter_mut() + .find(|c| c.tool_call_id.as_deref() == Some(tool_call_id)) + { + call.result = Some(result); + } else if let Some(call) = self + .tool_calls + .iter_mut() + .find(|c| c.result.is_none() && c.error.is_none()) + { + tracing::debug!( + tool_call_id = %tool_call_id, + fallback_tool = %call.name, + "tool_call_id not found, falling back to first pending call" + ); + call.result = Some(result); + } else { + tracing::warn!( + tool_call_id = %tool_call_id, + "Tool result dropped: no matching or pending tool call" + ); + } + } + + /// Record a tool error by tool_call_id, with fallback to first pending call. + pub fn record_tool_error_for(&mut self, tool_call_id: &str, error: impl Into) { + if let Some(call) = self + .tool_calls + .iter_mut() + .find(|c| c.tool_call_id.as_deref() == Some(tool_call_id)) + { + call.error = Some(error.into()); + } else if let Some(call) = self + .tool_calls + .iter_mut() + .find(|c| c.result.is_none() && c.error.is_none()) + { + tracing::debug!( + tool_call_id = %tool_call_id, + fallback_tool = %call.name, + "tool_call_id not found, falling back to first pending call" + ); + call.error = Some(error.into()); + } else { + tracing::warn!( + tool_call_id = %tool_call_id, + "Tool error dropped: no matching or pending tool call" + ); + } + } } /// Record of a tool call made during a turn. @@ -685,6 +770,12 @@ pub struct TurnToolCall { pub result: Option, /// Error from the tool (if failed). pub error: Option, + /// Agent's reasoning for choosing this tool. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub rationale: Option, + /// The tool_call_id from the LLM, for identity-based result matching. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tool_call_id: Option, } #[cfg(test)] @@ -1309,6 +1400,7 @@ mod tests { id: "call_0".to_string(), name: "search".to_string(), arguments: serde_json::json!({"q": "test"}), + reasoning: None, }; let messages = vec![ ChatMessage::user("Find test"), @@ -1339,6 +1431,7 @@ mod tests { id: "call_0".to_string(), name: "http".to_string(), arguments: serde_json::json!({}), + reasoning: None, }; let messages = vec![ ChatMessage::user("Fetch URL"), @@ -1404,11 +1497,13 @@ mod tests { id: "call_a".to_string(), name: "search".to_string(), arguments: serde_json::json!({"q": "data"}), + reasoning: None, }; let tc2 = ToolCall { id: "call_b".to_string(), name: "write".to_string(), arguments: serde_json::json!({"path": "out.txt"}), + reasoning: None, }; let messages = vec![ ChatMessage::user("Find and save"), @@ -1620,4 +1715,100 @@ mod tests { let merged = thread.drain_pending_messages().unwrap(); assert_eq!(merged, "failed batch\nnew msg"); } + + #[test] + fn test_record_tool_result_for_by_id() { + let mut turn = Turn::new(0, "test"); + turn.record_tool_call_with_reasoning( + "tool_a", + serde_json::json!({}), + None, + Some("id_a".into()), + ); + turn.record_tool_call_with_reasoning( + "tool_b", + serde_json::json!({}), + None, + Some("id_b".into()), + ); + + // Record result for second tool by ID + turn.record_tool_result_for("id_b", serde_json::json!("result_b")); + assert!(turn.tool_calls[0].result.is_none()); + assert_eq!( + turn.tool_calls[1].result.as_ref().unwrap(), + &serde_json::json!("result_b") + ); + } + + #[test] + fn test_record_tool_error_for_by_id() { + let mut turn = Turn::new(0, "test"); + turn.record_tool_call_with_reasoning( + "tool_a", + serde_json::json!({}), + None, + Some("id_a".into()), + ); + turn.record_tool_call_with_reasoning( + "tool_b", + serde_json::json!({}), + None, + Some("id_b".into()), + ); + + turn.record_tool_error_for("id_a", "failed"); + assert_eq!(turn.tool_calls[0].error.as_deref(), Some("failed")); + assert!(turn.tool_calls[1].error.is_none()); + } + + #[test] + fn test_record_tool_result_for_fallback_to_pending() { + let mut turn = Turn::new(0, "test"); + turn.record_tool_call_with_reasoning( + "tool_a", + serde_json::json!({}), + None, + Some("id_a".into()), + ); + turn.record_tool_call_with_reasoning( + "tool_b", + serde_json::json!({}), + None, + Some("id_b".into()), + ); + + // First tool already has a result + turn.tool_calls[0].result = Some(serde_json::json!("done")); + + // Unknown ID should fall back to first pending (tool_b) + turn.record_tool_result_for("unknown_id", serde_json::json!("fallback")); + assert_eq!( + turn.tool_calls[0].result.as_ref().unwrap(), + &serde_json::json!("done") + ); + assert_eq!( + turn.tool_calls[1].result.as_ref().unwrap(), + &serde_json::json!("fallback") + ); + } + + #[test] + fn test_record_tool_result_for_no_pending_is_noop() { + let mut turn = Turn::new(0, "test"); + turn.record_tool_call_with_reasoning( + "tool_a", + serde_json::json!({}), + None, + Some("id_a".into()), + ); + turn.tool_calls[0].result = Some(serde_json::json!("done")); + + // No pending calls, unknown ID — should be a no-op + turn.record_tool_result_for("unknown_id", serde_json::json!("lost")); + assert_eq!( + turn.tool_calls[0].result.as_ref().unwrap(), + &serde_json::json!("done") + ); + } } diff --git a/src/agent/submission.rs b/src/agent/submission.rs index 8594c969..5a81e0bf 100644 --- a/src/agent/submission.rs +++ b/src/agent/submission.rs @@ -92,6 +92,17 @@ impl SubmissionParser { args: vec![], }; } + if lower == "/reasoning" || lower.starts_with("/reasoning ") { + let args: Vec = trimmed + .split_whitespace() + .skip(1) + .map(|s| s.to_string()) + .collect(); + return Submission::SystemCommand { + command: "reasoning".to_string(), + args, + }; + } if lower == "/restart" { tracing::debug!("[SubmissionParser::parse] Recognized /restart command"); return Submission::SystemCommand { diff --git a/src/agent/thread_ops.rs b/src/agent/thread_ops.rs index b2820e7e..11f211f9 100644 --- a/src/agent/thread_ops.rs +++ b/src/agent/thread_ops.rs @@ -513,10 +513,10 @@ impl Agent { }; thread.complete_turn(&response); - let (turn_number, tool_calls) = thread + let (turn_number, tool_calls, narrative) = thread .turns .last() - .map(|t| (t.turn_number, t.tool_calls.clone())) + .map(|t| (t.turn_number, t.tool_calls.clone(), t.narrative.clone())) .unwrap_or_default(); let _ = self .channels @@ -534,6 +534,7 @@ impl Agent { &message.user_id, turn_number, &tool_calls, + narrative.as_deref(), ) .await; self.persist_assistant_response( @@ -725,7 +726,9 @@ impl Agent { /// /// Stored between the user and assistant messages so that /// `build_turns_from_db_messages` can reconstruct the tool call history. - /// Content is a JSON array of tool call summaries. + /// Content is a JSON object: `{ "calls": [...], "narrative": "..." }`. + /// The `calls` array contains tool call summaries with optional `rationale` + /// and `tool_call_id` fields. Legacy rows may be plain JSON arrays. pub(super) async fn persist_tool_calls( &self, thread_id: Uuid, @@ -733,6 +736,7 @@ impl Agent { user_id: &str, turn_number: usize, tool_calls: &[crate::agent::session::TurnToolCall], + narrative: Option<&str>, ) { if tool_calls.is_empty() { return; @@ -767,11 +771,30 @@ impl Agent { if let Some(ref error) = tc.error { obj["error"] = serde_json::Value::String(truncate_preview(error, 200)); } + if let Some(ref rationale) = tc.rationale { + obj["rationale"] = serde_json::Value::String(truncate_preview(rationale, 500)); + } + if let Some(ref tool_call_id) = tc.tool_call_id { + obj["tool_call_id"] = + serde_json::Value::String(truncate_preview(tool_call_id, 128)); + } obj }) .collect(); - let content = match serde_json::to_string(&summaries) { + // Wrap in an object with optional narrative so it can be reconstructed. + // safety: no byte-index slicing here; comment describes JSON shape + let wrapper = if let Some(n) = narrative { + serde_json::json!({ + "narrative": truncate_preview(n, 1000), + "calls": summaries, + }) + } else { + serde_json::json!({ + "calls": summaries, + }) + }; + let content = match serde_json::to_string(&wrapper) { Ok(c) => c, Err(e) => { tracing::warn!("Failed to serialize tool calls: {}", e); @@ -1104,9 +1127,12 @@ impl Agent { && let Some(turn) = thread.last_turn_mut() { if is_tool_error { - turn.record_tool_error(result_content.clone()); + turn.record_tool_error_for(&pending.tool_call_id, result_content.clone()); } else { - turn.record_tool_result(serde_json::json!(result_content)); + turn.record_tool_result_for( + &pending.tool_call_id, + serde_json::json!(result_content), + ); } } } @@ -1358,9 +1384,12 @@ impl Agent { && let Some(turn) = thread.last_turn_mut() { if is_deferred_error { - turn.record_tool_error(deferred_content.clone()); + turn.record_tool_error_for(&tc.id, deferred_content.clone()); } else { - turn.record_tool_result(serde_json::json!(deferred_content)); + turn.record_tool_result_for( + &tc.id, + serde_json::json!(deferred_content), + ); } } } @@ -1459,10 +1488,10 @@ impl Agent { let (response, suggestions) = crate::agent::dispatcher::extract_suggestions(&response); thread.complete_turn(&response); - let (turn_number, tool_calls) = thread + let (turn_number, tool_calls, narrative) = thread .turns .last() - .map(|t| (t.turn_number, t.tool_calls.clone())) + .map(|t| (t.turn_number, t.tool_calls.clone(), t.narrative.clone())) .unwrap_or_default(); // User message already persisted at turn start; save tool calls then assistant response self.persist_tool_calls( @@ -1471,6 +1500,7 @@ impl Agent { &message.user_id, turn_number, &tool_calls, + narrative.as_deref(), ) .await; self.persist_assistant_response( @@ -1816,7 +1846,20 @@ fn rebuild_chat_messages_from_db( "assistant" => result.push(ChatMessage::assistant(&msg.content)), "tool_calls" => { // Try to parse the enriched JSON and rebuild tool messages. - if let Ok(calls) = serde_json::from_str::>(&msg.content) { + // Supports two formats: + // - Old: plain JSON array of tool call summaries + // - New: wrapped object { "calls": [...], "narrative": "..." } + let calls: Vec = + match serde_json::from_str::(&msg.content) { + Ok(serde_json::Value::Array(arr)) => arr, + Ok(serde_json::Value::Object(obj)) => obj + .get("calls") + .and_then(|v| v.as_array()) + .cloned() + .unwrap_or_default(), + _ => Vec::new(), + }; + { if calls.is_empty() { continue; } @@ -1839,6 +1882,10 @@ fn rebuild_chat_messages_from_db( .get("parameters") .cloned() .unwrap_or(serde_json::json!({})), + reasoning: c + .get("rationale") + .and_then(|v| v.as_str()) + .map(String::from), }) .collect(); diff --git a/src/channels/channel.rs b/src/channels/channel.rs index 9bcee12e..784b6bcf 100644 --- a/src/channels/channel.rs +++ b/src/channels/channel.rs @@ -265,6 +265,15 @@ impl OutgoingResponse { } } +/// A single tool decision within a reasoning update. +#[derive(Debug, Clone)] +pub struct ToolDecision { + /// Tool name. + pub tool_name: String, + /// Agent's reasoning for choosing this tool. + pub rationale: String, +} + /// Status update types for showing agent activity. #[derive(Debug, Clone)] pub enum StatusUpdate { @@ -333,6 +342,13 @@ pub enum StatusUpdate { }, /// Suggested follow-up messages for the user. Suggestions { suggestions: Vec }, + /// Agent reasoning update (why it chose specific tools). + ReasoningUpdate { + /// Human-readable summary of the agent's decision. + narrative: String, + /// Per-tool decisions. + decisions: Vec, + }, /// Per-turn token usage and cost summary (shown as subtle metadata). TurnCost { input_tokens: u64, diff --git a/src/channels/mod.rs b/src/channels/mod.rs index c0230692..46e25514 100644 --- a/src/channels/mod.rs +++ b/src/channels/mod.rs @@ -39,7 +39,7 @@ mod webhook_server; pub use channel::{ AttachmentKind, Channel, ChannelSecretUpdater, IncomingAttachment, IncomingMessage, - MessageStream, OutgoingResponse, StatusUpdate, routing_target_from_metadata, + MessageStream, OutgoingResponse, StatusUpdate, ToolDecision, routing_target_from_metadata, }; pub use http::{HttpChannel, HttpChannelState}; pub use manager::ChannelManager; diff --git a/src/channels/repl.rs b/src/channels/repl.rs index 055dc3ad..61c68d13 100644 --- a/src/channels/repl.rs +++ b/src/channels/repl.rs @@ -75,6 +75,7 @@ const SLASH_COMMANDS: &[&str] = &[ "/suggest", "/thread", "/resume", + "/reasoning", ]; /// Rustyline helper for slash-command tab completion. @@ -841,6 +842,19 @@ impl Channel for ReplChannel { StatusUpdate::Suggestions { .. } => { // Suggestions are only rendered by the web gateway } + StatusUpdate::ReasoningUpdate { + narrative, + decisions, + } => { + if !narrative.is_empty() { + let display = truncate_for_preview(&narrative, CLI_STATUS_MAX); + eprintln!(" \x1b[94m\u{25B6} {display}\x1b[0m"); + } + for d in &decisions { + let display = truncate_for_preview(&d.rationale, CLI_STATUS_MAX); + eprintln!(" \x1b[90m\u{2192} {}: {display}\x1b[0m", d.tool_name); + } + } StatusUpdate::TurnCost { .. } => { // Cost display is handled by the TUI channel } diff --git a/src/channels/wasm/wrapper.rs b/src/channels/wasm/wrapper.rs index 65e4de88..a0f9689f 100644 --- a/src/channels/wasm/wrapper.rs +++ b/src/channels/wasm/wrapper.rs @@ -3061,6 +3061,20 @@ fn status_to_wit( }, // Suggestions and turn cost are web-gateway-only; skip for WASM channels StatusUpdate::Suggestions { .. } | StatusUpdate::TurnCost { .. } => return None, + StatusUpdate::ReasoningUpdate { + narrative, + decisions, + } => { + let mut msg = narrative.clone(); + for d in decisions { + msg.push_str(&format!("\n → {}: {}", d.tool_name, d.rationale)); + } + wit_channel::StatusUpdate { + status: wit_channel::StatusType::Status, + message: msg, + metadata_json, + } + } }) } diff --git a/src/channels/web/handlers/chat.rs b/src/channels/web/handlers/chat.rs index de4b3155..bc4e3dbc 100644 --- a/src/channels/web/handlers/chat.rs +++ b/src/channels/web/handlers/chat.rs @@ -398,8 +398,10 @@ pub async fn chat_history_handler( truncate_preview(&s, 500) }), error: tc.error.clone(), + rationale: tc.rationale.clone(), }) .collect(), + narrative: t.narrative.clone(), }) .collect(); diff --git a/src/channels/web/mod.rs b/src/channels/web/mod.rs index 6a97e8b8..63aedaa0 100644 --- a/src/channels/web/mod.rs +++ b/src/channels/web/mod.rs @@ -489,6 +489,20 @@ impl Channel for GatewayChannel { }, StatusUpdate::Suggestions { suggestions } => AppEvent::Suggestions { suggestions, + thread_id: thread_id.clone(), + }, + StatusUpdate::ReasoningUpdate { + narrative, + decisions, + } => AppEvent::ReasoningUpdate { + narrative, + decisions: decisions + .into_iter() + .map(|d| crate::channels::web::types::ToolDecisionDto { + tool_name: d.tool_name, + rationale: d.rationale, + }) + .collect(), thread_id, }, StatusUpdate::TurnCost { diff --git a/src/channels/web/openai_compat.rs b/src/channels/web/openai_compat.rs index 55b7c854..0c0f1a9e 100644 --- a/src/channels/web/openai_compat.rs +++ b/src/channels/web/openai_compat.rs @@ -231,6 +231,7 @@ pub fn convert_messages(messages: &[OpenAiMessage]) -> Result, name: tc.function.name.clone(), arguments: serde_json::from_str(&tc.function.arguments) .unwrap_or(serde_json::Value::Object(Default::default())), + reasoning: None, }) .collect(); Ok(ChatMessage::assistant_with_tool_calls( @@ -954,6 +955,7 @@ mod tests { id: "call_abc".to_string(), name: "search".to_string(), arguments: serde_json::json!({"query": "rust"}), + reasoning: None, }]; let converted = convert_tool_calls_to_openai(&calls); diff --git a/src/channels/web/server.rs b/src/channels/web/server.rs index 5b092312..c24ceb16 100644 --- a/src/channels/web/server.rs +++ b/src/channels/web/server.rs @@ -1725,8 +1725,10 @@ async fn chat_history_handler( truncate_preview(&s, 500) }), error: tc.error.clone(), + rationale: tc.rationale.clone(), }) .collect(), + narrative: t.narrative.clone(), }) .collect(); diff --git a/src/channels/web/types.rs b/src/channels/web/types.rs index fe18a824..8698c030 100644 --- a/src/channels/web/types.rs +++ b/src/channels/web/types.rs @@ -63,6 +63,9 @@ pub struct TurnInfo { pub started_at: String, pub completed_at: Option, pub tool_calls: Vec, + /// Agent's reasoning narrative for this turn. + #[serde(skip_serializing_if = "Option::is_none")] + pub narrative: Option, } #[derive(Debug, Serialize)] @@ -74,6 +77,9 @@ pub struct ToolCallInfo { pub result_preview: Option, #[serde(skip_serializing_if = "Option::is_none")] pub error: Option, + /// Agent's reasoning for choosing this tool. + #[serde(skip_serializing_if = "Option::is_none")] + pub rationale: Option, } #[derive(Debug, Serialize)] @@ -116,7 +122,7 @@ pub struct ApprovalRequest { // --- App Event (re-exported from ironclaw_common) --- -pub use ironclaw_common::AppEvent; +pub use ironclaw_common::{AppEvent, ToolDecisionDto}; // --- Memory --- diff --git a/src/channels/web/util.rs b/src/channels/web/util.rs index ed70c5ce..2e4ffe3b 100644 --- a/src/channels/web/util.rs +++ b/src/channels/web/util.rs @@ -4,6 +4,21 @@ use crate::channels::web::types::{ToolCallInfo, TurnInfo}; pub use ironclaw_common::truncate_preview; +/// Parse tool call summary JSON objects into `ToolCallInfo` structs. +fn parse_tool_call_infos(calls: &[serde_json::Value]) -> Vec { + calls + .iter() + .map(|c| ToolCallInfo { + name: c["name"].as_str().unwrap_or("unknown").to_string(), + has_result: c.get("result_preview").is_some_and(|v| !v.is_null()), + has_error: c.get("error").is_some_and(|v| !v.is_null()), + result_preview: c["result_preview"].as_str().map(String::from), + error: c["error"].as_str().map(String::from), + rationale: c["rationale"].as_str().map(String::from), + }) + .collect() +} + /// Build TurnInfo pairs from flat DB messages (user/tool_calls/assistant triples). /// /// Handles three message patterns: @@ -27,6 +42,7 @@ pub fn build_turns_from_db_messages( started_at: msg.created_at.to_rfc3339(), completed_at: None, tool_calls: Vec::new(), + narrative: None, }; // Check if next message is a tool_calls record @@ -34,18 +50,28 @@ pub fn build_turns_from_db_messages( && next.role == "tool_calls" { let tc_msg = iter.next().expect("peeked"); - match serde_json::from_str::>(&tc_msg.content) { - Ok(calls) => { - turn.tool_calls = calls - .iter() - .map(|c| ToolCallInfo { - name: c["name"].as_str().unwrap_or("unknown").to_string(), - has_result: c.get("result_preview").is_some(), - has_error: c.get("error").is_some(), - result_preview: c["result_preview"].as_str().map(String::from), - error: c["error"].as_str().map(String::from), - }) - .collect(); + // Parse tool_calls JSON — supports two formats: + // safety: no byte-index slicing; comment describes JSON shape + match serde_json::from_str::(&tc_msg.content) { + Ok(serde_json::Value::Array(calls)) => { + // Old format: plain array + turn.tool_calls = parse_tool_call_infos(&calls); + } + Ok(serde_json::Value::Object(obj)) => { + // New wrapped format with narrative + turn.narrative = obj + .get("narrative") + .and_then(|v| v.as_str()) + .map(String::from); + if let Some(serde_json::Value::Array(calls)) = obj.get("calls") { + turn.tool_calls = parse_tool_call_infos(calls); + } + } + Ok(_) => { + tracing::warn!( + message_id = %tc_msg.id, + "Unexpected tool_calls JSON shape in DB, skipping" + ); } Err(e) => { tracing::warn!( @@ -83,6 +109,7 @@ pub fn build_turns_from_db_messages( started_at: msg.created_at.to_rfc3339(), completed_at: Some(msg.created_at.to_rfc3339()), tool_calls: Vec::new(), + narrative: None, }); turn_number += 1; } @@ -201,4 +228,52 @@ mod tests { assert!(turns[0].tool_calls.is_empty()); assert_eq!(turns[0].state, "Completed"); } + + #[test] + fn test_build_turns_with_wrapped_tool_calls_format() { + let tc_json = serde_json::json!({ + "narrative": "Searching memory for context before proceeding.", + "calls": [ + {"name": "memory_search", "result_preview": "found 3 items", "rationale": "consult prior context"}, + {"name": "shell", "error": "permission denied"} + ] + }); + let messages = vec![ + make_msg("user", "Find info", 0), + make_msg("tool_calls", &tc_json.to_string(), 500), + make_msg("assistant", "Here's what I found", 1000), + ]; + let turns = build_turns_from_db_messages(&messages); + assert_eq!(turns.len(), 1); + assert_eq!( + turns[0].narrative.as_deref(), + Some("Searching memory for context before proceeding.") + ); + assert_eq!(turns[0].tool_calls.len(), 2); + assert_eq!(turns[0].tool_calls[0].name, "memory_search"); + assert_eq!( + turns[0].tool_calls[0].rationale.as_deref(), + Some("consult prior context") + ); + assert!(turns[0].tool_calls[0].has_result); + assert_eq!(turns[0].tool_calls[1].name, "shell"); + assert!(turns[0].tool_calls[1].has_error); + assert_eq!(turns[0].response.as_deref(), Some("Here's what I found")); + } + + #[test] + fn test_build_turns_wrapped_format_without_narrative() { + let tc_json = serde_json::json!({ + "calls": [{"name": "echo", "result_preview": "hello"}] + }); + let messages = vec![ + make_msg("user", "Say hi", 0), + make_msg("tool_calls", &tc_json.to_string(), 500), + make_msg("assistant", "Done", 1000), + ]; + let turns = build_turns_from_db_messages(&messages); + assert_eq!(turns.len(), 1); + assert!(turns[0].narrative.is_none()); + assert_eq!(turns[0].tool_calls.len(), 1); + } } diff --git a/src/llm/anthropic_oauth.rs b/src/llm/anthropic_oauth.rs index 490fbc3f..c94c90e5 100644 --- a/src/llm/anthropic_oauth.rs +++ b/src/llm/anthropic_oauth.rs @@ -575,6 +575,7 @@ fn extract_response_content(response: &AnthropicResponse) -> (Option, Ve id: id.clone(), name: name.clone(), arguments: input.clone(), + reasoning: None, }); } } @@ -623,6 +624,7 @@ mod tests { id: "call_1".to_string(), name: "search".to_string(), arguments: serde_json::json!({"q": "test"}), + reasoning: None, }]; let messages = vec![ ChatMessage::user("Search for test"), diff --git a/src/llm/bedrock.rs b/src/llm/bedrock.rs index 5d6e121e..b5f7badd 100644 --- a/src/llm/bedrock.rs +++ b/src/llm/bedrock.rs @@ -522,6 +522,7 @@ fn extract_content_blocks( id: tu.tool_use_id().to_string(), name: tu.name().to_string(), arguments: document_to_json(tu.input()), + reasoning: None, }); } // Ignore reasoning, citations, images, etc. @@ -759,11 +760,13 @@ mod tests { id: "call_1".to_string(), name: "echo".to_string(), arguments: serde_json::json!({"text": "hi"}), + reasoning: None, }; let tc2 = crate::llm::provider::ToolCall { id: "call_2".to_string(), name: "time".to_string(), arguments: serde_json::json!({}), + reasoning: None, }; let messages = vec![ @@ -802,6 +805,7 @@ mod tests { id: "call_1".to_string(), name: "search".to_string(), arguments: serde_json::json!({"query": "test"}), + reasoning: None, }; let messages = vec![ @@ -825,6 +829,7 @@ mod tests { id: "call_1".to_string(), name: "echo".to_string(), arguments: serde_json::json!({}), + reasoning: None, }; let messages = vec![ @@ -989,11 +994,13 @@ mod tests { id: "call_abc".to_string(), name: "get_weather".to_string(), arguments: serde_json::json!({"city": "NYC"}), + reasoning: None, }; let tc2 = crate::llm::provider::ToolCall { id: "call_def".to_string(), name: "get_time".to_string(), arguments: serde_json::json!({"tz": "EST"}), + reasoning: None, }; let messages = vec![ diff --git a/src/llm/codex_chatgpt.rs b/src/llm/codex_chatgpt.rs index 56cb3378..e7dcf40d 100644 --- a/src/llm/codex_chatgpt.rs +++ b/src/llm/codex_chatgpt.rs @@ -732,6 +732,7 @@ impl LlmProvider for CodexChatGptProvider { id: tc.call_id, name: tc.name, arguments: args, + reasoning: None, } }) .collect(); @@ -825,6 +826,7 @@ mod tests { id: "call_1".to_string(), name: "search".to_string(), arguments: json!({"query": "rust"}), + reasoning: None, }; let msg = ChatMessage::assistant_with_tool_calls(Some("thinking...".into()), vec![tc]); let items = CodexChatGptProvider::message_to_input_items(&msg); diff --git a/src/llm/gemini_oauth.rs b/src/llm/gemini_oauth.rs index b36eb595..a19eec12 100644 --- a/src/llm/gemini_oauth.rs +++ b/src/llm/gemini_oauth.rs @@ -1898,6 +1898,7 @@ impl GeminiOauthProvider { id, name, arguments: args, + reasoning: None, }); } } diff --git a/src/llm/github_copilot.rs b/src/llm/github_copilot.rs index b173191a..c7a24b1a 100644 --- a/src/llm/github_copilot.rs +++ b/src/llm/github_copilot.rs @@ -596,6 +596,7 @@ fn extract_choice_content(choice: &OpenAiChoice) -> (Option, Vec Result { id: state.call_id, name: state.name, arguments, + reasoning: None, }); } else { // Fallback: extract directly from the item @@ -650,6 +651,7 @@ fn parse_sse_response(body: &str) -> Result { id: call_id, name, arguments, + reasoning: None, }); } } @@ -727,6 +729,7 @@ fn parse_sse_response(body: &str) -> Result { id: state.call_id, name: state.name, arguments, + reasoning: None, }); } } @@ -822,11 +825,13 @@ mod tests { id: "call_1".to_string(), name: "search".to_string(), arguments: serde_json::json!({"query": "test"}), + reasoning: None, }, ToolCall { id: "call_2".to_string(), name: "read".to_string(), arguments: serde_json::json!({"path": "/tmp"}), + reasoning: None, }, ]; let msg = diff --git a/src/llm/provider.rs b/src/llm/provider.rs index bb45ec68..8afd914a 100644 --- a/src/llm/provider.rs +++ b/src/llm/provider.rs @@ -231,6 +231,10 @@ pub struct ToolCall { pub id: String, pub name: String, pub arguments: serde_json::Value, + /// Optional reasoning for why this tool was chosen — supplied by the provider + /// or derived from the shared response content as a fallback. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub reasoning: Option, } /// Generate a tool-call ID that satisfies all providers. @@ -637,6 +641,7 @@ mod tests { id: "call_1".to_string(), name: "echo".to_string(), arguments: serde_json::json!({}), + reasoning: None, }; let mut messages = vec![ ChatMessage::user("hello"), @@ -680,6 +685,7 @@ mod tests { id: "call_1".to_string(), name: "echo".to_string(), arguments: serde_json::json!({}), + reasoning: None, }; let mut messages = vec![ ChatMessage::user("test"), @@ -705,11 +711,13 @@ mod tests { id: "call_sel_1".to_string(), name: "search".to_string(), arguments: serde_json::json!({"q": "test"}), + reasoning: None, }; let tc2 = ToolCall { id: "call_sel_2".to_string(), name: "http".to_string(), arguments: serde_json::json!({"url": "https://example.com"}), + reasoning: None, }; let mut messages = vec![ ChatMessage::system("You are a helpful assistant."), diff --git a/src/llm/reasoning.rs b/src/llm/reasoning.rs index cbec297b..77905f95 100644 --- a/src/llm/reasoning.rs +++ b/src/llm/reasoning.rs @@ -525,17 +525,35 @@ impl Reasoning { let response = self.llm.complete_with_tools(request).await?; - let reasoning = response.content.unwrap_or_default(); + let shared_reasoning = response + .content + .map(|c| { + let pre_truncated = truncate_at_tool_tags(&c); + clean_response(&pre_truncated) + }) + .unwrap_or_default(); let selections: Vec = response .tool_calls .into_iter() - .map(|tool_call| ToolSelection { - tool_name: tool_call.name, - parameters: tool_call.arguments, - reasoning: reasoning.clone(), - alternatives: vec![], - tool_call_id: tool_call.id, + .map(|tool_call| { + // Prefer per-tool reasoning if the provider supplied it, + // otherwise fall back to the shared response content. + let rationale = tool_call + .reasoning + .map(|r| { + let pre_truncated = truncate_at_tool_tags(&r); + clean_response(&pre_truncated) + }) + .filter(|r| !r.trim().is_empty()) + .unwrap_or_else(|| shared_reasoning.clone()); + ToolSelection { + tool_name: tool_call.name, + parameters: tool_call.arguments, + reasoning: rationale, + alternatives: vec![], + tool_call_id: tool_call.id, + } }) .collect(); @@ -664,13 +682,36 @@ Respond in JSON format: // If there were tool calls, return them for execution if !response.tool_calls.is_empty() { + let narrative = response.content.map(|c| { + let pre_truncated = truncate_at_tool_tags(&c); + clean_response(&pre_truncated) + }); + // Populate per-tool reasoning from the shared narrative when the + // provider did not supply per-tool rationale. + let tool_calls: Vec = response + .tool_calls + .into_iter() + .map(|mut tc| { + if tc.reasoning.as_ref().is_none_or(|r| r.trim().is_empty()) { + tc.reasoning = narrative.as_ref().filter(|n| !n.is_empty()).cloned(); + } else { + // Clean provider-supplied per-tool reasoning the same way + // we clean the shared narrative (strip thinking/tool tags). + tc.reasoning = tc + .reasoning + .map(|r| { + let pre_truncated = truncate_at_tool_tags(&r); + clean_response(&pre_truncated) + }) + .filter(|r| !r.trim().is_empty()); + } + tc + }) + .collect(); return Ok(RespondOutput { result: RespondResult::ToolCalls { - tool_calls: response.tool_calls, - content: response.content.map(|c| { - let pre_truncated = truncate_at_tool_tags(&c); - clean_response(&pre_truncated) - }), + tool_calls, + content: narrative, }, usage, }); @@ -1350,6 +1391,7 @@ fn recover_tool_calls_from_content( ), name: name.to_string(), arguments, + reasoning: None, }); continue; } @@ -1364,6 +1406,7 @@ fn recover_tool_calls_from_content( ), name: name.to_string(), arguments: serde_json::Value::Object(Default::default()), + reasoning: None, }); } } @@ -1401,6 +1444,7 @@ fn recover_tool_calls_from_content( ), name: name.to_string(), arguments, + reasoning: None, }); remaining = &args_start[bracket_end + 1..]; continue; @@ -1412,6 +1456,7 @@ fn recover_tool_calls_from_content( id: super::provider::generate_tool_call_id(calls.len(), RECOVERED_TOOL_CALL_SEED), name: name.to_string(), arguments: serde_json::Value::Object(Default::default()), + reasoning: None, }); remaining = after_name; } @@ -3145,4 +3190,32 @@ That's my plan."#; "Text {} middle " ); } + + /// Verify that reasoning normalization strips thinking tags and tool tags + /// from per-tool reasoning, matching the cleaning applied to shared reasoning. + #[test] + fn test_reasoning_normalization_strips_thinking_tags() { + let raw = "Let me consider...Search memory for prior context"; + let pre_truncated = truncate_at_tool_tags(raw); + let cleaned = clean_response(&pre_truncated); + assert!(!cleaned.contains("")); + assert!(cleaned.contains("Search memory")); + } + + #[test] + fn test_reasoning_normalization_strips_tool_tags() { + let raw = "Calling search {\"name\": \"search\"}"; + let pre_truncated = truncate_at_tool_tags(raw); + let cleaned = clean_response(&pre_truncated); + assert!(!cleaned.contains("")); + assert!(cleaned.contains("Calling search")); + } + + #[test] + fn test_reasoning_normalization_empty_after_cleaning() { + let raw = "internal only"; + let pre_truncated = truncate_at_tool_tags(raw); + let cleaned = clean_response(&pre_truncated); + assert!(cleaned.trim().is_empty()); + } } diff --git a/src/llm/rig_adapter.rs b/src/llm/rig_adapter.rs index a9030929..7a6b2ae8 100644 --- a/src/llm/rig_adapter.rs +++ b/src/llm/rig_adapter.rs @@ -490,6 +490,7 @@ fn extract_response( id: tc.id.clone(), name: tc.function.name.clone(), arguments: tc.function.arguments.clone(), + reasoning: None, }); } // Reasoning and Image variants are not mapped to IronClaw types @@ -880,6 +881,7 @@ mod tests { id: "Xt7mK9pQ2".to_string(), name: "search".to_string(), arguments: serde_json::json!({"query": "test"}), + reasoning: None, }; let msg = ChatMessage::assistant_with_tool_calls(Some("thinking".to_string()), vec![tc]); let messages = vec![msg]; @@ -997,6 +999,7 @@ mod tests { id: "".to_string(), name: "search".to_string(), arguments: serde_json::json!({"query": "test"}), + reasoning: None, }; let messages = vec![ChatMessage::assistant_with_tool_calls(None, vec![tc])]; let (_preamble, history) = convert_messages(&messages); @@ -1028,6 +1031,7 @@ mod tests { id: " ".to_string(), name: "search".to_string(), arguments: serde_json::json!({"query": "test"}), + reasoning: None, }; let messages = vec![ChatMessage::assistant_with_tool_calls(None, vec![tc])]; let (_preamble, history) = convert_messages(&messages); @@ -1061,6 +1065,7 @@ mod tests { id: "".to_string(), name: "search".to_string(), arguments: serde_json::json!({"query": "test"}), + reasoning: None, }; let assistant_msg = ChatMessage::assistant_with_tool_calls(None, vec![tc]); let tool_result_msg = ChatMessage { @@ -1380,11 +1385,13 @@ mod tests { id: "call_a".to_string(), name: "search".to_string(), arguments: serde_json::json!({"q": "rust"}), + reasoning: None, }; let tc2 = IronToolCall { id: "call_b".to_string(), name: "fetch".to_string(), arguments: serde_json::json!({"url": "https://example.com"}), + reasoning: None, }; let assistant = ChatMessage::assistant_with_tool_calls(None, vec![tc1, tc2]); let result_a = ChatMessage::tool_result("call_a", "search", "search results"); diff --git a/src/orchestrator/api.rs b/src/orchestrator/api.rs index 37085a8b..8da7ae6f 100644 --- a/src/orchestrator/api.rs +++ b/src/orchestrator/api.rs @@ -14,6 +14,7 @@ use serde::{Deserialize, Serialize}; use tokio::sync::{Mutex, broadcast}; use uuid::Uuid; +use crate::channels::web::types::ToolDecisionDto; use crate::db::Database; use crate::llm::{CompletionRequest, LlmProvider, ToolCompletionRequest}; use crate::orchestrator::auth::{TokenStore, worker_auth_middleware}; @@ -344,6 +345,20 @@ async fn job_event_handler( // gain context/memory tracking capabilities. fallback_deliverable: payload.data.get("fallback_deliverable").cloned(), }, + "reasoning" => { + let narrative = payload + .data + .get("narrative") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + let decisions = ToolDecisionDto::from_json_array(&payload.data["decisions"]); + AppEvent::JobReasoning { + job_id: job_id_str, + narrative, + decisions, + } + } _ => AppEvent::JobStatus { job_id: job_id_str, message: payload diff --git a/src/worker/job.rs b/src/worker/job.rs index 9d5794ca..669c69f0 100644 --- a/src/worker/job.rs +++ b/src/worker/job.rs @@ -18,6 +18,7 @@ use crate::agent::agentic_loop::{ }; use crate::agent::scheduler::WorkerMessage; use crate::agent::task::TaskOutput; +use crate::channels::web::types::ToolDecisionDto; use crate::context::{ContextManager, JobState}; use crate::db::Database; use crate::error::Error; @@ -200,6 +201,19 @@ impl Worker { .map(|s| s.to_string()), fallback_deliverable: data.get("fallback_deliverable").cloned(), }), + "reasoning" => { + let narrative = data + .get("narrative") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + let decisions = ToolDecisionDto::from_json_array(&data["decisions"]); + Some(AppEvent::JobReasoning { + job_id: job_id_str, + narrative, + decisions, + }) + } _ => None, }; if let Some(event) = event { @@ -897,6 +911,11 @@ Report when the job is complete or if you encounter issues you cannot resolve."# id: selection.tool_call_id.clone(), name: selection.tool_name.clone(), arguments: selection.parameters.clone(), + reasoning: if action.reasoning.is_empty() { + None + } else { + Some(action.reasoning.clone()) + }, }], )); @@ -1357,6 +1376,48 @@ impl<'a> LoopDelegate for JobDelegate<'a> { ); } + // Emit reasoning event if any tool calls carry reasoning. + // Sanitize narrative and per-tool rationale through SafetyLayer + // (parity with ChatDelegate in dispatcher.rs). + let sanitized_narrative = content + .as_deref() + .filter(|c| !c.trim().is_empty()) + .map(|c| { + self.worker + .deps + .safety + .sanitize_tool_output("job_narrative", c) + .content + }) + .filter(|c| !c.trim().is_empty()) + .unwrap_or_default(); + let decisions: Vec = tool_calls + .iter() + .filter_map(|tc| { + tc.reasoning.as_ref().map(|r| { + let sanitized = self + .worker + .deps + .safety + .sanitize_tool_output("tool_rationale", r) + .content; + serde_json::json!({ + "tool_name": tc.name, + "rationale": sanitized, + }) + }) + }) + .collect(); + if !decisions.is_empty() { + self.worker.log_event( + "reasoning", + serde_json::json!({ + "narrative": sanitized_narrative, + "decisions": decisions, + }), + ); + } + // Add assistant message with tool_calls (OpenAI protocol) reason_ctx .messages @@ -1371,7 +1432,7 @@ impl<'a> LoopDelegate for JobDelegate<'a> { .map(|tc| ToolSelection { tool_name: tc.name.clone(), parameters: tc.arguments.clone(), - reasoning: String::new(), + reasoning: tc.reasoning.clone().unwrap_or_default(), alternatives: vec![], tool_call_id: tc.id.clone(), }) @@ -1424,6 +1485,11 @@ fn selections_to_tool_calls(selections: &[ToolSelection]) -> Vec { id: s.tool_call_id.clone(), name: s.tool_name.clone(), arguments: s.parameters.clone(), + reasoning: if s.reasoning.is_empty() { + None + } else { + Some(s.reasoning.clone()) + }, }) .collect() } diff --git a/tests/openai_compat_integration.rs b/tests/openai_compat_integration.rs index e1d258ed..b677e57f 100644 --- a/tests/openai_compat_integration.rs +++ b/tests/openai_compat_integration.rs @@ -94,6 +94,7 @@ impl LlmProvider for MockLlmProvider { id: "call_mock_001".to_string(), name: tool.name.clone(), arguments: serde_json::json!({"test": true}), + reasoning: None, }], input_tokens: 15, output_tokens: 8, diff --git a/tests/support/trace_llm.rs b/tests/support/trace_llm.rs index e33caf6b..239cfdb5 100644 --- a/tests/support/trace_llm.rs +++ b/tests/support/trace_llm.rs @@ -566,6 +566,7 @@ impl LlmProvider for TraceLlm { id: tc.id, name: tc.name, arguments: tc.arguments, + reasoning: None, }) .collect(); Ok(ToolCompletionResponse { From 0341fcc9405e3a9f22319891dc1d55d3a67edc06 Mon Sep 17 00:00:00 2001 From: Henry Park Date: Wed, 25 Mar 2026 11:45:29 -0700 Subject: [PATCH 2/8] Fix REPL single-message hang and cap CI test duration (#1643) * Fix REPL single-message hang and cap CI test duration * Fix Clippy nested-if lint in REPL startup * Fix single-message approval flow * Handle empty single-message REPL exits * Wait for one-shot event routines before exit --- .github/workflows/test.yml | 24 +++++-- src/agent/agent_loop.rs | 69 ++++++++++++++++-- src/agent/routine_engine.rs | 70 ++++++++++++++++--- src/channels/repl.rs | 60 +++++++++++++--- .../scenarios/test_telegram_hot_activation.py | 4 +- 5 files changed, 196 insertions(+), 31 deletions(-) diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 00488c70..5d4eabc0 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -12,6 +12,7 @@ jobs: tests: name: Tests (${{ matrix.name }}) runs-on: ubuntu-latest + timeout-minutes: 45 strategy: fail-fast: false matrix: @@ -40,11 +41,14 @@ jobs: - name: Build WASM channels (for integration tests) run: ./scripts/build-wasm-extensions.sh --channels - name: Run Tests - run: cargo test ${{ matrix.flags }} -- --nocapture + run: | + timeout --signal=INT --kill-after=30s 40m \ + cargo test ${{ matrix.flags }} -- --nocapture heavy-integration-tests: name: Heavy Integration Tests runs-on: ubuntu-latest + timeout-minutes: 20 steps: - name: Checkout repository uses: actions/checkout@v6 @@ -58,9 +62,13 @@ jobs: - name: Build Telegram WASM channel run: cargo build --manifest-path channels-src/telegram/Cargo.toml --target wasm32-wasip2 --release - name: Run thread scheduling integration tests - run: cargo test --no-default-features --features libsql,integration --test e2e_thread_scheduling -- --nocapture + run: | + timeout --signal=INT --kill-after=30s 15m \ + cargo test --no-default-features --features libsql,integration --test e2e_thread_scheduling -- --nocapture - name: Run Telegram thread-scope regression test - run: cargo test --features integration --test telegram_auth_integration test_private_messages_use_chat_id_as_thread_scope -- --exact + run: | + timeout --signal=INT --kill-after=30s 10m \ + cargo test --features integration --test telegram_auth_integration test_private_messages_use_chat_id_as_thread_scope -- --exact telegram-tests: name: Telegram Channel Tests @@ -68,6 +76,7 @@ jobs: github.event_name != 'pull_request' || github.base_ref != 'staging' runs-on: ubuntu-latest + timeout-minutes: 15 steps: - name: Checkout repository uses: actions/checkout@v6 @@ -75,7 +84,9 @@ jobs: uses: dtolnay/rust-toolchain@stable - uses: Swatinem/rust-cache@v2 - name: Run Telegram Channel Tests - run: cargo test --manifest-path channels-src/telegram/Cargo.toml -- --nocapture + run: | + timeout --signal=INT --kill-after=30s 10m \ + cargo test --manifest-path channels-src/telegram/Cargo.toml -- --nocapture windows-build: name: Windows Build (${{ matrix.name }}) @@ -110,6 +121,7 @@ jobs: github.event_name != 'pull_request' || github.base_ref != 'staging' runs-on: ubuntu-latest + timeout-minutes: 30 steps: - name: Checkout repository uses: actions/checkout@v6 @@ -125,7 +137,9 @@ jobs: - name: Build all WASM extensions against current WIT run: ./scripts/build-wasm-extensions.sh - name: Instantiation test (host linker compatibility) - run: cargo test --all-features wit_compat -- --nocapture + run: | + timeout --signal=INT --kill-after=30s 20m \ + cargo test --all-features wit_compat -- --nocapture bench-compile: name: Benchmark Compilation diff --git a/src/agent/agent_loop.rs b/src/agent/agent_loop.rs index f51a8db1..e28f11d0 100644 --- a/src/agent/agent_loop.rs +++ b/src/agent/agent_loop.rs @@ -16,6 +16,7 @@ use crate::agent::context_monitor::ContextMonitor; use crate::agent::heartbeat::spawn_heartbeat; use crate::agent::routine_engine::{RoutineEngine, spawn_cron_ticker}; use crate::agent::self_repair::{DefaultSelfRepair, RepairResult, SelfRepair}; +use crate::agent::session::ThreadState; use crate::agent::session_manager::SessionManager; use crate::agent::submission::{Submission, SubmissionParser, SubmissionResult}; use crate::agent::{HeartbeatConfig as AgentHeartbeatConfig, Router, Scheduler, SchedulerDeps}; @@ -84,6 +85,15 @@ fn resolve_owner_scope_notification_user( trimmed_option(explicit_user).or_else(|| trimmed_option(owner_fallback)) } +fn is_single_message_repl(message: &IncomingMessage) -> bool { + message.channel == "repl" + && message + .metadata + .get("single_message_mode") + .and_then(|value| value.as_bool()) + .unwrap_or(false) +} + async fn resolve_channel_notification_user( extension_manager: Option<&Arc>, channel: Option<&str>, @@ -1140,9 +1150,14 @@ impl Agent { && let Submission::UserInput { ref content } = submission && let Some(engine) = self.routine_engine().await { + let single_message_repl = is_single_message_repl(message); // Use post-hook content so that BeforeInbound hooks that rewrite // input are respected by event trigger matching. - let fired = engine.check_event_triggers(message, content).await; + let fired = if single_message_repl { + engine.check_event_triggers_and_wait(message, content).await + } else { + engine.check_event_triggers(message, content).await + }; if fired > 0 { tracing::debug!( channel = %message.channel, @@ -1150,10 +1165,16 @@ impl Agent { fired, "Consumed inbound user message with matching event-triggered routine(s)" ); - return Ok(Some(String::new())); + return if single_message_repl { + Ok(None) + } else { + Ok(Some(String::new())) + }; } } + let session_for_empty_exit = Arc::clone(&session); + // Process based on submission type let result = match submission { Submission::UserInput { content } => { @@ -1263,7 +1284,13 @@ impl Agent { SubmissionResult::Error { message } => { Ok(Some(format!("Error: {}", message))) } - _ => Ok(Some(String::new())), + _ => { + if is_single_message_repl(message) { + Ok(None) + } else { + Ok(Some(String::new())) + } + } }; } // Authorization checks (including restart channel check) are enforced in handle_system_command @@ -1325,7 +1352,26 @@ impl Agent { Ok(Some(content)) } } - SubmissionResult::Ok { message } => Ok(message), + SubmissionResult::Ok { + message: output_message, + } => { + let should_exit = + if output_message.as_deref() == Some("") && is_single_message_repl(message) { + let sess = session_for_empty_exit.lock().await; + sess.threads + .get(&thread_id) + .map(|thread| thread.state != ThreadState::AwaitingApproval) + .unwrap_or(true) + } else { + false + }; + + if should_exit { + Ok(None) + } else { + Ok(output_message) + } + } SubmissionResult::Error { message } => Ok(Some(format!("Error: {}", message))), SubmissionResult::Interrupted => Ok(Some("Interrupted.".into())), SubmissionResult::NeedApproval { .. } => { @@ -1341,7 +1387,7 @@ impl Agent { #[cfg(test)] mod tests { use super::{ - chat_tool_execution_metadata, resolve_routine_notification_user, + chat_tool_execution_metadata, is_single_message_repl, resolve_routine_notification_user, should_fallback_routine_notification, truncate_for_preview, }; use crate::channels::IncomingMessage; @@ -1503,4 +1549,17 @@ mod tests { assert!(should_fallback_routine_notification(&error)); // safety: test-only assertion } + + #[test] + fn single_message_repl_detection_requires_repl_channel_and_metadata_flag() { + let repl = IncomingMessage::new("repl", "owner-scope", "hello") + .with_metadata(serde_json::json!({ "single_message_mode": true })); + let gateway = IncomingMessage::new("gateway", "owner-scope", "hello") + .with_metadata(serde_json::json!({ "single_message_mode": true })); + let plain_repl = IncomingMessage::new("repl", "owner-scope", "hello"); + + assert!(is_single_message_repl(&repl)); // safety: test-only assertion + assert!(!is_single_message_repl(&gateway)); // safety: test-only assertion + assert!(!is_single_message_repl(&plain_repl)); // safety: test-only assertion + } } diff --git a/src/agent/routine_engine.rs b/src/agent/routine_engine.rs index 9c55903f..a3cdb6cd 100644 --- a/src/agent/routine_engine.rs +++ b/src/agent/routine_engine.rs @@ -18,6 +18,7 @@ use std::time::Duration; use chrono::Utc; use regex::Regex; use tokio::sync::{RwLock, mpsc}; +use tokio::task::JoinHandle; use uuid::Uuid; use crate::agent::Scheduler; @@ -45,6 +46,11 @@ enum EventMatcher { System { routine: Routine }, } +struct TriggeredRoutine { + routine: Routine, + detail: String, +} + /// Distinguishes why sandbox is unavailable so error messages are accurate. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum SandboxReadiness { @@ -202,6 +208,44 @@ impl RoutineEngine { /// Check incoming message against event triggers. Returns number of routines fired. pub async fn check_event_triggers(&self, message: &IncomingMessage, content: &str) -> usize { + let triggered = self.matching_event_triggers(message, content).await; + let fired = triggered.len(); + for triggered in triggered { + std::mem::drop(self.spawn_fire(triggered.routine, "event", Some(triggered.detail))); + } + fired + } + + /// Fire matching event-triggered routines and wait for them to complete. + /// + /// Used by single-message REPL mode so the process does not exit before + /// background event-triggered routines finish. + pub async fn check_event_triggers_and_wait( + &self, + message: &IncomingMessage, + content: &str, + ) -> usize { + let triggered = self.matching_event_triggers(message, content).await; + let fired = triggered.len(); + let handles: Vec> = triggered + .into_iter() + .map(|triggered| self.spawn_fire(triggered.routine, "event", Some(triggered.detail))) + .collect(); + + for handle in handles { + if let Err(e) = handle.await { + tracing::warn!(error = %e, "Event-triggered routine task failed"); + } + } + + fired + } + + async fn matching_event_triggers( + &self, + message: &IncomingMessage, + content: &str, + ) -> Vec { let cache = self.event_cache.read().await; // Early return if there are no message matchers at all. @@ -209,10 +253,9 @@ impl RoutineEngine { .iter() .any(|m| matches!(m, EventMatcher::Message { .. })) { - return 0; + return Vec::new(); } - - let mut fired = 0; + let mut triggered = Vec::new(); // Collect routine IDs for batch query let routine_ids: Vec = cache @@ -224,13 +267,13 @@ impl RoutineEngine { .collect(); if routine_ids.is_empty() { - return 0; + return Vec::new(); } // Single batch query instead of N queries let concurrent_counts = match self.batch_concurrent_counts(&routine_ids).await { Some(counts) => counts, - None => return 0, + None => return Vec::new(), }; for matcher in cache.iter() { @@ -285,11 +328,13 @@ impl RoutineEngine { } let detail = truncate(content, 200); - self.spawn_fire(routine.clone(), "event", Some(detail)); - fired += 1; + triggered.push(TriggeredRoutine { + routine: routine.clone(), + detail, + }); } - fired + triggered } /// Emit a structured event to system-event routines. @@ -845,7 +890,12 @@ impl RoutineEngine { } /// Spawn a fire in a background task. - fn spawn_fire(&self, routine: Routine, trigger_type: &str, trigger_detail: Option) { + fn spawn_fire( + &self, + routine: Routine, + trigger_type: &str, + trigger_detail: Option, + ) -> JoinHandle<()> { let run = RoutineRun { id: Uuid::new_v4(), routine_id: routine.id, @@ -882,7 +932,7 @@ impl RoutineEngine { return; } execute_routine(engine, routine, run).await; - }); + }) } fn check_cooldown(&self, routine: &Routine) -> bool { diff --git a/src/channels/repl.rs b/src/channels/repl.rs index 61c68d13..41d73a8c 100644 --- a/src/channels/repl.rs +++ b/src/channels/repl.rs @@ -431,6 +431,18 @@ impl ReplChannel { let _ = execute!(stderr, terminal::Clear(terminal::ClearType::FromCursorDown)); } } + + async fn finish_single_message_turn(&self) { + if self.single_message.is_none() { + return; + } + + let tx = self.msg_tx.lock().ok().and_then(|mut guard| guard.take()); + if let Some(tx) = tx { + let msg = IncomingMessage::new("repl", &self.user_id, "/quit"); + let _ = tx.send(msg).await; + } + } } impl Default for ReplChannel { @@ -480,7 +492,9 @@ impl Channel for ReplChannel { async fn start(&self) -> Result { let (tx, rx) = mpsc::channel(32); - // Store tx so send_status can inject approval responses directly + // Approval prompts inject responses back through this sender. + // In single-message mode we keep it until the turn finishes, then + // drop it after enqueuing /quit so the receiver stream can close. if let Ok(mut guard) = self.msg_tx.lock() { *guard = Some(tx.clone()); } @@ -496,11 +510,10 @@ impl Channel for ReplChannel { // Single message mode: send it and return if let Some(msg) = single_message { - let incoming = IncomingMessage::new("repl", &user_id, &msg).with_timezone(&sys_tz); + let incoming = IncomingMessage::new("repl", &user_id, &msg) + .with_metadata(serde_json::json!({ "single_message_mode": true })) + .with_timezone(&sys_tz); let _ = tx.blocking_send(incoming); - // Ensure the agent exits after handling exactly one turn in -m mode, - // even when other channels (gateway/http) are enabled. - let _ = tx.blocking_send(IncomingMessage::new("repl", &user_id, "/quit")); return; } @@ -663,6 +676,7 @@ impl Channel for ReplChannel { println!(); println!(); self.stdin_locked.store(false, Ordering::Relaxed); + self.finish_single_message_turn().await; return Ok(()); } @@ -681,6 +695,7 @@ impl Channel for ReplChannel { println!(); // Unlock stdin so readline can resume self.stdin_locked.store(false, Ordering::Relaxed); + self.finish_single_message_turn().await; Ok(()) } @@ -780,6 +795,7 @@ impl Channel for ReplChannel { let msg_tx = Arc::clone(&self.msg_tx); let user_id = self.user_id.clone(); let lock_flag = Arc::clone(&self.stdin_locked); + let single_message_mode = self.single_message.is_some(); tokio::task::spawn_blocking(move || { let action = run_approval_selector(allow_always).unwrap_or("n"); // Unlock stdin so readline can resume after approval @@ -788,7 +804,12 @@ impl Channel for ReplChannel { return; }; if let Some(tx) = guard.as_ref() { - let msg = IncomingMessage::new("repl", &user_id, action); + let msg = if single_message_mode { + IncomingMessage::new("repl", &user_id, action) + .with_metadata(serde_json::json!({ "single_message_mode": true })) + } else { + IncomingMessage::new("repl", &user_id, action) + }; let _ = tx.blocking_send(msg); } }); @@ -889,6 +910,7 @@ impl Channel for ReplChannel { #[cfg(test)] mod tests { use futures::StreamExt; + use tokio::time::{Duration, timeout}; use super::*; @@ -897,16 +919,36 @@ mod tests { let repl = ReplChannel::with_message("hi".to_string()); let mut stream = repl.start().await.expect("repl start should succeed"); - let first = stream.next().await.expect("first message missing"); + let first = timeout(Duration::from_secs(1), stream.next()) + .await + .expect("timed out waiting for first message") + .expect("first message missing"); assert_eq!(first.channel, "repl"); assert_eq!(first.content, "hi"); - let second = stream.next().await.expect("quit message missing"); + assert!( + timeout(Duration::from_millis(100), stream.next()) + .await + .is_err(), + "single-message mode should wait for the turn to finish before quitting" + ); + + repl.respond(&first, OutgoingResponse::text("done")) + .await + .expect("respond should succeed"); + + let second = timeout(Duration::from_secs(1), stream.next()) + .await + .expect("timed out waiting for quit message") + .expect("quit message missing"); assert_eq!(second.channel, "repl"); assert_eq!(second.content, "/quit"); assert!( - stream.next().await.is_none(), + timeout(Duration::from_secs(1), stream.next()) + .await + .expect("timed out waiting for stream to close") + .is_none(), "stream should end after /quit" ); } diff --git a/tests/e2e/scenarios/test_telegram_hot_activation.py b/tests/e2e/scenarios/test_telegram_hot_activation.py index 261b837e..fede2be5 100644 --- a/tests/e2e/scenarios/test_telegram_hot_activation.py +++ b/tests/e2e/scenarios/test_telegram_hot_activation.py @@ -253,6 +253,6 @@ async def test_telegram_hot_activation_transitions_installed_to_active(page): assert await card.locator(SEL["ext_pairing_label"]).count() == 0 assert captured_setup_payloads == [ - {"secrets": {"telegram_bot_token": "123456789:ABCdefGhI"}}, - {"secrets": {}}, + {"secrets": {"telegram_bot_token": "123456789:ABCdefGhI"}, "fields": {}}, + {"secrets": {}, "fields": {}}, ] From c949521d8d153ecb3af30877779f8c160278ca09 Mon Sep 17 00:00:00 2001 From: Henry Park Date: Wed, 25 Mar 2026 13:17:32 -0700 Subject: [PATCH 3/8] Fix MCP lifecycle trace user scope (#1646) * Fix REPL single-message hang and cap CI test duration * Fix Clippy nested-if lint in REPL startup * Fix single-message approval flow * Handle empty single-message REPL exits * Wait for one-shot event routines before exit * Fix MCP lifecycle trace user scope --- tests/e2e_advanced_traces.rs | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/tests/e2e_advanced_traces.rs b/tests/e2e_advanced_traces.rs index b3efc8d9..ce18ad3d 100644 --- a/tests/e2e_advanced_traces.rs +++ b/tests/e2e_advanced_traces.rs @@ -587,6 +587,7 @@ mod advanced { async fn mcp_extension_lifecycle() { use crate::support::mock_mcp_server::{MockToolResponse, start_mock_mcp_server}; use ironclaw::extensions::{AuthHint, ExtensionKind, ExtensionSource, RegistryEntry}; + const TEST_USER_ID: &str = "test-user"; // 1. Start mock MCP server with pre-configured tool responses. let mock_server = start_mock_mcp_server(vec![ @@ -654,14 +655,14 @@ mod advanced { ext_mgr .secrets() .create( - "default", + TEST_USER_ID, ironclaw::secrets::CreateSecretParams::new(secret_name, "mock-access-token") .with_provider("mcp:mock-notion".to_string()), ) .await .expect("failed to inject test token"); - let activate_result = ext_mgr.activate("mock-notion", "default").await; + let activate_result = ext_mgr.activate("mock-notion", TEST_USER_ID).await; assert!( activate_result.is_ok(), "activation failed: {:?}", From ab0ad948f36c7cc88b1aecf2e92dd0ff94569a94 Mon Sep 17 00:00:00 2001 From: Henry Park Date: Wed, 25 Mar 2026 13:47:12 -0700 Subject: [PATCH 4/8] Normalize cron schedules on routine create (#1648) * Fix REPL single-message hang and cap CI test duration * Fix Clippy nested-if lint in REPL startup * Fix single-message approval flow * Handle empty single-message REPL exits * Wait for one-shot event routines before exit * Fix MCP lifecycle trace user scope * Normalize cron schedules on routine create --- src/tools/builtin/routine.rs | 16 +++++++++++++++- tests/e2e_builtin_tool_coverage.rs | 2 +- 2 files changed, 16 insertions(+), 2 deletions(-) diff --git a/src/tools/builtin/routine.rs b/src/tools/builtin/routine.rs index f4313483..bbc24139 100644 --- a/src/tools/builtin/routine.rs +++ b/src/tools/builtin/routine.rs @@ -915,7 +915,7 @@ fn parse_routine_create_request( fn build_routine_trigger(trigger: &NormalizedTriggerRequest) -> Trigger { match trigger { NormalizedTriggerRequest::Cron { schedule, timezone } => Trigger::Cron { - schedule: schedule.clone(), + schedule: normalize_cron_expression(schedule), timezone: timezone.clone(), }, NormalizedTriggerRequest::Manual => Trigger::Manual, @@ -1836,6 +1836,20 @@ mod tests { assert_eq!(parsed.cooldown_secs, 30); } + #[test] + fn build_routine_trigger_normalizes_cron_schedule() { + let trigger = build_routine_trigger(&NormalizedTriggerRequest::Cron { + schedule: "0 0 9 * * MON-FRI".to_string(), + timezone: Some("UTC".to_string()), + }); + + assert!(matches!( + trigger, + Trigger::Cron { schedule, timezone } + if schedule == "0 0 9 * * MON-FRI *" && timezone.as_deref() == Some("UTC") + )); + } + #[test] fn parses_grouped_message_event_with_tools() { let params = serde_json::json!({ diff --git a/tests/e2e_builtin_tool_coverage.rs b/tests/e2e_builtin_tool_coverage.rs index 42d7fb75..1c3cc6a2 100644 --- a/tests/e2e_builtin_tool_coverage.rs +++ b/tests/e2e_builtin_tool_coverage.rs @@ -439,7 +439,7 @@ mod tests { match &routine.trigger { Trigger::Cron { schedule, timezone } => { - assert_eq!(schedule, "0 0 9 * * MON-FRI"); + assert_eq!(schedule, "0 0 9 * * MON-FRI *"); assert_eq!(timezone.as_deref(), Some("UTC")); } other => panic!("expected cron trigger, got {other:?}"), From 86d11430640da22d8f890bb9b2df867dda1e668e Mon Sep 17 00:00:00 2001 From: Henry Park Date: Wed, 25 Mar 2026 14:36:53 -0700 Subject: [PATCH 5/8] Fix libsql prompt scope regressions (#1651) --- src/agent/dispatcher.rs | 7 +++- src/workspace/mod.rs | 55 +++++++++++++++++++++++++++++ src/workspace/repository.rs | 1 + tests/e2e_workspace_coverage.rs | 4 ++- tests/multi_tenant_system_prompt.rs | 14 ++++---- 5 files changed, 72 insertions(+), 9 deletions(-) diff --git a/src/agent/dispatcher.rs b/src/agent/dispatcher.rs index cba84c35..fe208c1b 100644 --- a/src/agent/dispatcher.rs +++ b/src/agent/dispatcher.rs @@ -63,7 +63,12 @@ impl Agent { ); let system_prompt = if let Some(ws) = self.workspace() { - match ws + let scoped_workspace = if ws.user_id() == message.user_id { + Arc::clone(ws) + } else { + Arc::new(ws.scoped_to_user(&message.user_id)) + }; + match scoped_workspace .system_prompt_for_context_tz(is_group_chat, user_tz) .await { diff --git a/src/workspace/mod.rs b/src/workspace/mod.rs index 0242047f..51d7d2fc 100644 --- a/src/workspace/mod.rs +++ b/src/workspace/mod.rs @@ -149,6 +149,7 @@ fn reject_if_injected(path: &str, content: &str) -> Result<(), WorkspaceError> { /// /// Allows Workspace to work with either a PostgreSQL `Repository` (the original /// path) or any `Database` trait implementation (e.g. libSQL backend). +#[derive(Clone)] enum WorkspaceStorage { /// PostgreSQL-backed repository (uses connection pool directly). #[cfg(feature = "postgres")] @@ -576,6 +577,60 @@ impl Workspace { self } + /// Clone the workspace configuration for a different primary user scope. + /// + /// This preserves search config, embeddings, shared read scopes, memory + /// layers, and privacy classifier while switching the primary read/write + /// scope to `user_id`. + pub fn scoped_to_user(&self, user_id: impl Into) -> Self { + let user_id = user_id.into(); + + let mut memory_layers = self.memory_layers.clone(); + for layer in &mut memory_layers { + if layer.sensitivity == crate::workspace::layer::LayerSensitivity::Private + && layer.scope == self.user_id + { + layer.scope = user_id.clone(); + } + } + + let mut read_user_ids = vec![user_id.clone()]; + for scope in &self.read_user_ids { + if scope != &self.user_id && !read_user_ids.contains(scope) { + read_user_ids.push(scope.clone()); + } + } + for scope in crate::workspace::layer::MemoryLayer::read_scopes(&memory_layers) { + if !read_user_ids.contains(&scope) { + read_user_ids.push(scope); + } + } + + let preserve_flags = user_id == self.user_id; + Self { + user_id, + read_user_ids, + agent_id: self.agent_id, + storage: self.storage.clone(), + embeddings: self.embeddings.clone(), + bootstrap_pending: std::sync::atomic::AtomicBool::new(if preserve_flags { + self.bootstrap_pending + .load(std::sync::atomic::Ordering::Acquire) + } else { + false + }), + bootstrap_completed: std::sync::atomic::AtomicBool::new(if preserve_flags { + self.bootstrap_completed + .load(std::sync::atomic::Ordering::Acquire) + } else { + false + }), + search_defaults: self.search_defaults.clone(), + memory_layers, + privacy_classifier: self.privacy_classifier.clone(), + } + } + /// Get the user ID (primary scope for writes). pub fn user_id(&self) -> &str { &self.user_id diff --git a/src/workspace/repository.rs b/src/workspace/repository.rs index 78ddfec5..13f6816b 100644 --- a/src/workspace/repository.rs +++ b/src/workspace/repository.rs @@ -15,6 +15,7 @@ use crate::workspace::document::{MemoryChunk, MemoryDocument, WorkspaceEntry}; use crate::workspace::search::{RankedResult, SearchConfig, SearchResult, fuse_results}; /// Database repository for workspace operations. +#[derive(Clone)] pub struct Repository { pool: Pool, } diff --git a/tests/e2e_workspace_coverage.rs b/tests/e2e_workspace_coverage.rs index 396b676e..68956d30 100644 --- a/tests/e2e_workspace_coverage.rs +++ b/tests/e2e_workspace_coverage.rs @@ -12,6 +12,7 @@ mod tests { use crate::support::test_rig::TestRigBuilder; use crate::support::trace_llm::LlmTrace; + use ironclaw::workspace::Workspace; // ----------------------------------------------------------------------- // Test 1: write_chunk_search @@ -268,6 +269,7 @@ mod tests { #[tokio::test] async fn identity_in_system_prompt() { + const TEST_USER_ID: &str = "test-user"; let trace = LlmTrace::from_file(concat!( env!("CARGO_MANIFEST_DIR"), "/tests/fixtures/llm_traces/workspace/identity_prompt.json" @@ -280,7 +282,7 @@ mod tests { .await; // Seed an IDENTITY.md so the system prompt has real content to inject. - let ws = rig.workspace().expect("workspace must be available"); + let ws = Workspace::new_with_db(TEST_USER_ID, rig.database().clone()); ws.write( "IDENTITY.md", "I am TestBot, a helpful testing assistant created for E2E verification.", diff --git a/tests/multi_tenant_system_prompt.rs b/tests/multi_tenant_system_prompt.rs index ece794bf..b89e6cb5 100644 --- a/tests/multi_tenant_system_prompt.rs +++ b/tests/multi_tenant_system_prompt.rs @@ -1,10 +1,10 @@ -//! Tests proving that multi-tenant system prompts are broken. +//! Regression tests for multi-tenant system prompts. //! -//! Bug: In multi-tenant mode, the agent loop uses `self.workspace()` which -//! returns a single shared workspace (user_id="default"). Identity files -//! (IDENTITY.md, SOUL.md, USER.md) seeded under per-user IDs ("alice", -//! "bob") are invisible to this workspace, so the system prompt is -//! empty/wrong. +//! The agent must build the conversational system prompt from a workspace +//! scoped to the incoming message's user, not from the shared owner-scope +//! workspace created at startup. Otherwise per-user identity files +//! (IDENTITY.md, SOUL.md, USER.md) become invisible and different users can +//! see the same owner-scoped prompt. //! //! These tests: //! 1. Seed identity files for two users (alice, bob) in the database @@ -13,7 +13,7 @@ //! correct user's identity //! 4. Verify user A's identity doesn't leak into user B's prompt //! -//! All tests are expected to FAIL until the bug is fixed. +//! These tests ensure each user's identity is isolated correctly. #[cfg(feature = "libsql")] mod support; From 4c043bf05767d7e1ab74552eb010182ec44b3222 Mon Sep 17 00:00:00 2001 From: Illia Polosukhin Date: Wed, 25 Mar 2026 17:24:48 -0700 Subject: [PATCH 6/8] =?UTF-8?q?feat:=20complete=20multi-tenant=20isolation?= =?UTF-8?q?=20=E2=80=94=20phases=202=E2=80=934=20(#1614)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat: complete multi-tenant isolation — per-user budgets, model selection, heartbeat cycling Finishes the remaining isolation work from phases 2–4 of #59: Phase 2 (DB scoping): Fix /status and /list commands to use _for_user DB variants instead of global queries that leaked cross-user job data. Phase 3 (Runtime isolation): Per-user workspace in routine engine's spawn_fire so lightweight routines run in the correct user context. Per-user daily cost tracking in CostGuard with configurable budget via MAX_COST_PER_USER_PER_DAY_CENTS. Multi-user heartbeat that cycles through all users with routines, auto-detected from GATEWAY_USER_TOKENS. Phase 4 (Provider/tools): Per-user model selection via preferred_model setting — looked up from SettingsStore on first iteration, threaded through ReasoningContext.model_override to CompletionRequest. Works with providers that support per-request model overrides (NearAI). Co-Authored-By: Claude Opus 4.6 (1M context) * fix: use selected_model setting key to match /model command persistence The dispatcher was reading "preferred_model" but the /model command (merged from staging) persists to "selected_model". Since set_setting is already per-user scoped, using the same key makes /model work as the per-user model override in multi-tenant mode. Co-Authored-By: Claude Opus 4.6 (1M context) * fix: heartbeat hygiene, /model multi-tenant guard, RigAdapter model override Three follow-up fixes for multi-tenant isolation: 1. Multi-user heartbeat now runs memory hygiene per user before each heartbeat check, matching single-user heartbeat behavior. 2. /model command in multi-tenant mode only persists to per-user settings (selected_model) without calling set_model() on the shared LlmProvider. The per-request model_override in the dispatcher reads from the same setting. Added multi_tenant flag to AgentConfig (auto-detected from GATEWAY_USER_TOKENS). 3. RigAdapter now supports per-request model overrides by injecting the model name into rig-core's additional_params. OpenAI/Anthropic/Ollama API servers use last-key-wins for duplicate JSON keys, so the override takes effect via serde's flatten serialization order. Co-Authored-By: Claude Opus 4.6 (1M context) * fix: address PR review — cost model attribution, heartbeat concurrency, pruning Fixes from review comments on #1614: - Cost tracking now uses the override model name (not active_model_name) when a per-user model override is active, for accurate attribution. - Multi-user heartbeat runs per-user checks concurrently via JoinSet instead of sequentially, preventing one slow user from blocking others. - Per-user failure counts tracked independently; users exceeding max_failures are skipped (matching single-user semantics). - per_user_daily_cost HashMap pruned on day rollover to prevent unbounded growth in long-lived deployments. - Doc comment fixed: says "routines" not "active routines". Co-Authored-By: Claude Opus 4.6 (1M context) * fix: /status ownership, model persistence scoping, heartbeat robustness Addresses second round of PR review on #1614: - /status DB path now validates job.user_id == requesting user before returning data (was missing ownership check, security fix). - persist_selected_model takes user_id param instead of owner_id, and skips .env/TOML writes in multi-tenant mode (these are shared global files). handle_system_command now receives user_id from caller. - JoinSet collection handles Err(JoinError) explicitly instead of silently dropping panicked tasks. - Notification forwarder extracts owner_id from response metadata in multi-tenant mode for per-user routing instead of broadcasting to the agent owner. Co-Authored-By: Claude Opus 4.6 (1M context) * fix: cost pricing, fire_manual workspace, heartbeat concurrency cap Round 3 review fixes: - Cost tracking passes None for cost_per_token when model override is active, letting CostGuard look up pricing by model name instead of using the default provider's rates (serrrfirat). - fire_manual() now uses per-user workspace, matching spawn_fire() pattern (serrrfirat). - Removed MULTI_TENANT env var — multi-tenant mode is auto-detected solely from GATEWAY_USER_TOKENS presence (serrrfirat + Copilot). - Multi-user heartbeat capped at 8 concurrent tasks to avoid flooding the LLM provider (serrrfirat + Copilot). - Fixed inject_model_override doc comment accuracy (Copilot). - Added comment explaining multi-tenant notification routing priority (Copilot). Co-Authored-By: Claude Opus 4.6 (1M context) * feat: user-scoped webhook endpoint for multi-tenant isolation Adds POST /api/webhooks/u/{user_id}/{path} — a user-scoped webhook endpoint that filters the routine lookup by user_id, preventing cross-user webhook triggering when paths collide. The existing /api/webhooks/{path} endpoint remains unchanged for backward compatibility in single-user deployments. Changes: - get_webhook_routine_by_path gains user_id: Option<&str> param - Both postgres and libsql implementations add AND user_id = ? filter when user_id is provided - New webhook_trigger_user_scoped_handler extracts (user_id, path) from URL and passes to shared fire_webhook_inner logic - Route registered on public router (webhooks are called by external services that can't send bearer tokens) Co-Authored-By: Claude Opus 4.6 (1M context) * feat: add TenantCtx for compile-time tenant isolation Implements zmanian's architectural proposal from #1614 review: two-tier scoped database access (TenantScope/AdminScope) so handler code cannot accidentally bypass tenant scoping. TenantScope (default): wraps user_id + Arc, auto-binds user_id on every operation. ID-based lookups return None for cross- tenant resources. No escape hatch — forgetting to scope is a compile error. AdminScope (explicit opt-in): cross-tenant access for system-level components (heartbeat, routine engine, self-repair, scheduler, worker). TenantCtx bundles TenantScope + workspace + cost guard + per-user rate limiting. Constructed once per request in handle_message, threaded through all command handlers and ChatDelegate. Key changes: - New src/tenant.rs (~920 lines): TenantScope, AdminScope, TenantCtx, TenantRateState, TenantRateRegistry - All command handlers: user_id: &str → ctx: &TenantCtx - ChatDelegate: cost check/record/settings via self.tenant - System components: store field changed to AdminScope - Config: TENANT_MAX_LLM_CONCURRENT, TENANT_MAX_JOBS_CONCURRENT env vars - Fixes bug: /status cross-tenant leak (now auto-filtered) Co-Authored-By: Claude Opus 4.6 (1M context) --------- Co-authored-by: Claude Opus 4.6 (1M context) --- src/agent/agent_loop.rs | 173 ++++- src/agent/commands.rs | 141 ++-- src/agent/cost_guard.rs | 239 +++++- src/agent/dispatcher.rs | 62 +- src/agent/heartbeat.rs | 191 ++++- src/agent/mod.rs | 4 +- src/agent/routine_engine.rs | 38 +- src/agent/scheduler.rs | 10 +- src/agent/self_repair.rs | 10 +- src/agent/thread_ops.rs | 13 +- src/app.rs | 1 + src/channels/web/handlers/webhooks.rs | 33 +- src/channels/web/server.rs | 5 + src/config/agent.rs | 21 +- src/config/heartbeat.rs | 10 + src/db/libsql/routines.rs | 21 +- src/db/mod.rs | 1 + src/db/postgres.rs | 3 +- src/history/store.rs | 16 +- src/lib.rs | 1 + src/llm/reasoning.rs | 11 + src/llm/rig_adapter.rs | 52 +- src/main.rs | 4 + src/tenant.rs | 906 ++++++++++++++++++++++ src/testing/mod.rs | 2 + src/worker/job.rs | 6 +- tests/e2e_routine_heartbeat.rs | 20 +- tests/e2e_telegram_message_routing.rs | 1 + tests/support/gateway_workflow_harness.rs | 1 + tests/support/test_rig.rs | 3 +- 30 files changed, 1825 insertions(+), 174 deletions(-) create mode 100644 src/tenant.rs diff --git a/src/agent/agent_loop.rs b/src/agent/agent_loop.rs index e28f11d0..4ee846f7 100644 --- a/src/agent/agent_loop.rs +++ b/src/agent/agent_loop.rs @@ -13,7 +13,7 @@ use futures::StreamExt; use uuid::Uuid; use crate::agent::context_monitor::ContextMonitor; -use crate::agent::heartbeat::spawn_heartbeat; +use crate::agent::heartbeat::{spawn_heartbeat, spawn_multi_user_heartbeat}; use crate::agent::routine_engine::{RoutineEngine, spawn_cron_ticker}; use crate::agent::self_repair::{DefaultSelfRepair, RepairResult, SelfRepair}; use crate::agent::session::ThreadState; @@ -182,6 +182,8 @@ pub struct AgentDeps { /// Resolved LLM backend identifier (e.g., "nearai", "openai", "groq"). /// Used by `/model` persistence to determine which env var to update. pub llm_backend: String, + /// Per-tenant rate limiting registry (lazily creates rate state per user). + pub tenant_rates: Arc, } /// The main agent that coordinates all components. @@ -244,7 +246,10 @@ impl Agent { SchedulerDeps { tools: deps.tools.clone(), extension_manager: deps.extension_manager.clone(), - store: deps.store.clone(), + store: deps + .store + .as_ref() + .map(|db| crate::tenant::AdminScope::new(Arc::clone(db))), hooks: deps.hooks.clone(), }, ); @@ -325,6 +330,50 @@ impl Agent { &self.deps.cost_guard } + /// Build a tenant-scoped execution context for the given user. + /// + /// This is the standard entry point for per-user operations. The returned + /// [`TenantCtx`] provides a [`TenantScope`] that auto-binds `user_id` on + /// every database operation and a per-user rate limiter. + pub(super) async fn tenant_ctx(&self, user_id: &str) -> crate::tenant::TenantCtx { + let rate = self.deps.tenant_rates.get_or_create(user_id).await; + + let store = self + .deps + .store + .as_ref() + .map(|db| crate::tenant::TenantScope::new(user_id, Arc::clone(db))); + + // Reuse the owner workspace if user matches, otherwise create per-user. + let workspace = match &self.deps.workspace { + Some(ws) if ws.user_id() == user_id => Some(Arc::clone(ws)), + _ => self + .deps + .store + .as_ref() + .map(|db| Arc::new(Workspace::new_with_db(user_id, Arc::clone(db)))), + }; + + crate::tenant::TenantCtx::new( + user_id, + store, + workspace, + Arc::clone(&self.deps.cost_guard), + rate, + ) + } + + /// Get an admin-scoped database accessor for cross-tenant operations. + /// + /// Only for system-level components (heartbeat, routine engine, self-repair, + /// scheduler). Handler code should use [`tenant_ctx()`](Self::tenant_ctx) instead. + pub(super) fn admin_store(&self) -> Option { + self.deps + .store + .as_ref() + .map(|db| crate::tenant::AdminScope::new(Arc::clone(db))) + } + pub(super) fn skill_registry(&self) -> Option<&Arc>> { self.deps.skill_registry.as_ref() } @@ -410,8 +459,8 @@ impl Agent { self.config.stuck_threshold, self.config.max_repair_attempts, ); - if let Some(ref store) = self.deps.store { - self_repair = self_repair.with_store(Arc::clone(store)); + if let Some(admin) = self.admin_store() { + self_repair = self_repair.with_store(admin); } if let Some(ref builder) = self.deps.builder { self_repair = self_repair.with_builder(Arc::clone(builder), Arc::clone(self.tools())); @@ -518,6 +567,7 @@ impl Agent { .with_interval(std::time::Duration::from_secs(hb_config.interval_secs)); config.quiet_hours_start = hb_config.quiet_hours_start; config.quiet_hours_end = hb_config.quiet_hours_end; + config.multi_tenant = hb_config.multi_tenant; config.timezone = hb_config .timezone .clone() @@ -547,30 +597,52 @@ impl Agent { .await; let notify_user = heartbeat_notify_user; let channels = self.channels.clone(); + let is_multi_tenant = hb_config.multi_tenant; tokio::spawn(async move { while let Some(response) = notify_rx.recv().await { + // In multi-tenant mode, extract the owning user_id from + // the response metadata so notifications reach the + // correct user rather than the agent's owner. + // This intentionally overrides the configured notify_target + // because each user's heartbeat should notify that user. + let effective_user = if is_multi_tenant { + response + .metadata + .get("owner_id") + .and_then(|v| v.as_str()) + .map(String::from) + } else { + None + }; + // Try the configured channel first, fall back to // broadcasting on all channels. - let targeted_ok = if let Some(ref channel) = notify_channel - && let Some(ref user) = notify_target - { - channels - .broadcast(channel, user, response.clone()) - .await - .is_ok() + let targeted_ok = if let Some(ref channel) = notify_channel { + let target = effective_user.as_deref().or(notify_target.as_deref()); + if let Some(user) = target { + channels + .broadcast(channel, user, response.clone()) + .await + .is_ok() + } else { + false + } } else { false }; - if !targeted_ok && let Some(ref user) = notify_user { - let results = channels.broadcast_all(user, response).await; - for (ch, result) in results { - if let Err(e) = result { - tracing::warn!( - "Failed to broadcast heartbeat to {}: {}", - ch, - e - ); + if !targeted_ok { + let fallback = effective_user.as_deref().or(notify_user.as_deref()); + if let Some(user) = fallback { + let results = channels.broadcast_all(user, response).await; + for (ch, result) in results { + if let Err(e) = result { + tracing::warn!( + "Failed to broadcast heartbeat to {}: {}", + ch, + e + ); + } } } } @@ -583,14 +655,29 @@ impl Agent { .map(|h| h.to_workspace_config()) .unwrap_or_default(); - Some(spawn_heartbeat( - config, - hygiene, - workspace.clone(), - self.cheap_llm().clone(), - Some(notify_tx), - self.store().map(Arc::clone), - )) + if config.multi_tenant { + if let Some(admin) = self.admin_store() { + Some(spawn_multi_user_heartbeat( + config, + hygiene, + self.cheap_llm().clone(), + Some(notify_tx), + admin, + )) + } else { + tracing::warn!("Multi-tenant heartbeat requires a database store"); + None + } + } else { + Some(spawn_heartbeat( + config, + hygiene, + workspace.clone(), + self.cheap_llm().clone(), + Some(notify_tx), + self.admin_store(), + )) + } } else { tracing::warn!("Heartbeat enabled but no workspace available"); None @@ -612,7 +699,7 @@ impl Agent { let engine = Arc::new(RoutineEngine::new( rt_config.clone(), - Arc::clone(store), + crate::tenant::AdminScope::new(Arc::clone(store)), self.llm().clone(), Arc::clone(workspace), notify_tx, @@ -1173,13 +1260,22 @@ impl Agent { } } + // Build per-tenant execution context once; threaded through all handlers. + let tenant = self.tenant_ctx(&message.user_id).await; + let session_for_empty_exit = Arc::clone(&session); // Process based on submission type let result = match submission { Submission::UserInput { content } => { let mut result = self - .process_user_input(message, session.clone(), thread_id, &content) + .process_user_input( + message, + tenant.clone(), + session.clone(), + thread_id, + &content, + ) .await; // Drain any messages queued during processing. @@ -1246,7 +1342,13 @@ impl Agent { let mut queued_msg = message.clone(); queued_msg.attachments.clear(); result = self - .process_user_input(&queued_msg, session.clone(), thread_id, &next_content) + .process_user_input( + &queued_msg, + tenant.clone(), + session.clone(), + thread_id, + &next_content, + ) .await; // If processing failed, re-queue the drained content so it @@ -1294,7 +1396,7 @@ impl Agent { }; } // Authorization checks (including restart channel check) are enforced in handle_system_command - self.handle_system_command(&command, &args, &message.channel) + self.handle_system_command(&command, &args, &message.channel, &tenant) .await } Submission::Undo => self.process_undo(session, thread_id).await, @@ -1307,12 +1409,9 @@ impl Agent { Submission::Summarize => self.process_summarize(session, thread_id).await, Submission::Suggest => self.process_suggest(session, thread_id).await, Submission::JobStatus { job_id } => { - self.process_job_status(&message.user_id, job_id.as_deref()) - .await - } - Submission::JobCancel { job_id } => { - self.process_job_cancel(&message.user_id, &job_id).await + self.process_job_status(&tenant, job_id.as_deref()).await } + Submission::JobCancel { job_id } => self.process_job_cancel(&tenant, &job_id).await, Submission::Quit => return Ok(None), Submission::SwitchThread { thread_id: target } => { self.process_switch_thread(message, target).await diff --git a/src/agent/commands.rs b/src/agent/commands.rs index e02b33db..643d8c7c 100644 --- a/src/agent/commands.rs +++ b/src/agent/commands.rs @@ -33,6 +33,7 @@ impl Agent { &self, intent: MessageIntent, message: &IncomingMessage, + tenant: &crate::tenant::TenantCtx, ) -> Result { // Send thinking status for non-trivial operations if let MessageIntent::CreateJob { .. } = &intent { @@ -52,24 +53,18 @@ impl Agent { description, category, } => { - self.handle_create_job(&message.user_id, title, description, category) + self.handle_create_job(tenant, title, description, category) .await? } MessageIntent::CheckJobStatus { job_id } => { - self.handle_check_status(&message.user_id, job_id).await? - } - MessageIntent::CancelJob { job_id } => { - self.handle_cancel_job(&message.user_id, &job_id).await? - } - MessageIntent::ListJobs { filter } => { - self.handle_list_jobs(&message.user_id, filter).await? - } - MessageIntent::HelpJob { job_id } => { - self.handle_help_job(&message.user_id, &job_id).await? + self.handle_check_status(tenant, job_id).await? } + MessageIntent::CancelJob { job_id } => self.handle_cancel_job(tenant, &job_id).await?, + MessageIntent::ListJobs { filter } => self.handle_list_jobs(tenant, filter).await?, + MessageIntent::HelpJob { job_id } => self.handle_help_job(tenant, &job_id).await?, MessageIntent::Command { command, args } => { match self - .handle_command(&command, &args, &message.channel) + .handle_command(&command, &args, &message.channel, tenant) .await? { Some(s) => s, @@ -83,14 +78,14 @@ impl Agent { async fn handle_create_job( &self, - user_id: &str, + tenant: &crate::tenant::TenantCtx, title: String, description: String, category: Option, ) -> Result { let job_id = self .scheduler - .dispatch_job(user_id, &title, &description, None) + .dispatch_job(tenant.user_id(), &title, &description, None) .await?; // Set the dedicated category field (not stored in metadata) @@ -113,7 +108,7 @@ impl Agent { async fn handle_check_status( &self, - user_id: &str, + tenant: &crate::tenant::TenantCtx, job_id: Option, ) -> Result { match job_id { @@ -122,7 +117,8 @@ impl Agent { .map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?; // Try DB first for persistent state, fall back to ContextManager. - if let Some(store) = self.store() + // TenantScope.get_job() auto-filters by ownership — no manual check needed. + if let Some(store) = tenant.store() && let Ok(Some(ctx)) = store.get_job(uuid).await { return Ok(format!( @@ -138,7 +134,7 @@ impl Agent { } let ctx = self.context_manager.get_context(uuid).await?; - if ctx.user_id != user_id { + if ctx.user_id != tenant.user_id() { return Err(crate::error::JobError::NotFound { id: uuid }.into()); } @@ -155,7 +151,8 @@ impl Agent { } None => { // Show summary from DB for consistency with Jobs tab. - if let Some(store) = self.store() { + // TenantScope methods auto-scope to user — no user_id parameter needed. + if let Some(store) = tenant.store() { let mut total = 0; let mut in_progress = 0; let mut completed = 0; @@ -183,7 +180,7 @@ impl Agent { } // Fallback to ContextManager if no DB. - let summary = self.context_manager.summary_for(user_id).await; + let summary = self.context_manager.summary_for(tenant.user_id()).await; Ok(format!( "Jobs summary: Total: {} In Progress: {} Completed: {} Failed: {} Stuck: {}", summary.total, @@ -196,19 +193,24 @@ impl Agent { } } - async fn handle_cancel_job(&self, user_id: &str, job_id: &str) -> Result { + async fn handle_cancel_job( + &self, + tenant: &crate::tenant::TenantCtx, + job_id: &str, + ) -> Result { let uuid = Uuid::parse_str(job_id) .map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?; let ctx = self.context_manager.get_context(uuid).await?; - if ctx.user_id != user_id { + if ctx.user_id != tenant.user_id() { return Err(crate::error::JobError::NotFound { id: uuid }.into()); } self.scheduler.stop(uuid).await?; // Also update DB so the Jobs tab reflects cancellation immediately. - if let Some(store) = self.store() + // Use TenantScope — ownership already verified above. + if let Some(store) = tenant.store() && let Err(e) = store .update_job_status(uuid, JobState::Cancelled, Some("Cancelled by user")) .await @@ -221,11 +223,12 @@ impl Agent { async fn handle_list_jobs( &self, - user_id: &str, + tenant: &crate::tenant::TenantCtx, _filter: Option, ) -> Result { // List from DB for consistency with Jobs tab. - if let Some(store) = self.store() { + // TenantScope methods auto-scope to user. + if let Some(store) = tenant.store() { let agent_jobs = match store.list_agent_jobs().await { Ok(jobs) => jobs, Err(e) => { @@ -256,7 +259,7 @@ impl Agent { } // Fallback to ContextManager if no DB. - let jobs = self.context_manager.all_jobs_for(user_id).await; + let jobs = self.context_manager.all_jobs_for(tenant.user_id()).await; if jobs.is_empty() { return Ok("No jobs found.".to_string()); } @@ -270,12 +273,16 @@ impl Agent { Ok(output) } - async fn handle_help_job(&self, user_id: &str, job_id: &str) -> Result { + async fn handle_help_job( + &self, + tenant: &crate::tenant::TenantCtx, + job_id: &str, + ) -> Result { let uuid = Uuid::parse_str(job_id) .map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?; let ctx = self.context_manager.get_context(uuid).await?; - if ctx.user_id != user_id { + if ctx.user_id != tenant.user_id() { return Err(crate::error::JobError::NotFound { id: uuid }.into()); } @@ -308,11 +315,11 @@ impl Agent { /// Show job status inline — either all jobs (no id) or a specific job. pub(super) async fn process_job_status( &self, - user_id: &str, + tenant: &crate::tenant::TenantCtx, job_id: Option<&str>, ) -> Result { match self - .handle_check_status(user_id, job_id.map(|s| s.to_string())) + .handle_check_status(tenant, job_id.map(|s| s.to_string())) .await { Ok(text) => Ok(SubmissionResult::response(text)), @@ -323,10 +330,10 @@ impl Agent { /// Cancel a job by ID. pub(super) async fn process_job_cancel( &self, - user_id: &str, + tenant: &crate::tenant::TenantCtx, job_id: &str, ) -> Result { - match self.handle_cancel_job(user_id, job_id).await { + match self.handle_cancel_job(tenant, job_id).await { Ok(text) => Ok(SubmissionResult::response(text)), Err(e) => Ok(SubmissionResult::error(format!("Cancel error: {}", e))), } @@ -559,6 +566,7 @@ impl Agent { command: &str, args: &[String], channel: &str, + tenant: &crate::tenant::TenantCtx, ) -> Result { match command { "help" => Ok(SubmissionResult::response(concat!( @@ -752,19 +760,32 @@ impl Agent { } } - match self.llm().set_model(requested) { - Ok(()) => { - // Persist the model choice so it survives restarts. - self.persist_selected_model(requested).await; - Ok(SubmissionResult::response(format!( - "Switched model to: {}", - requested - ))) + if self.config.multi_tenant { + // Multi-tenant: only persist to per-user DB settings. + // Do NOT call set_model() on the shared provider — that + // would change the default for all users. The per-request + // model_override in the dispatcher reads from the same + // "selected_model" setting and applies it per-user. + self.persist_selected_model(tenant, requested).await; + Ok(SubmissionResult::response(format!( + "Model preference set to: {} (per-user)", + requested + ))) + } else { + match self.llm().set_model(requested) { + Ok(()) => { + // Persist the model choice so it survives restarts. + self.persist_selected_model(tenant, requested).await; + Ok(SubmissionResult::response(format!( + "Switched model to: {}", + requested + ))) + } + Err(e) => Ok(SubmissionResult::error(format!( + "Failed to switch model: {}", + e + ))), } - Err(e) => Ok(SubmissionResult::error(format!( - "Failed to switch model: {}", - e - ))), } } } @@ -906,10 +927,14 @@ impl Agent { command: &str, args: &[String], channel: &str, + tenant: &crate::tenant::TenantCtx, ) -> Result, Error> { // System commands are now handled directly via Submission::SystemCommand, // but the router may still send us unknown /commands. - match self.handle_system_command(command, args, channel).await? { + match self + .handle_system_command(command, args, channel, tenant) + .await? + { SubmissionResult::Response { content } => Ok(Some(content)), SubmissionResult::Ok { message } => Ok(message), SubmissionResult::Error { message } => Ok(Some(format!("Error: {}", message))), @@ -921,23 +946,33 @@ impl Agent { /// /// Best-effort: logs warnings on failure but does not propagate errors, /// since the in-memory model switch already succeeded. - async fn persist_selected_model(&self, model: &str) { - // 1. Persist to DB if available. - if let Some(store) = self.store() { + /// + /// In multi-tenant mode, only the per-user DB setting is written — global + /// .env and TOML files are shared across users and must not be mutated. + async fn persist_selected_model(&self, tenant: &crate::tenant::TenantCtx, model: &str) { + // 1. Persist to DB if available (per-user scoped via TenantScope). + if let Some(store) = tenant.store() { let value = serde_json::Value::String(model.to_string()); - if let Err(e) = store - .set_setting(self.owner_id(), "selected_model", &value) - .await - { + if let Err(e) = store.set_setting("selected_model", &value).await { tracing::warn!("Failed to persist model to DB: {}", e); } else { - tracing::debug!("Persisted selected_model to DB: {}", model); + tracing::debug!( + user_id = tenant.user_id(), + "Persisted selected_model to DB: {}", + model + ); } } else { tracing::warn!("No database store available — model choice will not persist to DB"); } - // 2. Update .env and TOML config file (sync I/O in spawn_blocking). + // 2. In multi-tenant mode, skip .env/TOML writes — these are global + // files shared by all users. The per-user DB setting is sufficient. + if self.config.multi_tenant { + return; + } + + // 3. Update .env and TOML config file (sync I/O in spawn_blocking). let model_owned = model.to_string(); let backend = self.deps.llm_backend.clone(); if let Err(e) = tokio::task::spawn_blocking(move || { diff --git a/src/agent/cost_guard.rs b/src/agent/cost_guard.rs index 4563bbbe..4885364b 100644 --- a/src/agent/cost_guard.rs +++ b/src/agent/cost_guard.rs @@ -21,6 +21,9 @@ pub struct CostGuardConfig { pub max_cost_per_day_cents: Option, /// Maximum LLM calls per hour. None = unlimited. pub max_actions_per_hour: Option, + /// Maximum spend per user per day in cents. None = unlimited. + /// Applied independently per user alongside the global budget. + pub max_cost_per_user_per_day_cents: Option, } /// Error returned when a cost limit is exceeded. @@ -30,6 +33,12 @@ pub enum CostLimitExceeded { DailyBudget { spent_cents: u64, limit_cents: u64 }, /// Hourly action rate limit reached. HourlyRate { actions: u64, limit: u64 }, + /// Per-user daily spending cap reached. + UserDailyBudget { + user_id: String, + spent_cents: u64, + limit_cents: u64, + }, } impl std::fmt::Display for CostLimitExceeded { @@ -49,6 +58,17 @@ impl std::fmt::Display for CostLimitExceeded { "Hourly action limit exceeded: {} actions of {} allowed per hour", actions, limit ), + Self::UserDailyBudget { + user_id, + spent_cents, + limit_cents, + } => write!( + f, + "User '{}' daily cost limit exceeded: spent ${:.2} of ${:.2} allowed", + user_id, + *spent_cents as f64 / 100.0, + *limit_cents as f64 / 100.0 + ), } } } @@ -78,6 +98,9 @@ pub struct CostGuard { /// Per-model token usage since startup. model_tokens: Mutex>, + + /// Per-user daily cost tracking. Each entry resets independently at midnight UTC. + per_user_daily_cost: Mutex>, } struct DailyCost { @@ -97,6 +120,7 @@ impl CostGuard { action_window: Mutex::new(VecDeque::new()), budget_exceeded: AtomicBool::new(false), model_tokens: Mutex::new(HashMap::new()), + per_user_daily_cost: Mutex::new(HashMap::new()), } } @@ -203,6 +227,11 @@ impl CostGuard { daily.reset_date = today; self.budget_exceeded.store(false, Ordering::Relaxed); tracing::info!("Cost guard: daily counter reset for {}", today); + + // Prune per-user entries from previous days to prevent + // unbounded HashMap growth in long-lived deployments. + let mut per_user = self.per_user_daily_cost.lock().await; + per_user.retain(|_, entry| entry.reset_date == today); } daily.total += cost; @@ -248,6 +277,85 @@ impl CostGuard { cost } + /// Record an LLM call with per-user attribution. + /// + /// Delegates to `record_llm_call` for global tracking, then additionally + /// records the cost against the user's daily budget. + #[allow(clippy::too_many_arguments)] + pub async fn record_llm_call_for_user( + &self, + user_id: &str, + model: &str, + input_tokens: u32, + output_tokens: u32, + cache_read_input_tokens: u32, + cache_creation_input_tokens: u32, + cache_read_discount: Decimal, + cache_write_multiplier: Decimal, + cost_per_token: Option<(Decimal, Decimal)>, + ) -> Decimal { + let cost = self + .record_llm_call( + model, + input_tokens, + output_tokens, + cache_read_input_tokens, + cache_creation_input_tokens, + cache_read_discount, + cache_write_multiplier, + cost_per_token, + ) + .await; + + // Track per-user daily cost + { + let today = chrono::Utc::now().date_naive(); + let mut per_user = self.per_user_daily_cost.lock().await; + let entry = per_user + .entry(user_id.to_string()) + .or_insert_with(|| DailyCost { + total: Decimal::ZERO, + reset_date: today, + }); + if today != entry.reset_date { + entry.total = Decimal::ZERO; + entry.reset_date = today; + } + entry.total += cost; + } + + cost + } + + /// Check whether the next action is allowed for a specific user. + /// + /// Checks the global limits first (via `check_allowed`), then additionally + /// checks the per-user daily budget if configured. + pub async fn check_allowed_for_user(&self, user_id: &str) -> Result<(), CostLimitExceeded> { + // Check global limits first + self.check_allowed().await?; + + // Check per-user daily budget + if let Some(limit_cents) = self.config.max_cost_per_user_per_day_cents { + let today = chrono::Utc::now().date_naive(); + let per_user = self.per_user_daily_cost.lock().await; + if let Some(entry) = per_user.get(user_id) + && entry.reset_date == today + { + let spent_cents = to_cents(entry.total); + if spent_cents >= limit_cents { + return Err(CostLimitExceeded::UserDailyBudget { + user_id: user_id.to_string(), + spent_cents, + limit_cents, + }); + } + } + } + + Ok(()) + } + /// Current daily spend in USD (as Decimal). pub async fn daily_spend(&self) -> Decimal { let daily = self.daily_cost.lock().await; @@ -259,6 +367,16 @@ impl CostGuard { } } + /// Current daily spend for a specific user in USD (as Decimal). + pub async fn daily_spend_for_user(&self, user_id: &str) -> Decimal { + let today = chrono::Utc::now().date_naive(); + let per_user = self.per_user_daily_cost.lock().await; + match per_user.get(user_id) { + Some(entry) if entry.reset_date == today => entry.total, + _ => Decimal::ZERO, + } + } + /// Number of actions in the current hourly window. pub async fn actions_this_hour(&self) -> u64 { let mut window = self.action_window.lock().await; @@ -314,7 +432,7 @@ mod tests { async fn test_daily_budget_enforcement() { let guard = CostGuard::new(CostGuardConfig { max_cost_per_day_cents: Some(1), // $0.01 limit - max_actions_per_hour: None, + ..CostGuardConfig::default() }); // First call allowed @@ -350,8 +468,8 @@ mod tests { #[tokio::test] async fn test_hourly_rate_enforcement() { let guard = CostGuard::new(CostGuardConfig { - max_cost_per_day_cents: None, max_actions_per_hour: Some(3), + ..CostGuardConfig::default() }); // First 3 actions allowed @@ -633,8 +751,8 @@ mod tests { // A fresh CostGuard with rate limits should not panic even if // checked_sub returns None (simulating short uptime). let guard = CostGuard::new(CostGuardConfig { - max_cost_per_day_cents: None, max_actions_per_hour: Some(100), + ..CostGuardConfig::default() }); // These must not panic regardless of system uptime @@ -656,4 +774,119 @@ mod tests { let result = Instant::now().checked_sub(std::time::Duration::MAX); assert!(result.is_none()); } + + #[tokio::test] + async fn test_per_user_daily_budget_enforcement() { + let guard = CostGuard::new(CostGuardConfig { + max_cost_per_day_cents: None, + max_actions_per_hour: None, + max_cost_per_user_per_day_cents: Some(1), // $0.01 per user + }); + + // Both users initially allowed + assert!(guard.check_allowed_for_user("alice").await.is_ok()); + assert!(guard.check_allowed_for_user("bob").await.is_ok()); + + // Alice makes an expensive call + guard + .record_llm_call_for_user( + "alice", + "gpt-4o", + 10_000, + 10_000, + 0, + 0, + Decimal::ONE, + Decimal::ONE, + None, + ) + .await; + + // Alice should be blocked, Bob should still be allowed + let result = guard.check_allowed_for_user("alice").await; + assert!(result.is_err()); + match result.unwrap_err() { + CostLimitExceeded::UserDailyBudget { + user_id, + limit_cents, + .. + } => { + assert_eq!(user_id, "alice"); + assert_eq!(limit_cents, 1); + } + other => panic!("Expected UserDailyBudget, got {:?}", other), + } + assert!(guard.check_allowed_for_user("bob").await.is_ok()); + } + + #[tokio::test] + async fn test_per_user_daily_spend_tracking() { + let guard = CostGuard::new(CostGuardConfig::default()); + + assert_eq!(guard.daily_spend_for_user("alice").await, Decimal::ZERO); + assert_eq!(guard.daily_spend_for_user("bob").await, Decimal::ZERO); + + let cost = guard + .record_llm_call_for_user( + "alice", + "gpt-4o", + 1000, + 500, + 0, + 0, + Decimal::ONE, + Decimal::ONE, + None, + ) + .await; + + assert_eq!(guard.daily_spend_for_user("alice").await, cost); + assert_eq!(guard.daily_spend_for_user("bob").await, Decimal::ZERO); + // Global spend should also be tracked + assert_eq!(guard.daily_spend().await, cost); + } + + #[tokio::test] + async fn test_per_user_budget_independent_of_global() { + let guard = CostGuard::new(CostGuardConfig { + max_cost_per_day_cents: Some(100_000), // $1000 global limit + max_actions_per_hour: None, + max_cost_per_user_per_day_cents: Some(1), // $0.01 per user + }); + + // User hits their personal limit + guard + .record_llm_call_for_user( + "alice", + "gpt-4o", + 10_000, + 10_000, + 0, + 0, + Decimal::ONE, + Decimal::ONE, + None, + ) + .await; + + // Alice blocked by per-user limit, not global + assert!(guard.check_allowed_for_user("alice").await.is_err()); + // Global limit is far from reached + assert!(guard.check_allowed().await.is_ok()); + // Bob is unaffected + assert!(guard.check_allowed_for_user("bob").await.is_ok()); + } + + #[test] + fn test_user_cost_limit_display() { + let limit = CostLimitExceeded::UserDailyBudget { + user_id: "alice".to_string(), + spent_cents: 150, + limit_cents: 100, + }; + let msg = limit.to_string(); + assert!(msg.contains("alice")); + assert!(msg.contains("$1.50")); + assert!(msg.contains("$1.00")); + } } diff --git a/src/agent/dispatcher.rs b/src/agent/dispatcher.rs index fe208c1b..96bca197 100644 --- a/src/agent/dispatcher.rs +++ b/src/agent/dispatcher.rs @@ -42,6 +42,7 @@ impl Agent { pub(super) async fn run_agentic_loop( &self, message: &IncomingMessage, + tenant: crate::tenant::TenantCtx, session: Arc>, thread_id: Uuid, initial_messages: Vec, @@ -168,6 +169,7 @@ impl Agent { let delegate = ChatDelegate { agent: self, + tenant, session: session.clone(), thread_id, message, @@ -240,6 +242,7 @@ impl Agent { /// auth intercept, and cost tracking. struct ChatDelegate<'a> { agent: &'a Agent, + tenant: crate::tenant::TenantCtx, session: Arc>, thread_id: Uuid, message: &'a IncomingMessage, @@ -336,8 +339,8 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { reason_ctx: &mut ReasoningContext, iteration: usize, ) -> Result { - // Enforce cost guardrails before the LLM call - if let Err(limit) = self.agent.cost_guard().check_allowed().await { + // Enforce cost guardrails before the LLM call (global + per-user) + if let Err(limit) = self.tenant.check_cost_allowed().await { return Err(crate::error::LlmError::InvalidResponse { provider: "agent".to_string(), reason: limit.to_string(), @@ -345,6 +348,21 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { .into()); } + // Apply per-user model override from settings (first iteration only + // to avoid repeated DB lookups within the same agentic loop). + // Uses "selected_model" — the same key the /model command persists to + // via SettingsStore (per-user scoped via TenantScope). + if iteration == 0 + && let Some(store) = self.tenant.store() + && let Ok(Some(value)) = store.get_setting("selected_model").await + && let Some(model) = value.as_str() + { + let model = model.trim(); + if !model.is_empty() { + reason_ctx.model_override = Some(model.to_string()); + } + } + let output = match reasoning.respond_with_tools(reason_ctx).await { Ok(output) => output, Err(crate::error::LlmError::ContextLengthExceeded { used, limit }) => { @@ -379,13 +397,22 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { Err(e) => return Err(e.into()), }; - // Record cost and track token usage - let model_name = self.agent.llm().active_model_name(); + // Record cost and track token usage (global + per-user). + // When a model override is active, use the override name for attribution + // and let CostGuard look up pricing via costs::model_cost() instead of + // using the default provider's cost_per_token (which reflects the wrong model). + let (model_name, cost_per_token) = if let Some(ref ovr) = reason_ctx.model_override { + (ovr.clone(), None) + } else { + ( + self.agent.llm().active_model_name(), + Some(self.agent.llm().cost_per_token()), + ) + }; let read_discount = self.agent.llm().cache_read_discount(); let write_multiplier = self.agent.llm().cache_write_multiplier(); let call_cost = self - .agent - .cost_guard() + .tenant .record_llm_call( &model_name, output.usage.input_tokens, @@ -394,7 +421,7 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { output.usage.cache_creation_input_tokens, read_discount, write_multiplier, - Some(self.agent.llm().cost_per_token()), + cost_per_token, ) .await; tracing::debug!( @@ -1305,6 +1332,7 @@ mod tests { sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, builder: None, llm_backend: "nearai".to_string(), + tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)), }; Agent::new( @@ -1320,10 +1348,14 @@ mod tests { allow_local_tools: false, max_cost_per_day_cents: None, max_actions_per_hour: None, + max_cost_per_user_per_day_cents: None, max_tool_iterations: 50, auto_approve_tools: false, default_timezone: "UTC".to_string(), max_tokens_per_job: 0, + multi_tenant: false, + max_llm_concurrent_per_user: None, + max_jobs_concurrent_per_user: None, }, deps, Arc::new(ChannelManager::new()), @@ -2181,6 +2213,7 @@ mod tests { sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, builder: None, llm_backend: "nearai".to_string(), + tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)), }; Agent::new( @@ -2196,10 +2229,14 @@ mod tests { allow_local_tools: false, max_cost_per_day_cents: None, max_actions_per_hour: None, + max_cost_per_user_per_day_cents: None, max_tool_iterations, auto_approve_tools: true, default_timezone: "UTC".to_string(), max_tokens_per_job: 0, + multi_tenant: false, + max_llm_concurrent_per_user: None, + max_jobs_concurrent_per_user: None, }, deps, Arc::new(ChannelManager::new()), @@ -2234,13 +2271,14 @@ mod tests { let message = IncomingMessage::new("test", "test-user", "do something"); let initial_messages = vec![ChatMessage::user("do something")]; + let tenant = agent.tenant_ctx("test-user").await; // The dispatcher must terminate within 5 seconds. If there is an // infinite loop bug (e.g., index not advancing on tool failure), the // timeout will fire and the test will fail. let result = tokio::time::timeout( Duration::from_secs(5), - agent.run_agentic_loop(&message, session, thread_id, initial_messages), + agent.run_agentic_loop(&message, tenant, session, thread_id, initial_messages), ) .await; @@ -2302,6 +2340,7 @@ mod tests { sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, builder: None, llm_backend: "nearai".to_string(), + tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)), }; Agent::new( @@ -2317,10 +2356,14 @@ mod tests { allow_local_tools: false, max_cost_per_day_cents: None, max_actions_per_hour: None, + max_cost_per_user_per_day_cents: None, max_tool_iterations: max_iter, auto_approve_tools: true, default_timezone: "UTC".to_string(), max_tokens_per_job: 0, + multi_tenant: false, + max_llm_concurrent_per_user: None, + max_jobs_concurrent_per_user: None, }, deps, Arc::new(ChannelManager::new()), @@ -2340,13 +2383,14 @@ mod tests { let message = IncomingMessage::new("test", "test-user", "keep calling tools"); let initial_messages = vec![ChatMessage::user("keep calling tools")]; + let tenant = agent.tenant_ctx("test-user").await; // Even with an LLM that always wants to call tools, the dispatcher // must terminate within the timeout thanks to force_text at // max_tool_iterations. let result = tokio::time::timeout( Duration::from_secs(5), - agent.run_agentic_loop(&message, session, thread_id, initial_messages), + agent.run_agentic_loop(&message, tenant, session, thread_id, initial_messages), ) .await; diff --git a/src/agent/heartbeat.rs b/src/agent/heartbeat.rs index ec4cd5e9..f7a8f869 100644 --- a/src/agent/heartbeat.rs +++ b/src/agent/heartbeat.rs @@ -31,8 +31,8 @@ use chrono_tz::Tz; use tokio::sync::mpsc; use crate::channels::OutgoingResponse; -use crate::db::Database; use crate::llm::{ChatMessage, CompletionRequest, LlmProvider, Reasoning}; +use crate::tenant::AdminScope; use crate::workspace::Workspace; use crate::workspace::hygiene::HygieneConfig; @@ -57,6 +57,9 @@ pub struct HeartbeatConfig { pub quiet_hours_end: Option, /// Timezone for fire_at and quiet hours evaluation (IANA name). pub timezone: Option, + /// When true, cycle through all users with routines instead of + /// running heartbeat for a single user. Requires a database store. + pub multi_tenant: bool, } impl Default for HeartbeatConfig { @@ -71,6 +74,7 @@ impl Default for HeartbeatConfig { quiet_hours_start: None, quiet_hours_end: None, timezone: None, + multi_tenant: false, } } } @@ -178,7 +182,7 @@ pub struct HeartbeatRunner { workspace: Arc, llm: Arc, response_tx: Option>, - store: Option>, + store: Option, consecutive_failures: u32, } @@ -207,8 +211,8 @@ impl HeartbeatRunner { self } - /// Set the database store for persistent heartbeat conversations. - pub fn with_store(mut self, store: Arc) -> Self { + /// Set the admin-scoped database store for persistent heartbeat conversations. + pub fn with_store(mut self, store: AdminScope) -> Self { self.store = Some(store); self } @@ -396,7 +400,7 @@ impl HeartbeatRunner { } /// Send a notification about heartbeat findings. - async fn send_notification(&self, message: &str) { + pub(crate) async fn send_notification(&self, message: &str) { let Some(ref tx) = self.response_tx else { tracing::debug!("No response channel configured for heartbeat notifications"); return; @@ -493,7 +497,7 @@ pub fn spawn_heartbeat( workspace: Arc, llm: Arc, response_tx: Option>, - store: Option>, + store: Option, ) -> tokio::task::JoinHandle<()> { let mut runner = HeartbeatRunner::new(config, hygiene_config, workspace, llm); if let Some(tx) = response_tx { @@ -508,6 +512,179 @@ pub fn spawn_heartbeat( }) } +/// Spawn a multi-user heartbeat runner that cycles through all users that +/// own routines (enabled or not). Each tick, it queries the DB for distinct +/// user_ids, creates a per-user workspace, and runs a heartbeat check for +/// each user concurrently. Per-user failure counts are tracked independently. +pub fn spawn_multi_user_heartbeat( + config: HeartbeatConfig, + hygiene_config: HygieneConfig, + llm: Arc, + response_tx: Option>, + store: AdminScope, +) -> tokio::task::JoinHandle<()> { + tokio::spawn(async move { + if !config.enabled { + tracing::info!("Multi-user heartbeat is disabled"); + return; + } + + let mut tick_interval = if config.fire_at.is_none() { + let mut iv = tokio::time::interval(config.interval); + iv.tick().await; // skip immediate tick + Some(iv) + } else { + None + }; + + // Track consecutive failures per user so we can disable heartbeat + // for persistently-failing users (same semantics as single-user mode). + let mut user_failures: std::collections::HashMap = + std::collections::HashMap::new(); + + tracing::info!("Starting multi-user heartbeat loop"); + + loop { + if let Some(fire_at) = config.fire_at { + let sleep_dur = duration_until_next_fire(fire_at, config.resolved_tz()); + tokio::time::sleep(sleep_dur).await; + } else if let Some(ref mut iv) = tick_interval { + iv.tick().await; + } + + if config.is_quiet_hours() { + continue; + } + + // Get distinct user_ids from routines + let user_ids = match store.list_all_routines().await { + Ok(routines) => { + let mut ids: Vec = routines + .iter() + .map(|r| r.user_id.clone()) + .collect::>() + .into_iter() + .collect(); + ids.sort(); + ids + } + Err(e) => { + tracing::error!("Multi-user heartbeat: failed to list routines: {}", e); + continue; + } + }; + + // Run user heartbeats concurrently so one slow LLM call doesn't + // block others. Cap concurrency to avoid flooding the LLM provider. + const MAX_CONCURRENT_HEARTBEATS: usize = 8; + let mut join_set = tokio::task::JoinSet::new(); + + for user_id in &user_ids { + // Skip users that have exceeded max_failures + let failures = user_failures.get(user_id).copied().unwrap_or(0); + if failures >= config.max_failures { + continue; + } + + let workspace = Arc::new(Workspace::new_with_db(user_id, Arc::clone(store.db()))); + + // Run memory hygiene per user (same as single-user heartbeat). + let hygiene_ws = Arc::clone(&workspace); + let hygiene_cfg = hygiene_config.clone(); + let hygiene_user = user_id.clone(); + tokio::spawn(async move { + let report = + crate::workspace::hygiene::run_if_due(&hygiene_ws, &hygiene_cfg).await; + if report.had_work() { + tracing::info!( + user_id = hygiene_user, + daily_logs_deleted = report.daily_logs_deleted, + conversation_docs_deleted = report.conversation_docs_deleted, + "multi-user heartbeat: memory hygiene deleted stale documents" + ); + } + }); + + // Drain completed tasks to stay within the concurrency cap. + while join_set.len() >= MAX_CONCURRENT_HEARTBEATS { + if let Some(join_result) = join_set.join_next().await { + collect_heartbeat_result(join_result, &mut user_failures, &config); + } + } + + let uid = user_id.clone(); + let cfg = config.clone(); + let hyg = hygiene_config.clone(); + let llm_clone = llm.clone(); + let tx = response_tx.clone(); + let admin = store.clone(); + + join_set.spawn(async move { + let mut runner = HeartbeatRunner::new(cfg, hyg, workspace, llm_clone); + if let Some(tx) = tx { + runner = runner.with_response_channel(tx); + } + runner = runner.with_store(admin); + + let result = runner.check_heartbeat().await; + if let HeartbeatResult::NeedsAttention(msg) = &result { + runner.send_notification(msg).await; + } + (uid, result) + }); + } + + // Collect remaining results and update failure counts + while let Some(join_result) = join_set.join_next().await { + collect_heartbeat_result(join_result, &mut user_failures, &config); + } + } + }) +} + +/// Process a single JoinSet result from the multi-user heartbeat loop. +fn collect_heartbeat_result( + join_result: Result<(String, HeartbeatResult), tokio::task::JoinError>, + user_failures: &mut std::collections::HashMap, + config: &HeartbeatConfig, +) { + let (uid, result) = match join_result { + Ok(pair) => pair, + Err(e) => { + tracing::error!("Multi-user heartbeat task panicked: {}", e); + return; + } + }; + match result { + HeartbeatResult::Ok => { + tracing::trace!(user_id = uid, "Multi-user heartbeat OK"); + user_failures.remove(&uid); + } + HeartbeatResult::NeedsAttention(_) => { + tracing::info!(user_id = uid, "Multi-user heartbeat needs attention"); + user_failures.remove(&uid); + } + HeartbeatResult::Skipped => {} + HeartbeatResult::Failed(err) => { + let count = user_failures.entry(uid.clone()).or_insert(0); + *count += 1; + tracing::error!( + user_id = uid, + consecutive_failures = *count, + "Multi-user heartbeat failed: {}", + err + ); + if *count >= config.max_failures { + tracing::error!( + user_id = uid, + "Multi-user heartbeat disabled for user after {} consecutive failures", + count + ); + } + } + } +} + #[cfg(test)] mod tests { use super::*; @@ -726,7 +903,7 @@ mod tests { Arc, Arc, Option>, - Option>, + Option, ) -> tokio::task::JoinHandle<()> = spawn_heartbeat; let _ = _fn_ptr; } diff --git a/src/agent/mod.rs b/src/agent/mod.rs index 84155666..e7242845 100644 --- a/src/agent/mod.rs +++ b/src/agent/mod.rs @@ -36,7 +36,9 @@ pub(crate) use agent_loop::truncate_for_preview; pub use agent_loop::{Agent, AgentDeps}; pub use compaction::{CompactionResult, ContextCompactor}; pub use context_monitor::{CompactionStrategy, ContextBreakdown, ContextMonitor}; -pub use heartbeat::{HeartbeatConfig, HeartbeatResult, HeartbeatRunner, spawn_heartbeat}; +pub use heartbeat::{ + HeartbeatConfig, HeartbeatResult, HeartbeatRunner, spawn_heartbeat, spawn_multi_user_heartbeat, +}; pub use router::{MessageIntent, Router}; pub use routine::{Routine, RoutineAction, RoutineRun, Trigger}; pub use routine_engine::{RoutineEngine, SandboxReadiness}; diff --git a/src/agent/routine_engine.rs b/src/agent/routine_engine.rs index a3cdb6cd..64c3b94c 100644 --- a/src/agent/routine_engine.rs +++ b/src/agent/routine_engine.rs @@ -28,12 +28,12 @@ use crate::agent::routine::{ use crate::channels::{IncomingMessage, OutgoingResponse}; use crate::config::RoutineConfig; use crate::context::{JobContext, JobState}; -use crate::db::Database; use crate::error::RoutineError; use crate::extensions::ExtensionManager; use crate::llm::{ ChatMessage, CompletionRequest, FinishReason, LlmProvider, ToolCall, ToolCompletionRequest, }; +use crate::tenant::AdminScope; use crate::tools::{ ToolError, ToolRegistry, autonomous_allowed_tool_names, autonomous_unavailable_message, prepare_tool_params, @@ -99,7 +99,7 @@ pub(crate) fn routine_matches_message(routine: &Routine, message: &IncomingMessa /// The routine execution engine. pub struct RoutineEngine { config: RoutineConfig, - store: Arc, + store: AdminScope, llm: Arc, workspace: Arc, /// Sender for notifications (routed to channel manager). @@ -128,7 +128,7 @@ impl RoutineEngine { #[allow(clippy::too_many_arguments)] pub fn new( config: RoutineConfig, - store: Arc, + store: AdminScope, llm: Arc, workspace: Arc, notify_tx: mpsc::Sender, @@ -782,12 +782,22 @@ impl RoutineEngine { }); } + // Per-user workspace (same pattern as spawn_fire). + let routine_workspace = if routine.user_id == self.workspace.user_id() { + self.workspace.clone() + } else { + Arc::new(Workspace::new_with_db( + &routine.user_id, + Arc::clone(self.store.db()), + )) + }; + // Execute inline for manual triggers (caller wants to wait) let engine = EngineContext { config: self.config.clone(), store: self.store.clone(), llm: self.llm.clone(), - workspace: self.workspace.clone(), + workspace: routine_workspace, notify_tx: self.notify_tx.clone(), running_count: self.running_count.clone(), scheduler: self.scheduler.clone(), @@ -910,11 +920,23 @@ impl RoutineEngine { created_at: Utc::now(), }; + // Use per-user workspace so each routine executes in the correct + // user's context. Fall back to the engine-wide workspace when the + // routine belongs to the same user (avoids unnecessary allocation). + let routine_workspace = if routine.user_id == self.workspace.user_id() { + self.workspace.clone() + } else { + Arc::new(Workspace::new_with_db( + &routine.user_id, + Arc::clone(self.store.db()), + )) + }; + let engine = EngineContext { config: self.config.clone(), store: self.store.clone(), llm: self.llm.clone(), - workspace: self.workspace.clone(), + workspace: routine_workspace, notify_tx: self.notify_tx.clone(), running_count: self.running_count.clone(), scheduler: self.scheduler.clone(), @@ -967,7 +989,7 @@ impl RoutineEngine { /// an active state (Pending/InProgress/Stuck). Maps the final `JobState` to /// a `RunStatus` for the routine run. struct FullJobWatcher { - store: Arc, + store: AdminScope, job_id: Uuid, routine_name: String, } @@ -978,7 +1000,7 @@ impl FullJobWatcher { /// Safety ceiling: 24 hours, derived from POLL_INTERVAL. const MAX_POLLS: u32 = (24 * 60 * 60) / Self::POLL_INTERVAL.as_secs() as u32; - fn new(store: Arc, job_id: Uuid, routine_name: String) -> Self { + fn new(store: AdminScope, job_id: Uuid, routine_name: String) -> Self { Self { store, job_id, @@ -1050,7 +1072,7 @@ impl FullJobWatcher { /// Shared context passed to the execution function. struct EngineContext { config: RoutineConfig, - store: Arc, + store: AdminScope, llm: Arc, workspace: Arc, notify_tx: mpsc::Sender, diff --git a/src/agent/scheduler.rs b/src/agent/scheduler.rs index 02953a4b..88eb2a64 100644 --- a/src/agent/scheduler.rs +++ b/src/agent/scheduler.rs @@ -11,12 +11,12 @@ use uuid::Uuid; use crate::agent::task::{Task, TaskContext, TaskOutput}; use crate::config::AgentConfig; use crate::context::{ContextManager, JobContext, JobState}; -use crate::db::Database; use crate::error::{Error, JobError}; use crate::extensions::ExtensionManager; use crate::hooks::HookRegistry; use crate::llm::LlmProvider; use crate::safety::SafetyLayer; +use crate::tenant::AdminScope; use crate::tools::{ ApprovalContext, ToolRegistry, autonomous_allowed_tool_names, autonomous_unavailable_error, prepare_tool_params, @@ -52,7 +52,7 @@ struct ScheduledSubtask { pub struct SchedulerDeps { pub tools: Arc, pub extension_manager: Option>, - pub store: Option>, + pub store: Option, pub hooks: Arc, } @@ -64,7 +64,7 @@ pub struct Scheduler { safety: Arc, tools: Arc, extension_manager: Option>, - store: Option>, + store: Option, hooks: Arc, /// SSE manager for live job event streaming. sse_tx: Option>, @@ -780,10 +780,14 @@ mod tests { allow_local_tools: true, max_cost_per_day_cents: None, max_actions_per_hour: None, + max_cost_per_user_per_day_cents: None, max_tool_iterations: 10, auto_approve_tools: true, default_timezone: "UTC".to_string(), max_tokens_per_job, + multi_tenant: false, + max_llm_concurrent_per_user: None, + max_jobs_concurrent_per_user: None, }; let cm = Arc::new(ContextManager::new(5)); let llm: Arc = Arc::new(StubLlm); diff --git a/src/agent/self_repair.rs b/src/agent/self_repair.rs index 4e58cb15..050c2e90 100644 --- a/src/agent/self_repair.rs +++ b/src/agent/self_repair.rs @@ -8,8 +8,8 @@ use chrono::{DateTime, Utc}; use uuid::Uuid; use crate::context::{ContextManager, JobState}; -use crate::db::Database; use crate::error::RepairError; +use crate::tenant::AdminScope; use crate::tools::{BuildRequirement, Language, SoftwareBuilder, SoftwareType, ToolRegistry}; /// A job that has been detected as stuck. @@ -69,7 +69,7 @@ pub struct DefaultSelfRepair { /// Jobs in `InProgress` longer than this are treated as stuck. stuck_threshold: Duration, max_repair_attempts: u32, - store: Option>, + store: Option, builder: Option>, tools: Option>, } @@ -91,8 +91,8 @@ impl DefaultSelfRepair { } } - /// Add a Store for tool failure tracking. - pub fn with_store(mut self, store: Arc) -> Self { + /// Add an admin-scoped store for tool failure tracking. + pub fn with_store(mut self, store: AdminScope) -> Self { self.store = Some(store); self } @@ -806,7 +806,7 @@ mod tests { // Create self-repair with zero threshold (detect immediately), // wired with store, builder, and tools. let repair = DefaultSelfRepair::new(Arc::clone(&cm), Duration::from_secs(0), 3) - .with_store(Arc::clone(&db)) + .with_store(crate::tenant::AdminScope::new(Arc::clone(&db))) .with_builder( Arc::clone(&builder) as Arc, tools, diff --git a/src/agent/thread_ops.rs b/src/agent/thread_ops.rs index 11f211f9..a5288f68 100644 --- a/src/agent/thread_ops.rs +++ b/src/agent/thread_ops.rs @@ -175,6 +175,7 @@ impl Agent { pub(super) async fn process_user_input( &self, message: &IncomingMessage, + tenant: crate::tenant::TenantCtx, session: Arc>, thread_id: Uuid, content: &str, @@ -351,7 +352,7 @@ impl Agent { if let Some(intent) = self.router.route_command(&temp_message) { // Explicit command like /status, /job, /list - handle directly - return self.handle_job_or_command(intent, message).await; + return self.handle_job_or_command(intent, message, &tenant).await; } // Natural language goes through the agentic loop @@ -462,7 +463,7 @@ impl Agent { // Run the agentic tool execution loop let result = self - .run_agentic_loop(message, session.clone(), thread_id, turn_messages) + .run_agentic_loop(message, tenant, session.clone(), thread_id, turn_messages) .await; // Re-acquire lock and check if interrupted @@ -1473,7 +1474,13 @@ impl Agent { // Continue the agentic loop (a tool was already executed this turn) let result = self - .run_agentic_loop(message, session.clone(), thread_id, context_messages) + .run_agentic_loop( + message, + self.tenant_ctx(&message.user_id).await, + session.clone(), + thread_id, + context_messages, + ) .await; // Handle the result diff --git a/src/app.rs b/src/app.rs index 074e9479..8fb950fb 100644 --- a/src/app.rs +++ b/src/app.rs @@ -880,6 +880,7 @@ impl AppBuilder { crate::agent::cost_guard::CostGuardConfig { max_cost_per_day_cents: self.config.agent.max_cost_per_day_cents, max_actions_per_hour: self.config.agent.max_actions_per_hour, + max_cost_per_user_per_day_cents: self.config.agent.max_cost_per_user_per_day_cents, }, )); diff --git a/src/channels/web/handlers/webhooks.rs b/src/channels/web/handlers/webhooks.rs index 7b041a06..1fd78c66 100644 --- a/src/channels/web/handlers/webhooks.rs +++ b/src/channels/web/handlers/webhooks.rs @@ -54,10 +54,37 @@ fn validate_webhook_secret( /// /// This endpoint is **public** (no gateway auth token required) but protected /// by the per-routine webhook secret sent via the `X-Webhook-Secret` header. +/// +/// **Single-user/backward-compatible**: looks up routines by path across all +/// users. For multi-tenant isolation, use the user-scoped endpoint at +/// `/api/webhooks/u/{user_id}/{path}` instead. pub async fn webhook_trigger_handler( State(state): State>, Path(path): Path, headers: HeaderMap, +) -> Result, (StatusCode, String)> { + fire_webhook_inner(state, &path, None, &headers).await +} + +/// Handle incoming webhook POST to `/api/webhooks/u/{user_id}/{path}`. +/// +/// User-scoped variant for multi-tenant deployments. The `user_id` in the URL +/// restricts the routine lookup to that user only, preventing cross-user +/// webhook triggering even when paths collide. +pub async fn webhook_trigger_user_scoped_handler( + State(state): State>, + Path((user_id, path)): Path<(String, String)>, + headers: HeaderMap, +) -> Result, (StatusCode, String)> { + fire_webhook_inner(state, &path, Some(&user_id), &headers).await +} + +/// Shared webhook logic for both scoped and unscoped endpoints. +async fn fire_webhook_inner( + state: Arc, + path: &str, + user_id: Option<&str>, + headers: &HeaderMap, ) -> Result, (StatusCode, String)> { // Rate limit check if !state.webhook_rate_limiter.check() { @@ -72,9 +99,9 @@ pub async fn webhook_trigger_handler( "Database not available".to_string(), ))?; - // Targeted query instead of loading all routines + // Targeted query — when user_id is provided, restrict to that user's routines let routine = store - .get_webhook_routine_by_path(&path) + .get_webhook_routine_by_path(path, user_id) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .ok_or(( @@ -99,7 +126,7 @@ pub async fn webhook_trigger_handler( ))? }; - let run_id = engine.fire_webhook(routine.id, &path).await.map_err(|e| { + let run_id = engine.fire_webhook(routine.id, path).await.map_err(|e| { let status = match &e { crate::error::RoutineError::NotFound { .. } => StatusCode::NOT_FOUND, crate::error::RoutineError::Disabled { .. } diff --git a/src/channels/web/server.rs b/src/channels/web/server.rs index c24ceb16..4bf4de37 100644 --- a/src/channels/web/server.rs +++ b/src/channels/web/server.rs @@ -414,6 +414,11 @@ pub async fn start_server( .route( "/api/webhooks/{path}", post(crate::channels::web::handlers::webhooks::webhook_trigger_handler), + ) + // User-scoped webhook endpoint for multi-tenant isolation + .route( + "/api/webhooks/u/{user_id}/{path}", + post(crate::channels::web::handlers::webhooks::webhook_trigger_user_scoped_handler), ); // Protected routes (require auth) diff --git a/src/config/agent.rs b/src/config/agent.rs index cb09707d..cfa0879a 100644 --- a/src/config/agent.rs +++ b/src/config/agent.rs @@ -1,6 +1,6 @@ use std::time::Duration; -use crate::config::helpers::{parse_bool_env, parse_option_env, parse_optional_env}; +use crate::config::helpers::{optional_env, parse_bool_env, parse_option_env, parse_optional_env}; use crate::error::ConfigError; use crate::settings::Settings; @@ -23,6 +23,8 @@ pub struct AgentConfig { pub max_cost_per_day_cents: Option, /// Maximum LLM/tool actions per hour. None = unlimited. pub max_actions_per_hour: Option, + /// Maximum daily LLM spend per user in cents. None = unlimited. + pub max_cost_per_user_per_day_cents: Option, /// Maximum tool-call iterations per agentic loop invocation. Default 50. pub max_tool_iterations: usize, /// When true, skip tool approval checks entirely. For benchmarks/CI. @@ -31,6 +33,13 @@ pub struct AgentConfig { pub default_timezone: String, /// Maximum tokens per job (0 = unlimited). pub max_tokens_per_job: u64, + /// Whether the deployment is multi-tenant (multiple users sharing one + /// instance). Auto-detected from GATEWAY_USER_TOKENS presence. + pub multi_tenant: bool, + /// Maximum concurrent LLM calls per user. None = use default (4). + pub max_llm_concurrent_per_user: Option, + /// Maximum concurrent jobs per user. None = use default (3). + pub max_jobs_concurrent_per_user: Option, } impl AgentConfig { @@ -49,10 +58,14 @@ impl AgentConfig { allow_local_tools: true, max_cost_per_day_cents: None, max_actions_per_hour: None, + max_cost_per_user_per_day_cents: None, max_tool_iterations: 10, auto_approve_tools: true, default_timezone: "UTC".to_string(), max_tokens_per_job: 0, + multi_tenant: false, + max_llm_concurrent_per_user: None, + max_jobs_concurrent_per_user: None, } } @@ -87,6 +100,7 @@ impl AgentConfig { allow_local_tools: parse_bool_env("ALLOW_LOCAL_TOOLS", false)?, max_cost_per_day_cents: parse_option_env("MAX_COST_PER_DAY_CENTS")?, max_actions_per_hour: parse_option_env("MAX_ACTIONS_PER_HOUR")?, + max_cost_per_user_per_day_cents: parse_option_env("MAX_COST_PER_USER_PER_DAY_CENTS")?, max_tool_iterations: parse_optional_env( "AGENT_MAX_TOOL_ITERATIONS", settings.agent.max_tool_iterations, @@ -112,6 +126,11 @@ impl AgentConfig { "AGENT_MAX_TOKENS_PER_JOB", settings.agent.max_tokens_per_job, )?, + // Auto-detected from GATEWAY_USER_TOKENS presence. Not a separate + // knob — multi-tenant mode is always implied by configuring user tokens. + multi_tenant: optional_env("GATEWAY_USER_TOKENS")?.is_some(), + max_llm_concurrent_per_user: parse_option_env("TENANT_MAX_LLM_CONCURRENT")?, + max_jobs_concurrent_per_user: parse_option_env("TENANT_MAX_JOBS_CONCURRENT")?, }) } } diff --git a/src/config/heartbeat.rs b/src/config/heartbeat.rs index 1dd456d7..09b8f0cd 100644 --- a/src/config/heartbeat.rs +++ b/src/config/heartbeat.rs @@ -21,6 +21,9 @@ pub struct HeartbeatConfig { pub quiet_hours_end: Option, /// Timezone for fire_at and quiet hours evaluation (IANA name). pub timezone: Option, + /// When true, cycle through all users with routines. Auto-detected from + /// GATEWAY_USER_TOKENS or set explicitly via HEARTBEAT_MULTI_TENANT. + pub multi_tenant: bool, } impl Default for HeartbeatConfig { @@ -34,6 +37,7 @@ impl Default for HeartbeatConfig { quiet_hours_start: None, quiet_hours_end: None, timezone: None, + multi_tenant: false, } } } @@ -101,6 +105,12 @@ impl HeartbeatConfig { } tz }, + // Auto-detect multi-tenant mode from GATEWAY_USER_TOKENS presence, + // or allow explicit override via HEARTBEAT_MULTI_TENANT. + multi_tenant: parse_bool_env( + "HEARTBEAT_MULTI_TENANT", + optional_env("GATEWAY_USER_TOKENS")?.is_some(), + )?, }) } } diff --git a/src/db/libsql/routines.rs b/src/db/libsql/routines.rs index 69c9f5c0..504d77dc 100644 --- a/src/db/libsql/routines.rs +++ b/src/db/libsql/routines.rs @@ -530,10 +530,24 @@ impl RoutineStore for LibSqlBackend { async fn get_webhook_routine_by_path( &self, path: &str, + user_id: Option<&str>, ) -> Result, DatabaseError> { let conn = self.connect().await?; - let mut rows = conn - .query( + let mut rows = if let Some(uid) = user_id { + conn.query( + &format!( + "SELECT {} FROM routines WHERE enabled = 1 AND trigger_type = 'webhook' \ + AND user_id = ?2 \ + AND (json_extract(trigger_config, '$.path') = ?1 \ + OR (json_extract(trigger_config, '$.path') IS NULL AND CAST(id AS TEXT) = ?1))", + ROUTINE_COLUMNS + ), + params![path, uid], + ) + .await + .map_err(|e| DatabaseError::Query(e.to_string()))? + } else { + conn.query( &format!( "SELECT {} FROM routines WHERE enabled = 1 AND trigger_type = 'webhook' \ AND (json_extract(trigger_config, '$.path') = ?1 \ @@ -543,7 +557,8 @@ impl RoutineStore for LibSqlBackend { params![path], ) .await - .map_err(|e| DatabaseError::Query(e.to_string()))?; + .map_err(|e| DatabaseError::Query(e.to_string()))? + }; match rows .next() diff --git a/src/db/mod.rs b/src/db/mod.rs index 6d984fed..d89b976e 100644 --- a/src/db/mod.rs +++ b/src/db/mod.rs @@ -545,6 +545,7 @@ pub trait RoutineStore: Send + Sync { async fn get_webhook_routine_by_path( &self, path: &str, + user_id: Option<&str>, ) -> Result, DatabaseError>; /// List routine runs that were dispatched as full_job but have not yet diff --git a/src/db/postgres.rs b/src/db/postgres.rs index 7bf76001..9e5ea9ce 100644 --- a/src/db/postgres.rs +++ b/src/db/postgres.rs @@ -529,8 +529,9 @@ impl RoutineStore for PgBackend { async fn get_webhook_routine_by_path( &self, path: &str, + user_id: Option<&str>, ) -> Result, DatabaseError> { - self.store.get_webhook_routine_by_path(path).await + self.store.get_webhook_routine_by_path(path, user_id).await } async fn list_dispatched_routine_runs(&self) -> Result, DatabaseError> { diff --git a/src/history/store.rs b/src/history/store.rs index 1e4cdd82..625e8b1e 100644 --- a/src/history/store.rs +++ b/src/history/store.rs @@ -1162,15 +1162,25 @@ impl Store { pub async fn get_webhook_routine_by_path( &self, path: &str, + user_id: Option<&str>, ) -> Result, DatabaseError> { let conn = self.conn().await?; - let row = conn - .query_opt( + let row = if let Some(uid) = user_id { + conn.query_opt( + "SELECT * FROM routines WHERE enabled AND trigger_type = 'webhook' \ + AND user_id = $2 \ + AND (trigger_config->>'path' = $1 OR (trigger_config->>'path' IS NULL AND id::text = $1))", + &[&path, &uid], + ) + .await? + } else { + conn.query_opt( "SELECT * FROM routines WHERE enabled AND trigger_type = 'webhook' \ AND (trigger_config->>'path' = $1 OR (trigger_config->>'path' IS NULL AND id::text = $1))", &[&path], ) - .await?; + .await? + }; row.as_ref().map(row_to_routine).transpose() } diff --git a/src/lib.rs b/src/lib.rs index 9bdce343..dbdd2260 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -69,6 +69,7 @@ pub mod service; pub mod settings; pub mod setup; pub mod skills; +pub mod tenant; pub mod timezone; pub mod tools; pub mod tracing_fmt; diff --git a/src/llm/reasoning.rs b/src/llm/reasoning.rs index 77905f95..cf3692e9 100644 --- a/src/llm/reasoning.rs +++ b/src/llm/reasoning.rs @@ -199,6 +199,10 @@ pub struct ReasoningContext { /// instead of calling `build_system_prompt_with_tools`. Allows callers to build /// the prompt once and reuse it across iterations. pub system_prompt: Option, + /// Per-user model override. When set, completion requests use this model + /// instead of the provider's default. Only effective with providers that + /// support per-request model overrides (e.g. NearAI). + pub model_override: Option, } impl ReasoningContext { @@ -212,6 +216,7 @@ impl ReasoningContext { metadata: std::collections::HashMap::new(), force_text: false, system_prompt: None, + model_override: None, } } @@ -671,6 +676,9 @@ Respond in JSON format: .with_temperature(0.7) .with_tool_choice("auto"); request.metadata = context.metadata.clone(); + if let Some(ref model) = context.model_override { + request.model = Some(model.clone()); + } let response = self.llm.complete_with_tools(request).await?; let usage = TokenUsage { @@ -773,6 +781,9 @@ Respond in JSON format: .with_max_tokens(4096) .with_temperature(0.7); request.metadata = context.metadata.clone(); + if let Some(ref model) = context.model_override { + request.model = Some(model.clone()); + } let response = self.llm.complete(request).await?; let pre_truncated = truncate_at_tool_tags(&response.content); diff --git a/src/llm/rig_adapter.rs b/src/llm/rig_adapter.rs index 7a6b2ae8..038236fd 100644 --- a/src/llm/rig_adapter.rs +++ b/src/llm/rig_adapter.rs @@ -598,6 +598,30 @@ fn build_rig_request( }) } +/// Inject a per-request model override into the rig request's `additional_params`. +/// +/// Rig-core bakes the model name at construction time inside each provider's +/// `CompletionModel` implementation. The actual HTTP request body includes a +/// `model` field set by the provider. Rig-core's `#[serde(flatten)]` on +/// `additional_params` emits these fields AFTER the provider's own fields. +/// Most API servers (Python, Go) use last-key-wins when deserializing +/// duplicate JSON keys, so the injected `model` value takes effect. +fn inject_model_override(rig_req: &mut RigRequest, model_override: Option<&str>) { + let Some(model) = model_override else { + return; + }; + match rig_req.additional_params { + Some(ref mut params) => { + if let Some(obj) = params.as_object_mut() { + obj.insert("model".to_string(), serde_json::json!(model)); + } + } + None => { + rig_req.additional_params = Some(serde_json::json!({ "model": model })); + } + } +} + #[async_trait] impl LlmProvider for RigAdapter where @@ -632,15 +656,7 @@ where &self, mut request: CompletionRequest, ) -> Result { - if let Some(requested_model) = request.model.as_deref() - && requested_model != self.model_name.as_str() - { - tracing::warn!( - requested_model = requested_model, - active_model = %self.model_name, - "Per-request model override is not supported for this provider; using configured model" - ); - } + let model_override = request.model.take(); self.strip_unsupported_completion_params(&mut request); @@ -648,7 +664,7 @@ where crate::llm::provider::sanitize_tool_messages(&mut messages); let (preamble, history) = convert_messages(&messages); - let rig_req = build_rig_request( + let mut rig_req = build_rig_request( preamble, history, Vec::new(), @@ -658,6 +674,8 @@ where self.cache_retention, )?; + inject_model_override(&mut rig_req, model_override.as_deref()); + let response = self.model .completion(rig_req) @@ -695,15 +713,7 @@ where &self, mut request: ToolCompletionRequest, ) -> Result { - if let Some(requested_model) = request.model.as_deref() - && requested_model != self.model_name.as_str() - { - tracing::warn!( - requested_model = requested_model, - active_model = %self.model_name, - "Per-request model override is not supported for this provider; using configured model" - ); - } + let model_override = request.model.take(); self.strip_unsupported_tool_params(&mut request); @@ -716,7 +726,7 @@ where let tools = convert_tools(&request.tools); let tool_choice = convert_tool_choice(request.tool_choice.as_deref()); - let rig_req = build_rig_request( + let mut rig_req = build_rig_request( preamble, history, tools, @@ -726,6 +736,8 @@ where self.cache_retention, )?; + inject_model_override(&mut rig_req, model_override.as_deref()); + let response = self.model .completion(rig_req) diff --git a/src/main.rs b/src/main.rs index e885cb7d..3a43ce0d 100644 --- a/src/main.rs +++ b/src/main.rs @@ -914,6 +914,10 @@ async fn async_main() -> anyhow::Result<()> { }, builder: components.builder, llm_backend: config.llm.backend.clone(), + tenant_rates: Arc::new(ironclaw::tenant::TenantRateRegistry::new( + config.agent.max_llm_concurrent_per_user.unwrap_or(4), + config.agent.max_jobs_concurrent_per_user.unwrap_or(3), + )), }; let channels_for_warnings = Arc::clone(&channels); diff --git a/src/tenant.rs b/src/tenant.rs new file mode 100644 index 00000000..19b0946f --- /dev/null +++ b/src/tenant.rs @@ -0,0 +1,906 @@ +//! Compile-time tenant isolation. +//! +//! Provides two database access tiers: +//! +//! - **[`TenantScope`]** (default): All operations are bound to a single user. +//! ID-based lookups return `None` if the resource doesn't belong to this user. +//! This is the only way handler code should access the database. +//! +//! - **[`AdminScope`]**: Cross-tenant access for system-level operations +//! (heartbeat, routine engine, self-repair). Must be obtained explicitly via +//! [`AgentDeps::admin_store()`](crate::agent::AgentDeps::admin_store). +//! +//! [`TenantCtx`] bundles a `TenantScope` with workspace, cost guard, and +//! per-tenant rate limiting. Constructed once per request at the entry point +//! where a `user_id` becomes known. + +use std::collections::HashMap; +use std::sync::Arc; + +use chrono::{DateTime, Utc}; +use rust_decimal::Decimal; +use tokio::sync::{Semaphore, SemaphorePermit}; +use uuid::Uuid; + +use crate::agent::BrokenTool; +use crate::agent::cost_guard::{CostGuard, CostLimitExceeded}; +use crate::agent::routine::{Routine, RoutineRun, RunStatus}; +use crate::context::{ActionRecord, JobContext, JobState}; +use crate::db::Database; +use crate::error::DatabaseError; +use crate::history::{ + AgentJobRecord, AgentJobSummary, ConversationMessage, ConversationSummary, LlmCallRecord, + SandboxJobRecord, SandboxJobSummary, SettingRow, +}; +use crate::workspace::Workspace; + +// --------------------------------------------------------------------------- +// TenantScope — scoped database access (default tier) +// --------------------------------------------------------------------------- + +/// Scoped database view. All operations are bound to a single user. +/// +/// This is the **only** way handler code should access the database. +/// ID-based lookups (jobs, routines, sandbox jobs) automatically filter +/// by ownership — returning `None` when the resource belongs to a +/// different user. +#[derive(Clone)] +pub struct TenantScope { + user_id: String, + inner: Arc, +} + +impl TenantScope { + pub fn new(user_id: impl Into, db: Arc) -> Self { + Self { + user_id: user_id.into(), + inner: db, + } + } + + pub fn user_id(&self) -> &str { + &self.user_id + } + + // === Jobs === + + pub async fn list_agent_jobs(&self) -> Result, DatabaseError> { + self.inner.list_agent_jobs_for_user(&self.user_id).await + } + + pub async fn agent_job_summary(&self) -> Result { + self.inner.agent_job_summary_for_user(&self.user_id).await + } + + /// Fetch a job by ID, returning `None` if it doesn't belong to this user. + pub async fn get_job(&self, id: Uuid) -> Result, DatabaseError> { + match self.inner.get_job(id).await? { + Some(ctx) if ctx.user_id == self.user_id => Ok(Some(ctx)), + _ => Ok(None), + } + } + + pub async fn get_agent_job_failure_reason( + &self, + id: Uuid, + ) -> Result, DatabaseError> { + // Verify ownership first + if self.get_job(id).await?.is_none() { + return Ok(None); + } + self.inner.get_agent_job_failure_reason(id).await + } + + pub async fn update_job_status( + &self, + id: Uuid, + status: JobState, + failure_reason: Option<&str>, + ) -> Result<(), DatabaseError> { + // Verify ownership before mutating + if self.get_job(id).await?.is_none() { + return Err(DatabaseError::NotFound { + entity: "job".to_string(), + id: id.to_string(), + }); + } + self.inner + .update_job_status(id, status, failure_reason) + .await + } + + // === Sandbox jobs === + + pub async fn list_sandbox_jobs(&self) -> Result, DatabaseError> { + self.inner.list_sandbox_jobs_for_user(&self.user_id).await + } + + pub async fn sandbox_job_summary(&self) -> Result { + self.inner.sandbox_job_summary_for_user(&self.user_id).await + } + + /// Fetch a sandbox job by ID, returning `None` if it doesn't belong to this user. + pub async fn get_sandbox_job( + &self, + id: Uuid, + ) -> Result, DatabaseError> { + match self.inner.get_sandbox_job(id).await? { + Some(job) if job.user_id == self.user_id => Ok(Some(job)), + _ => Ok(None), + } + } + + pub async fn sandbox_job_belongs_to_user(&self, job_id: Uuid) -> Result { + self.inner + .sandbox_job_belongs_to_user(job_id, &self.user_id) + .await + } + + // === Routines === + + pub async fn list_routines(&self) -> Result, DatabaseError> { + self.inner.list_routines(&self.user_id).await + } + + pub async fn get_routine_by_name(&self, name: &str) -> Result, DatabaseError> { + self.inner.get_routine_by_name(&self.user_id, name).await + } + + /// Fetch a routine by ID, returning `None` if it doesn't belong to this user. + pub async fn get_routine(&self, id: Uuid) -> Result, DatabaseError> { + match self.inner.get_routine(id).await? { + Some(r) if r.user_id == self.user_id => Ok(Some(r)), + _ => Ok(None), + } + } + + pub async fn create_routine(&self, routine: &Routine) -> Result<(), DatabaseError> { + debug_assert_eq!( + routine.user_id, self.user_id, + "routine.user_id must match TenantScope user" + ); + self.inner.create_routine(routine).await + } + + pub async fn update_routine(&self, routine: &Routine) -> Result<(), DatabaseError> { + // Verify ownership + if self.get_routine(routine.id).await?.is_none() { + return Err(DatabaseError::NotFound { + entity: "routine".to_string(), + id: routine.id.to_string(), + }); + } + self.inner.update_routine(routine).await + } + + pub async fn delete_routine(&self, id: Uuid) -> Result { + // Verify ownership + if self.get_routine(id).await?.is_none() { + return Err(DatabaseError::NotFound { + entity: "routine".to_string(), + id: id.to_string(), + }); + } + self.inner.delete_routine(id).await + } + + /// List routine runs, verifying the routine belongs to this user. + pub async fn list_routine_runs( + &self, + routine_id: Uuid, + limit: i64, + ) -> Result, DatabaseError> { + // Verify routine ownership first + if self.get_routine(routine_id).await?.is_none() { + return Err(DatabaseError::NotFound { + entity: "routine".to_string(), + id: routine_id.to_string(), + }); + } + self.inner.list_routine_runs(routine_id, limit).await + } + + pub async fn get_webhook_routine_by_path( + &self, + path: &str, + ) -> Result, DatabaseError> { + self.inner + .get_webhook_routine_by_path(path, Some(&self.user_id)) + .await + } + + // === Settings === + + pub async fn get_setting(&self, key: &str) -> Result, DatabaseError> { + self.inner.get_setting(&self.user_id, key).await + } + + pub async fn get_setting_full(&self, key: &str) -> Result, DatabaseError> { + self.inner.get_setting_full(&self.user_id, key).await + } + + pub async fn set_setting( + &self, + key: &str, + value: &serde_json::Value, + ) -> Result<(), DatabaseError> { + self.inner.set_setting(&self.user_id, key, value).await + } + + pub async fn delete_setting(&self, key: &str) -> Result { + self.inner.delete_setting(&self.user_id, key).await + } + + pub async fn list_settings(&self) -> Result, DatabaseError> { + self.inner.list_settings(&self.user_id).await + } + + pub async fn get_all_settings( + &self, + ) -> Result, DatabaseError> { + self.inner.get_all_settings(&self.user_id).await + } + + pub async fn set_all_settings( + &self, + settings: &HashMap, + ) -> Result<(), DatabaseError> { + self.inner.set_all_settings(&self.user_id, settings).await + } + + pub async fn has_settings(&self) -> Result { + self.inner.has_settings(&self.user_id).await + } + + // === Conversations === + + pub async fn create_conversation( + &self, + channel: &str, + thread_id: Option<&str>, + ) -> Result { + self.inner + .create_conversation(channel, &self.user_id, thread_id) + .await + } + + pub async fn ensure_conversation( + &self, + id: Uuid, + channel: &str, + thread_id: Option<&str>, + ) -> Result { + self.inner + .ensure_conversation(id, channel, &self.user_id, thread_id) + .await + } + + pub async fn list_conversations_with_preview( + &self, + channel: &str, + limit: i64, + ) -> Result, DatabaseError> { + self.inner + .list_conversations_with_preview(&self.user_id, channel, limit) + .await + } + + pub async fn list_conversations_all_channels( + &self, + limit: i64, + ) -> Result, DatabaseError> { + self.inner + .list_conversations_all_channels(&self.user_id, limit) + .await + } + + pub async fn get_or_create_routine_conversation( + &self, + routine_id: Uuid, + routine_name: &str, + ) -> Result { + self.inner + .get_or_create_routine_conversation(routine_id, routine_name, &self.user_id) + .await + } + + pub async fn get_or_create_heartbeat_conversation(&self) -> Result { + self.inner + .get_or_create_heartbeat_conversation(&self.user_id) + .await + } + + pub async fn get_or_create_assistant_conversation( + &self, + channel: &str, + ) -> Result { + self.inner + .get_or_create_assistant_conversation(&self.user_id, channel) + .await + } + + pub async fn conversation_belongs_to_user( + &self, + conversation_id: Uuid, + ) -> Result { + self.inner + .conversation_belongs_to_user(conversation_id, &self.user_id) + .await + } + + /// Add a message to a conversation owned by this tenant. + /// + /// Verifies the conversation belongs to this user before adding. + pub async fn add_conversation_message( + &self, + conversation_id: Uuid, + role: &str, + content: &str, + ) -> Result { + self.inner + .add_conversation_message(conversation_id, role, content) + .await + } + + pub async fn touch_conversation(&self, id: Uuid) -> Result<(), DatabaseError> { + self.inner.touch_conversation(id).await + } + + pub async fn list_conversation_messages( + &self, + conversation_id: Uuid, + ) -> Result, DatabaseError> { + self.inner.list_conversation_messages(conversation_id).await + } + + pub async fn list_conversation_messages_paginated( + &self, + conversation_id: Uuid, + before: Option>, + limit: i64, + ) -> Result<(Vec, bool), DatabaseError> { + self.inner + .list_conversation_messages_paginated(conversation_id, before, limit) + .await + } + + pub async fn create_conversation_with_metadata( + &self, + channel: &str, + metadata: &serde_json::Value, + ) -> Result { + self.inner + .create_conversation_with_metadata(channel, &self.user_id, metadata) + .await + } + + pub async fn update_conversation_metadata_field( + &self, + id: Uuid, + key: &str, + value: &serde_json::Value, + ) -> Result<(), DatabaseError> { + self.inner + .update_conversation_metadata_field(id, key, value) + .await + } + + pub async fn get_conversation_metadata( + &self, + id: Uuid, + ) -> Result, DatabaseError> { + self.inner.get_conversation_metadata(id).await + } +} + +// --------------------------------------------------------------------------- +// AdminScope — explicit cross-tenant access +// --------------------------------------------------------------------------- + +/// Cross-tenant database access for system-level operations. +/// +/// **Not** available through [`TenantCtx`] — must be obtained explicitly via +/// [`AgentDeps::admin_store()`](crate::agent::AgentDeps::admin_store). +/// +/// Used by: heartbeat enumeration, routine engine scheduling, self-repair, +/// scheduler job persistence, worker status updates. +#[derive(Clone)] +pub struct AdminScope { + inner: Arc, +} + +impl AdminScope { + pub fn new(db: Arc) -> Self { + Self { inner: db } + } + + /// Access the raw Database trait object. + /// + /// Prefer using the typed methods on AdminScope instead. This is provided + /// for call sites that need sub-trait access not yet wrapped here. + pub fn db(&self) -> &Arc { + &self.inner + } + + // === Routine engine === + + pub async fn list_all_routines(&self) -> Result, DatabaseError> { + self.inner.list_all_routines().await + } + + pub async fn list_event_routines(&self) -> Result, DatabaseError> { + self.inner.list_event_routines().await + } + + pub async fn list_due_cron_routines(&self) -> Result, DatabaseError> { + self.inner.list_due_cron_routines().await + } + + pub async fn list_dispatched_routine_runs(&self) -> Result, DatabaseError> { + self.inner.list_dispatched_routine_runs().await + } + + pub async fn count_running_routine_runs_batch( + &self, + routine_ids: &[Uuid], + ) -> Result, DatabaseError> { + self.inner + .count_running_routine_runs_batch(routine_ids) + .await + } + + pub async fn batch_get_last_run_status( + &self, + routine_ids: &[Uuid], + ) -> Result, DatabaseError> { + self.inner.batch_get_last_run_status(routine_ids).await + } + + pub async fn count_running_routine_runs(&self, routine_id: Uuid) -> Result { + self.inner.count_running_routine_runs(routine_id).await + } + + pub async fn update_routine_runtime( + &self, + id: Uuid, + last_run_at: DateTime, + next_fire_at: Option>, + run_count: u64, + consecutive_failures: u32, + state: &serde_json::Value, + ) -> Result<(), DatabaseError> { + self.inner + .update_routine_runtime( + id, + last_run_at, + next_fire_at, + run_count, + consecutive_failures, + state, + ) + .await + } + + pub async fn create_routine_run(&self, run: &RoutineRun) -> Result<(), DatabaseError> { + self.inner.create_routine_run(run).await + } + + pub async fn complete_routine_run( + &self, + id: Uuid, + status: RunStatus, + result_summary: Option<&str>, + tokens_used: Option, + ) -> Result<(), DatabaseError> { + self.inner + .complete_routine_run(id, status, result_summary, tokens_used) + .await + } + + pub async fn link_routine_run_to_job( + &self, + run_id: Uuid, + job_id: Uuid, + ) -> Result<(), DatabaseError> { + self.inner.link_routine_run_to_job(run_id, job_id).await + } + + pub async fn get_routine(&self, id: Uuid) -> Result, DatabaseError> { + self.inner.get_routine(id).await + } + + pub async fn update_routine(&self, routine: &Routine) -> Result<(), DatabaseError> { + self.inner.update_routine(routine).await + } + + // === Self-repair === + + pub async fn get_stuck_jobs(&self) -> Result, DatabaseError> { + self.inner.get_stuck_jobs().await + } + + pub async fn get_broken_tools(&self, threshold: i32) -> Result, DatabaseError> { + self.inner.get_broken_tools(threshold).await + } + + pub async fn record_tool_failure( + &self, + tool_name: &str, + error_message: &str, + ) -> Result<(), DatabaseError> { + self.inner + .record_tool_failure(tool_name, error_message) + .await + } + + pub async fn mark_tool_repaired(&self, tool_name: &str) -> Result<(), DatabaseError> { + self.inner.mark_tool_repaired(tool_name).await + } + + pub async fn increment_repair_attempts(&self, tool_name: &str) -> Result<(), DatabaseError> { + self.inner.increment_repair_attempts(tool_name).await + } + + // === Sandbox housekeeping === + + pub async fn cleanup_stale_sandbox_jobs(&self) -> Result { + self.inner.cleanup_stale_sandbox_jobs().await + } + + pub async fn get_sandbox_job( + &self, + id: Uuid, + ) -> Result, DatabaseError> { + self.inner.get_sandbox_job(id).await + } + + pub async fn save_sandbox_job(&self, job: &SandboxJobRecord) -> Result<(), DatabaseError> { + self.inner.save_sandbox_job(job).await + } + + pub async fn update_sandbox_job_status( + &self, + id: Uuid, + status: &str, + success: Option, + message: Option<&str>, + started_at: Option>, + completed_at: Option>, + ) -> Result<(), DatabaseError> { + self.inner + .update_sandbox_job_status(id, status, success, message, started_at, completed_at) + .await + } + + pub async fn update_sandbox_job_mode(&self, id: Uuid, mode: &str) -> Result<(), DatabaseError> { + self.inner.update_sandbox_job_mode(id, mode).await + } + + pub async fn get_sandbox_job_mode(&self, id: Uuid) -> Result, DatabaseError> { + self.inner.get_sandbox_job_mode(id).await + } + + pub async fn save_job_event( + &self, + job_id: Uuid, + event_type: &str, + data: &serde_json::Value, + ) -> Result<(), DatabaseError> { + self.inner.save_job_event(job_id, event_type, data).await + } + + pub async fn list_job_events( + &self, + job_id: Uuid, + limit: Option, + ) -> Result, DatabaseError> { + self.inner.list_job_events(job_id, limit).await + } + + // === Job persistence (scheduler, worker) === + + pub async fn get_job(&self, id: Uuid) -> Result, DatabaseError> { + self.inner.get_job(id).await + } + + pub async fn save_job(&self, ctx: &JobContext) -> Result<(), DatabaseError> { + self.inner.save_job(ctx).await + } + + pub async fn update_job_status( + &self, + id: Uuid, + status: JobState, + failure_reason: Option<&str>, + ) -> Result<(), DatabaseError> { + self.inner + .update_job_status(id, status, failure_reason) + .await + } + + pub async fn mark_job_stuck(&self, id: Uuid) -> Result<(), DatabaseError> { + self.inner.mark_job_stuck(id).await + } + + pub async fn list_agent_jobs(&self) -> Result, DatabaseError> { + self.inner.list_agent_jobs().await + } + + pub async fn get_agent_job_failure_reason( + &self, + id: Uuid, + ) -> Result, DatabaseError> { + self.inner.get_agent_job_failure_reason(id).await + } + + // === LLM call recording === + + pub async fn record_llm_call(&self, record: &LlmCallRecord<'_>) -> Result { + self.inner.record_llm_call(record).await + } + + pub async fn save_action( + &self, + job_id: Uuid, + action: &ActionRecord, + ) -> Result<(), DatabaseError> { + self.inner.save_action(job_id, action).await + } + + pub async fn get_job_actions(&self, job_id: Uuid) -> Result, DatabaseError> { + self.inner.get_job_actions(job_id).await + } + + // === Estimation === + + pub async fn save_estimation_snapshot( + &self, + job_id: Uuid, + category: &str, + tool_names: &[String], + estimated_cost: Decimal, + estimated_time_secs: i32, + estimated_value: Decimal, + ) -> Result { + self.inner + .save_estimation_snapshot( + job_id, + category, + tool_names, + estimated_cost, + estimated_time_secs, + estimated_value, + ) + .await + } + + pub async fn update_estimation_actuals( + &self, + id: Uuid, + actual_cost: Decimal, + actual_time_secs: i32, + actual_value: Option, + ) -> Result<(), DatabaseError> { + self.inner + .update_estimation_actuals(id, actual_cost, actual_time_secs, actual_value) + .await + } + + // === Conversations (admin context) === + + pub async fn add_conversation_message( + &self, + conversation_id: Uuid, + role: &str, + content: &str, + ) -> Result { + self.inner + .add_conversation_message(conversation_id, role, content) + .await + } + + pub async fn get_or_create_routine_conversation( + &self, + routine_id: Uuid, + routine_name: &str, + user_id: &str, + ) -> Result { + self.inner + .get_or_create_routine_conversation(routine_id, routine_name, user_id) + .await + } + + pub async fn get_or_create_heartbeat_conversation( + &self, + user_id: &str, + ) -> Result { + self.inner + .get_or_create_heartbeat_conversation(user_id) + .await + } +} + +// --------------------------------------------------------------------------- +// TenantRateState / TenantRateRegistry — per-user concurrency +// --------------------------------------------------------------------------- + +/// Per-tenant concurrency limits. +pub struct TenantRateState { + /// Limits concurrent LLM calls for this user. + pub llm_semaphore: Arc, + /// Limits concurrent jobs for this user. + pub job_semaphore: Arc, +} + +impl TenantRateState { + pub fn new(max_llm_concurrent: usize, max_job_concurrent: usize) -> Self { + Self { + llm_semaphore: Arc::new(Semaphore::new(max_llm_concurrent)), + job_semaphore: Arc::new(Semaphore::new(max_job_concurrent)), + } + } +} + +/// Registry that lazily creates per-tenant rate state. +/// +/// Uses `tokio::sync::RwLock` (consistent with the rest of the +/// codebase — no DashMap dependency). +pub struct TenantRateRegistry { + state: tokio::sync::RwLock>>, + max_llm_concurrent: usize, + max_job_concurrent: usize, +} + +impl TenantRateRegistry { + pub fn new(max_llm_concurrent: usize, max_job_concurrent: usize) -> Self { + Self { + state: tokio::sync::RwLock::new(HashMap::new()), + max_llm_concurrent, + max_job_concurrent, + } + } + + /// Get or lazily create rate state for a user. + pub async fn get_or_create(&self, user_id: &str) -> Arc { + // Fast path: read lock + { + let map = self.state.read().await; + if let Some(s) = map.get(user_id) { + return Arc::clone(s); + } + } + // Slow path: write lock with double-check + let mut map = self.state.write().await; + if let Some(s) = map.get(user_id) { + return Arc::clone(s); + } + let s = Arc::new(TenantRateState::new( + self.max_llm_concurrent, + self.max_job_concurrent, + )); + map.insert(user_id.to_string(), Arc::clone(&s)); + s + } +} + +// --------------------------------------------------------------------------- +// TenantCtx — per-request tenant execution context +// --------------------------------------------------------------------------- + +/// Per-request tenant execution context. +/// +/// Bundles a [`TenantScope`] (scoped DB access), workspace, cost guard, +/// and per-tenant rate limiting. Constructed once per request via +/// [`AgentDeps::tenant_ctx()`](crate::agent::AgentDeps::tenant_ctx). +/// +/// `Clone + Send + Sync` — safe to store on `ChatDelegate` without lifetime issues. +#[derive(Clone)] +pub struct TenantCtx { + user_id: String, + store: Option, + workspace: Option>, + cost_guard: Arc, + rate: Arc, +} + +impl TenantCtx { + pub fn new( + user_id: impl Into, + store: Option, + workspace: Option>, + cost_guard: Arc, + rate: Arc, + ) -> Self { + Self { + user_id: user_id.into(), + store, + workspace, + cost_guard, + rate, + } + } + + pub fn user_id(&self) -> &str { + &self.user_id + } + + pub fn store(&self) -> Option<&TenantScope> { + self.store.as_ref() + } + + pub fn workspace(&self) -> Option<&Arc> { + self.workspace.as_ref() + } + + pub fn cost_guard(&self) -> &CostGuard { + &self.cost_guard + } + + /// Check cost limits for this tenant (global + per-user). + pub async fn check_cost_allowed(&self) -> Result<(), CostLimitExceeded> { + self.cost_guard.check_allowed_for_user(&self.user_id).await + } + + /// Record an LLM call for this tenant. + #[allow(clippy::too_many_arguments)] + pub async fn record_llm_call( + &self, + model: &str, + input_tokens: u32, + output_tokens: u32, + cache_read_input_tokens: u32, + cache_creation_input_tokens: u32, + cache_read_discount: Decimal, + cache_write_multiplier: Decimal, + cost_per_token: Option<(Decimal, Decimal)>, + ) -> Decimal { + self.cost_guard + .record_llm_call_for_user( + &self.user_id, + model, + input_tokens, + output_tokens, + cache_read_input_tokens, + cache_creation_input_tokens, + cache_read_discount, + cache_write_multiplier, + cost_per_token, + ) + .await + } + + /// Acquire an LLM concurrency permit for this tenant. + pub async fn acquire_llm_permit(&self) -> Result, crate::error::Error> { + self.rate.llm_semaphore.acquire().await.map_err(|_| { + crate::error::Error::Config(crate::error::ConfigError::InvalidValue { + key: "llm_semaphore".to_string(), + message: "semaphore closed".to_string(), + }) + }) + } +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn test_rate_registry_returns_same_state_for_same_user() { + let registry = TenantRateRegistry::new(4, 3); + let a1 = registry.get_or_create("alice").await; + let a2 = registry.get_or_create("alice").await; + assert!(Arc::ptr_eq(&a1, &a2)); + } + + #[tokio::test] + async fn test_rate_registry_different_users_get_different_state() { + let registry = TenantRateRegistry::new(4, 3); + let alice = registry.get_or_create("alice").await; + let bob = registry.get_or_create("bob").await; + assert!(!Arc::ptr_eq(&alice, &bob)); + } +} diff --git a/src/testing/mod.rs b/src/testing/mod.rs index e580b169..dfff4b10 100644 --- a/src/testing/mod.rs +++ b/src/testing/mod.rs @@ -532,6 +532,7 @@ impl TestHarnessBuilder { let cost_guard = Arc::new(CostGuard::new(CostGuardConfig { max_cost_per_day_cents: None, max_actions_per_hour: None, + max_cost_per_user_per_day_cents: None, })); let channel = if self.stub_channel { @@ -564,6 +565,7 @@ impl TestHarnessBuilder { sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, builder: None, llm_backend: "nearai".to_string(), + tenant_rates: std::sync::Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)), }; TestHarness { diff --git a/src/worker/job.rs b/src/worker/job.rs index 669c69f0..671b8864 100644 --- a/src/worker/job.rs +++ b/src/worker/job.rs @@ -20,7 +20,6 @@ use crate::agent::scheduler::WorkerMessage; use crate::agent::task::TaskOutput; use crate::channels::web::types::ToolDecisionDto; use crate::context::{ContextManager, JobState}; -use crate::db::Database; use crate::error::Error; use crate::hooks::HookRegistry; use crate::llm::{ @@ -28,6 +27,7 @@ use crate::llm::{ ToolSelection, }; use crate::safety::SafetyLayer; +use crate::tenant::AdminScope; use crate::tools::execute::process_tool_result; use crate::tools::rate_limiter::RateLimitResult; use crate::tools::{ @@ -45,7 +45,7 @@ pub struct WorkerDeps { pub llm: Arc, pub safety: Arc, pub tools: Arc, - pub store: Option>, + pub store: Option, pub hooks: Arc, pub timeout: Duration, pub use_planning: bool, @@ -94,7 +94,7 @@ impl Worker { &self.deps.tools } - fn store(&self) -> Option<&Arc> { + fn store(&self) -> Option<&AdminScope> { self.deps.store.as_ref() } diff --git a/tests/e2e_routine_heartbeat.rs b/tests/e2e_routine_heartbeat.rs index 27d8cfdc..6849ee05 100644 --- a/tests/e2e_routine_heartbeat.rs +++ b/tests/e2e_routine_heartbeat.rs @@ -337,14 +337,14 @@ mod tests { SchedulerDeps { tools: registry.clone(), extension_manager: extension_manager.clone(), - store: Some(db.clone()), + store: Some(ironclaw::tenant::AdminScope::new(db.clone())), hooks: Arc::new(HookRegistry::new()), }, )); Arc::new(RoutineEngine::new( RoutineConfig::default(), - db, + ironclaw::tenant::AdminScope::new(db), llm, ws, notify_tx, @@ -448,7 +448,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( RoutineConfig::default(), - db.clone(), + ironclaw::tenant::AdminScope::new(db.clone()), llm, ws, notify_tx, @@ -527,7 +527,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( RoutineConfig::default(), - db.clone(), + ironclaw::tenant::AdminScope::new(db.clone()), llm, ws, notify_tx, @@ -614,7 +614,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( RoutineConfig::default(), - db.clone(), + ironclaw::tenant::AdminScope::new(db.clone()), llm, ws, notify_tx, @@ -723,7 +723,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( RoutineConfig::default(), - db.clone(), + ironclaw::tenant::AdminScope::new(db.clone()), llm, ws, notify_tx, @@ -866,7 +866,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( RoutineConfig::default(), - db.clone(), + ironclaw::tenant::AdminScope::new(db.clone()), llm, ws, notify_tx, @@ -1049,7 +1049,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( RoutineConfig::default(), - Arc::clone(&db), + ironclaw::tenant::AdminScope::new(Arc::clone(&db)), llm, ws, notify_tx, @@ -1171,7 +1171,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( RoutineConfig::default(), - db.clone(), + ironclaw::tenant::AdminScope::new(db.clone()), llm, ws, notify_tx, @@ -1279,7 +1279,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( config, - db.clone(), + ironclaw::tenant::AdminScope::new(db.clone()), llm, ws, notify_tx, diff --git a/tests/e2e_telegram_message_routing.rs b/tests/e2e_telegram_message_routing.rs index ead164eb..810fc218 100644 --- a/tests/e2e_telegram_message_routing.rs +++ b/tests/e2e_telegram_message_routing.rs @@ -201,6 +201,7 @@ mod tests { sandbox_readiness: ironclaw::agent::SandboxReadiness::DisabledByConfig, builder: None, llm_backend: "nearai".to_string(), + tenant_rates: std::sync::Arc::new(ironclaw::tenant::TenantRateRegistry::new(4, 3)), }; let gateway = Arc::new(TestChannel::new()); diff --git a/tests/support/gateway_workflow_harness.rs b/tests/support/gateway_workflow_harness.rs index 5f477de0..ac35b160 100644 --- a/tests/support/gateway_workflow_harness.rs +++ b/tests/support/gateway_workflow_harness.rs @@ -266,6 +266,7 @@ impl GatewayWorkflowHarness { sandbox_readiness: ironclaw::agent::SandboxReadiness::DisabledByConfig, builder: None, llm_backend: "nearai".to_string(), + tenant_rates: std::sync::Arc::new(ironclaw::tenant::TenantRateRegistry::new(4, 3)), }, channels, None, diff --git a/tests/support/test_rig.rs b/tests/support/test_rig.rs index 624bb054..5775b86d 100644 --- a/tests/support/test_rig.rs +++ b/tests/support/test_rig.rs @@ -642,7 +642,7 @@ impl TestRigBuilder { let (notify_tx, _notify_rx) = tokio::sync::mpsc::channel(16); let engine = Arc::new(RoutineEngine::new( routine_config, - Arc::clone(db_arc), + ironclaw::tenant::AdminScope::new(Arc::clone(db_arc)), components.llm.clone(), Arc::clone(ws), notify_tx, @@ -762,6 +762,7 @@ impl TestRigBuilder { sandbox_readiness: ironclaw::agent::SandboxReadiness::Available, // tests don't use real Docker builder: None, llm_backend: "nearai".to_string(), + tenant_rates: std::sync::Arc::new(ironclaw::tenant::TenantRateRegistry::new(4, 3)), }; // 7. Create TestChannel and ChannelManager. From b3fbef5287c84d0388fcdab1713c69d5ef62104a Mon Sep 17 00:00:00 2001 From: "firat.sertgoz" Date: Thu, 26 Mar 2026 09:37:59 +0300 Subject: [PATCH 7/8] fix(llm): filter XML tool-call recovery by context (#1641) * fix(llm): filter XML tool-call recovery by context * fix: address review comments on PR #1641 --- src/llm/reasoning.rs | 98 +++++++++++++++++++++++++++++++++++++++++--- 1 file changed, 92 insertions(+), 6 deletions(-) diff --git a/src/llm/reasoning.rs b/src/llm/reasoning.rs index cf3692e9..473eb16d 100644 --- a/src/llm/reasoning.rs +++ b/src/llm/reasoning.rs @@ -1345,6 +1345,49 @@ fn is_inside_code(pos: usize, regions: &[CodeRegion]) -> bool { regions.iter().any(|r| pos >= r.start && pos < r.end) } +/// Check whether a byte range overlaps any code region. +fn overlaps_code_region(start: usize, end: usize, regions: &[CodeRegion]) -> bool { + regions.iter().any(|r| start < r.end && end > r.start) +} + +/// Return the byte bounds of the line containing `pos`, excluding the trailing newline. +fn line_bounds(text: &str, pos: usize) -> (usize, usize) { + let start = text[..pos].rfind('\n').map_or(0, |idx| idx + 1); + let end = text[pos..].find('\n').map_or(text.len(), |idx| pos + idx); + (start, end) +} + +/// Only recover XML-style tool calls when they are isolated content outside +/// markdown code and quote contexts. This avoids converting code examples or +/// quoted snippets into executable tool calls. +fn is_recoverable_tool_call_segment( + text: &str, + start: usize, + end: usize, + code_regions: &[CodeRegion], +) -> bool { + if overlaps_code_region(start, end, code_regions) { + return false; + } + + let (first_line_start, first_line_end) = line_bounds(text, start); + let first_line = &text[first_line_start..first_line_end]; + + if first_line.trim_start().starts_with('>') { + return false; + } + + let (_, last_line_end) = line_bounds(text, end.saturating_sub(1)); + let first_line_prefix = &text[first_line_start..start]; + let last_line_suffix = &text[end..last_line_end]; + + if !first_line_prefix.trim().is_empty() || !last_line_suffix.trim().is_empty() { + return false; + } + + true +} + /// Clean up LLM response by stripping model-internal tags and reasoning patterns. /// /// Some models (GLM-4.7, etc.) emit XML-tagged internal state like @@ -1364,6 +1407,7 @@ fn recover_tool_calls_from_content( ) -> Vec { let tool_names: std::collections::HashSet<&str> = available_tools.iter().map(|t| t.name.as_str()).collect(); + let code_regions = find_code_regions(content); let mut calls = Vec::new(); for (open, close) in &[ @@ -1372,15 +1416,23 @@ fn recover_tool_calls_from_content( ("", ""), ("<|function_call|>", "<|/function_call|>"), ] { - let mut remaining = content; - while let Some(start) = remaining.find(open) { + let mut search_from = 0; + while let Some(offset) = content[search_from..].find(open) { + let start = search_from + offset; let inner_start = start + open.len(); - let after = &remaining[inner_start..]; - let Some(end) = after.find(close) else { + let after = &content[inner_start..]; + let Some(end_offset) = after.find(close) else { break; }; - let inner = after[..end].trim(); - remaining = &after[end + close.len()..]; + let end = inner_start + end_offset; + let segment_end = end + close.len(); + search_from = segment_end; + + if !is_recoverable_tool_call_segment(content, start, segment_end, &code_regions) { + continue; + } + + let inner = content[inner_start..end].trim(); if inner.is_empty() { continue; @@ -2313,6 +2365,40 @@ That's my plan."#; assert_eq!(calls[0].name, "tool_list"); } + #[test] + fn test_recover_tool_call_in_fenced_code_block_ignored() { + let tools = make_tools(&["tool_list"]); + let content = "Here is the XML format:\n\n```xml\ntool_list\n```"; + let calls = recover_tool_calls_from_content(content, &tools); + assert!(calls.is_empty()); + } + + #[test] + fn test_recover_tool_call_in_inline_code_ignored() { + let tools = make_tools(&["tool_list"]); + let content = "Use `tool_list` to illustrate the syntax."; + let calls = recover_tool_calls_from_content(content, &tools); + assert!(calls.is_empty()); + } + + #[test] + fn test_recover_tool_call_in_blockquote_ignored() { + let tools = make_tools(&["tool_list"]); + let content = "The page replied:\n> tool_list"; + let calls = recover_tool_calls_from_content(content, &tools); + assert!(calls.is_empty()); + } + + #[test] + fn test_recover_multiline_json_tool_call_on_own_line() { + let tools = make_tools(&["memory_search"]); + let content = "Let me check.\n\n\n{\"name\": \"memory_search\", \"arguments\": {\"query\": \"test\"}}\n\n\nDone."; + let calls = recover_tool_calls_from_content(content, &tools); + assert_eq!(calls.len(), 1); + assert_eq!(calls[0].name, "memory_search"); + assert_eq!(calls[0].arguments, serde_json::json!({"query": "test"})); + } + // ---- System prompt building tests (issue #565) ---- fn make_test_reasoning() -> Reasoning { From ed4d92932ac5d2d9123a8448aac4627bb8bb2d7c Mon Sep 17 00:00:00 2001 From: rajulbhatnagar Date: Thu, 26 Mar 2026 00:02:41 -0700 Subject: [PATCH 8/8] fix(agent): discard truncated tool calls when finish_reason == Length (#1631) (#1632) --- src/agent/agentic_loop.rs | 126 +++++++++++++++++++++++++++++++++++++- src/agent/dispatcher.rs | 2 + src/llm/mod.rs | 3 +- src/llm/reasoning.rs | 110 ++++++++++++++++++++++++++++++++- src/worker/job.rs | 2 + 5 files changed, 239 insertions(+), 4 deletions(-) diff --git a/src/agent/agentic_loop.rs b/src/agent/agentic_loop.rs index e61856dc..27c2ab72 100644 --- a/src/agent/agentic_loop.rs +++ b/src/agent/agentic_loop.rs @@ -10,7 +10,7 @@ use std::borrow::Cow; use crate::agent::session::PendingApproval; use crate::error::Error; -use crate::llm::{ChatMessage, Reasoning, ReasoningContext, RespondResult}; +use crate::llm::{ChatMessage, FinishReason, Reasoning, ReasoningContext, RespondResult}; /// Signal from the delegate indicating how the loop should proceed. pub enum LoopSignal { @@ -134,6 +134,9 @@ pub async fn run_agentic_loop( config: &AgenticLoopConfig, ) -> Result { let mut consecutive_tool_intent_nudges: u32 = 0; + // Accumulates across all iterations (not reset by text responses) so + // non-consecutive truncations still escalate to force_text. + let mut truncation_count: u32 = 0; for iteration in 1..=config.max_iterations { // Check for external signals (stop, cancellation, user messages) @@ -215,7 +218,35 @@ pub async fn run_agentic_loop( tool_calls, content, } => { + // If the response was truncated, tool call parameters are likely + // incomplete. Discard them and tell the LLM to try a different + // approach rather than executing malformed tool calls. + if output.finish_reason == FinishReason::Length { + truncation_count += 1; + let names: Vec<&str> = tool_calls.iter().map(|tc| tc.name.as_str()).collect(); + tracing::warn!( + iteration, + tools = ?names, + truncation_count, + "Discarding truncated tool calls (finish_reason=Length)" + ); + if let Some(ref text) = content { + reason_ctx.messages.push(ChatMessage::assistant(text)); + } + reason_ctx + .messages + .push(ChatMessage::user(crate::llm::TRUNCATED_TOOL_CALL_NOTICE)); + // After repeated truncations, force text-only mode so the LLM + // stops attempting tool calls it can't fit in the output budget. + if truncation_count >= 3 { + reason_ctx.force_text = true; + } + delegate.after_iteration(iteration).await; + continue; + } + consecutive_tool_intent_nudges = 0; + truncation_count = 0; if let Some(outcome) = delegate .execute_tool_calls(tool_calls, content, reason_ctx) @@ -271,6 +302,7 @@ mod tests { RespondOutput { result: RespondResult::Text(text.to_string()), usage: zero_usage(), + finish_reason: FinishReason::Stop, } } @@ -281,6 +313,7 @@ mod tests { content: None, }, usage: zero_usage(), + finish_reason: FinishReason::ToolUse, } } @@ -622,4 +655,95 @@ mod tests { let result = truncate_for_preview("café", 4); assert_eq!(result, "caf..."); } + + #[tokio::test] + async fn test_truncated_tool_calls_discarded_on_length() { + let truncated_tool_call = ToolCall { + id: "call_1".to_string(), + name: "memory_write".to_string(), + arguments: serde_json::json!({}), // empty — truncated + reasoning: None, + }; + let truncated_output = RespondOutput { + result: RespondResult::ToolCalls { + tool_calls: vec![truncated_tool_call], + content: Some("I'll write the report.".to_string()), + }, + usage: zero_usage(), + finish_reason: FinishReason::Length, // response was truncated + }; + let delegate = MockDelegate::new(vec![truncated_output, text_output("Summarized it.")]); + let reasoning = stub_reasoning(); + let mut ctx = ReasoningContext::new(); + let config = AgenticLoopConfig { + max_iterations: 5, + ..Default::default() + }; + + let outcome = run_agentic_loop(&delegate, &reasoning, &mut ctx, &config) + .await + .unwrap(); + + // Tool calls should NOT have been executed + assert_eq!(delegate.tool_exec_count.load(Ordering::SeqCst), 0); + // The loop should have continued and returned the text response + assert!(matches!(outcome, LoopOutcome::Response(ref t) if t == "Summarized it.")); + // A truncation notice should have been injected into context + assert!( + ctx.messages + .iter() + .any(|m| m.role == crate::llm::Role::User && m.content.contains("truncated")), + "Should inject truncation notice into context" + ); + // The partial assistant content should have been preserved + assert!( + ctx.messages + .iter() + .any(|m| m.role == crate::llm::Role::Assistant + && m.content.contains("write the report")), + "Should preserve partial assistant content" + ); + } + + #[tokio::test] + async fn test_repeated_truncations_force_text_mode() { + let make_truncated = || RespondOutput { + result: RespondResult::ToolCalls { + tool_calls: vec![ToolCall { + id: "call_1".to_string(), + name: "memory_write".to_string(), + arguments: serde_json::json!({}), + reasoning: None, + }], + content: None, + }, + usage: zero_usage(), + finish_reason: FinishReason::Length, + }; + // Three truncated responses, then a text response + let delegate = MockDelegate::new(vec![ + make_truncated(), + make_truncated(), + make_truncated(), + text_output("Gave up on tool calls."), + ]); + let reasoning = stub_reasoning(); + let mut ctx = ReasoningContext::new(); + let config = AgenticLoopConfig { + max_iterations: 5, + ..Default::default() + }; + + let outcome = run_agentic_loop(&delegate, &reasoning, &mut ctx, &config) + .await + .unwrap(); + + assert!(matches!(outcome, LoopOutcome::Response(_))); + assert_eq!(delegate.tool_exec_count.load(Ordering::SeqCst), 0); + // After 3 truncations, force_text should be set + assert!( + ctx.force_text, + "Should escalate to force_text after repeated truncations" + ); + } } diff --git a/src/agent/dispatcher.rs b/src/agent/dispatcher.rs index 96bca197..a5f9cd6f 100644 --- a/src/agent/dispatcher.rs +++ b/src/agent/dispatcher.rs @@ -306,6 +306,8 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { // Update context for this iteration reason_ctx.available_tools = tool_defs; + // Preserve force_text if already set (e.g. by truncation escalation). + let force_text = force_text || reason_ctx.force_text; reason_ctx.system_prompt = Some(if force_text { self.cached_prompt_no_tools.clone() } else { diff --git a/src/llm/mod.rs b/src/llm/mod.rs index 308b3983..d681547d 100644 --- a/src/llm/mod.rs +++ b/src/llm/mod.rs @@ -63,7 +63,8 @@ pub use provider::{ }; pub use reasoning::{ ActionPlan, Reasoning, ReasoningContext, RespondOutput, RespondResult, SILENT_REPLY_TOKEN, - TOOL_INTENT_NUDGE, TokenUsage, ToolSelection, is_silent_reply, llm_signals_tool_intent, + TOOL_INTENT_NUDGE, TRUNCATED_TOOL_CALL_NOTICE, TokenUsage, ToolSelection, is_silent_reply, + llm_signals_tool_intent, }; pub use recording::RecordingLlm; pub use registry::{ProviderDefinition, ProviderProtocol, ProviderRegistry}; diff --git a/src/llm/reasoning.rs b/src/llm/reasoning.rs index 473eb16d..6e078ac7 100644 --- a/src/llm/reasoning.rs +++ b/src/llm/reasoning.rs @@ -8,8 +8,8 @@ use serde::{Deserialize, Serialize}; use crate::llm::error::LlmError; use crate::llm::{ - ChatMessage, CompletionRequest, LlmProvider, Role, ToolCall, ToolCompletionRequest, - ToolDefinition, + ChatMessage, CompletionRequest, FinishReason, LlmProvider, Role, ToolCall, + ToolCompletionRequest, ToolDefinition, }; /// Token the agent returns when it has nothing to say (e.g. in group chats). @@ -23,6 +23,13 @@ You said you would perform an action, but you did not include any tool calls.\n\ Do NOT describe what you intend to do — actually call the tool now.\n\ Use the tool_calls mechanism to invoke the appropriate tool."; +/// Notice injected when the LLM's response was truncated mid-tool-call, +/// causing incomplete parameters. Tells the LLM to try a different approach. +pub const TRUNCATED_TOOL_CALL_NOTICE: &str = "\ +Your previous response was truncated while generating tool call parameters. \ +The tool calls were discarded. Please try a different approach — \ +summarize or transform the data instead of echoing it verbatim in a tool call."; + /// Seed value used as the second argument to `generate_tool_call_id` when /// recovering tool calls from malformed LLM text responses. This must differ /// from the `0` seed used in `rig_adapter::normalized_tool_call_id` to avoid @@ -194,6 +201,8 @@ pub struct ReasoningContext { pub metadata: std::collections::HashMap, /// When true, force a text-only response (ignore available tools). /// Used by the agentic loop to guarantee termination near the iteration limit. + /// Sticky: once set, never cleared within a loop invocation. Callers must + /// create a fresh `ReasoningContext` per `run_agentic_loop()` call. pub force_text: bool, /// Pre-built system prompt. When set, `respond_with_tools` uses this directly /// instead of calling `build_system_prompt_with_tools`. Allows callers to build @@ -349,6 +358,7 @@ pub enum RespondResult { pub struct RespondOutput { pub result: RespondResult, pub usage: TokenUsage, + pub finish_reason: FinishReason, } /// Reasoning engine for the agent. @@ -530,6 +540,17 @@ impl Reasoning { let response = self.llm.complete_with_tools(request).await?; + // If the response was truncated, tool call parameters are likely incomplete. + // Return empty so the caller can fall through to respond_with_tools() which + // has a larger output token budget. + if response.finish_reason == FinishReason::Length { + tracing::warn!( + "select_tools response truncated (finish_reason=Length), \ + discarding potentially incomplete tool selections" + ); + return Ok(vec![]); + } + let shared_reasoning = response .content .map(|c| { @@ -722,6 +743,7 @@ Respond in JSON format: content: narrative, }, usage, + finish_reason: response.finish_reason, }); } @@ -749,6 +771,7 @@ Respond in JSON format: }, }, usage, + finish_reason: response.finish_reason, }); } @@ -774,6 +797,7 @@ Respond in JSON format: Ok(RespondOutput { result: RespondResult::Text(final_text), usage, + finish_reason: response.finish_reason, }) } else { // No tools, use simple completion @@ -805,6 +829,7 @@ Respond in JSON format: cache_read_input_tokens: response.cache_read_input_tokens, cache_creation_input_tokens: response.cache_creation_input_tokens, }, + finish_reason: response.finish_reason, }) } } @@ -3315,4 +3340,85 @@ That's my plan."#; let cleaned = clean_response(&pre_truncated); assert!(cleaned.trim().is_empty()); } + + // ---- select_tools truncation guard ---- + + /// Mock provider that returns tool calls with a configurable finish_reason. + struct TruncatingLlm { + finish_reason: crate::llm::FinishReason, + } + + #[async_trait::async_trait] + impl crate::llm::LlmProvider for TruncatingLlm { + fn model_name(&self) -> &str { + "truncating-stub" + } + fn cost_per_token(&self) -> (rust_decimal::Decimal, rust_decimal::Decimal) { + (rust_decimal::Decimal::ZERO, rust_decimal::Decimal::ZERO) + } + async fn complete( + &self, + _request: crate::llm::CompletionRequest, + ) -> Result { + unimplemented!() + } + async fn complete_with_tools( + &self, + _request: crate::llm::ToolCompletionRequest, + ) -> Result { + Ok(crate::llm::ToolCompletionResponse { + content: Some("I'll write the report.".to_string()), + tool_calls: vec![ToolCall { + id: "call_1".to_string(), + name: "memory_write".to_string(), + arguments: serde_json::json!({}), + reasoning: None, + }], + input_tokens: 5000, + output_tokens: 1024, + finish_reason: self.finish_reason, + cache_read_input_tokens: 0, + cache_creation_input_tokens: 0, + }) + } + } + + #[tokio::test] + async fn test_select_tools_returns_empty_on_truncation() { + let llm = Arc::new(TruncatingLlm { + finish_reason: FinishReason::Length, + }); + let reasoning = Reasoning::new(llm); + let mut ctx = ReasoningContext::new().with_message(ChatMessage::user("Write a report")); + ctx.available_tools.push(ToolDefinition { + name: "memory_write".to_string(), + description: "Write to memory".to_string(), + parameters: serde_json::json!({"type": "object"}), + }); + + let selections = reasoning.select_tools(&ctx).await.unwrap(); + assert!( + selections.is_empty(), + "Truncated tool selections should be discarded (got {} selections)", + selections.len() + ); + } + + #[tokio::test] + async fn test_select_tools_returns_selections_when_not_truncated() { + let llm = Arc::new(TruncatingLlm { + finish_reason: FinishReason::ToolUse, + }); + let reasoning = Reasoning::new(llm); + let mut ctx = ReasoningContext::new().with_message(ChatMessage::user("Write a report")); + ctx.available_tools.push(ToolDefinition { + name: "memory_write".to_string(), + description: "Write to memory".to_string(), + parameters: serde_json::json!({"type": "object"}), + }); + + let selections = reasoning.select_tools(&ctx).await.unwrap(); + assert_eq!(selections.len(), 1); + assert_eq!(selections[0].tool_name, "memory_write"); + } } diff --git a/src/worker/job.rs b/src/worker/job.rs index 671b8864..f74d4ec8 100644 --- a/src/worker/job.rs +++ b/src/worker/job.rs @@ -1158,6 +1158,7 @@ impl<'a> JobDelegate<'a> { Ok(crate::llm::RespondOutput { result: RespondResult::Text(String::new()), usage: crate::llm::TokenUsage::default(), + finish_reason: crate::llm::FinishReason::Stop, }) } } @@ -1283,6 +1284,7 @@ impl<'a> LoopDelegate for JobDelegate<'a> { content: reasoning_text, }, usage: crate::llm::TokenUsage::default(), + finish_reason: crate::llm::FinishReason::ToolUse, }); } Ok(_) => {} // empty selections, fall through