From 656151783cb9aa165d9dc99e82d7855ed3943b11 Mon Sep 17 00:00:00 2001 From: Illia Polosukhin Date: Tue, 24 Mar 2026 23:01:19 -0700 Subject: [PATCH 01/11] feat(cli): show credential auth status in tool info (#1572) * feat(cli): show credential auth status in `tool info` `ironclaw tool info` now checks the secrets store and shows whether each required credential is configured or missing, consolidated into a single Auth section that deduplicates across http.credentials, auth, and setup.required_secrets. Secrets already shown in Auth are filtered from the Secrets section to avoid redundancy. Co-Authored-By: Claude Opus 4.6 (1M context) * fix(cli): address review feedback on tool info auth status - Fix clippy collapsible-if by using `if let` + `&&` - Use HashMap for O(1) dedup instead of HashSet + linear scan - Add --user flag to `tool info` for checking non-default user credentials - Show "? unknown" on secrets store errors instead of silently reporting missing - Surface secrets store init failure via eprintln instead of silent .ok() - Sort auth entries by secret name for deterministic output Co-Authored-By: Claude Opus 4.6 (1M context) * fix(cli): only filter secrets when auth section renders, add regression test When the secrets store fails to initialize, the Auth section is not rendered. Previously, secret names were still filtered from the Secrets section, causing credential names to disappear entirely. Now secrets are only filtered when the Auth section will actually be displayed. Adds test verifying auth secret deduplication across auth, setup, and http.credentials sections, plus secrets store existence checks. Co-Authored-By: Claude Opus 4.6 (1M context) * refactor(cli): extract collect_auth_secrets helper, always render Auth section Address review feedback: - Extract dedup logic into `collect_auth_secrets()` so the test exercises the same code path as production (not a re-implementation) - Always render the Auth section when auth secrets exist, showing "? unknown" status when the secrets store is unavailable instead of hiding credential names entirely - Lazily init secrets store only when capabilities contain auth secrets, avoiding spurious warnings for tools with no auth - Add test for empty capabilities edge case Co-Authored-By: Claude Opus 4.6 (1M context) * style(cli): move HashMap/HashSet imports to top of file Co-Authored-By: Claude Opus 4.6 (1M context) * fix(cli): use correct tagged JSON format for credential location in test The CredentialLocationSchema uses serde tagged enum format ({"type": "bearer"}), not a bare string ("AuthorizationBearer"). Co-Authored-By: Claude Opus 4.6 (1M context) --------- Co-authored-by: Claude Opus 4.6 (1M context) --- src/cli/tool.rs | 286 +++++++++++++++++++++++++++++++++++++++++++++--- 1 file changed, 270 insertions(+), 16 deletions(-) diff --git a/src/cli/tool.rs b/src/cli/tool.rs index be684580..9d39c492 100644 --- a/src/cli/tool.rs +++ b/src/cli/tool.rs @@ -2,6 +2,7 @@ //! //! Commands for installing, listing, removing, and authenticating WASM tools. +use std::collections::{HashMap, HashSet}; use std::io::Write; use std::path::{Path, PathBuf}; use std::sync::Arc; @@ -79,6 +80,10 @@ pub enum ToolCommand { /// Directory to look for tool (default: ~/.ironclaw/tools/) #[arg(short, long)] dir: Option, + + /// User ID for checking credential status (default: "default") + #[arg(short, long, default_value = "default")] + user: String, }, /// Configure authentication for a tool @@ -124,7 +129,11 @@ pub async fn run_tool_command(cmd: ToolCommand) -> anyhow::Result<()> { } => install_tool(path, name, capabilities, target, release, skip_build, force).await, ToolCommand::List { dir, verbose } => list_tools(dir, verbose).await, ToolCommand::Remove { name, dir } => remove_tool(name, dir).await, - ToolCommand::Info { name_or_path, dir } => show_tool_info(name_or_path, dir).await, + ToolCommand::Info { + name_or_path, + dir, + user, + } => show_tool_info(name_or_path, dir, user).await, ToolCommand::Auth { name, dir, user } => auth_tool(name, dir, user).await, ToolCommand::Setup { name, dir, user } => setup_tool(name, dir, user).await, } @@ -388,7 +397,11 @@ async fn remove_tool(name: String, dir: Option) -> anyhow::Result<()> { } /// Show information about a tool. -async fn show_tool_info(name_or_path: String, dir: Option) -> anyhow::Result<()> { +async fn show_tool_info( + name_or_path: String, + dir: Option, + user_id: String, +) -> anyhow::Result<()> { let wasm_path = if name_or_path.ends_with(".wasm") { PathBuf::from(&name_or_path) } else { @@ -423,7 +436,37 @@ async fn show_tool_info(name_or_path: String, dir: Option) -> anyhow::R println!("\nCapabilities ({}):", caps_path.display()); let content = fs::read_to_string(&caps_path).await?; match CapabilitiesFile::from_json(&content) { - Ok(caps) => print_capabilities_detail(&caps), + Ok(caps) => { + // Lazily init secrets store only when auth secrets need checking. + let has_auth = caps.auth.is_some() + || caps + .setup + .as_ref() + .is_some_and(|s| !s.required_secrets.is_empty()) + || caps + .http + .as_ref() + .is_some_and(|h| !h.credentials.is_empty()); + let secrets_store = if has_auth { + match init_secrets_store().await { + Ok(store) => Some(store), + Err(e) => { + eprintln!(" Warning: could not init secrets store: {}", e); + None + } + } + } else { + None + }; + print_capabilities_detail( + &caps, + secrets_store + .as_ref() + .map(|s| s.as_ref() as &(dyn SecretsStore + Send + Sync)), + &user_id, + ) + .await; + } Err(e) => println!(" Error parsing: {}", e), } } else { @@ -476,8 +519,89 @@ fn print_capabilities_summary(caps: &CapabilitiesFile) { } } +/// Per-secret info collected from all auth-related capability sections. +struct AuthSecretInfo { + secret_name: String, + /// Human-readable label (from auth.display_name or setup prompt). + description: Option, + /// Injection location (from http.credentials). + location: Option, +} + +/// Collected auth secrets and the set of secret names they cover. +struct CollectedAuthSecrets { + secrets: Vec, + /// Secret names present in `secrets`, for filtering the Secrets capability section. + seen_names: HashSet, +} + +/// Collect and deduplicate auth secrets from all auth-related capability sections. +/// +/// Priority for the description label: auth.display_name > setup.required_secrets.prompt. +/// Injection location is merged from http.credentials. +fn collect_auth_secrets(caps: &CapabilitiesFile) -> CollectedAuthSecrets { + let mut secrets: Vec = Vec::new(); + let mut seen: HashMap = HashMap::new(); + + // auth.display_name is the best label — seed first. + if let Some(ref auth) = caps.auth { + let index = secrets.len(); + seen.insert(auth.secret_name.clone(), index); + secrets.push(AuthSecretInfo { + secret_name: auth.secret_name.clone(), + description: auth.display_name.clone(), + location: None, + }); + } + + // setup.required_secrets.prompt is second-best label. + if let Some(ref setup) = caps.setup { + for secret in &setup.required_secrets { + if !seen.contains_key(&secret.name) { + let index = secrets.len(); + seen.insert(secret.name.clone(), index); + secrets.push(AuthSecretInfo { + secret_name: secret.name.clone(), + description: Some(secret.prompt.clone()), + location: None, + }); + } + } + } + + // Merge injection location from http.credentials. + if let Some(ref http) = caps.http { + for cred in http.credentials.values() { + let loc = format!("{:?}", cred.location); + if let Some(&index) = seen.get(&cred.secret_name) { + secrets[index].location = Some(loc); + } else { + let index = secrets.len(); + seen.insert(cred.secret_name.clone(), index); + secrets.push(AuthSecretInfo { + secret_name: cred.secret_name.clone(), + description: None, + location: Some(loc), + }); + } + } + } + + let seen_names = seen.into_keys().collect(); + CollectedAuthSecrets { + secrets, + seen_names, + } +} + /// Print detailed capabilities. -fn print_capabilities_detail(caps: &CapabilitiesFile) { +async fn print_capabilities_detail( + caps: &CapabilitiesFile, + secrets_store: Option<&(dyn SecretsStore + Send + Sync)>, + user_id: &str, +) { + let mut collected = collect_auth_secrets(caps); + if let Some(ref http) = caps.http { println!(" HTTP:"); for endpoint in &http.allowlist { @@ -490,13 +614,6 @@ fn print_capabilities_detail(caps: &CapabilitiesFile) { println!(" {} {} {}", methods, endpoint.host, path); } - if !http.credentials.is_empty() { - println!(" Credentials:"); - for (key, cred) in &http.credentials { - println!(" {}: {} -> {:?}", key, cred.secret_name, cred.location); - } - } - if let Some(ref rate) = http.rate_limit { println!( " Rate limit: {}/min, {}/hour", @@ -505,12 +622,24 @@ fn print_capabilities_detail(caps: &CapabilitiesFile) { } } + // Filter secrets already covered by the auth section (always rendered when non-empty). if let Some(ref secrets) = caps.secrets && !secrets.allowed_names.is_empty() { - println!(" Secrets (existence check only):"); - for name in &secrets.allowed_names { - println!(" {}", name); + let extra: Vec<_> = if collected.secrets.is_empty() { + secrets.allowed_names.iter().collect() + } else { + secrets + .allowed_names + .iter() + .filter(|name| !collected.seen_names.contains(name.as_str())) + .collect() + }; + if !extra.is_empty() { + println!(" Secrets (existence check only):"); + for name in extra { + println!(" {}", name); + } } } @@ -531,6 +660,38 @@ fn print_capabilities_detail(caps: &CapabilitiesFile) { println!(" {}", prefix); } } + + // Consolidated auth status — sorted by secret name for deterministic output. + if !collected.secrets.is_empty() { + collected + .secrets + .sort_by(|a, b| a.secret_name.cmp(&b.secret_name)); + println!(" Auth:"); + for info in &collected.secrets { + let (icon, label) = match secrets_store { + Some(store) => match store.exists(user_id, &info.secret_name).await { + Ok(true) => ("\u{2713}", "configured"), + Ok(false) => ("\u{2717}", "missing"), + Err(e) => { + eprintln!( + " Warning: failed to check secret `{}`: {}", + info.secret_name, e + ); + ("?", "unknown") + } + }, + None => ("?", "unknown"), + }; + let mut parts = info.secret_name.clone(); + if let Some(ref desc) = info.description { + parts = format!("{} ({})", parts, desc); + } + if let Some(ref loc) = info.location { + parts = format!("{} -> {}", parts, loc); + } + println!(" {} {} {}", parts, icon, label); + } + } } /// Validate a tool name to prevent path traversal. @@ -677,8 +838,7 @@ async fn combine_provider_scopes( secret_name: &str, base_oauth: &crate::tools::wasm::OAuthConfigSchema, ) -> crate::tools::wasm::OAuthConfigSchema { - let mut all_scopes: std::collections::HashSet = - base_oauth.scopes.iter().cloned().collect(); + let mut all_scopes: HashSet = base_oauth.scopes.iter().cloned().collect(); if let Ok(mut entries) = tokio::fs::read_dir(tools_dir).await { while let Ok(Some(entry)) = entries.next_entry().await { @@ -1127,6 +1287,8 @@ async fn setup_tool(name: String, dir: Option, user_id: String) -> anyh #[cfg(test)] mod tests { use super::*; + use crate::secrets::{CreateSecretParams, SecretsStore}; + use crate::testing::credentials::test_secrets_store; #[test] fn test_format_size() { @@ -1143,4 +1305,96 @@ mod tests { assert!(dir.to_string_lossy().contains(".ironclaw")); assert!(dir.to_string_lossy().contains("tools")); } + + /// Verify that auth secrets are deduplicated across auth, setup, and http.credentials, + /// and that credential status is checked against the secrets store. + #[tokio::test] + async fn test_auth_secret_dedup_and_status() { + let caps = CapabilitiesFile::from_json( + r#"{ + "auth": { + "secret_name": "gh_token", + "display_name": "GitHub" + }, + "setup": { + "required_secrets": [ + { "name": "gh_token", "prompt": "GitHub PAT" }, + { "name": "extra_key", "prompt": "Extra API Key" } + ] + }, + "http": { + "allowlist": [{ "host": "api.github.com" }], + "credentials": { + "github": { + "secret_name": "gh_token", + "location": { "type": "bearer" }, + "host_patterns": ["api.github.com"] + } + } + }, + "secrets": { + "allowed_names": ["gh_token", "gh_*"] + } + }"#, + ) + .unwrap(); + + let collected = collect_auth_secrets(&caps); + + // gh_token should appear once (from auth), with location merged from credentials. + // extra_key should appear once (from setup). + assert_eq!(collected.secrets.len(), 2); + let gh = collected + .secrets + .iter() + .find(|s| s.secret_name == "gh_token") + .unwrap(); + assert_eq!(gh.description.as_deref(), Some("GitHub")); + assert!( + gh.location.is_some(), + "location should be merged from http.credentials" + ); + + let extra = collected + .secrets + .iter() + .find(|s| s.secret_name == "extra_key") + .unwrap(); + assert_eq!(extra.description.as_deref(), Some("Extra API Key")); + assert!(extra.location.is_none()); + + // Secrets section should filter gh_token (in seen_names) but keep gh_* (wildcard). + let secrets = caps.secrets.as_ref().unwrap(); + let extra_secrets: Vec<_> = secrets + .allowed_names + .iter() + .filter(|name| !collected.seen_names.contains(name.as_str())) + .collect(); + assert_eq!(extra_secrets, vec!["gh_*"]); + + // Verify store check: missing secret -> exists returns false. + let store = test_secrets_store(); + assert!(!store.exists("default", "gh_token").await.unwrap()); + + // Store gh_token and verify it's found. + store + .create( + "default", + CreateSecretParams::new("gh_token", "ghp_test123"), + ) + .await + .unwrap(); + assert!(store.exists("default", "gh_token").await.unwrap()); + // extra_key still missing. + assert!(!store.exists("default", "extra_key").await.unwrap()); + } + + /// No auth sections → collect_auth_secrets returns empty. + #[test] + fn test_collect_auth_secrets_empty_caps() { + let caps = CapabilitiesFile::default(); + let collected = collect_auth_secrets(&caps); + assert!(collected.secrets.is_empty()); + assert!(collected.seen_names.is_empty()); + } } From 706c3a1b4747d0335fd45013deddde3239be2f7f Mon Sep 17 00:00:00 2001 From: Illia Polosukhin Date: Tue, 24 Mar 2026 23:02:46 -0700 Subject: [PATCH 02/11] refactor: extract AppEvent to crates/ironclaw_common (#1615) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * refactor: extract AppEvent to crates/ironclaw_common SseEvent was defined in src/channels/web/types.rs but imported by 12+ modules across agent, orchestrator, worker, tools, and extensions — it had become the application-wide event protocol, not a web transport concern. Create crates/ironclaw_common as a shared workspace crate and move the enum there as AppEvent. Also move the truncate_preview utility which was similarly leaked from the web gateway into agent modules. - New crate: crates/ironclaw_common (AppEvent, truncate_preview) - Rename SseEvent → AppEvent, from_sse_event → from_app_event - web/types.rs re-exports AppEvent for internal gateway use - web/util.rs re-exports truncate_preview - Wire format unchanged (serde renames are on variants, not the enum) Aligned with the event bus direction on refactor/architectural-hardening where DomainEvent (≡ AppEvent) is wrapped in a SystemEvent envelope. Co-Authored-By: Claude Opus 4.6 (1M context) * refactor: add AppEvent::event_type() helper, deduplicate match blocks Address Gemini review: extract the variant→string match into a single method on AppEvent, replacing the duplicated 22-arm matches in sse.rs and types.rs. Co-Authored-By: Claude Opus 4.6 (1M context) * refactor: rename leftover sse vars/tests to match AppEvent rename Address Copilot review: rename sse_event vars to app_event in orchestrator/api.rs and ws.rs, rename test functions from test_ws_server_from_sse_* to test_ws_server_from_app_event_*, and update stale SSE comments. Co-Authored-By: Claude Opus 4.6 (1M context) * refactor: add Deserialize to AppEvent, round-trip test, fix stale comments Address zmanian review: - Add Deserialize derive to AppEvent so downstream consumers can deserialize incoming events - Add event_type_matches_serde_type_field test that round-trips every variant through serde and asserts event_type() matches the serialized "type" field — catches drift between serde renames and the manual match - Add round_trip_deserialize test for basic Serialize/Deserialize parity - Update remaining "SSE" references in comments across server.rs, manager.rs, ws_gateway_integration.rs, and worker/job.rs Co-Authored-By: Claude Opus 4.6 (1M context) --------- Co-authored-by: Claude Opus 4.6 (1M context) --- Cargo.lock | 9 + Cargo.toml | 5 +- crates/ironclaw_common/Cargo.toml | 18 ++ crates/ironclaw_common/src/event.rs | 338 ++++++++++++++++++++++++++++ crates/ironclaw_common/src/lib.rs | 7 + crates/ironclaw_common/src/util.rs | 100 ++++++++ src/agent/job_monitor.rs | 48 ++-- src/agent/session.rs | 2 +- src/agent/thread_ops.rs | 2 +- src/channels/web/handlers/chat.rs | 6 +- src/channels/web/mod.rs | 32 +-- src/channels/web/server.rs | 24 +- src/channels/web/sse.rs | 61 ++--- src/channels/web/types.rs | 233 +++---------------- src/channels/web/util.rs | 106 +-------- src/channels/web/ws.rs | 8 +- src/extensions/manager.rs | 6 +- src/orchestrator/api.rs | 28 +-- src/orchestrator/mod.rs | 4 +- src/tools/builtin/job.rs | 6 +- src/tools/registry.rs | 6 +- src/worker/job.rs | 14 +- tests/multi_tenant_integration.rs | 46 ++-- tests/ws_gateway_integration.rs | 18 +- 24 files changed, 646 insertions(+), 481 deletions(-) create mode 100644 crates/ironclaw_common/Cargo.toml create mode 100644 crates/ironclaw_common/src/event.rs create mode 100644 crates/ironclaw_common/src/lib.rs create mode 100644 crates/ironclaw_common/src/util.rs diff --git a/Cargo.lock b/Cargo.lock index a813ef2b..27c258c1 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3428,6 +3428,7 @@ dependencies = [ "hyper-util", "iana-time-zone", "insta", + "ironclaw_common", "ironclaw_safety", "json5", "libsql", @@ -3485,6 +3486,14 @@ dependencies = [ "zip", ] +[[package]] +name = "ironclaw_common" +version = "0.1.0" +dependencies = [ + "serde", + "serde_json", +] + [[package]] name = "ironclaw_safety" version = "0.1.0" diff --git a/Cargo.toml b/Cargo.toml index 99992a40..395e42d3 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,5 +1,5 @@ [workspace] -members = [".", "crates/ironclaw_safety"] +members = [".", "crates/ironclaw_common", "crates/ironclaw_safety"] exclude = [ "channels-src/discord", "channels-src/telegram", @@ -100,6 +100,9 @@ tower-http = { version = "0.6", features = ["trace", "cors", "set-header"] } # Cron scheduling for routines cron = "0.13" +# Shared types +ironclaw_common = { path = "crates/ironclaw_common", version = "0.1.0" } + # Safety/sanitization ironclaw_safety = { path = "crates/ironclaw_safety", version = "0.1.0" } regex = "1" diff --git a/crates/ironclaw_common/Cargo.toml b/crates/ironclaw_common/Cargo.toml new file mode 100644 index 00000000..353ab747 --- /dev/null +++ b/crates/ironclaw_common/Cargo.toml @@ -0,0 +1,18 @@ +[package] +name = "ironclaw_common" +version = "0.1.0" +edition = "2024" +rust-version = "1.92" +description = "Shared types and utilities for the IronClaw workspace" +authors = ["NEAR AI "] +license = "MIT OR Apache-2.0" +homepage = "https://github.com/nearai/ironclaw" +repository = "https://github.com/nearai/ironclaw" +publish = false + +[package.metadata.dist] +dist = false + +[dependencies] +serde = { version = "1", features = ["derive"] } +serde_json = "1" diff --git a/crates/ironclaw_common/src/event.rs b/crates/ironclaw_common/src/event.rs new file mode 100644 index 00000000..83592c95 --- /dev/null +++ b/crates/ironclaw_common/src/event.rs @@ -0,0 +1,338 @@ +//! Application-wide event types. +//! +//! `AppEvent` is the real-time event protocol used across the entire +//! application. The web gateway serialises these to SSE / WebSocket +//! frames, but other subsystems (agent loop, orchestrator, extensions) +//! produce and consume them too. + +use serde::{Deserialize, Serialize}; + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "type")] +pub enum AppEvent { + #[serde(rename = "response")] + Response { content: String, thread_id: String }, + #[serde(rename = "thinking")] + Thinking { + message: String, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + #[serde(rename = "tool_started")] + ToolStarted { + name: String, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + #[serde(rename = "tool_completed")] + ToolCompleted { + name: String, + success: bool, + #[serde(skip_serializing_if = "Option::is_none")] + error: Option, + #[serde(skip_serializing_if = "Option::is_none")] + parameters: Option, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + #[serde(rename = "tool_result")] + ToolResult { + name: String, + preview: String, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + #[serde(rename = "stream_chunk")] + StreamChunk { + content: String, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + #[serde(rename = "status")] + Status { + message: String, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + #[serde(rename = "job_started")] + JobStarted { + job_id: String, + title: String, + browse_url: String, + }, + #[serde(rename = "approval_needed")] + ApprovalNeeded { + request_id: String, + tool_name: String, + description: String, + parameters: String, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + /// Whether the "always" auto-approve option should be shown. + allow_always: bool, + }, + #[serde(rename = "auth_required")] + AuthRequired { + extension_name: String, + #[serde(skip_serializing_if = "Option::is_none")] + instructions: Option, + #[serde(skip_serializing_if = "Option::is_none")] + auth_url: Option, + #[serde(skip_serializing_if = "Option::is_none")] + setup_url: Option, + }, + #[serde(rename = "auth_completed")] + AuthCompleted { + extension_name: String, + success: bool, + message: String, + }, + #[serde(rename = "error")] + Error { + message: String, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + #[serde(rename = "heartbeat")] + Heartbeat, + + // Sandbox job streaming events (worker + Claude Code bridge) + #[serde(rename = "job_message")] + JobMessage { + job_id: String, + role: String, + content: String, + }, + #[serde(rename = "job_tool_use")] + JobToolUse { + job_id: String, + tool_name: String, + input: serde_json::Value, + }, + #[serde(rename = "job_tool_result")] + JobToolResult { + job_id: String, + tool_name: String, + output: String, + }, + #[serde(rename = "job_status")] + JobStatus { job_id: String, message: String }, + #[serde(rename = "job_result")] + JobResult { + job_id: String, + status: String, + #[serde(skip_serializing_if = "Option::is_none")] + session_id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + fallback_deliverable: Option, + }, + + /// An image was generated by a tool. + #[serde(rename = "image_generated")] + ImageGenerated { + data_url: String, + #[serde(skip_serializing_if = "Option::is_none")] + path: Option, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + + /// Suggested follow-up messages for the user. + #[serde(rename = "suggestions")] + Suggestions { + suggestions: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + + /// Per-turn token usage and cost summary. + #[serde(rename = "turn_cost")] + TurnCost { + input_tokens: u64, + output_tokens: u64, + cost_usd: String, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + + /// Extension activation status change (WASM channels). + #[serde(rename = "extension_status")] + ExtensionStatus { + extension_name: String, + status: String, + #[serde(skip_serializing_if = "Option::is_none")] + message: Option, + }, +} + +impl AppEvent { + /// The wire-format event type string (matches the `#[serde(rename)]` value). + pub fn event_type(&self) -> &'static str { + match self { + Self::Response { .. } => "response", + Self::Thinking { .. } => "thinking", + Self::ToolStarted { .. } => "tool_started", + Self::ToolCompleted { .. } => "tool_completed", + Self::ToolResult { .. } => "tool_result", + Self::StreamChunk { .. } => "stream_chunk", + Self::Status { .. } => "status", + Self::JobStarted { .. } => "job_started", + Self::ApprovalNeeded { .. } => "approval_needed", + Self::AuthRequired { .. } => "auth_required", + Self::AuthCompleted { .. } => "auth_completed", + Self::Error { .. } => "error", + Self::Heartbeat => "heartbeat", + Self::JobMessage { .. } => "job_message", + Self::JobToolUse { .. } => "job_tool_use", + Self::JobToolResult { .. } => "job_tool_result", + Self::JobStatus { .. } => "job_status", + Self::JobResult { .. } => "job_result", + Self::ImageGenerated { .. } => "image_generated", + Self::Suggestions { .. } => "suggestions", + Self::TurnCost { .. } => "turn_cost", + Self::ExtensionStatus { .. } => "extension_status", + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + /// Verify that `event_type()` returns the same string as the serde + /// `"type"` field for every variant. This catches drift between the + /// `#[serde(rename)]` attributes and the manual match arms. + #[test] + fn event_type_matches_serde_type_field() { + let variants: Vec = vec![ + AppEvent::Response { + content: String::new(), + thread_id: String::new(), + }, + AppEvent::Thinking { + message: String::new(), + thread_id: None, + }, + AppEvent::ToolStarted { + name: String::new(), + thread_id: None, + }, + AppEvent::ToolCompleted { + name: String::new(), + success: true, + error: None, + parameters: None, + thread_id: None, + }, + AppEvent::ToolResult { + name: String::new(), + preview: String::new(), + thread_id: None, + }, + AppEvent::StreamChunk { + content: String::new(), + thread_id: None, + }, + AppEvent::Status { + message: String::new(), + thread_id: None, + }, + AppEvent::JobStarted { + job_id: String::new(), + title: String::new(), + browse_url: String::new(), + }, + AppEvent::ApprovalNeeded { + request_id: String::new(), + tool_name: String::new(), + description: String::new(), + parameters: String::new(), + thread_id: None, + allow_always: false, + }, + AppEvent::AuthRequired { + extension_name: String::new(), + instructions: None, + auth_url: None, + setup_url: None, + }, + AppEvent::AuthCompleted { + extension_name: String::new(), + success: true, + message: String::new(), + }, + AppEvent::Error { + message: String::new(), + thread_id: None, + }, + AppEvent::Heartbeat, + AppEvent::JobMessage { + job_id: String::new(), + role: String::new(), + content: String::new(), + }, + AppEvent::JobToolUse { + job_id: String::new(), + tool_name: String::new(), + input: serde_json::Value::Null, + }, + AppEvent::JobToolResult { + job_id: String::new(), + tool_name: String::new(), + output: String::new(), + }, + AppEvent::JobStatus { + job_id: String::new(), + message: String::new(), + }, + AppEvent::JobResult { + job_id: String::new(), + status: String::new(), + session_id: None, + fallback_deliverable: None, + }, + AppEvent::ImageGenerated { + data_url: String::new(), + path: None, + thread_id: None, + }, + AppEvent::Suggestions { + suggestions: vec![], + thread_id: None, + }, + AppEvent::TurnCost { + input_tokens: 0, + output_tokens: 0, + cost_usd: String::new(), + thread_id: None, + }, + AppEvent::ExtensionStatus { + extension_name: String::new(), + status: String::new(), + message: None, + }, + ]; + + for variant in &variants { + let json: serde_json::Value = serde_json::to_value(variant).unwrap(); + let serde_type = json["type"].as_str().unwrap(); + assert_eq!( + variant.event_type(), + serde_type, + "event_type() mismatch for variant: {:?}", + variant + ); + } + } + + #[test] + fn round_trip_deserialize() { + let original = AppEvent::Response { + content: "hello".to_string(), + thread_id: "t1".to_string(), + }; + let json = serde_json::to_string(&original).unwrap(); + let deserialized: AppEvent = serde_json::from_str(&json).unwrap(); + assert_eq!(deserialized.event_type(), "response"); + } +} diff --git a/crates/ironclaw_common/src/lib.rs b/crates/ironclaw_common/src/lib.rs new file mode 100644 index 00000000..6822bad1 --- /dev/null +++ b/crates/ironclaw_common/src/lib.rs @@ -0,0 +1,7 @@ +//! Shared types and utilities for the IronClaw workspace. + +mod event; +mod util; + +pub use event::AppEvent; +pub use util::truncate_preview; diff --git a/crates/ironclaw_common/src/util.rs b/crates/ironclaw_common/src/util.rs new file mode 100644 index 00000000..4f054671 --- /dev/null +++ b/crates/ironclaw_common/src/util.rs @@ -0,0 +1,100 @@ +//! Shared utility functions. + +/// Truncate a string to at most `max_bytes` bytes at a char boundary, appending "...". +/// +/// If the input is wrapped in `...` and truncation +/// removes the closing tag, the tag is re-appended so downstream XML parsers +/// never see an unclosed element. +pub fn truncate_preview(s: &str, max_bytes: usize) -> String { + if s.len() <= max_bytes { + return s.to_string(); + } + // Walk backwards from max_bytes to find a valid char boundary + let mut end = max_bytes; + while end > 0 && !s.is_char_boundary(end) { + end -= 1; + } + let mut result = format!("{}...", &s[..end]); + + // Re-close if truncation cut through the closing tag. + if s.starts_with("") { + result.push_str("\n"); + } + + result +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_truncate_preview_short_string() { + assert_eq!(truncate_preview("hello", 10), "hello"); + } + + #[test] + fn test_truncate_preview_exact_boundary() { + assert_eq!(truncate_preview("hello", 5), "hello"); + } + + #[test] + fn test_truncate_preview_truncates_ascii() { + assert_eq!(truncate_preview("hello world", 5), "hello..."); + } + + #[test] + fn test_truncate_preview_empty_string() { + assert_eq!(truncate_preview("", 10), ""); + } + + #[test] + fn test_truncate_preview_multibyte_char_boundary() { + let s = "a\u{20AC}b"; + let result = truncate_preview(s, 3); + assert_eq!(result, "a..."); + } + + #[test] + fn test_truncate_preview_emoji() { + let s = "hi\u{1F980}"; + let result = truncate_preview(s, 4); + assert_eq!(result, "hi..."); + } + + #[test] + fn test_truncate_preview_cjk() { + let s = "\u{4F60}\u{597D}\u{4E16}\u{754C}"; + let result = truncate_preview(s, 7); + assert_eq!(result, "\u{4F60}\u{597D}..."); + } + + #[test] + fn test_truncate_preview_zero_max_bytes() { + assert_eq!(truncate_preview("hello", 0), "..."); + } + + #[test] + fn test_truncate_preview_closes_tool_output_tag() { + let s = "\nSome very long content here\n"; + let result = truncate_preview(s, 60); + assert!(result.ends_with("")); + assert!(result.contains("...")); + } + + #[test] + fn test_truncate_preview_no_extra_close_when_intact() { + let s = "\nshort\n"; + let result = truncate_preview(s, 500); + assert_eq!(result, s); + assert_eq!(result.matches("").count(), 1); + } + + #[test] + fn test_truncate_preview_non_xml_unaffected() { + let s = "Just a plain long string that gets truncated"; + let result = truncate_preview(s, 10); + assert_eq!(result, "Just a pla..."); + assert!(!result.contains("")); + } +} diff --git a/src/agent/job_monitor.rs b/src/agent/job_monitor.rs index 02f5e3e2..e102dfbf 100644 --- a/src/agent/job_monitor.rs +++ b/src/agent/job_monitor.rs @@ -21,8 +21,8 @@ use tokio::task::JoinHandle; use uuid::Uuid; use crate::channels::IncomingMessage; -use crate::channels::web::types::SseEvent; use crate::context::{ContextManager, JobState}; +use ironclaw_common::AppEvent; /// Route context for forwarding job monitor events back to the user's channel. #[derive(Debug, Clone)] @@ -36,15 +36,15 @@ pub struct JobMonitorRoute { /// injects assistant messages into the agent loop. /// /// The monitor forwards: -/// - `SseEvent::JobMessage` (assistant role): injected as incoming messages so +/// - `AppEvent::JobMessage` (assistant role): injected as incoming messages so /// the main agent can read and relay to the user. -/// - `SseEvent::JobResult`: injected as a completion notice, then the task exits. +/// - `AppEvent::JobResult`: injected as a completion notice, then the task exits. /// /// Tool use/result and status events are intentionally skipped (too noisy for /// the main agent's context window). pub fn spawn_job_monitor( job_id: Uuid, - event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>, + event_rx: broadcast::Receiver<(Uuid, String, AppEvent)>, inject_tx: mpsc::Sender, route: JobMonitorRoute, ) -> JoinHandle<()> { @@ -56,7 +56,7 @@ pub fn spawn_job_monitor( /// jobs don't stay `InProgress` forever in the `ContextManager`. pub fn spawn_job_monitor_with_context( job_id: Uuid, - mut event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>, + mut event_rx: broadcast::Receiver<(Uuid, String, AppEvent)>, inject_tx: mpsc::Sender, route: JobMonitorRoute, context_manager: Option>, @@ -74,7 +74,7 @@ pub fn spawn_job_monitor_with_context( } match event { - SseEvent::JobMessage { role, content, .. } if role == "assistant" => { + AppEvent::JobMessage { role, content, .. } if role == "assistant" => { let mut msg = IncomingMessage::new( route.channel.clone(), route.user_id.clone(), @@ -92,7 +92,7 @@ pub fn spawn_job_monitor_with_context( break; } } - SseEvent::JobResult { status, .. } => { + AppEvent::JobResult { status, .. } => { // Transition in-memory state so the job frees its // max_jobs slot and query tools show the final state. if let Some(ref cm) = context_manager { @@ -162,7 +162,7 @@ pub fn spawn_job_monitor_with_context( /// inject messages into) but we still need to free the `max_jobs` slot. pub fn spawn_completion_watcher( job_id: Uuid, - mut event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>, + mut event_rx: broadcast::Receiver<(Uuid, String, AppEvent)>, context_manager: Arc, ) -> JoinHandle<()> { let short_id = job_id.to_string()[..8].to_string(); @@ -170,7 +170,7 @@ pub fn spawn_completion_watcher( tokio::spawn(async move { loop { match event_rx.recv().await { - Ok((ev_job_id, _user_id, SseEvent::JobResult { status, .. })) + Ok((ev_job_id, _user_id, AppEvent::JobResult { status, .. })) if ev_job_id == job_id => { let target = if status == "completed" { @@ -229,7 +229,7 @@ mod tests { #[tokio::test] async fn test_monitor_forwards_assistant_messages() { - let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let job_id = Uuid::new_v4(); @@ -240,7 +240,7 @@ mod tests { .send(( job_id, "test-user".to_string(), - SseEvent::JobMessage { + AppEvent::JobMessage { job_id: job_id.to_string(), role: "assistant".to_string(), content: "I found a bug".to_string(), @@ -262,7 +262,7 @@ mod tests { #[tokio::test] async fn test_monitor_ignores_other_jobs() { - let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let job_id = Uuid::new_v4(); @@ -274,7 +274,7 @@ mod tests { .send(( other_job_id, "test-user".to_string(), - SseEvent::JobMessage { + AppEvent::JobMessage { job_id: other_job_id.to_string(), role: "assistant".to_string(), content: "wrong job".to_string(), @@ -293,7 +293,7 @@ mod tests { #[tokio::test] async fn test_monitor_exits_on_job_result() { - let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let job_id = Uuid::new_v4(); @@ -304,7 +304,7 @@ mod tests { .send(( job_id, "test-user".to_string(), - SseEvent::JobResult { + AppEvent::JobResult { job_id: job_id.to_string(), status: "completed".to_string(), session_id: None, @@ -329,7 +329,7 @@ mod tests { #[tokio::test] async fn test_monitor_skips_tool_events() { - let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let job_id = Uuid::new_v4(); @@ -340,7 +340,7 @@ mod tests { .send(( job_id, "test-user".to_string(), - SseEvent::JobToolUse { + AppEvent::JobToolUse { job_id: job_id.to_string(), tool_name: "shell".to_string(), input: serde_json::json!({"command": "ls"}), @@ -353,7 +353,7 @@ mod tests { .send(( job_id, "test-user".to_string(), - SseEvent::JobMessage { + AppEvent::JobMessage { job_id: job_id.to_string(), role: "user".to_string(), content: "user prompt".to_string(), @@ -409,7 +409,7 @@ mod tests { .await .unwrap(); - let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let handle = spawn_job_monitor_with_context( @@ -425,7 +425,7 @@ mod tests { .send(( job_id, "test-user".to_string(), - SseEvent::JobResult { + AppEvent::JobResult { job_id: job_id.to_string(), status: "completed".to_string(), session_id: None, @@ -458,7 +458,7 @@ mod tests { .await .unwrap(); - let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let handle = spawn_job_monitor_with_context( @@ -474,7 +474,7 @@ mod tests { .send(( job_id, "test-user".to_string(), - SseEvent::JobResult { + AppEvent::JobResult { job_id: job_id.to_string(), status: "failed".to_string(), session_id: None, @@ -507,14 +507,14 @@ mod tests { .await .unwrap(); - let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let handle = spawn_completion_watcher(job_id, event_tx.subscribe(), Arc::clone(&cm)); event_tx .send(( job_id, "test-user".to_string(), - SseEvent::JobResult { + AppEvent::JobResult { job_id: job_id.to_string(), status: "completed".to_string(), session_id: None, diff --git a/src/agent/session.rs b/src/agent/session.rs index 45594922..7ec2023f 100644 --- a/src/agent/session.rs +++ b/src/agent/session.rs @@ -16,8 +16,8 @@ use chrono::{DateTime, TimeDelta, Utc}; use serde::{Deserialize, Serialize}; use uuid::Uuid; -use crate::channels::web::util::truncate_preview; use crate::llm::{ChatMessage, ToolCall, generate_tool_call_id}; +use ironclaw_common::truncate_preview; /// A session containing one or more threads. #[derive(Debug, Clone, Serialize, Deserialize)] diff --git a/src/agent/thread_ops.rs b/src/agent/thread_ops.rs index ddfd0c0f..b2820e7e 100644 --- a/src/agent/thread_ops.rs +++ b/src/agent/thread_ops.rs @@ -16,12 +16,12 @@ use crate::agent::dispatcher::{ }; use crate::agent::session::{MAX_PENDING_MESSAGES, PendingApproval, Session, ThreadState}; use crate::agent::submission::SubmissionResult; -use crate::channels::web::util::truncate_preview; use crate::channels::{IncomingMessage, StatusUpdate}; use crate::context::JobContext; use crate::error::Error; use crate::llm::{ChatMessage, ToolCall}; use crate::tools::redact_params; +use ironclaw_common::truncate_preview; const FORGED_THREAD_ID_ERROR: &str = "Invalid or unauthorized thread ID."; diff --git a/src/channels/web/handlers/chat.rs b/src/channels/web/handlers/chat.rs index 9753c015..de4b3155 100644 --- a/src/channels/web/handlers/chat.rs +++ b/src/channels/web/handlers/chat.rs @@ -175,7 +175,7 @@ pub async fn chat_auth_token_handler( if result.verification.is_some() { state.sse.broadcast_for_user( &user.user_id, - SseEvent::AuthRequired { + AppEvent::AuthRequired { extension_name: req.extension_name.clone(), instructions: Some(result.message), auth_url: None, @@ -187,7 +187,7 @@ pub async fn chat_auth_token_handler( state.sse.broadcast_for_user( &user.user_id, - SseEvent::AuthCompleted { + AppEvent::AuthCompleted { extension_name: req.extension_name.clone(), success: true, message: result.message, @@ -202,7 +202,7 @@ pub async fn chat_auth_token_handler( if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) { state.sse.broadcast_for_user( &user.user_id, - SseEvent::AuthRequired { + AppEvent::AuthRequired { extension_name: req.extension_name.clone(), instructions: Some(msg.clone()), auth_url: None, diff --git a/src/channels/web/mod.rs b/src/channels/web/mod.rs index a8b1ec41..6a97e8b8 100644 --- a/src/channels/web/mod.rs +++ b/src/channels/web/mod.rs @@ -58,7 +58,7 @@ use self::log_layer::{LogBroadcaster, LogLevelHandle}; use self::auth::MultiAuthState; use self::server::GatewayState; use self::sse::SseManager; -use self::types::SseEvent; +use self::types::AppEvent; /// Web gateway channel implementing the Channel trait. pub struct GatewayChannel { @@ -386,7 +386,7 @@ impl Channel for GatewayChannel { self.state.sse.broadcast_for_user( &msg.user_id, - SseEvent::Response { + AppEvent::Response { content: response.content, thread_id, }, @@ -405,11 +405,11 @@ impl Channel for GatewayChannel { .and_then(|v| v.as_str()) .map(String::from); let event = match status { - StatusUpdate::Thinking(msg) => SseEvent::Thinking { + StatusUpdate::Thinking(msg) => AppEvent::Thinking { message: msg, thread_id: thread_id.clone(), }, - StatusUpdate::ToolStarted { name } => SseEvent::ToolStarted { + StatusUpdate::ToolStarted { name } => AppEvent::ToolStarted { name, thread_id: thread_id.clone(), }, @@ -418,23 +418,23 @@ impl Channel for GatewayChannel { success, error, parameters, - } => SseEvent::ToolCompleted { + } => AppEvent::ToolCompleted { name, success, error, parameters, thread_id: thread_id.clone(), }, - StatusUpdate::ToolResult { name, preview } => SseEvent::ToolResult { + StatusUpdate::ToolResult { name, preview } => AppEvent::ToolResult { name, preview, thread_id: thread_id.clone(), }, - StatusUpdate::StreamChunk(content) => SseEvent::StreamChunk { + StatusUpdate::StreamChunk(content) => AppEvent::StreamChunk { content, thread_id: thread_id.clone(), }, - StatusUpdate::Status(msg) => SseEvent::Status { + StatusUpdate::Status(msg) => AppEvent::Status { message: msg, thread_id: thread_id.clone(), }, @@ -442,7 +442,7 @@ impl Channel for GatewayChannel { job_id, title, browse_url, - } => SseEvent::JobStarted { + } => AppEvent::JobStarted { job_id, title, browse_url, @@ -453,7 +453,7 @@ impl Channel for GatewayChannel { description, parameters, allow_always, - } => SseEvent::ApprovalNeeded { + } => AppEvent::ApprovalNeeded { request_id, tool_name, description, @@ -467,7 +467,7 @@ impl Channel for GatewayChannel { instructions, auth_url, setup_url, - } => SseEvent::AuthRequired { + } => AppEvent::AuthRequired { extension_name, instructions, auth_url, @@ -477,17 +477,17 @@ impl Channel for GatewayChannel { extension_name, success, message, - } => SseEvent::AuthCompleted { + } => AppEvent::AuthCompleted { extension_name, success, message, }, - StatusUpdate::ImageGenerated { data_url, path } => SseEvent::ImageGenerated { + StatusUpdate::ImageGenerated { data_url, path } => AppEvent::ImageGenerated { data_url, path, thread_id: thread_id.clone(), }, - StatusUpdate::Suggestions { suggestions } => SseEvent::Suggestions { + StatusUpdate::Suggestions { suggestions } => AppEvent::Suggestions { suggestions, thread_id, }, @@ -495,7 +495,7 @@ impl Channel for GatewayChannel { input_tokens, output_tokens, cost_usd, - } => SseEvent::TurnCost { + } => AppEvent::TurnCost { input_tokens, output_tokens, cost_usd, @@ -531,7 +531,7 @@ impl Channel for GatewayChannel { }; self.state.sse.broadcast_for_user( user_id, - SseEvent::Response { + AppEvent::Response { content: response.content, thread_id, }, diff --git a/src/channels/web/server.rs b/src/channels/web/server.rs index 31c2b296..5b092312 100644 --- a/src/channels/web/server.rs +++ b/src/channels/web/server.rs @@ -813,7 +813,7 @@ async fn oauth_callback_handler( if let Some(ref sse) = flow.sse_manager { sse.broadcast_for_user( &flow.user_id, - SseEvent::AuthCompleted { + AppEvent::AuthCompleted { extension_name: flow.extension_name.clone(), success: false, message: "OAuth flow expired. Please try again.".to_string(), @@ -951,11 +951,11 @@ async fn oauth_callback_handler( message }; - // Broadcast SSE event to notify the web UI + // Broadcast event to notify the web UI if let Some(ref sse) = flow.sse_manager { sse.broadcast_for_user( &flow.user_id, - SseEvent::AuthCompleted { + AppEvent::AuthCompleted { extension_name: flow.extension_name, success, message: final_message.clone(), @@ -1197,8 +1197,8 @@ async fn slack_relay_oauth_callback_handler( } }; - // Broadcast SSE event to notify the web UI - state.sse.broadcast(SseEvent::AuthCompleted { + // Broadcast event to notify the web UI + state.sse.broadcast(AppEvent::AuthCompleted { extension_name: DEFAULT_RELAY_NAME.to_string(), success, message: message.clone(), @@ -1471,7 +1471,7 @@ async fn chat_auth_token_handler( if result.verification.is_some() { state.sse.broadcast_for_user( &user.user_id, - SseEvent::AuthRequired { + AppEvent::AuthRequired { extension_name: req.extension_name.clone(), instructions: Some(result.message), auth_url: None, @@ -1484,7 +1484,7 @@ async fn chat_auth_token_handler( state.sse.broadcast_for_user( &user.user_id, - SseEvent::AuthCompleted { + AppEvent::AuthCompleted { extension_name: req.extension_name.clone(), success: true, message: result.message, @@ -1493,7 +1493,7 @@ async fn chat_auth_token_handler( } else { state.sse.broadcast_for_user( &user.user_id, - SseEvent::AuthCompleted { + AppEvent::AuthCompleted { extension_name: req.extension_name.clone(), success: false, message: result.message, @@ -1509,7 +1509,7 @@ async fn chat_auth_token_handler( if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) { state.sse.broadcast_for_user( &user.user_id, - SseEvent::AuthRequired { + AppEvent::AuthRequired { extension_name: req.extension_name.clone(), instructions: Some(msg.clone()), auth_url: None, @@ -2477,7 +2477,7 @@ async fn extensions_setup_submit_handler( // auth card or setup modal that was triggered by tool_auth/tool_activate. state.sse.broadcast_for_user( &user.user_id, - SseEvent::AuthCompleted { + AppEvent::AuthCompleted { extension_name: name.clone(), success: result.activated, message: resp.message.clone(), @@ -3169,7 +3169,7 @@ mod tests { Ok(Ok(scoped)) if matches!( scoped.event, - crate::channels::web::types::SseEvent::AuthRequired { .. } + crate::channels::web::types::AppEvent::AuthRequired { .. } ) => { panic!("verification responses should not emit auth_required SSE events") @@ -3451,7 +3451,7 @@ mod tests { assert_eq!(resp.status(), StatusCode::OK); match receiver.recv().await.expect("auth_completed event").event { - crate::channels::web::types::SseEvent::AuthCompleted { + crate::channels::web::types::AppEvent::AuthCompleted { extension_name, success, message, diff --git a/src/channels/web/sse.rs b/src/channels/web/sse.rs index 46841e19..e36cceab 100644 --- a/src/channels/web/sse.rs +++ b/src/channels/web/sse.rs @@ -11,7 +11,7 @@ use tokio::sync::broadcast; use tokio_stream::StreamExt; use tokio_stream::wrappers::BroadcastStream; -use crate::channels::web::types::SseEvent; +use crate::channels::web::types::AppEvent; /// Maximum number of concurrent SSE/WebSocket connections. /// Prevents resource exhaustion from connection flooding. @@ -25,7 +25,7 @@ const MAX_CONNECTIONS: u64 = 100; #[derive(Debug, Clone)] pub(crate) struct ScopedEvent { pub(crate) user_id: Option, - pub(crate) event: SseEvent, + pub(crate) event: AppEvent, } /// Manages SSE broadcast to all connected browser tabs. @@ -75,7 +75,7 @@ impl SseManager { } /// Broadcast an event to all connected clients (global/unscoped). - pub fn broadcast(&self, event: SseEvent) { + pub fn broadcast(&self, event: AppEvent) { let _ = self.tx.send(ScopedEvent { user_id: None, event, @@ -86,7 +86,7 @@ impl SseManager { /// /// Only subscribers for this user_id (or unscoped subscribers) will /// receive the event. - pub fn broadcast_for_user(&self, user_id: &str, event: SseEvent) { + pub fn broadcast_for_user(&self, user_id: &str, event: AppEvent) { let _ = self.tx.send(ScopedEvent { user_id: Some(user_id.to_string()), event, @@ -108,7 +108,7 @@ impl SseManager { pub fn subscribe_raw( &self, user_id: Option, - ) -> Option + Send + 'static + use<>> { + ) -> Option + Send + 'static + use<>> { // Atomically increment only if below the limit. This prevents // concurrent callers from overshooting max_connections. let counter = Arc::clone(&self.connection_count); @@ -186,30 +186,7 @@ impl SseManager { return None; } }; - let event_type = match &event { - SseEvent::Response { .. } => "response", - SseEvent::Thinking { .. } => "thinking", - SseEvent::ToolStarted { .. } => "tool_started", - SseEvent::ToolCompleted { .. } => "tool_completed", - SseEvent::ToolResult { .. } => "tool_result", - SseEvent::StreamChunk { .. } => "stream_chunk", - SseEvent::Status { .. } => "status", - SseEvent::ApprovalNeeded { .. } => "approval_needed", - SseEvent::AuthRequired { .. } => "auth_required", - SseEvent::AuthCompleted { .. } => "auth_completed", - SseEvent::Error { .. } => "error", - SseEvent::JobStarted { .. } => "job_started", - SseEvent::JobMessage { .. } => "job_message", - SseEvent::JobToolUse { .. } => "job_tool_use", - SseEvent::JobToolResult { .. } => "job_tool_result", - SseEvent::JobStatus { .. } => "job_status", - SseEvent::JobResult { .. } => "job_result", - SseEvent::Heartbeat => "heartbeat", - SseEvent::ImageGenerated { .. } => "image_generated", - SseEvent::Suggestions { .. } => "suggestions", - SseEvent::TurnCost { .. } => "turn_cost", - SseEvent::ExtensionStatus { .. } => "extension_status", - }; + let event_type = event.event_type(); Some(Ok(Event::default().event(event_type).data(data))) }); @@ -272,7 +249,7 @@ mod tests { fn test_broadcast_without_receivers() { let manager = SseManager::new(); // Should not panic even with no receivers - manager.broadcast(SseEvent::Heartbeat); + manager.broadcast(AppEvent::Heartbeat); } #[tokio::test] @@ -280,14 +257,14 @@ mod tests { let manager = SseManager::new(); let mut stream = Box::pin(manager.subscribe_raw(None).expect("should subscribe")); - manager.broadcast(SseEvent::Status { + manager.broadcast(AppEvent::Status { message: "test".to_string(), thread_id: None, }); let event = stream.next().await.unwrap(); match event { - SseEvent::Status { message, .. } => assert_eq!(message, "test"), + AppEvent::Status { message, .. } => assert_eq!(message, "test"), _ => panic!("unexpected event type"), } } @@ -299,14 +276,14 @@ mod tests { assert_eq!(manager.connection_count(), 1); - manager.broadcast(SseEvent::Thinking { + manager.broadcast(AppEvent::Thinking { message: "working".to_string(), thread_id: None, }); let event = stream.next().await.unwrap(); match event { - SseEvent::Thinking { message, .. } => assert_eq!(message, "working"), + AppEvent::Thinking { message, .. } => assert_eq!(message, "working"), _ => panic!("Expected Thinking event"), } } @@ -329,12 +306,12 @@ mod tests { let mut s2 = Box::pin(manager.subscribe_raw(None).expect("should subscribe")); assert_eq!(manager.connection_count(), 2); - manager.broadcast(SseEvent::Heartbeat); + manager.broadcast(AppEvent::Heartbeat); let e1 = s1.next().await.unwrap(); let e2 = s2.next().await.unwrap(); - assert!(matches!(e1, SseEvent::Heartbeat)); - assert!(matches!(e2, SseEvent::Heartbeat)); + assert!(matches!(e1, AppEvent::Heartbeat)); + assert!(matches!(e2, AppEvent::Heartbeat)); drop(s1); assert_eq!(manager.connection_count(), 1); @@ -373,25 +350,25 @@ mod tests { // Send event scoped to alice manager.broadcast_for_user( "alice", - SseEvent::Status { + AppEvent::Status { message: "alice only".to_string(), thread_id: None, }, ); // Send global event - manager.broadcast(SseEvent::Heartbeat); + manager.broadcast(AppEvent::Heartbeat); // Alice gets her scoped event let e = alice.next().await.unwrap(); - assert!(matches!(e, SseEvent::Status { .. })); + assert!(matches!(e, AppEvent::Status { .. })); // Alice also gets the global heartbeat let e = alice.next().await.unwrap(); - assert!(matches!(e, SseEvent::Heartbeat)); + assert!(matches!(e, AppEvent::Heartbeat)); // Bob only gets the global heartbeat (alice's event was filtered) let e = bob.next().await.unwrap(); // safety: test-only - assert!(matches!(e, SseEvent::Heartbeat)); // safety: test assertion + assert!(matches!(e, AppEvent::Heartbeat)); // safety: test assertion } } diff --git a/src/channels/web/types.rs b/src/channels/web/types.rs index 3ac4163c..fe18a824 100644 --- a/src/channels/web/types.rs +++ b/src/channels/web/types.rs @@ -114,165 +114,9 @@ pub struct ApprovalRequest { pub thread_id: Option, } -// --- SSE Event Types --- +// --- App Event (re-exported from ironclaw_common) --- -#[derive(Debug, Clone, Serialize)] -#[serde(tag = "type")] -pub enum SseEvent { - #[serde(rename = "response")] - Response { content: String, thread_id: String }, - #[serde(rename = "thinking")] - Thinking { - message: String, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - }, - #[serde(rename = "tool_started")] - ToolStarted { - name: String, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - }, - #[serde(rename = "tool_completed")] - ToolCompleted { - name: String, - success: bool, - #[serde(skip_serializing_if = "Option::is_none")] - error: Option, - #[serde(skip_serializing_if = "Option::is_none")] - parameters: Option, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - }, - #[serde(rename = "tool_result")] - ToolResult { - name: String, - preview: String, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - }, - #[serde(rename = "stream_chunk")] - StreamChunk { - content: String, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - }, - #[serde(rename = "status")] - Status { - message: String, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - }, - #[serde(rename = "job_started")] - JobStarted { - job_id: String, - title: String, - browse_url: String, - }, - #[serde(rename = "approval_needed")] - ApprovalNeeded { - request_id: String, - tool_name: String, - description: String, - parameters: String, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - /// Whether the "always" auto-approve option should be shown. - allow_always: bool, - }, - #[serde(rename = "auth_required")] - AuthRequired { - extension_name: String, - #[serde(skip_serializing_if = "Option::is_none")] - instructions: Option, - #[serde(skip_serializing_if = "Option::is_none")] - auth_url: Option, - #[serde(skip_serializing_if = "Option::is_none")] - setup_url: Option, - }, - #[serde(rename = "auth_completed")] - AuthCompleted { - extension_name: String, - success: bool, - message: String, - }, - #[serde(rename = "error")] - Error { - message: String, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - }, - #[serde(rename = "heartbeat")] - Heartbeat, - - // Sandbox job streaming events (worker + Claude Code bridge) - #[serde(rename = "job_message")] - JobMessage { - job_id: String, - role: String, - content: String, - }, - #[serde(rename = "job_tool_use")] - JobToolUse { - job_id: String, - tool_name: String, - input: serde_json::Value, - }, - #[serde(rename = "job_tool_result")] - JobToolResult { - job_id: String, - tool_name: String, - output: String, - }, - #[serde(rename = "job_status")] - JobStatus { job_id: String, message: String }, - #[serde(rename = "job_result")] - JobResult { - job_id: String, - status: String, - #[serde(skip_serializing_if = "Option::is_none")] - session_id: Option, - #[serde(skip_serializing_if = "Option::is_none")] - fallback_deliverable: Option, - }, - - /// An image was generated by a tool. - #[serde(rename = "image_generated")] - ImageGenerated { - data_url: String, - #[serde(skip_serializing_if = "Option::is_none")] - path: Option, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - }, - - /// Suggested follow-up messages for the user. - #[serde(rename = "suggestions")] - Suggestions { - suggestions: Vec, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - }, - - /// Per-turn token usage and cost summary. - #[serde(rename = "turn_cost")] - TurnCost { - input_tokens: u64, - output_tokens: u64, - cost_usd: String, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - }, - - /// Extension activation status change (WASM channels). - #[serde(rename = "extension_status")] - ExtensionStatus { - extension_name: String, - status: String, - #[serde(skip_serializing_if = "Option::is_none")] - message: Option, - }, -} +pub use ironclaw_common::AppEvent; // --- Memory --- @@ -784,32 +628,9 @@ pub enum WsServerMessage { } impl WsServerMessage { - /// Create a WsServerMessage from an SseEvent. - pub fn from_sse_event(event: &SseEvent) -> Self { - let event_type = match event { - SseEvent::Response { .. } => "response", - SseEvent::Thinking { .. } => "thinking", - SseEvent::ToolStarted { .. } => "tool_started", - SseEvent::ToolCompleted { .. } => "tool_completed", - SseEvent::ToolResult { .. } => "tool_result", - SseEvent::StreamChunk { .. } => "stream_chunk", - SseEvent::Status { .. } => "status", - SseEvent::JobStarted { .. } => "job_started", - SseEvent::ApprovalNeeded { .. } => "approval_needed", - SseEvent::AuthRequired { .. } => "auth_required", - SseEvent::AuthCompleted { .. } => "auth_completed", - SseEvent::Error { .. } => "error", - SseEvent::Heartbeat => "heartbeat", - SseEvent::JobMessage { .. } => "job_message", - SseEvent::JobToolUse { .. } => "job_tool_use", - SseEvent::JobToolResult { .. } => "job_tool_result", - SseEvent::JobStatus { .. } => "job_status", - SseEvent::JobResult { .. } => "job_result", - SseEvent::ImageGenerated { .. } => "image_generated", - SseEvent::Suggestions { .. } => "suggestions", - SseEvent::TurnCost { .. } => "turn_cost", - SseEvent::ExtensionStatus { .. } => "extension_status", - }; + /// Create a WsServerMessage from an AppEvent. + pub fn from_app_event(event: &AppEvent) -> Self { + let event_type = event.event_type(); let data = serde_json::to_value(event).unwrap_or(serde_json::Value::Null); WsServerMessage::Event { event_type: event_type.to_string(), @@ -1101,12 +922,12 @@ mod tests { } #[test] - fn test_ws_server_from_sse_response() { - let sse = SseEvent::Response { + fn test_ws_server_from_app_event_response() { + let event = AppEvent::Response { content: "hello".to_string(), thread_id: "t1".to_string(), }; - let ws = WsServerMessage::from_sse_event(&sse); + let ws = WsServerMessage::from_app_event(&event); match ws { WsServerMessage::Event { event_type, data } => { assert_eq!(event_type, "response"); @@ -1118,12 +939,12 @@ mod tests { } #[test] - fn test_ws_server_from_sse_thinking() { - let sse = SseEvent::Thinking { + fn test_ws_server_from_app_event_thinking() { + let event = AppEvent::Thinking { message: "reasoning...".to_string(), thread_id: None, }; - let ws = WsServerMessage::from_sse_event(&sse); + let ws = WsServerMessage::from_app_event(&event); match ws { WsServerMessage::Event { event_type, data } => { assert_eq!(event_type, "thinking"); @@ -1134,8 +955,8 @@ mod tests { } #[test] - fn test_ws_server_from_sse_approval_needed() { - let sse = SseEvent::ApprovalNeeded { + fn test_ws_server_from_app_event_approval_needed() { + let event = AppEvent::ApprovalNeeded { request_id: "r1".to_string(), tool_name: "shell".to_string(), description: "Run ls".to_string(), @@ -1143,7 +964,7 @@ mod tests { thread_id: Some("t1".to_string()), allow_always: true, }; - let ws = WsServerMessage::from_sse_event(&sse); + let ws = WsServerMessage::from_app_event(&event); match ws { WsServerMessage::Event { event_type, data } => { assert_eq!(event_type, "approval_needed"); @@ -1155,9 +976,9 @@ mod tests { } #[test] - fn test_ws_server_from_sse_heartbeat() { - let sse = SseEvent::Heartbeat; - let ws = WsServerMessage::from_sse_event(&sse); + fn test_ws_server_from_app_event_heartbeat() { + let event = AppEvent::Heartbeat; + let ws = WsServerMessage::from_app_event(&event); match ws { WsServerMessage::Event { event_type, .. } => { assert_eq!(event_type, "heartbeat"); @@ -1197,8 +1018,8 @@ mod tests { } #[test] - fn test_sse_auth_required_serialize() { - let event = SseEvent::AuthRequired { + fn test_app_event_auth_required_serialize() { + let event = AppEvent::AuthRequired { extension_name: "notion".to_string(), instructions: Some("Get your token from...".to_string()), auth_url: None, @@ -1214,8 +1035,8 @@ mod tests { } #[test] - fn test_sse_auth_completed_serialize() { - let event = SseEvent::AuthCompleted { + fn test_app_event_auth_completed_serialize() { + let event = AppEvent::AuthCompleted { extension_name: "notion".to_string(), success: true, message: "notion authenticated (3 tools loaded)".to_string(), @@ -1228,14 +1049,14 @@ mod tests { } #[test] - fn test_ws_server_from_sse_auth_required() { - let sse = SseEvent::AuthRequired { + fn test_ws_server_from_app_event_auth_required() { + let event = AppEvent::AuthRequired { extension_name: "openai".to_string(), instructions: Some("Enter API key".to_string()), auth_url: None, setup_url: None, }; - let ws = WsServerMessage::from_sse_event(&sse); + let ws = WsServerMessage::from_app_event(&event); match ws { WsServerMessage::Event { event_type, data } => { assert_eq!(event_type, "auth_required"); @@ -1246,13 +1067,13 @@ mod tests { } #[test] - fn test_ws_server_from_sse_auth_completed() { - let sse = SseEvent::AuthCompleted { + fn test_ws_server_from_app_event_auth_completed() { + let event = AppEvent::AuthCompleted { extension_name: "slack".to_string(), success: false, message: "Invalid token".to_string(), }; - let ws = WsServerMessage::from_sse_event(&sse); + let ws = WsServerMessage::from_app_event(&event); match ws { WsServerMessage::Event { event_type, data } => { assert_eq!(event_type, "auth_completed"); diff --git a/src/channels/web/util.rs b/src/channels/web/util.rs index 0debe6a9..ed70c5ce 100644 --- a/src/channels/web/util.rs +++ b/src/channels/web/util.rs @@ -2,29 +2,7 @@ use crate::channels::web::types::{ToolCallInfo, TurnInfo}; -/// Truncate a string to at most `max_bytes` bytes at a char boundary, appending "...". -/// -/// If the input is wrapped in `` and truncation -/// removes the closing tag, the tag is re-appended so downstream XML parsers -/// never see an unclosed element. -pub fn truncate_preview(s: &str, max_bytes: usize) -> String { - if s.len() <= max_bytes { - return s.to_string(); - } - // Walk backwards from max_bytes to find a valid char boundary - let mut end = max_bytes; - while end > 0 && !s.is_char_boundary(end) { - end -= 1; - } - let mut result = format!("{}...", &s[..end]); - - // Re-close if truncation cut through the closing tag. - if s.starts_with("") { - result.push_str("\n"); - } - - result -} +pub use ironclaw_common::truncate_preview; /// Build TurnInfo pairs from flat DB messages (user/tool_calls/assistant triples). /// @@ -118,88 +96,6 @@ mod tests { use super::*; use uuid::Uuid; - // ---- truncate_preview tests ---- - - #[test] - fn test_truncate_preview_short_string() { - assert_eq!(truncate_preview("hello", 10), "hello"); - } - - #[test] - fn test_truncate_preview_exact_boundary() { - assert_eq!(truncate_preview("hello", 5), "hello"); - } - - #[test] - fn test_truncate_preview_truncates_ascii() { - assert_eq!(truncate_preview("hello world", 5), "hello..."); - } - - #[test] - fn test_truncate_preview_empty_string() { - assert_eq!(truncate_preview("", 10), ""); - } - - #[test] - fn test_truncate_preview_multibyte_char_boundary() { - // '€' is 3 bytes (E2 82 AC). "a€b" = [61, E2, 82, AC, 62] = 5 bytes - // Truncating at max_bytes=3 should not split the euro sign. - let s = "a€b"; - let result = truncate_preview(s, 3); - // max_bytes=3 lands mid-€, so it walks back to byte 1 ("a") - assert_eq!(result, "a..."); - } - - #[test] - fn test_truncate_preview_emoji() { - // '🦀' is 4 bytes. "hi🦀" = 6 bytes - let s = "hi🦀"; - let result = truncate_preview(s, 4); - // max_bytes=4 lands mid-🦀, walks back to byte 2 ("hi") - assert_eq!(result, "hi..."); - } - - #[test] - fn test_truncate_preview_cjk() { - // CJK characters are 3 bytes each. "你好世界" = 12 bytes - let s = "你好世界"; - let result = truncate_preview(s, 7); - // max_bytes=7 lands mid-character (byte 7 is inside 世), walks back to 6 ("你好") - assert_eq!(result, "你好..."); - } - - #[test] - fn test_truncate_preview_zero_max_bytes() { - assert_eq!(truncate_preview("hello", 0), "..."); - } - - #[test] - fn test_truncate_preview_closes_tool_output_tag() { - let s = "\nSome very long content here\n"; - // Truncate so it cuts before the closing tag - let result = truncate_preview(s, 60); - assert!(result.ends_with("")); - assert!(result.contains("...")); - } - - #[test] - fn test_truncate_preview_no_extra_close_when_intact() { - let s = "\nshort\n"; - // The string is short enough not to be truncated - let result = truncate_preview(s, 500); - assert_eq!(result, s); - // Should not have a duplicate closing tag - assert_eq!(result.matches("").count(), 1); - } - - #[test] - fn test_truncate_preview_non_xml_unaffected() { - let s = "Just a plain long string that gets truncated"; - let result = truncate_preview(s, 10); - assert_eq!(result, "Just a pla..."); - assert!(!result.contains("")); - } - // ---- build_turns_from_db_messages tests ---- fn make_msg(role: &str, content: &str, offset_ms: i64) -> crate::history::ConversationMessage { diff --git a/src/channels/web/ws.rs b/src/channels/web/ws.rs index 9d4e919c..51beaafd 100644 --- a/src/channels/web/ws.rs +++ b/src/channels/web/ws.rs @@ -97,7 +97,7 @@ pub async fn handle_ws_connection( let msg = tokio::select! { event = event_stream.next() => { match event { - Some(sse_event) => WsServerMessage::from_sse_event(&sse_event), + Some(app_event) => WsServerMessage::from_app_event(&app_event), None => break, // Broadcast channel closed } } @@ -275,7 +275,7 @@ async fn handle_client_message( if result.verification.is_some() { state.sse.broadcast_for_user( user_id, - crate::channels::web::types::SseEvent::AuthRequired { + crate::channels::web::types::AppEvent::AuthRequired { extension_name: extension_name.clone(), instructions: Some(result.message), auth_url: None, @@ -286,7 +286,7 @@ async fn handle_client_message( crate::channels::web::server::clear_auth_mode(state, user_id).await; state.sse.broadcast_for_user( user_id, - crate::channels::web::types::SseEvent::AuthCompleted { + crate::channels::web::types::AppEvent::AuthCompleted { extension_name, success: true, message: result.message, @@ -299,7 +299,7 @@ async fn handle_client_message( if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) { state.sse.broadcast_for_user( user_id, - crate::channels::web::types::SseEvent::AuthRequired { + crate::channels::web::types::AppEvent::AuthRequired { extension_name: extension_name.clone(), instructions: Some(msg.clone()), auth_url: None, diff --git a/src/extensions/manager.rs b/src/extensions/manager.rs index 0f308352..90920767 100644 --- a/src/extensions/manager.rs +++ b/src/extensions/manager.rs @@ -1118,7 +1118,7 @@ impl ExtensionManager { /// Broadcast an extension status change to the web UI via SSE. async fn broadcast_extension_status(&self, name: &str, status: &str, message: Option<&str>) { if let Some(ref sse) = *self.sse_manager.read().await { - sse.broadcast(crate::channels::web::types::SseEvent::ExtensionStatus { + sse.broadcast(ironclaw_common::AppEvent::ExtensionStatus { extension_name: name.to_string(), status: status.to_string(), message: message.map(|m| m.to_string()), @@ -3288,7 +3288,7 @@ impl ExtensionManager { } .await; - // Broadcast SSE event + // Broadcast auth result event let (success, message) = match result { Ok(()) => (true, format!("{} authenticated successfully", display_name)), Err(ref e) => ( @@ -3314,7 +3314,7 @@ impl ExtensionManager { } if let Some(ref sse) = sse_manager { - sse.broadcast(crate::channels::web::types::SseEvent::AuthCompleted { + sse.broadcast(ironclaw_common::AppEvent::AuthCompleted { extension_name: ext_name, success, message, diff --git a/src/orchestrator/api.rs b/src/orchestrator/api.rs index 00f8a4da..37085a8b 100644 --- a/src/orchestrator/api.rs +++ b/src/orchestrator/api.rs @@ -14,7 +14,6 @@ use serde::{Deserialize, Serialize}; use tokio::sync::{Mutex, broadcast}; use uuid::Uuid; -use crate::channels::web::types::SseEvent; use crate::db::Database; use crate::llm::{CompletionRequest, LlmProvider, ToolCompletionRequest}; use crate::orchestrator::auth::{TokenStore, worker_auth_middleware}; @@ -25,6 +24,7 @@ use crate::worker::api::{ CompletionReport, CredentialResponse, JobDescription, ProxyCompletionRequest, ProxyCompletionResponse, ProxyToolCompletionRequest, ProxyToolCompletionResponse, StatusUpdate, }; +use ironclaw_common::AppEvent; /// A follow-up prompt queued for a Claude Code bridge. #[derive(Debug, Clone, Serialize, Deserialize)] @@ -41,7 +41,7 @@ pub struct OrchestratorState { pub token_store: TokenStore, /// Broadcast channel for job events (consumed by the web gateway SSE). /// Tuple: (job_id, user_id, event). - pub job_event_tx: Option>, + pub job_event_tx: Option>, /// Buffered follow-up prompts for sandbox jobs, keyed by job_id. pub prompt_queue: Arc>>>, /// Database handle for persisting job events. @@ -277,10 +277,10 @@ async fn job_event_handler( }); } - // Convert to SSE event and broadcast + // Convert to app event and broadcast let job_id_str = job_id.to_string(); - let sse_event = match payload.event_type.as_str() { - "message" => SseEvent::JobMessage { + let app_event = match payload.event_type.as_str() { + "message" => AppEvent::JobMessage { job_id: job_id_str, role: payload .data @@ -295,7 +295,7 @@ async fn job_event_handler( .unwrap_or("") .to_string(), }, - "tool_use" => SseEvent::JobToolUse { + "tool_use" => AppEvent::JobToolUse { job_id: job_id_str, tool_name: payload .data @@ -309,7 +309,7 @@ async fn job_event_handler( .cloned() .unwrap_or(serde_json::Value::Null), }, - "tool_result" => SseEvent::JobToolResult { + "tool_result" => AppEvent::JobToolResult { job_id: job_id_str, tool_name: payload .data @@ -324,7 +324,7 @@ async fn job_event_handler( .unwrap_or("") .to_string(), }, - "result" => SseEvent::JobResult { + "result" => AppEvent::JobResult { job_id: job_id_str, status: payload .data @@ -344,7 +344,7 @@ async fn job_event_handler( // gain context/memory tracking capabilities. fallback_deliverable: payload.data.get("fallback_deliverable").cloned(), }, - _ => SseEvent::JobStatus { + _ => AppEvent::JobStatus { job_id: job_id_str, message: payload .data @@ -390,9 +390,9 @@ async fn job_event_handler( }; if user_id.is_empty() { - let _ = tx.send((job_id, String::new(), sse_event)); + let _ = tx.send((job_id, String::new(), app_event)); } else { - let _ = tx.send((job_id, user_id, sse_event)); + let _ = tx.send((job_id, user_id, app_event)); } } @@ -817,7 +817,7 @@ mod tests { // No store configured, so user_id falls back to empty string. assert_eq!(recv_uid, ""); match event { - SseEvent::JobMessage { + AppEvent::JobMessage { job_id: jid, role, content, @@ -872,7 +872,7 @@ mod tests { let (_recv_id, _recv_uid, event) = rx.recv().await.unwrap(); match event { - SseEvent::JobToolUse { tool_name, .. } => { + AppEvent::JobToolUse { tool_name, .. } => { assert_eq!(tool_name, "shell"); } other => panic!("Expected JobToolUse, got {:?}", other), @@ -918,7 +918,7 @@ mod tests { let (_recv_id, _recv_uid, event) = rx.recv().await.unwrap(); // Unknown event types fall through to JobStatus - assert!(matches!(event, SseEvent::JobStatus { .. })); + assert!(matches!(event, AppEvent::JobStatus { .. })); } // -- Status update test -- diff --git a/src/orchestrator/mod.rs b/src/orchestrator/mod.rs index 896b5648..8d09dc53 100644 --- a/src/orchestrator/mod.rs +++ b/src/orchestrator/mod.rs @@ -46,10 +46,10 @@ use std::sync::Arc; use tokio::sync::{Mutex, broadcast}; use uuid::Uuid; -use crate::channels::web::types::SseEvent; use crate::db::Database; use crate::llm::LlmProvider; use crate::secrets::SecretsStore; +use ironclaw_common::AppEvent; /// Resolve the orchestrator port from the `ORCHESTRATOR_PORT` environment /// variable, falling back to 50051. @@ -63,7 +63,7 @@ fn resolve_orchestrator_port() -> u16 { /// Result of orchestrator setup, containing all handles needed by the agent. pub struct OrchestratorSetup { pub container_job_manager: Option>, - pub job_event_tx: Option>, + pub job_event_tx: Option>, pub prompt_queue: Arc>>>, pub docker_status: crate::sandbox::DockerStatus, } diff --git a/src/tools/builtin/job.rs b/src/tools/builtin/job.rs index 86d7e44d..4c711e69 100644 --- a/src/tools/builtin/job.rs +++ b/src/tools/builtin/job.rs @@ -17,7 +17,6 @@ use uuid::Uuid; use crate::bootstrap::ironclaw_base_dir; use crate::channels::IncomingMessage; -use crate::channels::web::types::SseEvent; use crate::context::{ContextManager, JobContext, JobState}; use crate::db::Database; use crate::history::SandboxJobRecord; @@ -25,6 +24,7 @@ use crate::orchestrator::auth::CredentialGrant; use crate::orchestrator::job_manager::{ContainerJobManager, JobMode}; use crate::secrets::SecretsStore; use crate::tools::tool::{ApprovalRequirement, Tool, ToolError, ToolOutput, require_str}; +use ironclaw_common::AppEvent; /// Lazy scheduler reference, filled after Agent::new creates the Scheduler. /// @@ -85,7 +85,7 @@ pub struct CreateJobTool { job_manager: Option>, store: Option>, /// Broadcast sender for job events (used to subscribe a monitor). - event_tx: Option>, + event_tx: Option>, /// Injection channel for pushing messages into the agent loop. inject_tx: Option>, /// Encrypted secrets store for validating credential grants. @@ -120,7 +120,7 @@ impl CreateJobTool { /// monitor that forwards Claude Code output to the main agent loop. pub fn with_monitor_deps( mut self, - event_tx: tokio::sync::broadcast::Sender<(Uuid, String, SseEvent)>, + event_tx: tokio::sync::broadcast::Sender<(Uuid, String, AppEvent)>, inject_tx: tokio::sync::mpsc::Sender, ) -> Self { self.event_tx = Some(event_tx); diff --git a/src/tools/registry.rs b/src/tools/registry.rs index bc3be144..8c08633b 100644 --- a/src/tools/registry.rs +++ b/src/tools/registry.rs @@ -383,11 +383,7 @@ impl ToolRegistry { job_manager: Option>, store: Option>, job_event_tx: Option< - tokio::sync::broadcast::Sender<( - uuid::Uuid, - String, - crate::channels::web::types::SseEvent, - )>, + tokio::sync::broadcast::Sender<(uuid::Uuid, String, ironclaw_common::AppEvent)>, >, inject_tx: Option>, prompt_queue: Option, diff --git a/src/worker/job.rs b/src/worker/job.rs index b2e3f7e6..ed261039 100644 --- a/src/worker/job.rs +++ b/src/worker/job.rs @@ -18,7 +18,6 @@ use crate::agent::agentic_loop::{ }; use crate::agent::scheduler::WorkerMessage; use crate::agent::task::TaskOutput; -use crate::channels::web::types::SseEvent; use crate::context::{ContextManager, JobState}; use crate::db::Database; use crate::error::Error; @@ -33,6 +32,7 @@ use crate::tools::rate_limiter::RateLimitResult; use crate::tools::{ ApprovalContext, ToolRegistry, autonomous_unavailable_error, prepare_tool_params, redact_params, }; +use ironclaw_common::AppEvent; /// Shared dependencies for worker execution. /// @@ -48,7 +48,7 @@ pub struct WorkerDeps { pub hooks: Arc, pub timeout: Duration, pub use_planning: bool, - /// SSE manager for live job event streaming to the web gateway. + /// Broadcast sender for live job event streaming to the web gateway. pub sse_tx: Option>, /// Approval context for tool execution. When `None`, all non-`Never` tools are /// blocked (legacy behavior). When `Some`, the context determines which tools @@ -141,7 +141,7 @@ impl Worker { if let Some(ref sse) = self.deps.sse_tx { let job_id_str = job_id.to_string(); let event = match event_type { - "message" => Some(SseEvent::JobMessage { + "message" => Some(AppEvent::JobMessage { job_id: job_id_str, role: data .get("role") @@ -154,7 +154,7 @@ impl Worker { .unwrap_or("") .to_string(), }), - "tool_use" => Some(SseEvent::JobToolUse { + "tool_use" => Some(AppEvent::JobToolUse { job_id: job_id_str, tool_name: data .get("tool_name") @@ -166,7 +166,7 @@ impl Worker { .cloned() .unwrap_or(serde_json::Value::Null), }), - "tool_result" => Some(SseEvent::JobToolResult { + "tool_result" => Some(AppEvent::JobToolResult { job_id: job_id_str, tool_name: data .get("tool_name") @@ -179,7 +179,7 @@ impl Worker { .unwrap_or("") .to_string(), }), - "status" => Some(SseEvent::JobStatus { + "status" => Some(AppEvent::JobStatus { job_id: job_id_str, message: data .get("message") @@ -187,7 +187,7 @@ impl Worker { .unwrap_or("") .to_string(), }), - "result" => Some(SseEvent::JobResult { + "result" => Some(AppEvent::JobResult { job_id: job_id_str, status: data .get("status") diff --git a/tests/multi_tenant_integration.rs b/tests/multi_tenant_integration.rs index f2529866..227fa721 100644 --- a/tests/multi_tenant_integration.rs +++ b/tests/multi_tenant_integration.rs @@ -307,7 +307,7 @@ fn per_user_rate_limiter_single_user_mode() { #[tokio::test] async fn sse_scoped_event_only_delivered_to_target_user() { - use ironclaw::channels::web::types::SseEvent; + use ironclaw_common::AppEvent; use tokio_stream::StreamExt; let manager = SseManager::new(); @@ -325,34 +325,34 @@ async fn sse_scoped_event_only_delivered_to_target_user() { // Send event scoped to alice manager.broadcast_for_user( ALICE_USER_ID, - SseEvent::Status { + AppEvent::Status { message: "alice's event".to_string(), thread_id: None, }, ); // Send global heartbeat (both should get it) - manager.broadcast(SseEvent::Heartbeat); + manager.broadcast(AppEvent::Heartbeat); // Alice gets her scoped event first let e = alice_stream.next().await.unwrap(); match &e { - SseEvent::Status { message, .. } => assert_eq!(message, "alice's event"), + AppEvent::Status { message, .. } => assert_eq!(message, "alice's event"), _ => panic!("Expected Status, got {:?}", e), } // Alice also gets heartbeat let e = alice_stream.next().await.unwrap(); - assert!(matches!(e, SseEvent::Heartbeat)); + assert!(matches!(e, AppEvent::Heartbeat)); // Bob only gets the heartbeat (alice's event was filtered) let e = bob_stream.next().await.unwrap(); - assert!(matches!(e, SseEvent::Heartbeat)); + assert!(matches!(e, AppEvent::Heartbeat)); } #[tokio::test] async fn sse_global_event_delivered_to_all_users() { - use ironclaw::channels::web::types::SseEvent; + use ironclaw_common::AppEvent; use tokio_stream::StreamExt; let manager = SseManager::new(); @@ -367,7 +367,7 @@ async fn sse_global_event_delivered_to_all_users() { .expect("subscribe"), ); - manager.broadcast(SseEvent::Status { + manager.broadcast(AppEvent::Status { message: "global announcement".to_string(), thread_id: None, }); @@ -375,7 +375,7 @@ async fn sse_global_event_delivered_to_all_users() { let ea = alice.next().await.unwrap(); let eb = bob.next().await.unwrap(); match (&ea, &eb) { - (SseEvent::Status { message: a, .. }, SseEvent::Status { message: b, .. }) => { + (AppEvent::Status { message: a, .. }, AppEvent::Status { message: b, .. }) => { assert_eq!(a, "global announcement"); assert_eq!(b, "global announcement"); } @@ -385,7 +385,7 @@ async fn sse_global_event_delivered_to_all_users() { #[tokio::test] async fn sse_user_b_event_not_visible_to_user_a() { - use ironclaw::channels::web::types::SseEvent; + use ironclaw_common::AppEvent; use tokio_stream::StreamExt; let manager = SseManager::new(); @@ -398,19 +398,19 @@ async fn sse_user_b_event_not_visible_to_user_a() { // Send event for bob only manager.broadcast_for_user( BOB_USER_ID, - SseEvent::Response { + AppEvent::Response { content: "bob's secret".to_string(), thread_id: "t1".to_string(), }, ); // Send heartbeat so alice has something to receive - manager.broadcast(SseEvent::Heartbeat); + manager.broadcast(AppEvent::Heartbeat); // Alice should only get heartbeat, not bob's response let e = alice.next().await.unwrap(); assert!( - matches!(e, SseEvent::Heartbeat), + matches!(e, AppEvent::Heartbeat), "Expected Heartbeat, got {:?}", e ); @@ -418,7 +418,7 @@ async fn sse_user_b_event_not_visible_to_user_a() { #[tokio::test] async fn sse_unscoped_subscriber_receives_all_events() { - use ironclaw::channels::web::types::SseEvent; + use ironclaw_common::AppEvent; use tokio_stream::StreamExt; let manager = SseManager::new(); @@ -427,19 +427,19 @@ async fn sse_unscoped_subscriber_receives_all_events() { manager.broadcast_for_user( ALICE_USER_ID, - SseEvent::Status { + AppEvent::Status { message: "alice only".to_string(), thread_id: None, }, ); manager.broadcast_for_user( BOB_USER_ID, - SseEvent::Status { + AppEvent::Status { message: "bob only".to_string(), thread_id: None, }, ); - manager.broadcast(SseEvent::Heartbeat); + manager.broadcast(AppEvent::Heartbeat); // Unscoped subscriber gets ALL three events let e1 = stream.next().await.unwrap(); @@ -447,14 +447,14 @@ async fn sse_unscoped_subscriber_receives_all_events() { let e3 = stream.next().await.unwrap(); match &e1 { - SseEvent::Status { message, .. } => assert_eq!(message, "alice only"), + AppEvent::Status { message, .. } => assert_eq!(message, "alice only"), _ => panic!("Expected alice's Status"), } match &e2 { - SseEvent::Status { message, .. } => assert_eq!(message, "bob only"), + AppEvent::Status { message, .. } => assert_eq!(message, "bob only"), _ => panic!("Expected bob's Status"), } - assert!(matches!(e3, SseEvent::Heartbeat)); + assert!(matches!(e3, AppEvent::Heartbeat)); } // =========================================================================== @@ -881,7 +881,7 @@ async fn full_server_jobs_endpoint_rejected_without_auth() { #[tokio::test] async fn full_server_ws_multi_user_event_isolation() { use futures::StreamExt; - use ironclaw::channels::web::types::SseEvent; + use ironclaw_common::AppEvent; use tokio_tungstenite::tungstenite::Message; use tokio_tungstenite::tungstenite::client::IntoClientRequest; @@ -914,14 +914,14 @@ async fn full_server_ws_multi_user_event_isolation() { // Broadcast an event scoped to Alice only state.sse.broadcast_for_user( ALICE_USER_ID, - SseEvent::Status { + AppEvent::Status { message: "alice-only-event".to_string(), thread_id: None, }, ); // Broadcast a global heartbeat so Bob has something to receive - state.sse.broadcast(SseEvent::Heartbeat); + state.sse.broadcast(AppEvent::Heartbeat); // Alice should get her scoped event let alice_msg = tokio::time::timeout(Duration::from_secs(2), alice_ws.next()) diff --git a/tests/ws_gateway_integration.rs b/tests/ws_gateway_integration.rs index a6db5af7..0ec5c929 100644 --- a/tests/ws_gateway_integration.rs +++ b/tests/ws_gateway_integration.rs @@ -5,7 +5,7 @@ //! - WebSocket upgrade with auth //! - Ping/pong //! - Client message → agent msg_tx -//! - Broadcast SSE event → WebSocket client +//! - Broadcast AppEvent → WebSocket client //! - Connection tracking (counter increment/decrement) //! - Gateway status endpoint @@ -22,8 +22,8 @@ use tokio_tungstenite::tungstenite::client::IntoClientRequest; use ironclaw::channels::IncomingMessage; use ironclaw::channels::web::server::{GatewayState, start_server}; use ironclaw::channels::web::sse::SseManager; -use ironclaw::channels::web::types::SseEvent; use ironclaw::channels::web::ws::WsConnectionTracker; +use ironclaw_common::AppEvent; const AUTH_TOKEN: &str = "test-token-12345"; const TIMEOUT: Duration = Duration::from_secs(5); @@ -164,8 +164,8 @@ async fn test_ws_broadcast_event_received() { // Give the connection a moment to fully establish tokio::time::sleep(Duration::from_millis(50)).await; - // Broadcast an SSE event (simulates agent sending a response) - state.sse.broadcast(SseEvent::Response { + // Broadcast an event (simulates agent sending a response) + state.sse.broadcast(AppEvent::Response { content: "agent says hi".to_string(), thread_id: "t1".to_string(), }); @@ -186,7 +186,7 @@ async fn test_ws_thinking_event() { let mut ws = connect_ws(addr).await; tokio::time::sleep(Duration::from_millis(50)).await; - state.sse.broadcast(SseEvent::Thinking { + state.sse.broadcast(AppEvent::Thinking { message: "analyzing...".to_string(), thread_id: None, }); @@ -311,22 +311,22 @@ async fn test_ws_multiple_events_in_sequence() { tokio::time::sleep(Duration::from_millis(50)).await; // Broadcast multiple events rapidly - state.sse.broadcast(SseEvent::Thinking { + state.sse.broadcast(AppEvent::Thinking { message: "step 1".to_string(), thread_id: None, }); - state.sse.broadcast(SseEvent::ToolStarted { + state.sse.broadcast(AppEvent::ToolStarted { name: "shell".to_string(), thread_id: None, }); - state.sse.broadcast(SseEvent::ToolCompleted { + state.sse.broadcast(AppEvent::ToolCompleted { name: "shell".to_string(), success: true, error: None, parameters: None, thread_id: None, }); - state.sse.broadcast(SseEvent::Response { + state.sse.broadcast(AppEvent::Response { content: "done".to_string(), thread_id: "t1".to_string(), }); From 6daa2f155f2683cf93669cac5844b6d85400b7a5 Mon Sep 17 00:00:00 2001 From: Jacob Lasky Date: Wed, 25 Mar 2026 03:31:44 -0400 Subject: [PATCH 03/11] fix: ensure LLM calls always end with user message (closes #763) (#1259) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix: ensure LLM calls always end with user message (closes #763) Claude 4.6 models (claude-sonnet-4-6, claude-opus-4-6) no longer support assistant message prefill — any LLM call where the conversation ends on an assistant message is rejected with HTTP 400 "This model does not support assistant message prefill". The same root cause also triggers NEAR AI's "No user query found in messages" 400 error for the routine engine path. Two fixes: 1. src/worker/container.rs — before_llm_call() After poll_and_inject_prompt(), if no user follow-up arrived and handle_text_response() left an assistant message at the end of the conversation, inject a sentinel "Continue." user message before the next LLM call. 2. src/agent/routine_engine.rs — execute_lightweight_with_tools() Before the force_text final completion call, ensure messages end with a user-role message. Tool result messages (Role::Tool) satisfy Anthropic but not NEAR AI; assistant messages satisfy neither. Also updates the worker system prompt to instruct the agent to include the phrase "The job is complete" in its final message, so the agentic loop can detect termination reliably. Tested with claude-sonnet-4-6 and claude-opus-4-6. Workaround: ANTHROPIC_MODEL=claude-sonnet-4-20250514 (still supports prefill). * fix: broaden sentinel guard to any non-user message (per review) Gemini suggested the Role::Assistant check in before_llm_call() is too specific. Changed to !Role::User to match the routine_engine.rs fix and cover tool results too. * fix: address zmanian review — JobDelegate sentinel, shared helper, NearAI complete() flattening - Extract ensure_ends_with_user_message() to src/util.rs with 4 unit tests (empty list, after assistant, after tool result, no-op when already user) - Add sentinel guard to JobDelegate::before_llm_call() in src/worker/job.rs so scheduler jobs (CreateJob / /job path) no longer hit Claude 4.6 / NEAR AI 400s - Replace inline guards in ContainerDelegate and routine_engine.rs with the shared helper — all 3 call sites now use one implementation - Fix complete() in nearai_chat.rs to apply flatten_tool_messages when flatten_tool_messages=true — previously only complete_with_tools() flattened, so force_text paths could still send role:"tool" messages to NEAR AI - Update stale comment in container.rs: "assistant message" → "non-user message" - Add flatten tests in nearai_chat.rs covering the complete() path Co-Authored-By: Claude Sonnet 4.6 * ci: fix fmt and tar advisory --------- Co-authored-by: Jacob Lasky Co-authored-by: Claude Sonnet 4.6 Co-authored-by: Illia Polosukhin Co-authored-by: firat.sertgoz --- src/agent/routine_engine.rs | 5 ++- src/llm/nearai_chat.rs | 70 +++++++++++++++++++++++++++++++++++-- src/util.rs | 52 ++++++++++++++++++++++++++- src/worker/container.rs | 6 +++- src/worker/job.rs | 5 +++ 5 files changed, 133 insertions(+), 5 deletions(-) diff --git a/src/agent/routine_engine.rs b/src/agent/routine_engine.rs index 39acb83d..9c55903f 100644 --- a/src/agent/routine_engine.rs +++ b/src/agent/routine_engine.rs @@ -1541,7 +1541,10 @@ async fn execute_lightweight_with_tools( let force_text = iteration >= max_iterations; if force_text { - // Final iteration: no tools, just get text response + // Final iteration: no tools, just get text response. + // Claude 4.6 rejects assistant prefill; NEAR AI rejects any non-user-ending + // conversation. Ensure the last message is user-role. + crate::util::ensure_ends_with_user_message(&mut messages); let request = CompletionRequest::new(messages) .with_max_tokens(effective_max_tokens) .with_temperature(0.3); diff --git a/src/llm/nearai_chat.rs b/src/llm/nearai_chat.rs index acbff6ad..5372d76d 100644 --- a/src/llm/nearai_chat.rs +++ b/src/llm/nearai_chat.rs @@ -463,8 +463,15 @@ impl LlmProvider for NearAiChatProvider { let model = req.model.unwrap_or_else(|| self.active_model_name()); let mut raw_messages = req.messages; crate::llm::provider::sanitize_tool_messages(&mut raw_messages); - let messages: Vec = - raw_messages.into_iter().map(|m| m.into()).collect(); + let raw: Vec = raw_messages.into_iter().map(|m| m.into()).collect(); + + // NEAR AI rejects `role:"tool"` messages even on text-only completion paths. + // Apply the same flattening used by complete_with_tools(). + let messages = if self.flatten_tool_messages { + flatten_tool_messages(raw) + } else { + raw + }; let request = ChatCompletionRequest { model, @@ -2193,6 +2200,65 @@ mod tests { assert_eq!(deserialized.function.arguments, r#"{"city":"London"}"#); } + // -- flatten_tool_messages in complete() path ---------------------------- + + #[test] + fn test_flatten_applied_on_text_only_path() { + // Verify that flatten_tool_messages converts tool-role messages to user + // messages (mirrors the complete_with_tools path). + let messages = vec![ + ChatCompletionMessage { + role: "user".to_string(), + content: Some(MessageContent::Text("run it".to_string())), + tool_call_id: None, + name: None, + tool_calls: None, + }, + ChatCompletionMessage { + role: "tool".to_string(), + content: Some(MessageContent::Text("ok".to_string())), + tool_call_id: Some("call_1".to_string()), + name: Some("run_cmd".to_string()), + tool_calls: None, + }, + ]; + let flattened = flatten_tool_messages(messages); + assert_eq!(flattened.len(), 2); + assert_eq!(flattened[1].role, "user"); + let text = flattened[1] + .content + .as_ref() + .and_then(|c| c.as_text()) + .unwrap(); + assert!(text.contains("run_cmd"), "should reference tool name"); + assert!(text.contains("ok"), "should include tool result"); + } + + #[test] + fn test_no_flatten_when_no_tool_messages() { + // When there are no tool-role messages, flatten_tool_messages is a no-op. + let messages = vec![ + ChatCompletionMessage { + role: "user".to_string(), + content: Some(MessageContent::Text("hi".to_string())), + tool_call_id: None, + name: None, + tool_calls: None, + }, + ChatCompletionMessage { + role: "assistant".to_string(), + content: Some(MessageContent::Text("hello".to_string())), + tool_call_id: None, + name: None, + tool_calls: None, + }, + ]; + let result = flatten_tool_messages(messages); + // No tool messages → unchanged roles + assert_eq!(result[0].role, "user"); + assert_eq!(result[1].role, "assistant"); + } + // -- api_url edge cases --------------------------------------------------- #[test] diff --git a/src/util.rs b/src/util.rs index 866f623c..a76f3b27 100644 --- a/src/util.rs +++ b/src/util.rs @@ -1,5 +1,7 @@ //! Shared utility functions used across the codebase. +use crate::llm::{ChatMessage, Role}; + /// Find the largest valid UTF-8 char boundary at or before `pos`. /// /// Polyfill for `str::floor_char_boundary` (nightly-only). Use when @@ -16,6 +18,17 @@ pub fn floor_char_boundary(s: &str, pos: usize) -> usize { i } +/// Ensure the last message in `messages` is a user-role message. +/// +/// NEAR AI rejects conversations that don't end with a user message; +/// Claude 4.6 rejects assistant prefill. Call this before any LLM +/// completion request to satisfy both requirements. +pub fn ensure_ends_with_user_message(messages: &mut Vec) { + if !matches!(messages.last(), Some(m) if m.role == Role::User) { + messages.push(ChatMessage::user("Continue.")); + } +} + /// Check if an LLM response explicitly signals that a job/task is complete. /// /// Uses phrase-level matching to avoid false positives from bare words like @@ -72,7 +85,8 @@ pub fn llm_signals_completion(response: &str) -> bool { #[cfg(test)] mod tests { - use crate::util::{floor_char_boundary, llm_signals_completion}; + use crate::llm::ChatMessage; + use crate::util::{ensure_ends_with_user_message, floor_char_boundary, llm_signals_completion}; // ── floor_char_boundary ── @@ -103,6 +117,42 @@ mod tests { assert_eq!(floor_char_boundary("", 5), 0); } + // ── ensure_ends_with_user_message ── + + #[test] + fn ensure_user_message_injects_when_empty() { + let mut msgs: Vec = vec![]; + ensure_ends_with_user_message(&mut msgs); + assert_eq!(msgs.len(), 1); + assert_eq!(msgs[0].role, crate::llm::Role::User); + } + + #[test] + fn ensure_user_message_injects_after_assistant() { + let mut msgs = vec![ChatMessage::user("hi"), ChatMessage::assistant("hello")]; + ensure_ends_with_user_message(&mut msgs); + assert_eq!(msgs.len(), 3); + assert_eq!(msgs[2].role, crate::llm::Role::User); + } + + #[test] + fn ensure_user_message_injects_after_tool_result() { + let mut msgs = vec![ + ChatMessage::user("run tool"), + ChatMessage::tool_result("call_1", "my_tool", "result"), + ]; + ensure_ends_with_user_message(&mut msgs); + assert_eq!(msgs.len(), 3); + assert_eq!(msgs[2].role, crate::llm::Role::User); + } + + #[test] + fn ensure_user_message_no_op_when_already_user() { + let mut msgs = vec![ChatMessage::user("hello")]; + ensure_ends_with_user_message(&mut msgs); + assert_eq!(msgs.len(), 1); + } + // ── llm_signals_completion ── #[test] diff --git a/src/worker/container.rs b/src/worker/container.rs index e0933975..5d8e03b5 100644 --- a/src/worker/container.rs +++ b/src/worker/container.rs @@ -151,7 +151,7 @@ Job: {} Description: {} You have tools for shell commands, file operations, and code editing. -Work independently to complete this job. Report when done."#, +Work independently to complete this job. When finished, your final message MUST include the phrase "The job is complete" to signal termination."#, job.title, job.description ))); @@ -373,6 +373,10 @@ impl LoopDelegate for ContainerDelegate { // Poll for follow-up prompts from the user self.poll_and_inject_prompt(reason_ctx).await; + // Claude 4.6 rejects assistant prefill; NEAR AI rejects any non-user-ending + // conversation. Ensure the last message is user-role before calling the LLM. + crate::util::ensure_ends_with_user_message(&mut reason_ctx.messages); + // Refresh tools (in case WASM tools were built) reason_ctx.available_tools = self.tools.tool_definitions().await; diff --git a/src/worker/job.rs b/src/worker/job.rs index ed261039..9d5794ca 100644 --- a/src/worker/job.rs +++ b/src/worker/job.rs @@ -1232,6 +1232,11 @@ impl<'a> LoopDelegate for JobDelegate<'a> { ) -> Option { // Refresh tool definitions so newly built tools become visible reason_ctx.available_tools = self.worker.tools().tool_definitions().await; + + // Claude 4.6 rejects assistant prefill; NEAR AI rejects any non-user-ending + // conversation. Ensure the last message is user-role before calling the LLM. + crate::util::ensure_ends_with_user_message(&mut reason_ctx.messages); + None } From 41ed0a0f9814d754c17df80c14d263ae10e09b45 Mon Sep 17 00:00:00 2001 From: Illia Polosukhin Date: Wed, 25 Mar 2026 08:35:41 -0700 Subject: [PATCH 04/11] 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 05/11] 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 06/11] 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 07/11] 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 08/11] 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 09/11] =?UTF-8?q?feat:=20complete=20multi-tenant=20isolati?= =?UTF-8?q?on=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 10/11] 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 11/11] 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