From dea789cca9853ee814ae05c565de4e84684801b5 Mon Sep 17 00:00:00 2001 From: Henry Park Date: Mon, 23 Mar 2026 11:01:26 -0700 Subject: [PATCH 01/20] Default new lightweight routines to tools-enabled (#1573) * Default new lightweight routines to tools-enabled * Fix fmt and clippy on lightweight routine PR * Use grouped execution field in routine no-tools fixture * Align CLI routine defaults with tools-enabled lightweight mode --- src/agent/routine_engine.rs | 1 + src/cli/routines.rs | 49 ++++++++++- src/tools/builtin/routine.rs | 87 +++++++++++++++++-- tests/e2e_builtin_tool_coverage.rs | 53 +++++++++-- .../tools/routine_manual_create_no_tools.json | 39 +++++++++ 5 files changed, 213 insertions(+), 16 deletions(-) create mode 100644 tests/fixtures/llm_traces/tools/routine_manual_create_no_tools.json diff --git a/src/agent/routine_engine.rs b/src/agent/routine_engine.rs index de2879b4..9b554582 100644 --- a/src/agent/routine_engine.rs +++ b/src/agent/routine_engine.rs @@ -1440,6 +1440,7 @@ fn handle_text_response( /// This is a simplified version of the full dispatcher loop: /// - Max 3-5 iterations (configurable) /// - Sequential tool execution (not parallel) +/// - Uses the owner's live autonomous tool scope when lightweight tools are enabled /// - Auto-approval of non-Always tools /// - No hooks or approval dialogs async fn execute_lightweight_with_tools( diff --git a/src/cli/routines.rs b/src/cli/routines.rs index dd8a2fa3..ebef8839 100644 --- a/src/cli/routines.rs +++ b/src/cli/routines.rs @@ -340,8 +340,8 @@ async fn create( prompt: prompt.to_string(), context_paths: Vec::new(), max_tokens: 4096, - use_tools: false, - max_tool_rounds: 0, + use_tools: true, + max_tool_rounds: 3, }, guardrails: RoutineGuardrails { cooldown: std::time::Duration::from_secs(cooldown_secs), @@ -685,6 +685,7 @@ fn truncate(s: &str, max_chars: usize) -> String { #[cfg(test)] mod tests { use super::*; + use crate::agent::routine::RoutineAction; #[test] fn format_relative_future() { @@ -743,4 +744,48 @@ mod tests { assert!(notify.on_failure); // safety: test-only assertion assert!(!notify.on_success); // safety: test-only assertion } + + #[cfg(feature = "libsql")] + #[tokio::test] + async fn cli_create_defaults_lightweight_routines_to_tools_enabled() { + let harness = crate::testing::TestHarnessBuilder::new().build().await; + let db = harness.db.clone(); + + run_routines_command( + RoutinesCommand::Create { + name: "cli-digest".to_string(), + schedule: "0 0 9 * * *".to_string(), + prompt: "Prepare the morning digest.".to_string(), + description: "CLI created routine".to_string(), + timezone: Some("UTC".to_string()), + cooldown: 300, + notify_channel: None, + }, + db.clone(), + "user1", + ) + .await + .expect("create routine"); + + let routine = db + .get_routine_by_name("user1", "cli-digest") + .await + .expect("get routine by name") + .expect("cli-digest should exist"); + + match routine.action { + RoutineAction::Lightweight { + use_tools, + max_tool_rounds, + .. + } => { + assert!( + use_tools, + "CLI-created lightweight routines should default to tools" + ); + assert_eq!(max_tool_rounds, 3); + } + other => panic!("expected lightweight action, got {other:?}"), + } + } } diff --git a/src/tools/builtin/routine.rs b/src/tools/builtin/routine.rs index c197fe25..17e17ba1 100644 --- a/src/tools/builtin/routine.rs +++ b/src/tools/builtin/routine.rs @@ -140,7 +140,8 @@ fn execution_properties() -> Value { }, "use_tools": { "type": "boolean", - "description": "Only applies to lightweight mode. When true, safe non-approval tools are available." + "default": true, + "description": "Only applies to lightweight mode. New lightweight routines default this to true; when enabled, the routine can use the owner's live autonomous tool scope." }, "max_tool_rounds": { "type": "integer", @@ -290,7 +291,7 @@ fn routine_request_discovery_schema() -> Value { fn lightweight_execution_variant() -> Value { serde_json::json!({ "type": "object", - "description": "Default lightweight execution. Applies when execution is omitted or execution.mode='lightweight'.", + "description": "Default lightweight execution. Applies when execution is omitted or execution.mode='lightweight'. New lightweight routines default to tools enabled unless execution.use_tools=false is set.", "properties": { "mode": { "type": "string", @@ -304,7 +305,8 @@ fn lightweight_execution_variant() -> Value { }, "use_tools": { "type": "boolean", - "description": "When true, safe non-approval tools are available." + "default": true, + "description": "Defaults to true for new lightweight routines. When enabled, the routine can use the owner's live autonomous tool scope." }, "max_tool_rounds": { "type": "integer", @@ -335,7 +337,7 @@ fn full_job_execution_variant() -> Value { fn execution_discovery_schema() -> Value { serde_json::json!({ "type": "object", - "description": "Optional execution settings. Omit this block for the default lightweight mode.", + "description": "Optional execution settings. Omit this block for the default lightweight mode with tools enabled.", "properties": execution_properties(), "oneOf": [ lightweight_execution_variant(), @@ -408,7 +410,8 @@ fn routine_create_tool_summary() -> ToolDiscoverySummary { "execution.mode='full_job' uses the owner's live autonomous tool scope and ignores use_tools, max_tool_rounds, and context_paths.".into(), ], notes: vec![ - "Omitting execution defaults to lightweight mode.".into(), + "Omitting execution defaults to lightweight mode with tools enabled.".into(), + "Set execution.use_tools=false to keep a new lightweight routine text-only.".into(), "Omitting delivery.user falls back to the owner's last-seen notification target.".into(), "advanced.cooldown_secs defaults to 300.".into(), "Legacy flat aliases are still accepted for compatibility, but grouped fields are preferred.".into(), @@ -852,11 +855,15 @@ fn parse_execution_mode(value: Option) -> Result Result { +fn parse_routine_execution( + params: &Value, + default_use_tools: bool, +) -> Result { let mode = parse_execution_mode(string_field(params, "execution", "mode", &["action_type"]))?; let context_paths = string_array_field(params, "execution", "context_paths", &["context_paths"]); - let use_tools = bool_field(params, "execution", "use_tools", &["use_tools"]).unwrap_or(false); + let use_tools = + bool_field(params, "execution", "use_tools", &["use_tools"]).unwrap_or(default_use_tools); let max_tool_rounds = u64_field(params, "execution", "max_tool_rounds", &["max_tool_rounds"]) .unwrap_or(3) .clamp(1, crate::agent::routine::MAX_TOOL_ROUNDS_LIMIT as u64) @@ -888,7 +895,7 @@ fn parse_routine_create_request( .unwrap_or("") .to_string(); let trigger = parse_routine_trigger(params)?; - let execution = parse_routine_execution(params)?; + let execution = parse_routine_execution(params, true)?; let delivery = parse_routine_delivery(params); let cooldown_secs = u64_field(params, "advanced", "cooldown_secs", &["cooldown_secs"]).unwrap_or(300); @@ -1863,6 +1870,56 @@ mod tests { ); } + #[test] + fn parses_lightweight_create_with_tools_enabled_by_default() { + let params = serde_json::json!({ + "name": "manual-check", + "prompt": "Inspect the repo for issues.", + "request": { + "kind": "manual" + } + }); + + let parsed = parse_routine_create_request(¶ms).expect("parse default lightweight"); + + assert!( + matches!(parsed.execution.mode, NormalizedExecutionMode::Lightweight), + "expected lightweight execution mode", + ); + assert!( + parsed.execution.use_tools, + "new lightweight routines should default use_tools=true", + ); + assert_eq!(parsed.execution.max_tool_rounds, 3); + } + + #[test] + fn parses_lightweight_create_with_explicit_tools_disabled() { + let params = serde_json::json!({ + "name": "manual-check", + "prompt": "Inspect the repo for issues.", + "request": { + "kind": "manual" + }, + "execution": { + "use_tools": false + } + }); + + let parsed = + parse_routine_create_request(¶ms).expect("parse lightweight with tools disabled"); + + assert!( + matches!(parsed.execution.mode, NormalizedExecutionMode::Lightweight), + "expected lightweight execution mode", + ); + assert!( + !parsed.execution.use_tools, + "explicit use_tools=false should be preserved", + ); + assert_eq!(parsed.execution.max_tool_rounds, 3); + } + #[test] fn parses_context_paths_with_trim_drop_empty_and_stable_dedupe() { let params = serde_json::json!({ @@ -2201,6 +2258,20 @@ mod tests { .any(|rule| rule.contains("request.kind='cron'")), "summary should explain cron requirement", ); + assert!( + summary + .notes + .iter() + .any(|note| note.contains("lightweight mode with tools enabled")), + "summary should mention the new lightweight default", + ); + assert!( + summary + .notes + .iter() + .any(|note| note.contains("execution.use_tools=false")), + "summary should mention the text-only opt-out", + ); assert!( summary .notes diff --git a/tests/e2e_builtin_tool_coverage.rs b/tests/e2e_builtin_tool_coverage.rs index 69982b84..42d7fb75 100644 --- a/tests/e2e_builtin_tool_coverage.rs +++ b/tests/e2e_builtin_tool_coverage.rs @@ -205,11 +205,11 @@ mod tests { } // ----------------------------------------------------------------------- - // Test 5: routine_manual_create + // Test 5: routine_manual_create_defaults_to_tools_enabled // ----------------------------------------------------------------------- #[tokio::test] - async fn routine_manual_create() { + async fn routine_manual_create_defaults_to_tools_enabled() { let trace = LlmTrace::from_file(concat!( env!("CARGO_MANIFEST_DIR"), "/tests/fixtures/llm_traces/tools/routine_manual_create.json" @@ -237,8 +237,8 @@ mod tests { assert!(matches!(routine.trigger, Trigger::Manual)); assert!( - matches!(&routine.action, RoutineAction::Lightweight { use_tools, .. } if !*use_tools), - "manual routine should default to lightweight without tools: {:?}", + matches!(&routine.action, RoutineAction::Lightweight { use_tools, .. } if *use_tools), + "manual routine should default to lightweight with tools enabled: {:?}", routine.action ); @@ -246,7 +246,48 @@ mod tests { } // ----------------------------------------------------------------------- - // Test 6: routine_history + // Test 6: routine_manual_create_explicit_no_tools + // ----------------------------------------------------------------------- + + #[tokio::test] + async fn routine_manual_create_explicit_no_tools() { + let trace = LlmTrace::from_file(concat!( + env!("CARGO_MANIFEST_DIR"), + "/tests/fixtures/llm_traces/tools/routine_manual_create_no_tools.json" + )) + .expect("failed to load routine_manual_create_no_tools.json"); + + let rig = TestRigBuilder::new() + .with_trace(trace.clone()) + .with_auto_approve_tools(true) + .build() + .await; + + rig.send_message("Create a manual routine for quiet text-only bug triage") + .await; + let responses = rig.wait_for_responses(1, Duration::from_secs(15)).await; + + rig.verify_trace_expects(&trace, &responses); + + let routine = rig + .database() + .get_routine_by_name("test-user", "manual-triage-no-tools") + .await + .expect("get_routine_by_name") + .expect("manual-triage-no-tools should exist"); + + assert!(matches!(routine.trigger, Trigger::Manual)); + assert!( + matches!(&routine.action, RoutineAction::Lightweight { use_tools, .. } if !*use_tools), + "manual routine should preserve explicit use_tools=false: {:?}", + routine.action + ); + + rig.shutdown(); + } + + // ----------------------------------------------------------------------- + // Test 7: routine_history // ----------------------------------------------------------------------- #[tokio::test] @@ -283,7 +324,7 @@ mod tests { } // ----------------------------------------------------------------------- - // Test 7: routine_system_event_emit + // Test 8: routine_system_event_emit // ----------------------------------------------------------------------- #[tokio::test] diff --git a/tests/fixtures/llm_traces/tools/routine_manual_create_no_tools.json b/tests/fixtures/llm_traces/tools/routine_manual_create_no_tools.json new file mode 100644 index 00000000..275f2269 --- /dev/null +++ b/tests/fixtures/llm_traces/tools/routine_manual_create_no_tools.json @@ -0,0 +1,39 @@ +{ + "model_name": "test-routine-manual-create-no-tools", + "expects": { + "tools_used": ["routine_create"], + "all_tools_succeeded": true, + "min_responses": 1 + }, + "steps": [ + { + "response": { + "type": "tool_calls", + "tool_calls": [ + { + "id": "call_rc_manual_2", + "name": "routine_create", + "arguments": { + "name": "manual-triage-no-tools", + "trigger_type": "manual", + "prompt": "Summarize the latest bug reports when this routine is fired.", + "execution": { + "use_tools": false + } + } + } + ], + "input_tokens": 90, + "output_tokens": 24 + } + }, + { + "response": { + "type": "text", + "content": "Created the manual-triage-no-tools routine. It will only run when explicitly fired and stay text-only.", + "input_tokens": 140, + "output_tokens": 18 + } + } + ] +} From fa51b9f52dde0727f5dd65f134b93095832de959 Mon Sep 17 00:00:00 2001 From: Illia Polosukhin Date: Mon, 23 Mar 2026 14:50:15 -0700 Subject: [PATCH 02/20] =?UTF-8?q?fix:=20post-merge=20review=20sweep=20?= =?UTF-8?q?=E2=80=94=208=20fixes=20across=20security,=20perf,=20and=20corr?= =?UTF-8?q?ectness=20(#1550)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix: post-merge review sweep — 8 fixes across security, perf, and correctness 1. Fix code fence detection in extract_suggestions() (issue #1180) - rfind("```") couldn't handle odd fence counts (unclosed blocks) - Now counts all fence positions and checks parity 2. Cache routine parameters_schema() with OnceLock (issue #1361) - routine_create_parameters_schema() and event_emit_parameters_schema() were regenerating JSON on every LLM call 3. Replace O(n) LRU eviction with lru crate (issue #1430) - Embedding cache now uses lru::LruCache for O(1) eviction - Removes manual HashMap + last_accessed tracking 4. Fix WASM router secret_validated semantics (issue #1281) - Now reflects whether any auth (secret/Ed25519/HMAC) was performed - Previously only checked if a secret was configured 5. Sanitize channel/user in routine prompt interpolation (issue #1364) - Defense-in-depth: strip newlines, replace backticks, truncate to 128 chars before injecting into LLM prompt 6. Remove duplicate 401 retry in github_copilot.rs (PR #1512 review) - Internal retry conflicted with outer RetryProvider causing nested retries; now invalidates token and lets RetryProvider handle retry 7. Fix token error classification in github_copilot.rs (PR #1512 review) - AccessDenied/Expired errors now map to AuthFailed (non-retryable) - Transient errors remain RequestFailed (retryable) 8. Fix parse_extra_headers() hardcoded env var name (PR #1512 review) - Error messages now report the actual env var being parsed instead of always saying LLM_EXTRA_HEADERS Co-Authored-By: Claude Opus 4.6 (1M context) * fix: address PR review comments and fix formatting - sanitize_prompt_field: single-pass with map() instead of collect+replace - embed(): re-check cache under lock before cloning (thundering herd) - embed_batch(): limit caching to cache capacity, skip overflow entries - router: thread did_authenticate bool instead of re-calling async methods - github_copilot 401: use generic error message, avoid leaking response body - cargo fmt: fix two formatting violations caught by CI Co-Authored-By: Claude Opus 4.6 (1M context) * chore: trigger CI re-run with updated refs [skip-regression-check] Co-Authored-By: Claude Opus 4.6 (1M context) --------- Co-authored-by: Claude Opus 4.6 (1M context) --- src/agent/dispatcher.rs | 26 +++++- src/agent/routine_engine.rs | 19 ++++- src/channels/wasm/router.rs | 11 ++- src/config/llm.rs | 16 +++- src/llm/github_copilot.rs | 71 +++++----------- src/tools/builtin/routine.rs | 6 +- src/workspace/embedding_cache.rs | 135 +++++++------------------------ 7 files changed, 114 insertions(+), 170 deletions(-) diff --git a/src/agent/dispatcher.rs b/src/agent/dispatcher.rs index 03548219..5d39866b 100644 --- a/src/agent/dispatcher.rs +++ b/src/agent/dispatcher.rs @@ -1098,15 +1098,23 @@ pub(crate) fn extract_suggestions(text: &str) -> (String, Vec) { Regex::new(r"(?s)\s*(.*?)\s*").expect("valid regex") // safety: constant pattern }); - // Find the position of the last closing code fence to avoid matching inside code blocks - let last_code_fence = text.rfind("```").unwrap_or(0); + // Build a sorted list of code fence positions to determine open/close pairing. + // A position is "inside" a fenced block when it falls between an odd-numbered + // fence (opening) and the next even-numbered fence (closing). + let fence_positions: Vec = text.match_indices("```").map(|(pos, _)| pos).collect(); - // Find all matches, take the last one that's after the last code fence + let is_inside_fence = |pos: usize| -> bool { + // Count how many fences appear before `pos`. If odd, we're inside a fence. + let count = fence_positions.iter().take_while(|&&fp| fp <= pos).count(); + count % 2 == 1 + }; + + // Find all matches, take the last one that's outside any code fence let mut best_match: Option> = None; let mut best_capture: Option = None; for caps in RE.captures_iter(text) { if let (Some(full), Some(inner)) = (caps.get(0), caps.get(1)) - && full.start() >= last_code_fence + && !is_inside_fence(full.start()) { best_match = Some(full); best_capture = Some(inner.as_str().to_string()); @@ -2345,6 +2353,16 @@ mod tests { assert!(suggestions.is_empty()); // safety: test } + #[test] + fn test_extract_suggestions_inside_unclosed_code_fence() { + // Regression: odd number of fences (unclosed fence) must still be + // treated as "inside a code block". + let input = "```\ncode\n[\"bar\"]"; + let (text, suggestions) = super::extract_suggestions(input); + assert_eq!(text, input); // safety: test + assert!(suggestions.is_empty()); // safety: test + } + #[test] fn test_extract_suggestions_after_code_fence() { let input = "```\ncode\n```\nAnswer.\n[\"foo\"]"; diff --git a/src/agent/routine_engine.rs b/src/agent/routine_engine.rs index 9b554582..7c7ef5f3 100644 --- a/src/agent/routine_engine.rs +++ b/src/agent/routine_engine.rs @@ -1305,6 +1305,19 @@ async fn execute_lightweight( } } +/// Sanitize a user-controlled string before interpolation into an LLM prompt. +/// Strips newlines (which could break prompt structure) and truncates to a +/// reasonable length to limit abuse surface. +fn sanitize_prompt_field(value: &str) -> String { + const MAX_LEN: usize = 128; + value + .chars() + .filter(|&c| c != '\n' && c != '\r') + .take(MAX_LEN) + .map(|c| if c == '`' { '\'' } else { c }) + .collect() +} + fn build_lightweight_prompt( prompt: &str, context_parts: &[String], @@ -1323,14 +1336,16 @@ fn build_lightweight_prompt( ); if let Some(channel) = notify.channel.as_deref() { + let sanitized = sanitize_prompt_field(channel); full_prompt.push_str(&format!( - "The configured delivery channel for this routine is `{channel}`.\n" + "The configured delivery channel for this routine is `{sanitized}`.\n" )); } if let Some(user) = notify.user.as_deref() { + let sanitized = sanitize_prompt_field(user); full_prompt.push_str(&format!( - "The configured delivery target for this routine is `{user}`.\n" + "The configured delivery target for this routine is `{sanitized}`.\n" )); } diff --git a/src/channels/wasm/router.rs b/src/channels/wasm/router.rs index 8005ccea..510bc461 100644 --- a/src/channels/wasm/router.rs +++ b/src/channels/wasm/router.rs @@ -333,6 +333,9 @@ async fn webhook_handler( let channel_name = channel.channel_name(); + // Track whether any authentication was performed and passed. + let mut did_authenticate = false; + // Check if secret is required if state.router.requires_secret(channel_name).await { // Get the secret header name for this channel (from capabilities or default) @@ -382,6 +385,7 @@ async fn webhook_handler( ); } tracing::debug!(channel = %channel_name, "Webhook secret validated"); + did_authenticate = true; } None => { tracing::warn!( @@ -433,6 +437,7 @@ async fn webhook_handler( ); } tracing::debug!(channel = %channel_name, "Ed25519 signature verified"); + did_authenticate = true; } _ => { tracing::warn!( @@ -484,6 +489,7 @@ async fn webhook_handler( ); } tracing::debug!(channel = %channel_name, "HMAC-SHA256 signature verified"); + did_authenticate = true; } _ => { tracing::warn!( @@ -510,8 +516,9 @@ async fn webhook_handler( }) .collect(); - // Call the WASM channel - let secret_validated = state.router.requires_secret(channel_name).await; + // Call the WASM channel. `did_authenticate` was set above by whichever + // auth guard (secret / Ed25519 / HMAC) successfully validated the request. + let secret_validated = did_authenticate; tracing::info!( channel = %channel_name, diff --git a/src/config/llm.rs b/src/config/llm.rs index 87e4daa5..ed4b8a05 100644 --- a/src/config/llm.rs +++ b/src/config/llm.rs @@ -406,7 +406,7 @@ impl LlmConfig { // Resolve extra headers let extra_headers = if let Some(env_var) = extra_headers_env { optional_env(env_var)? - .map(|val| parse_extra_headers(&val)) + .map(|val| parse_extra_headers_with_key(&val, env_var)) .transpose()? .unwrap_or_default() } else { @@ -475,7 +475,10 @@ impl LlmConfig { /// /// Format: `Key1:Value1,Key2:Value2` (colon-separated, not `=`, because /// header values often contain `=`). -fn parse_extra_headers(val: &str) -> Result, ConfigError> { +fn parse_extra_headers_with_key( + val: &str, + env_var_name: &str, +) -> Result, ConfigError> { if val.trim().is_empty() { return Ok(Vec::new()); } @@ -488,14 +491,14 @@ fn parse_extra_headers(val: &str) -> Result, ConfigError> } let Some((key, value)) = pair.split_once(':') else { return Err(ConfigError::InvalidValue { - key: "LLM_EXTRA_HEADERS".to_string(), + key: env_var_name.to_string(), message: format!("malformed header entry '{}', expected Key:Value", pair), }); }; let key = key.trim(); if key.is_empty() { return Err(ConfigError::InvalidValue { - key: "LLM_EXTRA_HEADERS".to_string(), + key: env_var_name.to_string(), message: format!("empty header name in entry '{}'", pair), }); } @@ -536,6 +539,11 @@ mod tests { use crate::settings::Settings; use crate::testing::credentials::*; + /// Convenience wrapper for tests — uses "TEST_HEADERS" as the env var name. + fn parse_extra_headers(val: &str) -> Result, ConfigError> { + parse_extra_headers_with_key(val, "TEST_HEADERS") + } + /// Clear all openai-compatible-related env vars. fn clear_openai_compatible_env() { // SAFETY: Only called under ENV_MUTEX in tests. diff --git a/src/llm/github_copilot.rs b/src/llm/github_copilot.rs index 9baf6c74..b173191a 100644 --- a/src/llm/github_copilot.rs +++ b/src/llm/github_copilot.rs @@ -107,14 +107,21 @@ impl GithubCopilotProvider { body: &impl Serialize, ) -> Result { let url = self.api_url(); - // Map token exchange failures to RequestFailed (retryable) rather than - // AuthFailed (non-retryable), since transient network errors during - // exchange should be retried by RetryProvider. + // Distinguish permanent auth errors (non-retryable) from transient + // network failures (retryable) so RetryProvider handles them correctly. let token = self.token_manager.get_token().await.map_err(|e| { tracing::warn!(error = %e, "Copilot: token exchange failed"); - LlmError::RequestFailed { - provider: "github_copilot".to_string(), - reason: format!("Token exchange failed: {e}"), + match &e { + crate::llm::github_copilot_auth::GithubCopilotAuthError::AccessDenied + | crate::llm::github_copilot_auth::GithubCopilotAuthError::Expired => { + LlmError::AuthFailed { + provider: "github_copilot".to_string(), + } + } + _ => LlmError::RequestFailed { + provider: "github_copilot".to_string(), + reason: format!("Token exchange failed: {e}"), + }, } })?; @@ -157,54 +164,14 @@ impl GithubCopilotProvider { ); if status.as_u16() == 401 { - // Invalidate the cached session token and retry once with a - // fresh exchange — stale tokens are the most common 401 cause. - tracing::warn!("Copilot: 401 Unauthorized — invalidating session token, retrying"); + // Invalidate the cached session token so the next attempt + // (driven by RetryProvider) gets a fresh one. We don't retry + // inline to avoid nested retries with the outer RetryProvider. + tracing::warn!("Copilot: 401 Unauthorized — invalidating session token for retry"); self.token_manager.invalidate().await; - let fresh = self.token_manager.get_token().await.map_err(|e| { - tracing::warn!(error = %e, "Copilot: re-exchange after 401 failed"); - LlmError::RequestFailed { - provider: "github_copilot".to_string(), - reason: format!("Token re-exchange after 401 failed: {e}"), - } - })?; - let mut retry_req = self - .client - .post(&url) - .bearer_auth(fresh.expose_secret()) - .header("Content-Type", "application/json"); - for (key, value) in &self.extra_headers { - retry_req = retry_req.header(key.as_str(), value.as_str()); - } - let retry = - retry_req - .json(body) - .send() - .await - .map_err(|e| LlmError::RequestFailed { - provider: "github_copilot".to_string(), - reason: format!("Retry after 401 failed: {e}"), - })?; - if retry.status().is_success() { - let text = retry.text().await.map_err(|e| LlmError::RequestFailed { - provider: "github_copilot".to_string(), - reason: format!("Failed to read retry response body: {e}"), - })?; - return serde_json::from_str(&text).map_err(|e| { - let truncated = crate::agent::truncate_for_preview(&text, 512); - LlmError::InvalidResponse { - provider: "github_copilot".to_string(), - reason: format!("JSON parse error: {e}. Raw: {truncated}"), - } - }); - } - let retry_status = retry.status(); - tracing::warn!( - status = %retry_status, - "Copilot: 401 retry also failed" - ); - return Err(LlmError::AuthFailed { + return Err(LlmError::RequestFailed { provider: "github_copilot".to_string(), + reason: "HTTP 401 Unauthorized".to_string(), }); } if status.as_u16() == 429 { diff --git a/src/tools/builtin/routine.rs b/src/tools/builtin/routine.rs index 17e17ba1..f4313483 100644 --- a/src/tools/builtin/routine.rs +++ b/src/tools/builtin/routine.rs @@ -608,7 +608,8 @@ fn routine_create_schema(include_compatibility_aliases: bool) -> Value { } pub(crate) fn routine_create_parameters_schema() -> Value { - routine_create_schema(false) + static CACHE: OnceLock = OnceLock::new(); + CACHE.get_or_init(|| routine_create_schema(false)).clone() } fn routine_create_discovery_schema() -> Value { @@ -1014,7 +1015,8 @@ fn event_emit_schema(include_source_alias: bool) -> Value { } pub(crate) fn event_emit_parameters_schema() -> Value { - event_emit_schema(false) + static CACHE: OnceLock = OnceLock::new(); + CACHE.get_or_init(|| event_emit_schema(false)).clone() } fn event_emit_discovery_schema() -> Value { diff --git a/src/workspace/embedding_cache.rs b/src/workspace/embedding_cache.rs index 21d3c7c3..60c2eb08 100644 --- a/src/workspace/embedding_cache.rs +++ b/src/workspace/embedding_cache.rs @@ -3,14 +3,13 @@ //! Avoids redundant HTTP calls for identical texts by caching embeddings //! in memory keyed by `SHA-256(model_name + "\0" + text)`. //! -//! Follows the same cache pattern as `llm::response_cache::CachedProvider`: -//! `HashMap` + `last_accessed` tracking + manual LRU eviction. +//! Uses `lru::LruCache` for O(1) insertion, lookup, and eviction. -use std::collections::HashMap; +use std::num::NonZeroUsize; use std::sync::{Arc, Mutex}; -use std::time::Instant; use async_trait::async_trait; +use lru::LruCache; use sha2::{Digest, Sha256}; use crate::workspace::embeddings::{EmbeddingError, EmbeddingProvider}; @@ -22,8 +21,7 @@ pub struct EmbeddingCacheConfig { /// /// Approximate raw embedding payload: `max_entries × dimension × 4 bytes`. /// At 10,000 entries × 1536 floats ≈ 58 MB (payload only; actual memory - /// is higher due to HashMap buckets, `[u8; 32]` hash keys, `Vec`/`Instant` - /// per-entry overhead). + /// is higher due to per-entry overhead in the linked-list LRU). pub max_entries: usize, } @@ -35,11 +33,6 @@ impl Default for EmbeddingCacheConfig { } } -struct CacheEntry { - embedding: Vec, - last_accessed: Instant, -} - /// Embedding provider wrapper that caches results in memory. /// /// Thread-safe via `std::sync::Mutex`. The lock is **never held** @@ -47,8 +40,7 @@ struct CacheEntry { /// so a synchronous mutex is cheaper than `tokio::sync::Mutex`. pub struct CachedEmbeddingProvider { inner: Arc, - cache: Mutex>, - config: EmbeddingCacheConfig, + cache: Mutex>>, } impl CachedEmbeddingProvider { @@ -56,19 +48,18 @@ impl CachedEmbeddingProvider { /// /// `config.max_entries` is clamped to at least 1. pub fn new(inner: Arc, config: EmbeddingCacheConfig) -> Self { - let config = EmbeddingCacheConfig { - max_entries: config.max_entries.max(1), - }; - if config.max_entries > 100_000 { + let max_entries = config.max_entries.max(1); + if max_entries > 100_000 { tracing::warn!( - max_entries = config.max_entries, + max_entries, "Embedding cache size exceeds 100,000 entries; memory usage may be significant" ); } + // safety: max_entries >= 1 due to .max(1) above + let cap = NonZeroUsize::new(max_entries).expect("clamped to >= 1"); // safety: always >= 1 Self { inner, - cache: Mutex::new(HashMap::with_capacity(config.max_entries.min(1024))), - config, + cache: Mutex::new(LruCache::new(cap)), } } @@ -100,49 +91,6 @@ impl CachedEmbeddingProvider { hasher.update(text.as_bytes()); hasher.finalize().into() } - - /// Evict the least-recently-used entry if at capacity (single-entry path). - // TODO: O(n) scan per eviction. If max_entries grows large, switch to - // an ordered data structure (e.g. `IndexMap` with swap_remove, or a - // linked-list LRU like the `lru` crate). - fn evict_lru(cache: &mut HashMap<[u8; 32], CacheEntry>, max_entries: usize) { - while cache.len() >= max_entries { - let oldest_key = cache - .iter() - .min_by_key(|(_, entry)| entry.last_accessed) - .map(|(k, _)| *k); - - if let Some(k) = oldest_key { - cache.remove(&k); - } else { - break; - } - } - } - - /// Evict the `k` oldest entries in O(n) average time via partial selection. - /// - /// Used by `embed_batch` to avoid the O(n×m) cost of calling - /// `evict_lru` per insert. - fn evict_k_oldest(cache: &mut HashMap<[u8; 32], CacheEntry>, k: usize) { - if k == 0 || cache.is_empty() { - return; - } - if k >= cache.len() { - cache.clear(); - return; - } - // Partial selection: find the k oldest in O(n) average via - // select_nth_unstable_by_key, then remove the first k entries. - let mut entries: Vec<([u8; 32], Instant)> = cache - .iter() - .map(|(key, entry)| (*key, entry.last_accessed)) - .collect(); - entries.select_nth_unstable_by_key(k - 1, |(_, t)| *t); - for (key, _) in entries.into_iter().take(k) { - cache.remove(&key); - } - } } #[async_trait] @@ -162,39 +110,32 @@ impl EmbeddingProvider for CachedEmbeddingProvider { async fn embed(&self, text: &str) -> Result, EmbeddingError> { let key = self.cache_key(text); - // Check cache (short critical section) + // Check cache (short critical section). LruCache::get promotes the + // entry to most-recently-used automatically. { let mut guard = self.cache.lock().unwrap_or_else(|e| e.into_inner()); - if let Some(entry) = guard.get_mut(&key) { - entry.last_accessed = Instant::now(); + if let Some(embedding) = guard.get(&key) { tracing::trace!("embedding cache hit"); - return Ok(entry.embedding.clone()); + return Ok(embedding.clone()); } } // Lock released before HTTP call. // NOTE: Thundering herd — multiple concurrent callers with the same // uncached key will each call the inner provider. This is acceptable: - // embeddings are idempotent and the last writer wins in the HashMap. + // embeddings are idempotent and the last writer wins in the LruCache. let embedding = self.inner.embed(text).await?; - // Store result. Re-check under lock: another concurrent caller may - // have inserted this key while the lock was released for the HTTP call. + // Store result under lock. Re-check first: another concurrent caller + // may have already cached this key while the lock was released. { let mut guard = self.cache.lock().unwrap_or_else(|e| e.into_inner()); - if let Some(entry) = guard.get_mut(&key) { - // Thundering herd — another caller already cached it. - // Just touch timestamp; skip the clone. - entry.last_accessed = Instant::now(); + if guard.get(&key).is_some() { + // Thundering herd — another caller beat us. LruCache::get + // already promoted it to most-recently-used; skip the clone. + tracing::trace!("embedding cache: concurrent insert, skipping clone"); } else { - Self::evict_lru(&mut guard, self.config.max_entries); - guard.insert( - key, - CacheEntry { - embedding: embedding.clone(), - last_accessed: Instant::now(), - }, - ); + guard.push(key, embedding.clone()); } } @@ -214,11 +155,9 @@ impl EmbeddingProvider for CachedEmbeddingProvider { { let mut guard = self.cache.lock().unwrap_or_else(|e| e.into_inner()); - let now = Instant::now(); for (i, key) in keys.iter().enumerate() { - if let Some(entry) = guard.get_mut(key) { - entry.last_accessed = now; - results[i] = Some(entry.embedding.clone()); + if let Some(embedding) = guard.get(key) { + results[i] = Some(embedding.clone()); } else { miss_indices.push(i); } @@ -228,7 +167,6 @@ impl EmbeddingProvider for CachedEmbeddingProvider { if miss_indices.is_empty() { tracing::trace!(count = texts.len(), "embedding batch: all cache hits"); - // All slots populated from cache hits return results .into_iter() .enumerate() @@ -260,29 +198,18 @@ impl EmbeddingProvider for CachedEmbeddingProvider { "embedding batch: partial cache" ); - // Cache FIRST (clone only the cacheable subset), then move originals - // into results. This avoids cloning capacity-skipped embeddings entirely. + // Cache only the last `cap` new embeddings — caching more than the + // cache capacity wastes clone work on entries that are immediately evicted. { let mut guard = self.cache.lock().unwrap_or_else(|e| e.into_inner()); - let cacheable = miss_indices.len().min(self.config.max_entries); - let skip = miss_indices.len() - cacheable; - let need_to_evict = (guard.len() + cacheable).saturating_sub(self.config.max_entries); - if need_to_evict > 0 { - Self::evict_k_oldest(&mut guard, need_to_evict); - } - let now = Instant::now(); + let cap = guard.cap().get(); + let skip = miss_indices.len().saturating_sub(cap); for (&orig_idx, emb) in miss_indices[skip..].iter().zip(&new_embeddings[skip..]) { - guard.insert( - keys[orig_idx], - CacheEntry { - embedding: emb.clone(), - last_accessed: now, - }, - ); + guard.push(keys[orig_idx], emb.clone()); } } - // Move originals into results (zero-copy for all, including cached ones). + // Move originals into results (zero-copy). for (orig_idx, emb) in miss_indices.iter().copied().zip(new_embeddings) { results[orig_idx] = Some(emb); } From b441ebec02bdedf650abbcc89c6321b477247504 Mon Sep 17 00:00:00 2001 From: standardtoaster Date: Tue, 24 Mar 2026 04:50:05 +0100 Subject: [PATCH 03/20] feat: multi-tenant auth with per-user workspace isolation (#1118) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat: multi-tenant auth with per-user scoping Multi-user authentication and authorization for IronClaw gateway: - Token-based auth mapping tokens to user IDs via GATEWAY_USER_TOKENS - Per-user SSE broadcast scoping - Per-user rate limiting with poisoned lock recovery - Handler auth and ownership checks for jobs, settings, routines - Extension secrets scoped per-user - Chat handlers use authenticated identity - Reverse proxy deployment documentation - Comprehensive integration tests for auth, SSE, rate limiting, and job isolation * fix: scope memory tools per-user in multi-tenant mode Memory tools (search, write, read, tree) held a single workspace created at startup with GATEWAY_USER_ID. In multi-tenant mode, all users' tool calls searched the default user's scope. Add WorkspaceResolver trait that resolves workspaces per-request using JobContext.user_id. In single-user mode, returns the startup workspace. In multi-tenant mode (GATEWAY_USER_TOKENS configured), creates and caches per-user workspaces on demand. Includes regression tests for workspace resolution and user isolation. Co-Authored-By: Claude Opus 4.6 (1M context) * fix: comprehensive multi-tenant isolation audit Address all review findings from @serrrfirat plus 7 additional gaps found via full security audit: Reviewer findings (5): - WorkspacePool now applies search config, memory layers, embedding cache, identity read scopes, and global config scopes (was bare) - jobs_summary_handler uses per-user queries instead of global counters - jobs_prompt_handler restructured to not 404 agent jobs + ownership check - jobs_restart_handler agent branch now verifies user ownership - agent_job_summary_for_user added to Database trait + both backends Audit findings (7): - Delete dead handlers/memory.rs (stale copies with no auth) - Add AuthenticatedUser to logs_events, logs_level_get, logs_level_set - Add AuthenticatedUser to extensions_tools_handler, gateway_status_handler - Add auth + ownership checks to all 6 routines handlers - Add auth to all 4 skills handlers with audit logging on mutations - Scope extension setup SSE broadcast to user (broadcast_for_user) - Fix pre-existing test compilation errors in extensions/manager.rs 17 new multi-tenant isolation tests covering: - WorkspacePool config propagation and scope merging - Jobs handler per-user isolation (summary, restart, prompt, cancel) - Routines handler auth enforcement and cross-user rejection - Auth middleware enforcement on logs, skills, status endpoints Co-Authored-By: Claude Opus 4.6 (1M context) * fix: second-pass multi-tenant audit — scope SSE broadcasts, DB queries, dead handlers Second audit pass applying learned patterns across the codebase: - OAuth callback SSE broadcasts now use broadcast_for_user (lines 773, 912) - jobs_list_handler uses list_agent_jobs_for_user instead of fetching all users' jobs and filtering in Rust - list_agent_jobs_for_user added to Database trait + postgres + libsql - Dead handler files (extensions.rs, static_files.rs) hardened with AuthenticatedUser to prevent auth regression if migrated Co-Authored-By: Claude Opus 4.6 (1M context) * fix: address review findings — token hashing, broadcast scoping, error handling Security fixes: - Hash tokens with SHA-256 at construction time so authentication compares fixed-size 32-byte digests, eliminating length-oracle timing leaks - Scope auth SSE broadcasts per-user in chat_auth_token_handler — AuthRequired/AuthCompleted events were leaking across tenants - Propagate DB errors in restart handlers instead of silently swallowing via `if let Ok(Some(...))` pattern Code quality: - Log SSE serialization failures instead of silently producing empty strings via unwrap_or_default() - Remove dead `pub type AuthState = MultiAuthState` alias - Replace `.unwrap()` with `Arc::clone(db)` in app.rs multi-tenant workspace setup (db is guaranteed Some in context, but unwrap violates project convention) - Fix telegram setup test to inject UserIdentity into request extensions (handler now requires AuthenticatedUser) - Add safety comments on test-only expect/unwrap calls for CI - Apply cargo fmt to fix pre-existing formatting Co-Authored-By: Claude Opus 4.6 (1M context) * fix: address review findings — unify workspace pool, fix SSE regression, cache job owners - Unify WorkspacePool and PerUserWorkspaceResolver: WorkspacePool now implements WorkspaceResolver, eliminating duplicate per-user workspace construction logic. app.rs uses WorkspacePool directly. - Fix sse_tx: None scheduler regression: change scheduler/worker SSE broadcasting from broadcast::Sender to Arc, restoring SSE event delivery for scheduled agent jobs. - Cache job owner in orchestrator: add job_owner_cache to OrchestratorState so job_event_handler avoids a DB round-trip on every event after the first per job. - Deduplicate ext_user_id computation in main.rs. - Remove unused _gateway_state variable. - Fix pre-existing test: first_token() returns None in multi-user mode by design; align test assertion. Co-Authored-By: Claude Opus 4.6 (1M context) * style: fix formatting in app.rs Co-Authored-By: Claude Opus 4.6 (1M context) * refactor: extract memory handlers back into handlers/memory.rs Move memory API handlers out of server.rs into their own module, consistent with how jobs, routines, and skills handlers are organized. The resolve_workspace() helper moves with them since it is only used by memory handlers. Co-Authored-By: Claude Opus 4.6 (1M context) --------- Co-authored-by: Claude Opus 4.6 (1M context) Co-authored-by: ilblackdragon@gmail.com --- src/agent/agent_loop.rs | 8 +- src/agent/job_monitor.rs | 34 +- src/agent/scheduler.rs | 11 +- src/agent/thread_ops.rs | 2 +- src/app.rs | 32 +- src/channels/web/auth.rs | 405 +++++++- src/channels/web/handlers/chat.rs | 97 +- src/channels/web/handlers/extensions.rs | 11 +- src/channels/web/handlers/jobs.rs | 677 +++++++------ src/channels/web/handlers/memory.rs | 113 ++- src/channels/web/handlers/mod.rs | 15 +- src/channels/web/handlers/routines.rs | 45 +- src/channels/web/handlers/settings.rs | 19 +- src/channels/web/handlers/skills.rs | 9 + src/channels/web/handlers/static_files.rs | 3 + src/channels/web/mod.rs | 112 ++- src/channels/web/openai_compat.rs | 3 +- src/channels/web/server.rs | 891 +++++++++-------- src/channels/web/sse.rs | 158 ++- src/channels/web/test_helpers.rs | 29 +- src/channels/web/tests/mod.rs | 3 + src/channels/web/tests/multi_tenant.rs | 796 ++++++++++++++++ src/channels/web/ws.rs | 37 +- src/cli/oauth_defaults.rs | 4 +- src/config/channels.rs | 142 +++ src/db/libsql/jobs.rs | 69 ++ src/db/mod.rs | 8 + src/db/postgres.rs | 14 + src/extensions/manager.rs | 559 ++++++----- src/history/store.rs | 53 ++ src/main.rs | 75 +- src/orchestrator/api.rs | 59 +- src/orchestrator/mod.rs | 3 +- src/tools/builtin/extension_tools.rs | 34 +- src/tools/builtin/job.rs | 4 +- src/tools/builtin/memory.rs | 328 ++++++- src/tools/builtin/mod.rs | 2 +- src/tools/registry.rs | 38 +- src/worker/job.rs | 8 +- tests/e2e_advanced_traces.rs | 2 +- tests/module_init_integration.rs | 2 +- tests/multi_tenant_integration.rs | 1059 +++++++++++++++++++++ tests/multi_tenant_system_prompt.rs | 240 +++++ tests/openai_compat_integration.rs | 35 +- tests/support/gateway_workflow_harness.rs | 17 +- tests/ws_gateway_integration.rs | 13 +- 46 files changed, 5074 insertions(+), 1204 deletions(-) create mode 100644 src/channels/web/tests/mod.rs create mode 100644 src/channels/web/tests/multi_tenant.rs create mode 100644 tests/multi_tenant_integration.rs create mode 100644 tests/multi_tenant_system_prompt.rs diff --git a/src/agent/agent_loop.rs b/src/agent/agent_loop.rs index 5cbd8166..ee91ea9a 100644 --- a/src/agent/agent_loop.rs +++ b/src/agent/agent_loop.rs @@ -157,8 +157,8 @@ pub struct AgentDeps { pub hooks: Arc, /// Cost enforcement guardrails (daily budget, hourly rate limits). pub cost_guard: Arc, - /// SSE broadcast sender for live job event streaming to the web gateway. - pub sse_tx: Option>, + /// SSE manager for live job event streaming to the web gateway. + pub sse_tx: Option>, /// HTTP interceptor for trace recording/replay. pub http_interceptor: Option>, /// Audio transcription middleware for voice messages. @@ -235,8 +235,8 @@ impl Agent { hooks: deps.hooks.clone(), }, ); - if let Some(ref tx) = deps.sse_tx { - scheduler.set_sse_sender(tx.clone()); + if let Some(ref sse) = deps.sse_tx { + scheduler.set_sse_sender(Arc::clone(sse)); } if let Some(ref interceptor) = deps.http_interceptor { scheduler.set_http_interceptor(Arc::clone(interceptor)); diff --git a/src/agent/job_monitor.rs b/src/agent/job_monitor.rs index 675d0426..02f5e3e2 100644 --- a/src/agent/job_monitor.rs +++ b/src/agent/job_monitor.rs @@ -44,7 +44,7 @@ pub struct JobMonitorRoute { /// the main agent's context window). pub fn spawn_job_monitor( job_id: Uuid, - event_rx: broadcast::Receiver<(Uuid, SseEvent)>, + event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>, 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, SseEvent)>, + mut event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>, inject_tx: mpsc::Sender, route: JobMonitorRoute, context_manager: Option>, @@ -68,7 +68,7 @@ pub fn spawn_job_monitor_with_context( loop { match event_rx.recv().await { - Ok((ev_job_id, event)) => { + Ok((ev_job_id, _user_id, event)) => { if ev_job_id != job_id { continue; } @@ -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, SseEvent)>, + mut event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>, context_manager: Arc, ) -> JoinHandle<()> { let short_id = job_id.to_string()[..8].to_string(); @@ -170,7 +170,9 @@ pub fn spawn_completion_watcher( tokio::spawn(async move { loop { match event_rx.recv().await { - Ok((ev_job_id, SseEvent::JobResult { status, .. })) if ev_job_id == job_id => { + Ok((ev_job_id, _user_id, SseEvent::JobResult { status, .. })) + if ev_job_id == job_id => + { let target = if status == "completed" { JobState::Completed } else { @@ -227,7 +229,7 @@ mod tests { #[tokio::test] async fn test_monitor_forwards_assistant_messages() { - let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let job_id = Uuid::new_v4(); @@ -237,6 +239,7 @@ mod tests { event_tx .send(( job_id, + "test-user".to_string(), SseEvent::JobMessage { job_id: job_id.to_string(), role: "assistant".to_string(), @@ -259,7 +262,7 @@ mod tests { #[tokio::test] async fn test_monitor_ignores_other_jobs() { - let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let job_id = Uuid::new_v4(); @@ -270,6 +273,7 @@ mod tests { event_tx .send(( other_job_id, + "test-user".to_string(), SseEvent::JobMessage { job_id: other_job_id.to_string(), role: "assistant".to_string(), @@ -289,7 +293,7 @@ mod tests { #[tokio::test] async fn test_monitor_exits_on_job_result() { - let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let job_id = Uuid::new_v4(); @@ -299,6 +303,7 @@ mod tests { event_tx .send(( job_id, + "test-user".to_string(), SseEvent::JobResult { job_id: job_id.to_string(), status: "completed".to_string(), @@ -324,7 +329,7 @@ mod tests { #[tokio::test] async fn test_monitor_skips_tool_events() { - let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let job_id = Uuid::new_v4(); @@ -334,6 +339,7 @@ mod tests { event_tx .send(( job_id, + "test-user".to_string(), SseEvent::JobToolUse { job_id: job_id.to_string(), tool_name: "shell".to_string(), @@ -346,6 +352,7 @@ mod tests { event_tx .send(( job_id, + "test-user".to_string(), SseEvent::JobMessage { job_id: job_id.to_string(), role: "user".to_string(), @@ -402,7 +409,7 @@ mod tests { .await .unwrap(); - let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let handle = spawn_job_monitor_with_context( @@ -417,6 +424,7 @@ mod tests { event_tx .send(( job_id, + "test-user".to_string(), SseEvent::JobResult { job_id: job_id.to_string(), status: "completed".to_string(), @@ -450,7 +458,7 @@ mod tests { .await .unwrap(); - let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let handle = spawn_job_monitor_with_context( @@ -465,6 +473,7 @@ mod tests { event_tx .send(( job_id, + "test-user".to_string(), SseEvent::JobResult { job_id: job_id.to_string(), status: "failed".to_string(), @@ -498,12 +507,13 @@ mod tests { .await .unwrap(); - let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(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 { job_id: job_id.to_string(), status: "completed".to_string(), diff --git a/src/agent/scheduler.rs b/src/agent/scheduler.rs index 1c4a7fde..02953a4b 100644 --- a/src/agent/scheduler.rs +++ b/src/agent/scheduler.rs @@ -9,7 +9,6 @@ use tokio::task::JoinHandle; use uuid::Uuid; use crate::agent::task::{Task, TaskContext, TaskOutput}; -use crate::channels::web::types::SseEvent; use crate::config::AgentConfig; use crate::context::{ContextManager, JobContext, JobState}; use crate::db::Database; @@ -67,8 +66,8 @@ pub struct Scheduler { extension_manager: Option>, store: Option>, hooks: Arc, - /// SSE broadcast sender for live job event streaming. - sse_tx: Option>, + /// SSE manager for live job event streaming. + sse_tx: Option>, /// HTTP interceptor for trace recording/replay (propagated to workers). http_interceptor: Option>, /// Running jobs (main LLM-driven jobs). @@ -102,9 +101,9 @@ impl Scheduler { } } - /// Set the SSE broadcast sender for live job event streaming. - pub fn set_sse_sender(&mut self, tx: tokio::sync::broadcast::Sender) { - self.sse_tx = Some(tx); + /// Set the SSE manager for live job event streaming. + pub fn set_sse_sender(&mut self, sse: Arc) { + self.sse_tx = Some(sse); } /// Set the HTTP interceptor for trace recording/replay. diff --git a/src/agent/thread_ops.rs b/src/agent/thread_ops.rs index eec29099..ddfd0c0f 100644 --- a/src/agent/thread_ops.rs +++ b/src/agent/thread_ops.rs @@ -1646,7 +1646,7 @@ impl Agent { }; match ext_mgr - .configure_token(&pending.extension_name, token) + .configure_token(&pending.extension_name, token, &message.user_id) .await { Ok(result) if result.activated => { diff --git a/src/app.rs b/src/app.rs index 94d949be..edd547d3 100644 --- a/src/app.rs +++ b/src/app.rs @@ -327,7 +327,7 @@ impl AppBuilder { .with_search_config(&self.config.search); if let Some(ref emb) = embeddings { - ws = ws.with_embeddings_cached(emb.clone(), emb_cache_config); + ws = ws.with_embeddings_cached(emb.clone(), emb_cache_config.clone()); } // Wire workspace-level settings (read scopes, memory layers) @@ -341,7 +341,35 @@ impl AppBuilder { } ws = ws.with_memory_layers(self.config.workspace.memory_layers.clone()); let ws = Arc::new(ws); - tools.register_memory_tools(Arc::clone(&ws)); + + // Detect multi-tenant mode: when GATEWAY_USER_TOKENS is configured, + // each authenticated user needs their own workspace scope. Use + // WorkspacePool (which implements WorkspaceResolver) to create + // per-user workspaces on demand instead of sharing the startup + // workspace across all users. + let is_multi_tenant = self + .config + .channels + .gateway + .as_ref() + .is_some_and(|gw| gw.user_tokens.is_some()); + + if is_multi_tenant { + let pool = Arc::new(crate::channels::web::server::WorkspacePool::new( + Arc::clone(db), + embeddings.clone(), + emb_cache_config, + self.config.search.clone(), + self.config.workspace.clone(), + )); + tools.register_memory_tools_with_resolver(pool); + tracing::info!( + "Memory tools configured with per-user workspace resolver (multi-tenant mode)" + ); + } else { + tools.register_memory_tools(Arc::clone(&ws)); + } + Some(ws) } else { None diff --git a/src/channels/web/auth.rs b/src/channels/web/auth.rs index b2fa4e4f..7dc8adb4 100644 --- a/src/channels/web/auth.rs +++ b/src/channels/web/auth.rs @@ -1,17 +1,133 @@ //! Bearer token authentication middleware for the web gateway. +//! +//! Supports multi-user mode: each token maps to a `UserIdentity` that carries +//! the user_id. The identity is inserted into request extensions so downstream +//! handlers can extract it via `AuthenticatedUser`. + +use std::collections::HashMap; use axum::{ - extract::{Request, State}, - http::{HeaderMap, Method, StatusCode}, + extract::{FromRequestParts, Request, State}, + http::{HeaderMap, Method, StatusCode, request::Parts}, middleware::Next, response::{IntoResponse, Response}, }; +use sha2::{Digest, Sha256}; use subtle::ConstantTimeEq; -/// Shared auth state injected via axum middleware state. +/// Identity resolved from a bearer token. +#[derive(Debug, Clone)] +pub struct UserIdentity { + pub user_id: String, + /// Additional user scopes this identity can read from. + pub workspace_read_scopes: Vec, +} + +/// Hash a token with SHA-256 for constant-size, timing-safe storage. +fn hash_token(token: &str) -> [u8; 32] { + let mut hasher = Sha256::new(); + hasher.update(token.as_bytes()); + hasher.finalize().into() +} + +/// Multi-user auth state: maps token hashes to user identities. +/// +/// Tokens are SHA-256 hashed on construction so they are never stored in +/// plaintext. Authentication compares fixed-size (32-byte) digests using +/// constant-time comparison, eliminating both length-oracle timing leaks +/// and accidental token exposure in memory dumps. +/// +/// In single-user mode (the default), contains exactly one entry. #[derive(Clone)] -pub struct AuthState { - pub token: String, +pub struct MultiAuthState { + /// Maps SHA-256(token) → identity. Tokens are never stored in cleartext. + hashed_tokens: Vec<([u8; 32], UserIdentity)>, + /// Original first token kept only for single-user startup printing. + /// Not used for authentication. + display_token: Option, +} + +impl MultiAuthState { + /// Create a single-user auth state (backwards compatible). + pub fn single(token: String, user_id: String) -> Self { + let hash = hash_token(&token); + Self { + hashed_tokens: vec![( + hash, + UserIdentity { + user_id, + workspace_read_scopes: Vec::new(), + }, + )], + display_token: Some(token), + } + } + + /// Create a multi-user auth state from a map of tokens to identities. + pub fn multi(tokens: HashMap) -> Self { + let hashed_tokens: Vec<([u8; 32], UserIdentity)> = tokens + .into_iter() + .map(|(tok, identity)| (hash_token(&tok), identity)) + .collect(); + Self { + hashed_tokens, + display_token: None, + } + } + + /// Authenticate a token, returning the associated identity if valid. + /// + /// Uses SHA-256 hashing + constant-time comparison (`subtle::ConstantTimeEq`) + /// to prevent timing side-channels. Both the candidate and stored tokens are + /// hashed to 32-byte digests, eliminating length-oracle leaks. Iterates all + /// entries regardless of match to avoid early-exit timing differences. + /// O(n) in the number of configured users — negligible for typical + /// deployments (< 10 users). + pub fn authenticate(&self, candidate: &str) -> Option<&UserIdentity> { + let candidate_hash = hash_token(candidate); + let mut matched: Option<&UserIdentity> = None; + for (stored_hash, identity) in &self.hashed_tokens { + if bool::from(candidate_hash.ct_eq(stored_hash)) { + matched = Some(identity); + } + } + matched + } + + /// Get the first token for backwards-compatible printing at startup. + /// + /// Only available in single-user mode; returns `None` in multi-user mode + /// to avoid exposing tokens. + pub fn first_token(&self) -> Option<&str> { + self.display_token.as_deref() + } + + /// Get the first user identity (for single-user fallback). + pub fn first_identity(&self) -> Option<&UserIdentity> { + self.hashed_tokens.first().map(|(_, id)| id) + } +} + +/// Axum extractor that provides the authenticated user identity. +/// +/// Only available on routes behind `auth_middleware`. Extracts the +/// `UserIdentity` that the middleware inserted into request extensions. +pub struct AuthenticatedUser(pub UserIdentity); + +impl FromRequestParts for AuthenticatedUser +where + S: Send + Sync, +{ + type Rejection = (StatusCode, &'static str); + + async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result { + parts + .extensions + .get::() + .cloned() + .map(AuthenticatedUser) + .ok_or((StatusCode::UNAUTHORIZED, "Not authenticated")) + } } /// Whether query-string token auth is allowed for this request. @@ -51,29 +167,34 @@ fn query_token(request: &Request) -> Option { /// Auth middleware that validates bearer token from header or query param. /// /// SSE connections can't set headers from `EventSource`, so we also accept -/// `?token=xxx` as a query parameter, but only on SSE endpoints. +/// `?token=xxx` as a query parameter, but only on SSE/WS endpoints. +/// +/// On successful authentication, inserts the matching `UserIdentity` into +/// request extensions for downstream extraction via `AuthenticatedUser`. pub async fn auth_middleware( - State(auth): State, + State(auth): State, headers: HeaderMap, - request: Request, + mut request: Request, next: Next, ) -> Response { - // Try Authorization header first (constant-time comparison). + // Try Authorization header first. // RFC 6750 Section 2.1: auth-scheme comparison is case-insensitive. if let Some(auth_header) = headers.get("authorization") && let Ok(value) = auth_header.to_str() && value.len() > 7 && value[..7].eq_ignore_ascii_case("Bearer ") - && bool::from(value.as_bytes()[7..].ct_eq(auth.token.as_bytes())) + && let Some(identity) = auth.authenticate(&value[7..]) { + request.extensions_mut().insert(identity.clone()); return next.run(request).await; } - // Fall back to query parameter, but only for SSE endpoints (constant-time comparison). + // Fall back to query parameter, but only for SSE/WS endpoints. if allows_query_token_auth(&request) && let Some(token) = query_token(&request) - && bool::from(token.as_bytes().ct_eq(auth.token.as_bytes())) + && let Some(identity) = auth.authenticate(&token) { + request.extensions_mut().insert(identity.clone()); return next.run(request).await; } @@ -83,15 +204,61 @@ pub async fn auth_middleware( #[cfg(test)] mod tests { use super::*; - use crate::testing::credentials::{TEST_AUTH_SECRET_TOKEN, TEST_BEARER_TOKEN}; + use crate::testing::credentials::TEST_AUTH_SECRET_TOKEN; #[test] - fn test_auth_state_clone() { - let state = AuthState { - token: TEST_BEARER_TOKEN.to_string(), - }; - let cloned = state.clone(); - assert_eq!(cloned.token, TEST_BEARER_TOKEN); + fn test_multi_auth_state_single() { + let state = MultiAuthState::single("tok-123".to_string(), "alice".to_string()); + let identity = state.authenticate("tok-123"); + assert!(identity.is_some()); + assert_eq!(identity.unwrap().user_id, "alice"); + } + + #[test] + fn test_multi_auth_state_reject_wrong_token() { + let state = MultiAuthState::single("tok-123".to_string(), "alice".to_string()); + assert!(state.authenticate("wrong-token").is_none()); + } + + #[test] + fn test_multi_auth_state_multi_users() { + let mut tokens = HashMap::new(); + tokens.insert( + "tok-alice".to_string(), + UserIdentity { + user_id: "alice".to_string(), + workspace_read_scopes: Vec::new(), + }, + ); + tokens.insert( + "tok-bob".to_string(), + UserIdentity { + user_id: "bob".to_string(), + workspace_read_scopes: Vec::new(), + }, + ); + let state = MultiAuthState::multi(tokens); + + let alice = state.authenticate("tok-alice").unwrap(); + assert_eq!(alice.user_id, "alice"); + + let bob = state.authenticate("tok-bob").unwrap(); + assert_eq!(bob.user_id, "bob"); + + assert!(state.authenticate("tok-charlie").is_none()); + } + + #[test] + fn test_multi_auth_state_first_token() { + let state = MultiAuthState::single("my-token".to_string(), "user1".to_string()); + assert_eq!(state.first_token(), Some("my-token")); + } + + #[test] + fn test_multi_auth_state_first_identity() { + let state = MultiAuthState::single("my-token".to_string(), "user1".to_string()); + let identity = state.first_identity().unwrap(); + assert_eq!(identity.user_id, "user1"); } use axum::Router; @@ -107,9 +274,7 @@ mod tests { /// Router with streaming endpoints (query auth allowed) and regular /// endpoints (query auth rejected). fn test_app(token: &str) -> Router { - let state = AuthState { - token: token.to_string(), - }; + let state = MultiAuthState::single(token.to_string(), "test-user".to_string()); Router::new() .route("/api/chat/events", get(dummy_handler)) .route("/api/logs/events", get(dummy_handler)) @@ -306,4 +471,200 @@ mod tests { let resp = app.oneshot(req).await.unwrap(); assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); } + + // --- Multi-tenant auth integration tests --- + + /// Handler that extracts `AuthenticatedUser` and returns the resolved user_id. + async fn identity_handler(AuthenticatedUser(identity): AuthenticatedUser) -> String { + identity.user_id + } + + /// Handler that extracts `AuthenticatedUser` and returns workspace_read_scopes as JSON. + async fn scopes_handler(AuthenticatedUser(identity): AuthenticatedUser) -> String { + serde_json::to_string(&identity.workspace_read_scopes).unwrap() + } + + /// Build a multi-user router where each token maps to a distinct identity. + fn multi_user_app(tokens: HashMap) -> Router { + let state = MultiAuthState::multi(tokens); + Router::new() + .route("/api/chat/events", get(identity_handler)) + .route("/api/chat/send", post(identity_handler)) + .route("/api/scopes", get(scopes_handler)) + .layer(middleware::from_fn_with_state(state, auth_middleware)) + } + + fn two_user_tokens() -> HashMap { + let mut tokens = HashMap::new(); + tokens.insert( + "tok-alice".to_string(), + UserIdentity { + user_id: "alice".to_string(), + workspace_read_scopes: vec!["shared".to_string()], + }, + ); + tokens.insert( + "tok-bob".to_string(), + UserIdentity { + user_id: "bob".to_string(), + workspace_read_scopes: vec!["shared".to_string(), "alice".to_string()], + }, + ); + tokens + } + + #[tokio::test] + async fn test_multi_user_alice_token_resolves_to_alice() { + let app = multi_user_app(two_user_tokens()); + let req = Request::builder() + .uri("/api/chat/events") + .header("Authorization", "Bearer tok-alice") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + assert_eq!(body, "alice"); + } + + #[tokio::test] + async fn test_multi_user_bob_token_resolves_to_bob() { + let app = multi_user_app(two_user_tokens()); + let req = Request::builder() + .uri("/api/chat/events") + .header("Authorization", "Bearer tok-bob") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + assert_eq!(body, "bob"); + } + + #[tokio::test] + async fn test_multi_user_sequential_tokens_resolve_independently() { + // Send both alice and bob tokens sequentially and verify each gets + // the correct identity — guards against token map corruption. + let tokens = two_user_tokens(); + + let app1 = multi_user_app(tokens.clone()); + let req = Request::builder() + .uri("/api/chat/events") + .header("Authorization", "Bearer tok-alice") + .body(Body::empty()) + .unwrap(); + let resp = app1.oneshot(req).await.unwrap(); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + assert_eq!(body, "alice"); + + let app2 = multi_user_app(tokens); + let req = Request::builder() + .uri("/api/chat/events") + .header("Authorization", "Bearer tok-bob") + .body(Body::empty()) + .unwrap(); + let resp = app2.oneshot(req).await.unwrap(); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + assert_eq!(body, "bob"); + } + + #[tokio::test] + async fn test_multi_user_unknown_token_rejected() { + let app = multi_user_app(two_user_tokens()); + let req = Request::builder() + .uri("/api/chat/events") + .header("Authorization", "Bearer tok-charlie") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn test_multi_user_workspace_read_scopes_propagated() { + let app = multi_user_app(two_user_tokens()); + + // Alice has ["shared"] + let req = Request::builder() + .uri("/api/scopes") + .header("Authorization", "Bearer tok-alice") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + let scopes: Vec = serde_json::from_slice(&body).unwrap(); + assert_eq!(scopes, vec!["shared"]); + } + + #[tokio::test] + async fn test_multi_user_bob_has_two_scopes() { + let app = multi_user_app(two_user_tokens()); + + // Bob has ["shared", "alice"] + let req = Request::builder() + .uri("/api/scopes") + .header("Authorization", "Bearer tok-bob") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + let scopes: Vec = serde_json::from_slice(&body).unwrap(); + assert_eq!(scopes, vec!["shared", "alice"]); + } + + #[tokio::test] + async fn test_multi_user_query_param_resolves_correct_identity() { + let app = multi_user_app(two_user_tokens()); + let req = Request::builder() + .uri("/api/chat/events?token=tok-bob") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + assert_eq!(body, "bob"); + } + + #[tokio::test] + async fn test_multi_user_post_with_bearer_resolves_identity() { + let app = multi_user_app(two_user_tokens()); + let req = Request::builder() + .method(Method::POST) + .uri("/api/chat/send") + .header("Authorization", "Bearer tok-alice") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + assert_eq!(body, "alice"); + } + + #[tokio::test] + async fn test_multi_user_empty_scopes_for_single_user() { + // Single-user mode creates identity with empty workspace_read_scopes. + let state = MultiAuthState::single("tok-only".to_string(), "solo".to_string()); + let app = Router::new() + .route("/api/scopes", get(scopes_handler)) + .layer(middleware::from_fn_with_state(state, auth_middleware)); + let req = Request::builder() + .uri("/api/scopes") + .header("Authorization", "Bearer tok-only") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + let scopes: Vec = serde_json::from_slice(&body).unwrap(); + assert!(scopes.is_empty()); + } + + #[tokio::test] + async fn test_prefix_and_extension_tokens_rejected() { + // Verifies that prefix/suffix variants of valid tokens are rejected. + // Note: the constant-time property is enforced structurally by use of + // subtle::ConstantTimeEq and cannot be verified via outcome testing. + let state = MultiAuthState::single("long-secret-token".to_string(), "user".to_string()); + assert!(state.authenticate("long-secret").is_none()); + assert!(state.authenticate("long-secret-token-extra").is_none()); + } } diff --git a/src/channels/web/handlers/chat.rs b/src/channels/web/handlers/chat.rs index 5cb2b9ea..9753c015 100644 --- a/src/channels/web/handlers/chat.rs +++ b/src/channels/web/handlers/chat.rs @@ -12,22 +12,24 @@ use serde::Deserialize; use uuid::Uuid; use crate::channels::IncomingMessage; +use crate::channels::web::auth::AuthenticatedUser; use crate::channels::web::server::GatewayState; use crate::channels::web::types::*; use crate::channels::web::util::{build_turns_from_db_messages, truncate_preview}; pub async fn chat_send_handler( State(state): State>, + AuthenticatedUser(identity): AuthenticatedUser, Json(req): Json, ) -> Result<(StatusCode, Json), (StatusCode, String)> { - if !state.chat_rate_limiter.check() { + if !state.chat_rate_limiter.check(&identity.user_id) { return Err(( StatusCode::TOO_MANY_REQUESTS, "Rate limit exceeded. Try again shortly.".to_string(), )); } - let mut msg = IncomingMessage::new("gateway", &state.user_id, &req.content); + let mut msg = IncomingMessage::new("gateway", &identity.user_id, &req.content); if let Some(ref thread_id) = req.thread_id { msg = msg.with_thread(thread_id); @@ -74,6 +76,7 @@ pub async fn chat_send_handler( pub async fn chat_approval_handler( State(state): State>, + AuthenticatedUser(identity): AuthenticatedUser, Json(req): Json, ) -> Result<(StatusCode, Json), (StatusCode, String)> { let (approved, always) = match req.action.as_str() { @@ -109,7 +112,7 @@ pub async fn chat_approval_handler( ) })?; - let mut msg = IncomingMessage::new("gateway", &state.user_id, content); + let mut msg = IncomingMessage::new("gateway", &identity.user_id, content); if let Some(ref thread_id) = req.thread_id { msg = msg.with_thread(thread_id); @@ -150,6 +153,7 @@ pub async fn chat_approval_handler( /// The token never touches the LLM, chat history, or SSE stream. pub async fn chat_auth_token_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Json(req): Json, ) -> Result, (StatusCode, String)> { let ext_mgr = state.extension_manager.as_ref().ok_or(( @@ -158,7 +162,7 @@ pub async fn chat_auth_token_handler( ))?; match ext_mgr - .configure_token(&req.extension_name, &req.token) + .configure_token(&req.extension_name, &req.token, &user.user_id) .await { Ok(result) => { @@ -169,20 +173,26 @@ pub async fn chat_auth_token_handler( resp.instructions = result.verification.as_ref().map(|v| v.instructions.clone()); if result.verification.is_some() { - state.sse.broadcast(SseEvent::AuthRequired { - extension_name: req.extension_name.clone(), - instructions: Some(result.message), - auth_url: None, - setup_url: None, - }); + state.sse.broadcast_for_user( + &user.user_id, + SseEvent::AuthRequired { + extension_name: req.extension_name.clone(), + instructions: Some(result.message), + auth_url: None, + setup_url: None, + }, + ); } else { - clear_auth_mode(&state).await; + clear_auth_mode(&state, &user.user_id).await; - state.sse.broadcast(SseEvent::AuthCompleted { - extension_name: req.extension_name.clone(), - success: true, - message: result.message, - }); + state.sse.broadcast_for_user( + &user.user_id, + SseEvent::AuthCompleted { + extension_name: req.extension_name.clone(), + success: true, + message: result.message, + }, + ); } Ok(Json(resp)) @@ -190,12 +200,15 @@ pub async fn chat_auth_token_handler( Err(e) => { let msg = e.to_string(); if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) { - state.sse.broadcast(SseEvent::AuthRequired { - extension_name: req.extension_name.clone(), - instructions: Some(msg.clone()), - auth_url: None, - setup_url: None, - }); + state.sse.broadcast_for_user( + &user.user_id, + SseEvent::AuthRequired { + extension_name: req.extension_name.clone(), + instructions: Some(msg.clone()), + auth_url: None, + setup_url: None, + }, + ); } Ok(Json(ActionResponse::fail(msg))) } @@ -205,16 +218,17 @@ pub async fn chat_auth_token_handler( /// Cancel an in-progress auth flow. pub async fn chat_auth_cancel_handler( State(state): State>, + AuthenticatedUser(identity): AuthenticatedUser, Json(_req): Json, ) -> Result, (StatusCode, String)> { - clear_auth_mode(&state).await; + clear_auth_mode(&state, &identity.user_id).await; Ok(Json(ActionResponse::ok("Auth cancelled"))) } /// Clear pending auth mode on the active thread. -pub async fn clear_auth_mode(state: &GatewayState) { +pub async fn clear_auth_mode(state: &GatewayState, user_id: &str) { if let Some(ref sm) = state.session_manager { - let session = sm.get_or_create_session(&state.user_id).await; + let session = sm.get_or_create_session(user_id).await; let mut sess = session.lock().await; if let Some(thread_id) = sess.active_thread && let Some(thread) = sess.threads.get_mut(&thread_id) @@ -226,8 +240,9 @@ pub async fn clear_auth_mode(state: &GatewayState) { pub async fn chat_events_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result { - state.sse.subscribe().ok_or(( + state.sse.subscribe(Some(user.user_id)).ok_or(( StatusCode::SERVICE_UNAVAILABLE, "Too many connections".to_string(), )) @@ -237,6 +252,7 @@ pub async fn chat_ws_handler( headers: axum::http::HeaderMap, ws: WebSocketUpgrade, State(state): State>, + AuthenticatedUser(identity): AuthenticatedUser, ) -> Result { // Validate Origin header to prevent cross-site WebSocket hijacking. let origin = headers @@ -262,7 +278,9 @@ pub async fn chat_ws_handler( "WebSocket origin not allowed".to_string(), )); } - Ok(ws.on_upgrade(move |socket| crate::channels::web::ws::handle_ws_connection(socket, state))) + Ok(ws.on_upgrade(move |socket| { + crate::channels::web::ws::handle_ws_connection(socket, state, identity) + })) } #[derive(Deserialize)] @@ -274,6 +292,7 @@ pub struct HistoryQuery { pub async fn chat_history_handler( State(state): State>, + AuthenticatedUser(identity): AuthenticatedUser, Query(query): Query, ) -> Result, (StatusCode, String)> { let session_manager = state.session_manager.as_ref().ok_or(( @@ -281,7 +300,9 @@ pub async fn chat_history_handler( "Session manager not available".to_string(), ))?; - let session = session_manager.get_or_create_session(&state.user_id).await; + let session = session_manager + .get_or_create_session(&identity.user_id) + .await; let limit = query.limit.unwrap_or(50); let before_cursor = query @@ -314,7 +335,7 @@ pub async fn chat_history_handler( && let Some(ref store) = state.store { let owned = store - .conversation_belongs_to_user(thread_id, &state.user_id) + .conversation_belongs_to_user(thread_id, &identity.user_id) .await .unwrap_or(false); if !owned { @@ -434,24 +455,27 @@ pub async fn chat_history_handler( pub async fn chat_threads_handler( State(state): State>, + AuthenticatedUser(identity): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let session_manager = state.session_manager.as_ref().ok_or(( StatusCode::SERVICE_UNAVAILABLE, "Session manager not available".to_string(), ))?; - let session = session_manager.get_or_create_session(&state.user_id).await; + let session = session_manager + .get_or_create_session(&identity.user_id) + .await; // Try DB first for persistent thread list if let Some(ref store) = state.store { // Auto-create assistant thread if it doesn't exist let assistant_id = store - .get_or_create_assistant_conversation(&state.user_id, "gateway") + .get_or_create_assistant_conversation(&identity.user_id, "gateway") .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; if let Ok(summaries) = store - .list_conversations_all_channels(&state.user_id, 50) + .list_conversations_all_channels(&identity.user_id, 50) .await { let mut assistant_thread = None; @@ -534,13 +558,16 @@ pub async fn chat_threads_handler( pub async fn chat_new_thread_handler( State(state): State>, + AuthenticatedUser(identity): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let session_manager = state.session_manager.as_ref().ok_or(( StatusCode::SERVICE_UNAVAILABLE, "Session manager not available".to_string(), ))?; - let session = session_manager.get_or_create_session(&state.user_id).await; + let session = session_manager + .get_or_create_session(&identity.user_id) + .await; let (thread_id, info) = { let mut sess = session.lock().await; let thread = sess.create_thread(); @@ -562,12 +589,12 @@ pub async fn chat_new_thread_handler( // so that the subsequent loadThreads() call from the frontend sees it. if let Some(ref store) = state.store { match store - .ensure_conversation(thread_id, "gateway", &state.user_id, None) + .ensure_conversation(thread_id, "gateway", &identity.user_id, None) .await { Ok(true) => {} Ok(false) => tracing::warn!( - user = %state.user_id, + user = %identity.user_id, thread_id = %thread_id, "Skipped persisting new thread due to ownership/channel conflict" ), diff --git a/src/channels/web/handlers/extensions.rs b/src/channels/web/handlers/extensions.rs index 855fba3e..d705591e 100644 --- a/src/channels/web/handlers/extensions.rs +++ b/src/channels/web/handlers/extensions.rs @@ -8,11 +8,13 @@ use axum::{ http::StatusCode, }; +use crate::channels::web::auth::AuthenticatedUser; use crate::channels::web::server::GatewayState; use crate::channels::web::types::*; pub async fn extensions_list_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let ext_mgr = state.extension_manager.as_ref().ok_or(( StatusCode::NOT_IMPLEMENTED, @@ -20,7 +22,7 @@ pub async fn extensions_list_handler( ))?; let installed = ext_mgr - .list(None, false) + .list(None, false, &user.user_id) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; @@ -80,6 +82,7 @@ pub async fn extensions_list_handler( pub async fn extensions_tools_handler( State(state): State>, + AuthenticatedUser(_user): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let registry = state.tool_registry.as_ref().ok_or(( StatusCode::SERVICE_UNAVAILABLE, @@ -100,6 +103,7 @@ pub async fn extensions_tools_handler( pub async fn extensions_install_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Json(req): Json, ) -> Result, (StatusCode, String)> { let ext_mgr = state.extension_manager.as_ref().ok_or(( @@ -116,7 +120,7 @@ pub async fn extensions_install_handler( }); match ext_mgr - .install(&req.name, req.url.as_deref(), kind_hint) + .install(&req.name, req.url.as_deref(), kind_hint, &user.user_id) .await { Ok(result) => Ok(Json(ActionResponse::ok(result.message))), @@ -126,6 +130,7 @@ pub async fn extensions_install_handler( pub async fn extensions_remove_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(name): Path, ) -> Result, (StatusCode, String)> { let ext_mgr = state.extension_manager.as_ref().ok_or(( @@ -133,7 +138,7 @@ pub async fn extensions_remove_handler( "Extension manager not available (secrets store required)".to_string(), ))?; - match ext_mgr.remove(&name).await { + match ext_mgr.remove(&name, &user.user_id).await { Ok(message) => Ok(Json(ActionResponse::ok(message))), Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))), } diff --git a/src/channels/web/handlers/jobs.rs b/src/channels/web/handlers/jobs.rs index 5a94e055..35adeec6 100644 --- a/src/channels/web/handlers/jobs.rs +++ b/src/channels/web/handlers/jobs.rs @@ -11,11 +11,13 @@ use axum::{ use serde::Deserialize; use uuid::Uuid; +use crate::channels::web::auth::AuthenticatedUser; use crate::channels::web::server::GatewayState; use crate::channels::web::types::*; pub async fn jobs_list_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let store = state.store.as_ref().ok_or(( StatusCode::SERVICE_UNAVAILABLE, @@ -25,8 +27,8 @@ pub async fn jobs_list_handler( let mut jobs: Vec = Vec::new(); let mut seen_ids: HashSet = HashSet::new(); - // Fetch sandbox jobs from database. - match store.list_sandbox_jobs().await { + // Fetch sandbox jobs scoped to this user. + match store.list_sandbox_jobs_for_user(&user.user_id).await { Ok(sandbox_jobs) => { for j in &sandbox_jobs { let ui_state = match j.status.as_str() { @@ -50,8 +52,8 @@ pub async fn jobs_list_handler( } } - // Fetch agent (non-sandbox) jobs from database, deduplicating by ID. - match store.list_agent_jobs().await { + // Fetch agent (non-sandbox) jobs scoped to this user, deduplicating by ID. + match store.list_agent_jobs_for_user(&user.user_id).await { Ok(agent_jobs) => { for j in &agent_jobs { if seen_ids.contains(&j.id) { @@ -80,6 +82,7 @@ pub async fn jobs_list_handler( pub async fn jobs_summary_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let store = state.store.as_ref().ok_or(( StatusCode::SERVICE_UNAVAILABLE, @@ -93,8 +96,8 @@ pub async fn jobs_summary_handler( let mut failed = 0; let mut stuck = 0; - // Sandbox job counts. - match store.sandbox_job_summary().await { + // Sandbox job counts scoped to this user. + match store.sandbox_job_summary_for_user(&user.user_id).await { Ok(s) => { total += s.total; pending += s.creating; @@ -107,8 +110,8 @@ pub async fn jobs_summary_handler( } } - // Agent job counts. - match store.agent_job_summary().await { + // Agent job counts scoped to this user. + match store.agent_job_summary_for_user(&user.user_id).await { Ok(s) => { total += s.total; pending += s.pending; @@ -134,6 +137,7 @@ pub async fn jobs_summary_handler( pub async fn jobs_detail_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(id): Path, ) -> Result, (StatusCode, String)> { let store = state.store.as_ref().ok_or(( @@ -145,169 +149,213 @@ pub async fn jobs_detail_handler( .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?; // Try sandbox job from DB first. - if let Ok(Some(job)) = store.get_sandbox_job(job_id).await { - let browse_id = std::path::Path::new(&job.project_dir) - .file_name() - .map(|n| n.to_string_lossy().to_string()) - .unwrap_or_else(|| job.id.to_string()); + match store.get_sandbox_job(job_id).await { + Ok(Some(job)) => { + if job.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + let browse_id = std::path::Path::new(&job.project_dir) + .file_name() + .map(|n| n.to_string_lossy().to_string()) + .unwrap_or_else(|| job.id.to_string()); - let ui_state = match job.status.as_str() { - "creating" => "pending", - "running" => "in_progress", - s => s, - }; + let ui_state = match job.status.as_str() { + "creating" => "pending", + "running" => "in_progress", + s => s, + }; - let elapsed_secs = job.started_at.map(|start| { - let end = job.completed_at.unwrap_or_else(chrono::Utc::now); - (end - start).num_seconds().max(0) as u64 - }); - - // Synthesize transitions from timestamps. - let mut transitions = Vec::new(); - if let Some(started) = job.started_at { - transitions.push(TransitionInfo { - from: "creating".to_string(), - to: "running".to_string(), - timestamp: started.to_rfc3339(), - reason: None, + let elapsed_secs = job.started_at.map(|start| { + let end = job.completed_at.unwrap_or_else(chrono::Utc::now); + (end - start).num_seconds().max(0) as u64 }); - } - if let Some(completed) = job.completed_at { - transitions.push(TransitionInfo { - from: "running".to_string(), - to: job.status.clone(), - timestamp: completed.to_rfc3339(), - reason: job.failure_reason.clone(), - }); - } - let mode = store.get_sandbox_job_mode(job.id).await.ok().flatten(); - let is_claude_code = mode.as_deref() == Some("claude_code"); + // Synthesize transitions from timestamps. + let mut transitions = Vec::new(); + if let Some(started) = job.started_at { + transitions.push(TransitionInfo { + from: "creating".to_string(), + to: "running".to_string(), + timestamp: started.to_rfc3339(), + reason: None, + }); + } + if let Some(completed) = job.completed_at { + transitions.push(TransitionInfo { + from: "running".to_string(), + to: job.status.clone(), + timestamp: completed.to_rfc3339(), + reason: job.failure_reason.clone(), + }); + } - return Ok(Json(JobDetailResponse { - id: job.id, - title: job.task.clone(), - description: String::new(), - state: ui_state.to_string(), - user_id: job.user_id.clone(), - created_at: job.created_at.to_rfc3339(), - started_at: job.started_at.map(|dt| dt.to_rfc3339()), - completed_at: job.completed_at.map(|dt| dt.to_rfc3339()), - elapsed_secs, - project_dir: Some(job.project_dir.clone()), - browse_url: Some(format!("/projects/{}/", browse_id)), - job_mode: mode.filter(|m| m != "worker"), - transitions, - can_restart: state.job_manager.is_some(), - can_prompt: is_claude_code && state.prompt_queue.is_some(), - job_kind: Some("sandbox".to_string()), - })); + let mode = store.get_sandbox_job_mode(job.id).await.ok().flatten(); + let is_claude_code = mode.as_deref() == Some("claude_code"); + + return Ok(Json(JobDetailResponse { + id: job.id, + title: job.task.clone(), + description: String::new(), + state: ui_state.to_string(), + user_id: job.user_id.clone(), + created_at: job.created_at.to_rfc3339(), + started_at: job.started_at.map(|dt| dt.to_rfc3339()), + completed_at: job.completed_at.map(|dt| dt.to_rfc3339()), + elapsed_secs, + project_dir: Some(job.project_dir.clone()), + browse_url: Some(format!("/projects/{}/", browse_id)), + job_mode: mode.filter(|m| m != "worker"), + transitions, + can_restart: state.job_manager.is_some(), + can_prompt: is_claude_code && state.prompt_queue.is_some(), + job_kind: Some("sandbox".to_string()), + })); + } + Ok(None) => {} + Err(e) => { + return Err(( + StatusCode::INTERNAL_SERVER_ERROR, + format!("Database error: {}", e), + )); + } } // Fall back to agent job from DB. - if let Ok(Some(ctx)) = store.get_job(job_id).await { - let elapsed_secs = ctx.started_at.map(|start| { - let end = ctx.completed_at.unwrap_or_else(chrono::Utc::now); - (end - start).num_seconds().max(0) as u64 - }); + match store.get_job(job_id).await { + Ok(Some(ctx)) => { + if ctx.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + let elapsed_secs = ctx.started_at.map(|start| { + let end = ctx.completed_at.unwrap_or_else(chrono::Utc::now); + (end - start).num_seconds().max(0) as u64 + }); - // Only show prompt bar for jobs that have a running worker (Pending/InProgress). - // Stuck jobs have no active worker loop, so messages would be silently dropped. - let is_promptable = matches!( - ctx.state, - crate::context::JobState::Pending | crate::context::JobState::InProgress - ); - return Ok(Json(JobDetailResponse { - id: ctx.job_id, - title: ctx.title.clone(), - description: ctx.description.clone(), - state: ctx.state.to_string(), - user_id: ctx.user_id.clone(), - created_at: ctx.created_at.to_rfc3339(), - started_at: ctx.started_at.map(|dt| dt.to_rfc3339()), - completed_at: ctx.completed_at.map(|dt| dt.to_rfc3339()), - elapsed_secs, - project_dir: None, - browse_url: None, - job_mode: None, - transitions: Vec::new(), - can_restart: state.scheduler.is_some(), - can_prompt: is_promptable && state.scheduler.is_some(), - job_kind: Some("agent".to_string()), - })); + // Only show prompt bar for jobs that have a running worker (Pending/InProgress). + // Stuck jobs have no active worker loop, so messages would be silently dropped. + let is_promptable = matches!( + ctx.state, + crate::context::JobState::Pending | crate::context::JobState::InProgress + ); + Ok(Json(JobDetailResponse { + id: ctx.job_id, + title: ctx.title.clone(), + description: ctx.description.clone(), + state: ctx.state.to_string(), + user_id: ctx.user_id.clone(), + created_at: ctx.created_at.to_rfc3339(), + started_at: ctx.started_at.map(|dt| dt.to_rfc3339()), + completed_at: ctx.completed_at.map(|dt| dt.to_rfc3339()), + elapsed_secs, + project_dir: None, + browse_url: None, + job_mode: None, + transitions: Vec::new(), + can_restart: state.scheduler.is_some(), + can_prompt: is_promptable && state.scheduler.is_some(), + job_kind: Some("agent".to_string()), + })) + } + Ok(None) => Err((StatusCode::NOT_FOUND, "Job not found".to_string())), + Err(e) => Err(( + StatusCode::INTERNAL_SERVER_ERROR, + format!("Database error: {}", e), + )), } - - Err((StatusCode::NOT_FOUND, "Job not found".to_string())) } pub async fn jobs_cancel_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(id): Path, ) -> Result, (StatusCode, String)> { let job_id = Uuid::parse_str(&id) .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?; // Try sandbox job cancellation. - if let Some(ref store) = state.store - && let Ok(Some(job)) = store.get_sandbox_job(job_id).await - { - if job.status == "running" || job.status == "creating" { - // Stop the container if we have a job manager. - if let Some(ref jm) = state.job_manager - && let Err(e) = jm.stop_job(job_id).await - { - tracing::warn!(job_id = %job_id, error = %e, "Failed to stop container during cancellation"); + if let Some(ref store) = state.store { + match store.get_sandbox_job(job_id).await { + Ok(Some(job)) => { + if job.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + if job.status == "running" || job.status == "creating" { + if let Some(ref jm) = state.job_manager + && let Err(e) = jm.stop_job(job_id).await + { + tracing::warn!(job_id = %job_id, error = %e, "Failed to stop container during cancellation"); + } + store + .update_sandbox_job_status( + job_id, + "failed", + Some(false), + Some("Cancelled by user"), + None, + Some(chrono::Utc::now()), + ) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + } + return Ok(Json(serde_json::json!({ + "status": "cancelled", + "job_id": job_id, + }))); + } + Ok(None) => {} + Err(e) => { + return Err(( + StatusCode::INTERNAL_SERVER_ERROR, + format!("Database error: {}", e), + )); } - store - .update_sandbox_job_status( - job_id, - "failed", - Some(false), - Some("Cancelled by user"), - None, - Some(chrono::Utc::now()), - ) - .await - .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; } - return Ok(Json(serde_json::json!({ - "status": "cancelled", - "job_id": job_id, - }))); } // Fall back to agent job cancellation: stop the worker via the scheduler // (which updates the in-memory ContextManager AND aborts the task handle), // then persist the status to the DB as a fallback. - if let Some(ref store) = state.store - && let Ok(Some(job)) = store.get_job(job_id).await - { - if job.state.is_active() { - // Try to stop via scheduler (aborts the worker task + updates - // in-memory ContextManager). This is best-effort — the job may - // not be in the scheduler map if it already finished. - if let Some(ref slot) = state.scheduler - && let Some(ref scheduler) = *slot.read().await - { - let _ = scheduler.stop(job_id).await; - } + if let Some(ref store) = state.store { + match store.get_job(job_id).await { + Ok(Some(job)) => { + if job.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + if job.state.is_active() { + // Try to stop via scheduler (aborts the worker task + updates + // in-memory ContextManager). This is best-effort — the job may + // not be in the scheduler map if it already finished. + if let Some(ref slot) = state.scheduler + && let Some(ref scheduler) = *slot.read().await + { + let _ = scheduler.stop(job_id).await; + } - // Always persist cancellation to the DB so the state is - // consistent even if the scheduler wasn't available or the - // job wasn't in its in-memory map. - store - .update_job_status( - job_id, - crate::context::JobState::Cancelled, - Some("Cancelled by user"), - ) - .await - .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + // Always persist cancellation to the DB so the state is + // consistent even if the scheduler wasn't available or the + // job wasn't in its in-memory map. + store + .update_job_status( + job_id, + crate::context::JobState::Cancelled, + Some("Cancelled by user"), + ) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + } + return Ok(Json(serde_json::json!({ + "status": "cancelled", + "job_id": job_id, + }))); + } + Ok(None) => {} + Err(e) => { + return Err(( + StatusCode::INTERNAL_SERVER_ERROR, + format!("Database error: {}", e), + )); + } } - return Ok(Json(serde_json::json!({ - "status": "cancelled", - "job_id": job_id, - }))); } Err((StatusCode::NOT_FOUND, "Job not found".to_string())) @@ -315,6 +363,7 @@ pub async fn jobs_cancel_handler( pub async fn jobs_restart_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(id): Path, ) -> Result, (StatusCode, String)> { let store = state.store.as_ref().ok_or(( @@ -326,146 +375,166 @@ pub async fn jobs_restart_handler( .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?; // Try sandbox job restart first. - if let Ok(Some(old_job)) = store.get_sandbox_job(old_job_id).await { - if old_job.status != "interrupted" && old_job.status != "failed" { + match store.get_sandbox_job(old_job_id).await { + Ok(Some(old_job)) => { + if old_job.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + if old_job.status != "interrupted" && old_job.status != "failed" { + return Err(( + StatusCode::CONFLICT, + format!("Cannot restart job in state '{}'", old_job.status), + )); + } + + let jm = state.job_manager.as_ref().ok_or(( + StatusCode::SERVICE_UNAVAILABLE, + "Sandbox not enabled".to_string(), + ))?; + + // Enrich the task with failure context. + let task = if let Some(ref reason) = old_job.failure_reason { + format!( + "Previous attempt failed: {}. Retry: {}", + reason, old_job.task + ) + } else { + old_job.task.clone() + }; + + let new_job_id = Uuid::new_v4(); + let now = chrono::Utc::now(); + + let record = crate::history::SandboxJobRecord { + id: new_job_id, + task: task.clone(), + status: "creating".to_string(), + user_id: old_job.user_id.clone(), + project_dir: old_job.project_dir.clone(), + success: None, + failure_reason: None, + created_at: now, + started_at: None, + completed_at: None, + credential_grants_json: old_job.credential_grants_json.clone(), + }; + store + .save_sandbox_job(&record) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + + let mode = match store.get_sandbox_job_mode(old_job_id).await { + Ok(Some(m)) if m == "claude_code" => { + crate::orchestrator::job_manager::JobMode::ClaudeCode + } + _ => crate::orchestrator::job_manager::JobMode::Worker, + }; + + let credential_grants: Vec = + serde_json::from_str(&old_job.credential_grants_json).unwrap_or_else(|e| { + tracing::warn!( + job_id = %old_job.id, + "Failed to deserialize credential grants from stored job: {}. \ + Restarted job will have no credentials.", + e + ); + vec![] + }); + + let project_dir = std::path::PathBuf::from(&old_job.project_dir); + let _token = jm + .create_job( + new_job_id, + &task, + Some(project_dir), + mode, + credential_grants, + ) + .await + .map_err(|e| { + ( + StatusCode::INTERNAL_SERVER_ERROR, + format!("Failed to create container: {}", e), + ) + })?; + + store + .update_sandbox_job_status(new_job_id, "running", None, None, Some(now), None) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + + return Ok(Json(serde_json::json!({ + "status": "restarted", + "old_job_id": old_job_id, + "new_job_id": new_job_id, + }))); + } + Ok(None) => {} + Err(e) => { return Err(( - StatusCode::CONFLICT, - format!("Cannot restart job in state '{}'", old_job.status), + StatusCode::INTERNAL_SERVER_ERROR, + format!("Database error: {}", e), )); } - - let jm = state.job_manager.as_ref().ok_or(( - StatusCode::SERVICE_UNAVAILABLE, - "Sandbox not enabled".to_string(), - ))?; - - // Enrich the task with failure context. - let task = if let Some(ref reason) = old_job.failure_reason { - format!( - "Previous attempt failed: {}. Retry: {}", - reason, old_job.task - ) - } else { - old_job.task.clone() - }; - - let new_job_id = Uuid::new_v4(); - let now = chrono::Utc::now(); - - let record = crate::history::SandboxJobRecord { - id: new_job_id, - task: task.clone(), - status: "creating".to_string(), - user_id: old_job.user_id.clone(), - project_dir: old_job.project_dir.clone(), - success: None, - failure_reason: None, - created_at: now, - started_at: None, - completed_at: None, - credential_grants_json: old_job.credential_grants_json.clone(), - }; - store - .save_sandbox_job(&record) - .await - .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; - - let mode = match store.get_sandbox_job_mode(old_job_id).await { - Ok(Some(m)) if m == "claude_code" => { - crate::orchestrator::job_manager::JobMode::ClaudeCode - } - _ => crate::orchestrator::job_manager::JobMode::Worker, - }; - - let credential_grants: Vec = - serde_json::from_str(&old_job.credential_grants_json).unwrap_or_else(|e| { - tracing::warn!( - job_id = %old_job.id, - "Failed to deserialize credential grants from stored job: {}. \ - Restarted job will have no credentials.", - e - ); - vec![] - }); - - let project_dir = std::path::PathBuf::from(&old_job.project_dir); - let _token = jm - .create_job( - new_job_id, - &task, - Some(project_dir), - mode, - credential_grants, - ) - .await - .map_err(|e| { - ( - StatusCode::INTERNAL_SERVER_ERROR, - format!("Failed to create container: {}", e), - ) - })?; - - store - .update_sandbox_job_status(new_job_id, "running", None, None, Some(now), None) - .await - .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; - - return Ok(Json(serde_json::json!({ - "status": "restarted", - "old_job_id": old_job_id, - "new_job_id": new_job_id, - }))); } // Try agent job restart: dispatch a new job via the scheduler. - if let Ok(Some(old_job)) = store.get_job(old_job_id).await { - if old_job.state.is_active() { - return Err(( - StatusCode::CONFLICT, - format!("Cannot restart job in state '{}'", old_job.state), - )); + match store.get_job(old_job_id).await { + Ok(Some(old_job)) => { + if old_job.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + if old_job.state.is_active() { + return Err(( + StatusCode::CONFLICT, + format!("Cannot restart job in state '{}'", old_job.state), + )); + } + + let slot = state.scheduler.as_ref().ok_or(( + StatusCode::SERVICE_UNAVAILABLE, + "Scheduler not available".to_string(), + ))?; + let scheduler_guard = slot.read().await; + let scheduler = scheduler_guard.as_ref().ok_or(( + StatusCode::SERVICE_UNAVAILABLE, + "Agent not started yet".to_string(), + ))?; + + // Look up failure reason (O(1) point lookup). + let failure_reason = store + .get_agent_job_failure_reason(old_job_id) + .await + .ok() + .flatten() + .unwrap_or_default(); + + let title = if !failure_reason.is_empty() { + format!( + "Previous attempt failed: {}. Retry: {}", + failure_reason, old_job.title + ) + } else { + old_job.title.clone() + }; + + let new_job_id = scheduler + .dispatch_job(&old_job.user_id, &title, &old_job.description, None) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + + Ok(Json(serde_json::json!({ + "status": "restarted", + "old_job_id": old_job_id, + "new_job_id": new_job_id, + }))) } - - let slot = state.scheduler.as_ref().ok_or(( - StatusCode::SERVICE_UNAVAILABLE, - "Scheduler not available".to_string(), - ))?; - let scheduler_guard = slot.read().await; - let scheduler = scheduler_guard.as_ref().ok_or(( - StatusCode::SERVICE_UNAVAILABLE, - "Agent not started yet".to_string(), - ))?; - - // Look up failure reason (O(1) point lookup). - let failure_reason = store - .get_agent_job_failure_reason(old_job_id) - .await - .ok() - .flatten() - .unwrap_or_default(); - - let title = if !failure_reason.is_empty() { - format!( - "Previous attempt failed: {}. Retry: {}", - failure_reason, old_job.title - ) - } else { - old_job.title.clone() - }; - - let new_job_id = scheduler - .dispatch_job(&old_job.user_id, &title, &old_job.description, None) - .await - .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; - - return Ok(Json(serde_json::json!({ - "status": "restarted", - "old_job_id": old_job_id, - "new_job_id": new_job_id, - }))); + Ok(None) => Err((StatusCode::NOT_FOUND, "Job not found".to_string())), + Err(e) => Err(( + StatusCode::INTERNAL_SERVER_ERROR, + format!("Database error: {}", e), + )), } - - Err((StatusCode::NOT_FOUND, "Job not found".to_string())) } /// Submit a follow-up prompt to a running job. @@ -476,6 +545,7 @@ pub async fn jobs_restart_handler( /// - Worker-mode sandbox jobs → not supported (no mechanism to inject) pub async fn jobs_prompt_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(id): Path, Json(body): Json, ) -> Result, (StatusCode, String)> { @@ -494,10 +564,15 @@ pub async fn jobs_prompt_handler( let done = body.get("done").and_then(|v| v.as_bool()).unwrap_or(false); - // Try sandbox job path: check if we have a sandbox record for this ID. + // Try sandbox job path first: verify ownership, then route to Claude Code or reject. if let Some(ref s) = state.store - && let Ok(Some(_)) = s.get_sandbox_job(job_id).await + && let Ok(Some(sandbox_job)) = s.get_sandbox_job(job_id).await { + // Verify ownership. + if sandbox_job.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + // It's a sandbox job. Check if Claude Code mode. let mode = s.get_sandbox_job_mode(job_id).await.ok().flatten(); if mode.as_deref() == Some("claude_code") { @@ -522,7 +597,26 @@ pub async fn jobs_prompt_handler( } } - // Try agent job path: send via scheduler. + // Try agent job path: verify ownership, then send via scheduler. + if let Some(ref store) = state.store { + match store.get_job(job_id).await { + Ok(Some(agent_job)) => { + if agent_job.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + } + Ok(None) => { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + Err(e) => { + return Err(( + StatusCode::INTERNAL_SERVER_ERROR, + format!("Database error: {}", e), + )); + } + } + } + let slot = state.scheduler.as_ref().ok_or(( StatusCode::NOT_IMPLEMENTED, "Agent job prompts require the scheduler to be configured".to_string(), @@ -550,6 +644,7 @@ pub async fn jobs_prompt_handler( /// Load persisted job events for a job (for history replay on page open). pub async fn jobs_events_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(id): Path, ) -> Result, (StatusCode, String)> { let store = state.store.as_ref().ok_or(( @@ -561,6 +656,24 @@ pub async fn jobs_events_handler( .parse() .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?; + // Verify ownership before returning events. + match store.get_sandbox_job(job_id).await { + Ok(Some(job)) => { + if job.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + } + Ok(None) => { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + Err(e) => { + return Err(( + StatusCode::INTERNAL_SERVER_ERROR, + format!("Database error: {}", e), + )); + } + } + let events = store .list_job_events(job_id, None) .await @@ -593,6 +706,7 @@ pub struct FilePathQuery { pub async fn job_files_list_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(id): Path, Query(query): Query, ) -> Result, (StatusCode, String)> { @@ -610,6 +724,10 @@ pub async fn job_files_list_handler( .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .ok_or((StatusCode::NOT_FOUND, "Job not found".to_string()))?; + if job.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + let base = std::path::PathBuf::from(&job.project_dir); let rel_path = query.path.as_deref().unwrap_or(""); let target = base.join(rel_path); @@ -656,6 +774,7 @@ pub async fn job_files_list_handler( pub async fn job_files_read_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(id): Path, Query(query): Query, ) -> Result, (StatusCode, String)> { @@ -673,6 +792,10 @@ pub async fn job_files_read_handler( .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .ok_or((StatusCode::NOT_FOUND, "Job not found".to_string()))?; + if job.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + let path = query.path.as_deref().ok_or(( StatusCode::BAD_REQUEST, "path parameter required".to_string(), diff --git a/src/channels/web/handlers/memory.rs b/src/channels/web/handlers/memory.rs index fc0e1fe4..ff0fac16 100644 --- a/src/channels/web/handlers/memory.rs +++ b/src/channels/web/handlers/memory.rs @@ -9,8 +9,27 @@ use axum::{ }; use serde::Deserialize; +use crate::channels::web::auth::{AuthenticatedUser, UserIdentity}; use crate::channels::web::server::GatewayState; use crate::channels::web::types::*; +use crate::workspace::Workspace; + +/// Resolve the workspace for the authenticated user. +/// +/// Prefers `workspace_pool` (multi-user mode) when available, falling back +/// to the single-user `state.workspace`. +pub(crate) async fn resolve_workspace( + state: &GatewayState, + user: &UserIdentity, +) -> Result, (StatusCode, String)> { + if let Some(ref pool) = state.workspace_pool { + return Ok(pool.get_or_create(user).await); + } + state.workspace.as_ref().cloned().ok_or(( + StatusCode::SERVICE_UNAVAILABLE, + "Workspace not available".to_string(), + )) +} #[derive(Deserialize)] pub struct TreeQuery { @@ -20,12 +39,10 @@ pub struct TreeQuery { pub async fn memory_tree_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Query(_query): Query, ) -> Result, (StatusCode, String)> { - let workspace = state.workspace.as_ref().ok_or(( - StatusCode::SERVICE_UNAVAILABLE, - "Workspace not available".to_string(), - ))?; + let workspace = resolve_workspace(&state, &user).await?; // Build tree from list_all (flat list of all paths) let all_paths = workspace @@ -68,12 +85,10 @@ pub struct ListQuery { pub async fn memory_list_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Query(query): Query, ) -> Result, (StatusCode, String)> { - let workspace = state.workspace.as_ref().ok_or(( - StatusCode::SERVICE_UNAVAILABLE, - "Workspace not available".to_string(), - ))?; + let workspace = resolve_workspace(&state, &user).await?; let path = query.path.as_deref().unwrap_or(""); let entries = workspace @@ -104,12 +119,10 @@ pub struct ReadQuery { pub async fn memory_read_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Query(query): Query, ) -> Result, (StatusCode, String)> { - let workspace = state.workspace.as_ref().ok_or(( - StatusCode::SERVICE_UNAVAILABLE, - "Workspace not available".to_string(), - ))?; + let workspace = resolve_workspace(&state, &user).await?; let doc = workspace .read(&query.path) @@ -123,17 +136,75 @@ pub async fn memory_read_handler( })) } -// memory_write_handler lives in server.rs (layer-aware version with append, -// privacy redirect, and proper error status codes). +pub async fn memory_write_handler( + State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, + Json(req): Json, +) -> Result, (StatusCode, String)> { + let workspace = resolve_workspace(&state, &user).await?; + + // Route through layer-aware methods when a layer is specified. + // + // Note: unlike MemoryWriteTool, this endpoint does NOT block writes to + // identity files (IDENTITY.md, SOUL.md, etc.). The HTTP API is an + // authenticated admin interface; the supervisor uses it to seed identity + // files at startup. Identity-file protection is enforced at the tool + // layer (LLM-facing) where the write originates from an untrusted agent. + if let Some(ref layer_name) = req.layer { + let result = if req.append { + workspace + .append_to_layer(layer_name, &req.path, &req.content, req.force) + .await + } else { + workspace + .write_to_layer(layer_name, &req.path, &req.content, req.force) + .await + } + .map_err(|e| { + use crate::error::WorkspaceError; + let status = match &e { + WorkspaceError::LayerNotFound { .. } => StatusCode::BAD_REQUEST, + WorkspaceError::LayerReadOnly { .. } => StatusCode::FORBIDDEN, + WorkspaceError::PrivacyRedirectFailed => StatusCode::UNPROCESSABLE_ENTITY, + _ => StatusCode::INTERNAL_SERVER_ERROR, + }; + (status, e.to_string()) + })?; + return Ok(Json(MemoryWriteResponse { + path: req.path, + status: "written", + redirected: Some(result.redirected), + actual_layer: Some(result.actual_layer), + })); + } + + // Non-layer path: honor the append field + if req.append { + workspace + .append(&req.path, &req.content) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + } else { + workspace + .write(&req.path, &req.content) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + } + + Ok(Json(MemoryWriteResponse { + path: req.path, + status: "written", + redirected: None, + actual_layer: None, + })) +} pub async fn memory_search_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Json(req): Json, ) -> Result, (StatusCode, String)> { - let workspace = state.workspace.as_ref().ok_or(( - StatusCode::SERVICE_UNAVAILABLE, - "Workspace not available".to_string(), - ))?; + let workspace = resolve_workspace(&state, &user).await?; let limit = req.limit.unwrap_or(10); let results = workspace @@ -142,10 +213,10 @@ pub async fn memory_search_handler( .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; let hits: Vec = results - .into_iter() + .iter() .map(|r| SearchHit { - path: r.document_path, - content: r.content, + path: r.document_id.to_string(), + content: r.content.clone(), score: r.score as f64, }) .collect(); diff --git a/src/channels/web/handlers/mod.rs b/src/channels/web/handlers/mod.rs index 2f942058..50c7a0b9 100644 --- a/src/channels/web/handlers/mod.rs +++ b/src/channels/web/handlers/mod.rs @@ -1,13 +1,10 @@ //! Handler modules for the web gateway API. //! //! Each module groups related endpoint handlers by domain. -//! -//! # Migration status -//! -//! `skills` is the canonical implementation used by `server.rs`. -//! The remaining modules are in-progress migrations from inline server.rs -//! handlers; their functions are not yet wired up, hence the `dead_code` allow. +pub mod jobs; +pub mod memory; +pub mod routines; pub mod skills; // Modules not yet wired into server.rs router -- suppress dead_code until @@ -17,12 +14,6 @@ pub mod chat; #[allow(dead_code)] pub mod extensions; #[allow(dead_code)] -pub mod jobs; -#[allow(dead_code)] -pub mod memory; -#[allow(dead_code)] -pub mod routines; -#[allow(dead_code)] pub mod settings; #[allow(dead_code)] pub mod static_files; diff --git a/src/channels/web/handlers/routines.rs b/src/channels/web/handlers/routines.rs index 368a28ae..d27adca2 100644 --- a/src/channels/web/handlers/routines.rs +++ b/src/channels/web/handlers/routines.rs @@ -11,12 +11,14 @@ use serde::Deserialize; use uuid::Uuid; use crate::agent::routine::{Trigger, next_cron_fire}; +use crate::channels::web::auth::AuthenticatedUser; use crate::channels::web::server::GatewayState; use crate::channels::web::types::*; use crate::error::RoutineError; pub async fn routines_list_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let store = state.store.as_ref().ok_or(( StatusCode::SERVICE_UNAVAILABLE, @@ -24,7 +26,7 @@ pub async fn routines_list_handler( ))?; let routines = store - .list_all_routines() + .list_routines(&user.user_id) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; @@ -35,6 +37,7 @@ pub async fn routines_list_handler( pub async fn routines_summary_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let store = state.store.as_ref().ok_or(( StatusCode::SERVICE_UNAVAILABLE, @@ -42,7 +45,7 @@ pub async fn routines_summary_handler( ))?; let routines = store - .list_all_routines() + .list_routines(&user.user_id) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; @@ -78,6 +81,7 @@ pub async fn routines_summary_handler( pub async fn routines_detail_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(id): Path, ) -> Result, (StatusCode, String)> { let store = state.store.as_ref().ok_or(( @@ -94,6 +98,10 @@ pub async fn routines_detail_handler( .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?; + if routine.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Routine not found".to_string())); + } + let runs = store .list_routine_runs(routine_id, 20) .await @@ -137,6 +145,7 @@ pub async fn routines_detail_handler( pub async fn routines_trigger_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(id): Path, ) -> Result, (StatusCode, String)> { // Clone the Arc out of the lock to avoid holding the RwLock across .await. @@ -152,7 +161,7 @@ pub async fn routines_trigger_handler( .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?; let run_id = engine - .fire_manual(routine_id, Some(&state.user_id)) + .fire_manual(routine_id, Some(&user.user_id)) .await .map_err(|e| (routine_error_status(&e), e.to_string()))?; @@ -170,6 +179,7 @@ pub struct ToggleRequest { pub async fn routines_toggle_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(id): Path, body: Option>, ) -> Result, (StatusCode, String)> { @@ -187,6 +197,10 @@ pub async fn routines_toggle_handler( .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?; + if routine.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Routine not found".to_string())); + } + let was_enabled = routine.enabled; // If a specific value was provided, use it; otherwise toggle. routine.enabled = match body { @@ -230,6 +244,7 @@ pub async fn routines_toggle_handler( pub async fn routines_delete_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(id): Path, ) -> Result, (StatusCode, String)> { let store = state.store.as_ref().ok_or(( @@ -240,6 +255,17 @@ pub async fn routines_delete_handler( let routine_id = Uuid::parse_str(&id) .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?; + // Verify ownership before deleting. + let routine = store + .get_routine(routine_id) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? + .ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?; + + if routine.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Routine not found".to_string())); + } + let deleted = store .delete_routine(routine_id) .await @@ -261,8 +287,10 @@ pub async fn routines_delete_handler( } } +#[allow(dead_code)] // Used by server.rs inline version; kept in sync here for future migration. pub async fn routines_runs_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(id): Path, ) -> Result, (StatusCode, String)> { let store = state.store.as_ref().ok_or(( @@ -273,6 +301,17 @@ pub async fn routines_runs_handler( let routine_id = Uuid::parse_str(&id) .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?; + // Verify ownership before listing runs. + let routine = store + .get_routine(routine_id) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? + .ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?; + + if routine.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Routine not found".to_string())); + } + let runs = store .list_routine_runs(routine_id, 50) .await diff --git a/src/channels/web/handlers/settings.rs b/src/channels/web/handlers/settings.rs index dd66027b..4dd7299a 100644 --- a/src/channels/web/handlers/settings.rs +++ b/src/channels/web/handlers/settings.rs @@ -8,17 +8,19 @@ use axum::{ http::StatusCode, }; +use crate::channels::web::auth::AuthenticatedUser; use crate::channels::web::server::GatewayState; use crate::channels::web::types::*; pub async fn settings_list_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result, StatusCode> { let store = state .store .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; - let rows = store.list_settings(&state.user_id).await.map_err(|e| { + let rows = store.list_settings(&user.user_id).await.map_err(|e| { tracing::error!("Failed to list settings: {}", e); StatusCode::INTERNAL_SERVER_ERROR })?; @@ -37,6 +39,7 @@ pub async fn settings_list_handler( pub async fn settings_get_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(key): Path, ) -> Result, StatusCode> { let store = state @@ -44,7 +47,7 @@ pub async fn settings_get_handler( .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; let row = store - .get_setting_full(&state.user_id, &key) + .get_setting_full(&user.user_id, &key) .await .map_err(|e| { tracing::error!("Failed to get setting '{}': {}", key, e); @@ -61,6 +64,7 @@ pub async fn settings_get_handler( pub async fn settings_set_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(key): Path, Json(body): Json, ) -> Result { @@ -69,7 +73,7 @@ pub async fn settings_set_handler( .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; store - .set_setting(&state.user_id, &key, &body.value) + .set_setting(&user.user_id, &key, &body.value) .await .map_err(|e| { tracing::error!("Failed to set setting '{}': {}", key, e); @@ -81,6 +85,7 @@ pub async fn settings_set_handler( pub async fn settings_delete_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(key): Path, ) -> Result { let store = state @@ -88,7 +93,7 @@ pub async fn settings_delete_handler( .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; store - .delete_setting(&state.user_id, &key) + .delete_setting(&user.user_id, &key) .await .map_err(|e| { tracing::error!("Failed to delete setting '{}': {}", key, e); @@ -100,12 +105,13 @@ pub async fn settings_delete_handler( pub async fn settings_export_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result, StatusCode> { let store = state .store .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; - let settings = store.get_all_settings(&state.user_id).await.map_err(|e| { + let settings = store.get_all_settings(&user.user_id).await.map_err(|e| { tracing::error!("Failed to export settings: {}", e); StatusCode::INTERNAL_SERVER_ERROR })?; @@ -115,6 +121,7 @@ pub async fn settings_export_handler( pub async fn settings_import_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Json(body): Json, ) -> Result { let store = state @@ -122,7 +129,7 @@ pub async fn settings_import_handler( .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; store - .set_all_settings(&state.user_id, &body.settings) + .set_all_settings(&user.user_id, &body.settings) .await .map_err(|e| { tracing::error!("Failed to import settings: {}", e); diff --git a/src/channels/web/handlers/skills.rs b/src/channels/web/handlers/skills.rs index 400d179a..c8ecaf9f 100644 --- a/src/channels/web/handlers/skills.rs +++ b/src/channels/web/handlers/skills.rs @@ -8,11 +8,13 @@ use axum::{ http::StatusCode, }; +use crate::channels::web::auth::AuthenticatedUser; use crate::channels::web::server::GatewayState; use crate::channels::web::types::*; pub async fn skills_list_handler( State(state): State>, + AuthenticatedUser(_user): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let registry = state.skill_registry.as_ref().ok_or(( StatusCode::NOT_IMPLEMENTED, @@ -45,6 +47,7 @@ pub async fn skills_list_handler( pub async fn skills_search_handler( State(state): State>, + AuthenticatedUser(_user): AuthenticatedUser, Json(req): Json, ) -> Result, (StatusCode, String)> { let registry = state.skill_registry.as_ref().ok_or(( @@ -119,6 +122,7 @@ pub async fn skills_search_handler( pub async fn skills_install_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, headers: axum::http::HeaderMap, Json(req): Json, ) -> Result, (StatusCode, String)> { @@ -135,6 +139,8 @@ pub async fn skills_install_handler( )); } + tracing::info!(user_id = %user.user_id, skill = %req.name, "skill install requested"); + let registry = state.skill_registry.as_ref().ok_or(( StatusCode::NOT_IMPLEMENTED, "Skills system not enabled".to_string(), @@ -219,6 +225,7 @@ pub async fn skills_install_handler( pub async fn skills_remove_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, headers: axum::http::HeaderMap, Path(name): Path, ) -> Result, (StatusCode, String)> { @@ -234,6 +241,8 @@ pub async fn skills_remove_handler( )); } + tracing::info!(user_id = %user.user_id, skill = %name, "skill remove requested"); + let registry = state.skill_registry.as_ref().ok_or(( StatusCode::NOT_IMPLEMENTED, "Skills system not enabled".to_string(), diff --git a/src/channels/web/handlers/static_files.rs b/src/channels/web/handlers/static_files.rs index c198d95e..effc7037 100644 --- a/src/channels/web/handlers/static_files.rs +++ b/src/channels/web/handlers/static_files.rs @@ -7,6 +7,7 @@ use axum::{ }; use crate::bootstrap::ironclaw_base_dir; +use crate::channels::web::auth::AuthenticatedUser; use crate::channels::web::types::*; // --- Static file handlers --- @@ -113,6 +114,7 @@ use crate::channels::web::server::GatewayState; pub async fn logs_events_handler( State(state): State>, + AuthenticatedUser(_user): AuthenticatedUser, ) -> Result< Sse> + Send + 'static>, (StatusCode, String), @@ -152,6 +154,7 @@ pub async fn logs_events_handler( pub async fn gateway_status_handler( State(state): State>, + AuthenticatedUser(_user): AuthenticatedUser, ) -> Json { let sse_connections = state.sse.connection_count(); let ws_connections = state diff --git a/src/channels/web/mod.rs b/src/channels/web/mod.rs index f40834cb..b26a7829 100644 --- a/src/channels/web/mod.rs +++ b/src/channels/web/mod.rs @@ -31,6 +31,9 @@ pub mod ws; /// [`TestGatewayBuilder`](test_helpers::TestGatewayBuilder). pub mod test_helpers; +#[cfg(test)] +mod tests; + use std::net::SocketAddr; use std::sync::Arc; @@ -52,6 +55,7 @@ use crate::workspace::Workspace; use self::log_layer::{LogBroadcaster, LogLevelHandle}; +use self::auth::MultiAuthState; use self::server::GatewayState; use self::sse::SseManager; use self::types::SseEvent; @@ -60,14 +64,15 @@ use self::types::SseEvent; pub struct GatewayChannel { config: GatewayConfig, state: Arc, - /// The actual auth token in use (generated or from config). - auth_token: String, + /// Multi-user auth state (replaces bare auth_token). + auth: MultiAuthState, } impl GatewayChannel { /// Create a new gateway channel. /// /// If no auth token is configured, generates a random one and prints it. + /// Builds a single-user `MultiAuthState` from the config. pub fn new(config: GatewayConfig) -> Self { let auth_token = config.auth_token.clone().unwrap_or_else(|| { use rand::RngCore; @@ -77,10 +82,13 @@ impl GatewayChannel { bytes.iter().map(|b| format!("{b:02x}")).collect() }); + let auth = MultiAuthState::single(auth_token, config.user_id.clone()); + let state = Arc::new(GatewayState { msg_tx: tokio::sync::RwLock::new(None), - sse: SseManager::new(), + sse: Arc::new(SseManager::new()), workspace: None, + workspace_pool: None, session_manager: None, log_broadcaster: None, log_level_handle: None, @@ -90,13 +98,13 @@ impl GatewayChannel { job_manager: None, prompt_queue: None, scheduler: None, - user_id: config.user_id.clone(), + default_user_id: config.user_id.clone(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())), llm_provider: None, skill_registry: None, skill_catalog: None, - chat_rate_limiter: server::RateLimiter::new(30, 60), + chat_rate_limiter: server::PerUserRateLimiter::new(30, 60), oauth_rate_limiter: server::RateLimiter::new(10, 60), webhook_rate_limiter: server::RateLimiter::new(10, 60), registry_entries: Vec::new(), @@ -109,7 +117,46 @@ impl GatewayChannel { Self { config, state, - auth_token, + auth, + } + } + + /// Create a gateway channel with a pre-built multi-user auth state. + pub fn new_multi_auth(config: GatewayConfig, auth: MultiAuthState) -> Self { + let state = Arc::new(GatewayState { + msg_tx: tokio::sync::RwLock::new(None), + sse: Arc::new(SseManager::new()), + workspace: None, + workspace_pool: None, + session_manager: None, + log_broadcaster: None, + log_level_handle: None, + extension_manager: None, + tool_registry: None, + store: None, + job_manager: None, + prompt_queue: None, + scheduler: None, + default_user_id: config.user_id.clone(), + shutdown_tx: tokio::sync::RwLock::new(None), + ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())), + llm_provider: None, + skill_registry: None, + skill_catalog: None, + chat_rate_limiter: server::PerUserRateLimiter::new(30, 60), + oauth_rate_limiter: server::RateLimiter::new(10, 60), + registry_entries: Vec::new(), + cost_guard: None, + routine_engine: Arc::new(tokio::sync::RwLock::new(None)), + startup_time: std::time::Instant::now(), + webhook_rate_limiter: server::RateLimiter::new(10, 60), + active_config: server::ActiveConfigSnapshot::default(), + }); + + Self { + config, + state, + auth, } } @@ -118,8 +165,9 @@ impl GatewayChannel { let mut new_state = GatewayState { msg_tx: tokio::sync::RwLock::new(None), // Preserve the existing broadcast channel so sender handles remain valid. - sse: SseManager::from_sender(self.state.sse.sender()), + sse: Arc::new(SseManager::from_sender(self.state.sse.sender())), workspace: self.state.workspace.clone(), + workspace_pool: self.state.workspace_pool.clone(), session_manager: self.state.session_manager.clone(), log_broadcaster: self.state.log_broadcaster.clone(), log_level_handle: self.state.log_level_handle.clone(), @@ -129,13 +177,13 @@ impl GatewayChannel { job_manager: self.state.job_manager.clone(), prompt_queue: self.state.prompt_queue.clone(), scheduler: self.state.scheduler.clone(), - user_id: self.state.user_id.clone(), + default_user_id: self.state.default_user_id.clone(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: self.state.ws_tracker.clone(), llm_provider: self.state.llm_provider.clone(), skill_registry: self.state.skill_registry.clone(), skill_catalog: self.state.skill_catalog.clone(), - chat_rate_limiter: server::RateLimiter::new(30, 60), + chat_rate_limiter: server::PerUserRateLimiter::new(30, 60), oauth_rate_limiter: server::RateLimiter::new(10, 60), webhook_rate_limiter: server::RateLimiter::new(10, 60), registry_entries: self.state.registry_entries.clone(), @@ -260,9 +308,15 @@ impl GatewayChannel { self } - /// Get the auth token (for printing to console on startup). + /// Inject the per-user workspace pool for multi-user mode. + pub fn with_workspace_pool(mut self, pool: Arc) -> Self { + self.rebuild_state(|s| s.workspace_pool = Some(pool)); + self + } + + /// Get the first auth token (for printing to console on startup). pub fn auth_token(&self) -> &str { - &self.auth_token + self.auth.first_token().unwrap_or("") } /// Get a reference to the shared gateway state (for the agent to push SSE events). @@ -291,7 +345,7 @@ impl Channel for GatewayChannel { ), })?; - server::start_server(addr, self.state.clone(), self.auth_token.clone()).await?; + server::start_server(addr, self.state.clone(), self.auth.clone()).await?; Ok(Box::pin(ReceiverStream::new(rx))) } @@ -311,10 +365,13 @@ impl Channel for GatewayChannel { } }; - self.state.sse.broadcast(SseEvent::Response { - content: response.content, - thread_id, - }); + self.state.sse.broadcast_for_user( + &msg.user_id, + SseEvent::Response { + content: response.content, + thread_id, + }, + ); Ok(()) } @@ -427,13 +484,21 @@ impl Channel for GatewayChannel { }, }; - self.state.sse.broadcast(event); + // Scope events to the user when user_id is available in metadata. + // When user_id is missing (heartbeat, routines), events go to all + // subscribers. In multi-tenant mode this leaks status across users. + if let Some(uid) = metadata.get("user_id").and_then(|v| v.as_str()) { + self.state.sse.broadcast_for_user(uid, event); + } else { + tracing::debug!("Status event missing user_id in metadata; broadcasting globally"); + self.state.sse.broadcast(event); + } Ok(()) } async fn broadcast( &self, - _user_id: &str, + user_id: &str, response: OutgoingResponse, ) -> Result<(), ChannelError> { let thread_id = match response.thread_id { @@ -445,10 +510,13 @@ impl Channel for GatewayChannel { return Ok(()); } }; - self.state.sse.broadcast(SseEvent::Response { - content: response.content, - thread_id, - }); + self.state.sse.broadcast_for_user( + user_id, + SseEvent::Response { + content: response.content, + thread_id, + }, + ); Ok(()) } diff --git a/src/channels/web/openai_compat.rs b/src/channels/web/openai_compat.rs index 51577e06..55b7c854 100644 --- a/src/channels/web/openai_compat.rs +++ b/src/channels/web/openai_compat.rs @@ -463,9 +463,10 @@ fn build_tool_request( pub async fn chat_completions_handler( State(state): State>, + super::auth::AuthenticatedUser(user): super::auth::AuthenticatedUser, Json(req): Json, ) -> Result)> { - if !state.chat_rate_limiter.check() { + if !state.chat_rate_limiter.check(&user.user_id) { return Err(openai_error( StatusCode::TOO_MANY_REQUESTS, "Rate limit exceeded. Please try again later.", diff --git a/src/channels/web/server.rs b/src/channels/web/server.rs index 7edaad67..aaa479fa 100644 --- a/src/channels/web/server.rs +++ b/src/channels/web/server.rs @@ -30,12 +30,18 @@ use crate::agent::SessionManager; use crate::bootstrap::ironclaw_base_dir; use crate::channels::IncomingMessage; use crate::channels::relay::DEFAULT_RELAY_NAME; -use crate::channels::web::auth::{AuthState, auth_middleware}; +use crate::channels::web::auth::{ + AuthenticatedUser, MultiAuthState, UserIdentity, auth_middleware, +}; use crate::channels::web::handlers::jobs::{ job_files_list_handler, job_files_read_handler, jobs_cancel_handler, jobs_detail_handler, jobs_events_handler, jobs_list_handler, jobs_prompt_handler, jobs_restart_handler, jobs_summary_handler, }; +use crate::channels::web::handlers::memory::{ + memory_list_handler, memory_read_handler, memory_search_handler, memory_tree_handler, + memory_write_handler, +}; use crate::channels::web::handlers::routines::{ routines_delete_handler, routines_detail_handler, routines_list_handler, routines_summary_handler, routines_toggle_handler, routines_trigger_handler, @@ -80,7 +86,6 @@ fn redact_oauth_state_for_logs(state: &str) -> String { /// Simple sliding-window rate limiter. /// /// Tracks the number of requests in the current window. Resets when the window expires. -/// Not per-IP (since this is a single-user gateway with auth), but prevents flooding. pub struct RateLimiter { /// Requests remaining in the current window. remaining: AtomicU64, @@ -108,6 +113,12 @@ impl RateLimiter { } /// Try to consume one request. Returns `true` if allowed, `false` if rate limited. + /// + /// Note: There is a benign TOCTOU race between checking `window_start` and + /// resetting it — two concurrent threads may both see an expired window + /// and reset it, granting a few extra requests at the window boundary. + /// This is acceptable for chat rate limiting where approximate enforcement + /// is sufficient, and avoids the cost of a Mutex. pub fn check(&self) -> bool { let now = std::time::SystemTime::now() .duration_since(std::time::UNIX_EPOCH) @@ -148,14 +159,176 @@ pub struct ActiveConfigSnapshot { pub enabled_channels: Vec, } +/// Per-user rate limiter that maintains a separate sliding window per user_id. +/// +/// Prevents one user from exhausting the rate limit for all users in multi-tenant mode. +pub struct PerUserRateLimiter { + limiters: std::sync::RwLock>, + max_requests: u64, + window_secs: u64, +} + +impl PerUserRateLimiter { + pub fn new(max_requests: u64, window_secs: u64) -> Self { + Self { + limiters: std::sync::RwLock::new(std::collections::HashMap::new()), + max_requests, + window_secs, + } + } + + /// Try to consume one request for the given user. Returns `true` if allowed. + pub fn check(&self, user_id: &str) -> bool { + // Fast path: check existing limiter under read lock. + // On lock poisoning (another thread panicked while holding the lock), + // allow the request rather than crashing the server. + { + let map = match self.limiters.read() { + Ok(m) => m, + Err(e) => { + tracing::warn!("PerUserRateLimiter read lock poisoned; recovering"); + e.into_inner() + } + }; + if let Some(limiter) = map.get(user_id) { + return limiter.check(); + } + } + // Slow path: create limiter under write lock. + let mut map = match self.limiters.write() { + Ok(m) => m, + Err(e) => { + tracing::warn!("PerUserRateLimiter write lock poisoned; recovering"); + e.into_inner() + } + }; + let limiter = map + .entry(user_id.to_string()) + .or_insert_with(|| RateLimiter::new(self.max_requests, self.window_secs)); + limiter.check() + } +} + +/// Per-user workspace pool: lazily creates and caches workspaces keyed by user_id. +/// +/// In single-user mode, exactly one workspace is cached. In multi-user mode, +/// each authenticated user gets their own workspace with appropriate scopes, +/// search config, memory layers, and embedding cache settings. +/// +/// Also implements [`WorkspaceResolver`] so it can be shared with memory tools, +/// avoiding a separate `PerUserWorkspaceResolver` with duplicated logic. +pub struct WorkspacePool { + db: Arc, + embeddings: Option>, + embedding_cache_config: crate::workspace::EmbeddingCacheConfig, + search_config: crate::config::WorkspaceSearchConfig, + workspace_config: crate::config::WorkspaceConfig, + cache: tokio::sync::RwLock>>, +} + +impl WorkspacePool { + pub fn new( + db: Arc, + embeddings: Option>, + embedding_cache_config: crate::workspace::EmbeddingCacheConfig, + search_config: crate::config::WorkspaceSearchConfig, + workspace_config: crate::config::WorkspaceConfig, + ) -> Self { + Self { + db, + embeddings, + embedding_cache_config, + search_config, + workspace_config, + cache: tokio::sync::RwLock::new(std::collections::HashMap::new()), + } + } + + /// Build a workspace for a user, applying search config, embeddings, + /// global read scopes, and memory layers. + fn build_workspace(&self, user_id: &str) -> Workspace { + let mut ws = Workspace::new_with_db(user_id, Arc::clone(&self.db)) + .with_search_config(&self.search_config); + + if let Some(ref emb) = self.embeddings { + ws = ws.with_embeddings_cached(Arc::clone(emb), self.embedding_cache_config.clone()); + } + + if !self.workspace_config.read_scopes.is_empty() { + ws = ws.with_additional_read_scopes(self.workspace_config.read_scopes.clone()); + } + + ws = ws.with_memory_layers(self.workspace_config.memory_layers.clone()); + ws + } + + /// Get or create a workspace for the given user identity. + /// + /// Applies search config, memory layers, embedding cache, and read scopes + /// (both from global config and from the token's `workspace_read_scopes`). + pub async fn get_or_create(&self, identity: &UserIdentity) -> Arc { + // Fast path: check read lock + { + let cache = self.cache.read().await; + if let Some(ws) = cache.get(&identity.user_id) { + return Arc::clone(ws); + } + } + + // Slow path: create workspace under write lock + let mut cache = self.cache.write().await; + // Double-check after acquiring write lock + if let Some(ws) = cache.get(&identity.user_id) { + return Arc::clone(ws); + } + + let mut ws = self.build_workspace(&identity.user_id); + + // Apply per-token read scopes from identity. + if !identity.workspace_read_scopes.is_empty() { + ws = ws.with_additional_read_scopes(identity.workspace_read_scopes.clone()); + } + + let ws = Arc::new(ws); + cache.insert(identity.user_id.clone(), Arc::clone(&ws)); + ws + } +} + +#[async_trait::async_trait] +impl crate::tools::builtin::memory::WorkspaceResolver for WorkspacePool { + async fn resolve(&self, user_id: &str) -> Arc { + // Fast path: check read lock + { + let cache = self.cache.read().await; + if let Some(ws) = cache.get(user_id) { + return Arc::clone(ws); + } + } + + // Slow path: create workspace under write lock + let mut cache = self.cache.write().await; + if let Some(ws) = cache.get(user_id) { + return Arc::clone(ws); + } + + let ws = Arc::new(self.build_workspace(user_id)); + cache.insert(user_id.to_string(), Arc::clone(&ws)); + tracing::debug!(user_id = user_id, "Created per-user workspace"); + ws + } +} + /// Shared state for all gateway handlers. pub struct GatewayState { /// Channel to send messages to the agent loop. pub msg_tx: tokio::sync::RwLock>>, - /// SSE broadcast manager. - pub sse: SseManager, - /// Workspace for memory API. + /// SSE broadcast manager (Arc-wrapped so extension manager can hold a reference). + pub sse: Arc, + /// Workspace for memory API (single-user fallback). pub workspace: Option>, + /// Per-user workspace pool for multi-user mode. + pub workspace_pool: Option>, /// Session manager for thread info. pub session_manager: Option>, /// Log broadcaster for the logs SSE endpoint. @@ -172,8 +345,8 @@ pub struct GatewayState { pub job_manager: Option>, /// Prompt queue for Claude Code follow-up prompts. pub prompt_queue: Option, - /// User ID for this gateway. - pub user_id: String, + /// Default user ID (fallback for non-request contexts like heartbeat/routines). + pub default_user_id: String, /// Shutdown signal sender. pub shutdown_tx: tokio::sync::RwLock>>, /// WebSocket connection tracker. @@ -186,8 +359,8 @@ pub struct GatewayState { pub skill_catalog: Option>, /// Scheduler for sending follow-up messages to running agent jobs. pub scheduler: Option, - /// Rate limiter for chat endpoints (30 messages per 60 seconds). - pub chat_rate_limiter: RateLimiter, + /// Per-user rate limiter for chat endpoints (30 messages per 60 seconds per user). + pub chat_rate_limiter: PerUserRateLimiter, /// Rate limiter for OAuth callback endpoints (10 requests per 60 seconds). pub oauth_rate_limiter: RateLimiter, /// Rate limiter for webhook trigger endpoints (10 requests per 60 seconds). @@ -211,7 +384,7 @@ pub struct GatewayState { pub async fn start_server( addr: SocketAddr, state: Arc, - auth_token: String, + auth: MultiAuthState, ) -> Result { let listener = tokio::net::TcpListener::bind(addr).await.map_err(|e| { crate::error::ChannelError::StartupFailed { @@ -242,7 +415,7 @@ pub async fn start_server( ); // Protected routes (require auth) - let auth_state = AuthState { token: auth_token }; + let auth_state = auth; let protected = Router::new() // Chat .route("/api/chat/send", post(chat_send_handler)) @@ -568,14 +741,12 @@ async fn oauth_callback_handler( .get("error_description") .cloned() .unwrap_or_else(|| error.clone()); - clear_auth_mode(&state).await; return oauth_error_page(&description); } let state_param = match params.get("state") { Some(s) if !s.is_empty() => s.clone(), _ => { - clear_auth_mode(&state).await; return oauth_error_page("IronClaw"); } }; @@ -583,7 +754,6 @@ async fn oauth_callback_handler( let code = match params.get("code") { Some(c) if !c.is_empty() => c.clone(), _ => { - clear_auth_mode(&state).await; return oauth_error_page("IronClaw"); } }; @@ -592,7 +762,6 @@ async fn oauth_callback_handler( let ext_mgr = match state.extension_manager.as_ref() { Some(mgr) => mgr, None => { - clear_auth_mode(&state).await; return oauth_error_page("IronClaw"); } }; @@ -606,7 +775,7 @@ async fn oauth_callback_handler( error = %error, "OAuth callback received with malformed state" ); - clear_auth_mode(&state).await; + clear_auth_mode(&state, &state.default_user_id).await; return oauth_error_page("IronClaw"); } }; @@ -628,7 +797,6 @@ async fn oauth_callback_handler( lookup_key = %redacted_lookup_key, "OAuth callback received with unknown or expired state" ); - clear_auth_mode(&state).await; return oauth_error_page("IronClaw"); } }; @@ -640,14 +808,17 @@ async fn oauth_callback_handler( "OAuth flow expired" ); // Notify UI so auth card can show error instead of staying stuck - if let Some(ref sender) = flow.sse_sender { - let _ = sender.send(SseEvent::AuthCompleted { - extension_name: flow.extension_name.clone(), - success: false, - message: "OAuth flow expired. Please try again.".to_string(), - }); + if let Some(ref sse) = flow.sse_manager { + sse.broadcast_for_user( + &flow.user_id, + SseEvent::AuthCompleted { + extension_name: flow.extension_name.clone(), + success: false, + message: "OAuth flow expired. Please try again.".to_string(), + }, + ); } - clear_auth_mode(&state).await; + clear_auth_mode(&state, &flow.user_id).await; return oauth_error_page(&flow.display_name); } @@ -753,14 +924,14 @@ async fn oauth_callback_handler( // Clear auth mode regardless of outcome so the next user message goes // through to the LLM instead of being intercepted as a token. - clear_auth_mode(&state).await; + clear_auth_mode(&state, &flow.user_id).await; // After successful OAuth, auto-activate the extension so it moves // from "Installed (Authenticate)" → "Active" without a second click. // OAuth success is independent of activation — tokens are already stored. // Report auth as successful and attempt activation as a bonus step. let final_message = if success { - match ext_mgr.activate(&flow.extension_name).await { + match ext_mgr.activate(&flow.extension_name, &flow.user_id).await { Ok(result) => result.message, Err(e) => { tracing::warn!( @@ -779,12 +950,15 @@ async fn oauth_callback_handler( }; // Broadcast SSE event to notify the web UI - if let Some(ref sender) = flow.sse_sender { - let _ = sender.send(SseEvent::AuthCompleted { - extension_name: flow.extension_name, - success, - message: final_message.clone(), - }); + if let Some(ref sse) = flow.sse_manager { + sse.broadcast_for_user( + &flow.user_id, + SseEvent::AuthCompleted { + extension_name: flow.extension_name, + success, + message: final_message.clone(), + }, + ); } let html = oauth_defaults::landing_html(&flow.display_name, success); @@ -962,7 +1136,7 @@ async fn slack_relay_oauth_callback_handler( let state_key = format!("relay:{}:oauth_state", DEFAULT_RELAY_NAME); let stored_state = match ext_mgr .secrets() - .get_decrypted(&state.user_id, &state_key) + .get_decrypted(&state.default_user_id, &state_key) .await { Ok(secret) => secret.expose().to_string(), @@ -986,7 +1160,10 @@ async fn slack_relay_oauth_callback_handler( } // Delete the nonce (one-time use) - let _ = ext_mgr.secrets().delete(&state.user_id, &state_key).await; + let _ = ext_mgr + .secrets() + .delete(&state.default_user_id, &state_key) + .await; let result: Result<(), String> = async { let store = state.store.as_ref().ok_or_else(|| { @@ -997,12 +1174,16 @@ async fn slack_relay_oauth_callback_handler( // Store team_id in settings let team_id_key = format!("relay:{}:team_id", DEFAULT_RELAY_NAME); let _ = store - .set_setting(&state.user_id, &team_id_key, &serde_json::json!(team_id)) + .set_setting( + &state.default_user_id, + &team_id_key, + &serde_json::json!(team_id), + ) .await; // Activate the relay channel ext_mgr - .activate_stored_relay(DEFAULT_RELAY_NAME) + .activate_stored_relay(DEFAULT_RELAY_NAME, &state.default_user_id) .await .map_err(|e| format!("Failed to activate relay channel: {}", e))?; @@ -1104,6 +1285,7 @@ fn mime_to_ext(mime: &str) -> &str { async fn chat_send_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, headers: axum::http::HeaderMap, Json(req): Json, ) -> Result<(StatusCode, Json), (StatusCode, String)> { @@ -1113,14 +1295,14 @@ async fn chat_send_handler( req.thread_id ); - if !state.chat_rate_limiter.check() { + if !state.chat_rate_limiter.check(&user.user_id) { return Err(( StatusCode::TOO_MANY_REQUESTS, "Rate limit exceeded. Try again shortly.".to_string(), )); } - let mut msg = IncomingMessage::new("gateway", &state.user_id, &req.content); + let mut msg = IncomingMessage::new("gateway", &user.user_id, &req.content); // Prefer timezone from JSON body, fall back to X-Timezone header let tz = req .timezone @@ -1130,10 +1312,13 @@ async fn chat_send_handler( msg = msg.with_timezone(tz); } + // Always include user_id in metadata so downstream SSE broadcasts can scope events. + let mut meta = serde_json::json!({"user_id": &user.user_id}); if let Some(ref thread_id) = req.thread_id { msg = msg.with_thread(thread_id); - msg = msg.with_metadata(serde_json::json!({"thread_id": thread_id})); + meta["thread_id"] = serde_json::json!(thread_id); } + msg = msg.with_metadata(meta); // Convert uploaded images to IncomingAttachments if !req.images.is_empty() { @@ -1182,6 +1367,7 @@ async fn chat_send_handler( async fn chat_approval_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Json(req): Json, ) -> Result<(StatusCode, Json), (StatusCode, String)> { let (approved, always) = match req.action.as_str() { @@ -1217,7 +1403,7 @@ async fn chat_approval_handler( ) })?; - let mut msg = IncomingMessage::new("gateway", &state.user_id, content); + let mut msg = IncomingMessage::new("gateway", &user.user_id, content); if let Some(ref thread_id) = req.thread_id { msg = msg.with_thread(thread_id); @@ -1258,6 +1444,7 @@ async fn chat_approval_handler( /// The token never touches the LLM, chat history, or SSE stream. async fn chat_auth_token_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Json(req): Json, ) -> Result, (StatusCode, String)> { let ext_mgr = state.extension_manager.as_ref().ok_or(( @@ -1266,7 +1453,7 @@ async fn chat_auth_token_handler( ))?; match ext_mgr - .configure_token(&req.extension_name, &req.token) + .configure_token(&req.extension_name, &req.token, &user.user_id) .await { Ok(result) => { @@ -1281,27 +1468,36 @@ async fn chat_auth_token_handler( resp.instructions = result.verification.as_ref().map(|v| v.instructions.clone()); if result.verification.is_some() { - state.sse.broadcast(SseEvent::AuthRequired { - extension_name: req.extension_name.clone(), - instructions: Some(result.message), - auth_url: None, - setup_url: None, - }); + state.sse.broadcast_for_user( + &user.user_id, + SseEvent::AuthRequired { + extension_name: req.extension_name.clone(), + instructions: Some(result.message), + auth_url: None, + setup_url: None, + }, + ); } else if result.activated { // Clear auth mode on the active thread - clear_auth_mode(&state).await; + clear_auth_mode(&state, &user.user_id).await; - state.sse.broadcast(SseEvent::AuthCompleted { - extension_name: req.extension_name.clone(), - success: true, - message: result.message, - }); + state.sse.broadcast_for_user( + &user.user_id, + SseEvent::AuthCompleted { + extension_name: req.extension_name.clone(), + success: true, + message: result.message, + }, + ); } else { - state.sse.broadcast(SseEvent::AuthCompleted { - extension_name: req.extension_name.clone(), - success: false, - message: result.message, - }); + state.sse.broadcast_for_user( + &user.user_id, + SseEvent::AuthCompleted { + extension_name: req.extension_name.clone(), + success: false, + message: result.message, + }, + ); } Ok(Json(resp)) @@ -1310,12 +1506,15 @@ async fn chat_auth_token_handler( let msg = e.to_string(); // Re-emit auth_required for retry on validation errors if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) { - state.sse.broadcast(SseEvent::AuthRequired { - extension_name: req.extension_name.clone(), - instructions: Some(msg.clone()), - auth_url: None, - setup_url: None, - }); + state.sse.broadcast_for_user( + &user.user_id, + SseEvent::AuthRequired { + extension_name: req.extension_name.clone(), + instructions: Some(msg.clone()), + auth_url: None, + setup_url: None, + }, + ); } Ok(Json(ActionResponse::fail(msg))) } @@ -1325,16 +1524,17 @@ async fn chat_auth_token_handler( /// Cancel an in-progress auth flow. async fn chat_auth_cancel_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Json(_req): Json, ) -> Result, (StatusCode, String)> { - clear_auth_mode(&state).await; + clear_auth_mode(&state, &user.user_id).await; Ok(Json(ActionResponse::ok("Auth cancelled"))) } /// Clear pending auth mode on the active thread. -pub async fn clear_auth_mode(state: &GatewayState) { +pub async fn clear_auth_mode(state: &GatewayState, user_id: &str) { if let Some(ref sm) = state.session_manager { - let session = sm.get_or_create_session(&state.user_id).await; + let session = sm.get_or_create_session(user_id).await; let mut sess = session.lock().await; if let Some(thread_id) = sess.active_thread && let Some(thread) = sess.threads.get_mut(&thread_id) @@ -1346,8 +1546,9 @@ pub async fn clear_auth_mode(state: &GatewayState) { async fn chat_events_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result { - let sse = state.sse.subscribe().ok_or(( + let sse = state.sse.subscribe(Some(user.user_id)).ok_or(( StatusCode::SERVICE_UNAVAILABLE, "Too many connections".to_string(), ))?; @@ -1357,7 +1558,31 @@ async fn chat_events_handler( )) } +/// Check whether an Origin header value points to a local address. +/// +/// Extracts the host from the origin (handling both IPv4/hostname and IPv6 +/// literal formats) and compares it against known local addresses. Used to +/// prevent cross-site WebSocket hijacking while allowing localhost access. +fn is_local_origin(origin: &str) -> bool { + let host = origin + .strip_prefix("http://") + .or_else(|| origin.strip_prefix("https://")) + .and_then(|rest| { + if rest.starts_with('[') { + // IPv6 literal: extract "[::1]" up to and including ']' + rest.find(']').map(|i| &rest[..=i]) + } else { + // IPv4 or hostname: take up to the first ':' (port) or '/' (path) + rest.split(':').next()?.split('/').next() + } + }) + .unwrap_or(""); + + matches!(host, "localhost" | "127.0.0.1" | "[::1]") +} + async fn chat_ws_handler( + AuthenticatedUser(user): AuthenticatedUser, headers: axum::http::HeaderMap, ws: WebSocketUpgrade, State(state): State>, @@ -1375,23 +1600,16 @@ async fn chat_ws_handler( ) })?; - // Extract the host from the origin and compare exactly, so that - // crafted origins like "http://localhost.evil.com" are rejected. - // Origin format is "scheme://host[:port]". - let host = origin - .strip_prefix("http://") - .or_else(|| origin.strip_prefix("https://")) - .and_then(|rest| rest.split(':').next()?.split('/').next()) - .unwrap_or(""); - - let is_local = matches!(host, "localhost" | "127.0.0.1" | "[::1]"); + let is_local = is_local_origin(origin); if !is_local { return Err(( StatusCode::FORBIDDEN, "WebSocket origin not allowed".to_string(), )); } - Ok(ws.on_upgrade(move |socket| crate::channels::web::ws::handle_ws_connection(socket, state))) + Ok(ws.on_upgrade(move |socket| { + crate::channels::web::ws::handle_ws_connection(socket, state, user) + })) } #[derive(Deserialize)] @@ -1403,6 +1621,7 @@ struct HistoryQuery { async fn chat_history_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Query(query): Query, ) -> Result, (StatusCode, String)> { let session_manager = state.session_manager.as_ref().ok_or(( @@ -1410,7 +1629,7 @@ async fn chat_history_handler( "Session manager not available".to_string(), ))?; - let session = session_manager.get_or_create_session(&state.user_id).await; + let session = session_manager.get_or_create_session(&user.user_id).await; let sess = session.lock().await; let limit = query.limit.unwrap_or(50); @@ -1445,9 +1664,12 @@ async fn chat_history_handler( && let Some(ref store) = state.store { let owned = store - .conversation_belongs_to_user(thread_id, &state.user_id) + .conversation_belongs_to_user(thread_id, &user.user_id) .await - .unwrap_or(false); + .map_err(|e| { + tracing::error!(thread_id = %thread_id, error = %e, "DB error during thread ownership check"); + (StatusCode::INTERNAL_SERVER_ERROR, "Database error".to_string()) + })?; if !owned && !sess.threads.contains_key(&thread_id) { return Err((StatusCode::NOT_FOUND, "Thread not found".to_string())); } @@ -1558,68 +1780,74 @@ async fn chat_history_handler( async fn chat_threads_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let session_manager = state.session_manager.as_ref().ok_or(( StatusCode::SERVICE_UNAVAILABLE, "Session manager not available".to_string(), ))?; - let session = session_manager.get_or_create_session(&state.user_id).await; + let session = session_manager.get_or_create_session(&user.user_id).await; let sess = session.lock().await; // Try DB first for persistent thread list if let Some(ref store) = state.store { // Auto-create assistant thread if it doesn't exist let assistant_id = store - .get_or_create_assistant_conversation(&state.user_id, "gateway") + .get_or_create_assistant_conversation(&user.user_id, "gateway") .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; - if let Ok(summaries) = store - .list_conversations_all_channels(&state.user_id, 50) + match store + .list_conversations_all_channels(&user.user_id, 50) .await { - let mut assistant_thread = None; - let mut threads = Vec::new(); + Ok(summaries) => { + let mut assistant_thread = None; + let mut threads = Vec::new(); - for s in &summaries { - let info = ThreadInfo { - id: s.id, - state: "Idle".to_string(), - turn_count: s.message_count.max(0) as usize, - created_at: s.started_at.to_rfc3339(), - updated_at: s.last_activity.to_rfc3339(), - title: s.title.clone(), - thread_type: s.thread_type.clone(), - channel: Some(s.channel.clone()), - }; + for s in &summaries { + let info = ThreadInfo { + id: s.id, + state: "Idle".to_string(), + turn_count: s.message_count.max(0) as usize, + created_at: s.started_at.to_rfc3339(), + updated_at: s.last_activity.to_rfc3339(), + title: s.title.clone(), + thread_type: s.thread_type.clone(), + channel: Some(s.channel.clone()), + }; - if s.id == assistant_id { - assistant_thread = Some(info); - } else { - threads.push(info); + if s.id == assistant_id { + assistant_thread = Some(info); + } else { + threads.push(info); + } } - } - // If assistant wasn't in the list (0 messages), synthesize it - if assistant_thread.is_none() { - assistant_thread = Some(ThreadInfo { - id: assistant_id, - state: "Idle".to_string(), - turn_count: 0, - created_at: chrono::Utc::now().to_rfc3339(), - updated_at: chrono::Utc::now().to_rfc3339(), - title: None, - thread_type: Some("assistant".to_string()), - channel: Some("gateway".to_string()), - }); - } + // If assistant wasn't in the list (0 messages), synthesize it + if assistant_thread.is_none() { + assistant_thread = Some(ThreadInfo { + id: assistant_id, + state: "Idle".to_string(), + turn_count: 0, + created_at: chrono::Utc::now().to_rfc3339(), + updated_at: chrono::Utc::now().to_rfc3339(), + title: None, + thread_type: Some("assistant".to_string()), + channel: Some("gateway".to_string()), + }); + } - return Ok(Json(ThreadListResponse { - assistant_thread, - threads, - active_thread: sess.active_thread, - })); + return Ok(Json(ThreadListResponse { + assistant_thread, + threads, + active_thread: sess.active_thread, + })); + } + Err(e) => { + tracing::error!(user_id = %user.user_id, error = %e, "DB error listing threads; falling back to in-memory"); + } } } @@ -1649,13 +1877,14 @@ async fn chat_threads_handler( async fn chat_new_thread_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let session_manager = state.session_manager.as_ref().ok_or(( StatusCode::SERVICE_UNAVAILABLE, "Session manager not available".to_string(), ))?; - let session = session_manager.get_or_create_session(&state.user_id).await; + let session = session_manager.get_or_create_session(&user.user_id).await; let (thread_id, info) = { let mut sess = session.lock().await; let thread = sess.create_thread(); @@ -1677,12 +1906,12 @@ async fn chat_new_thread_handler( // so that the subsequent loadThreads() call from the frontend sees it. if let Some(ref store) = state.store { match store - .ensure_conversation(thread_id, "gateway", &state.user_id, None) + .ensure_conversation(thread_id, "gateway", &user.user_id, None) .await { Ok(true) => {} Ok(false) => tracing::warn!( - user = %state.user_id, + user = %user.user_id, thread_id = %thread_id, "Skipped persisting new thread due to ownership/channel conflict" ), @@ -1700,216 +1929,12 @@ async fn chat_new_thread_handler( Ok(Json(info)) } -// --- Memory handlers --- - -#[derive(Deserialize)] -struct TreeQuery { - #[allow(dead_code)] - depth: Option, -} - -async fn memory_tree_handler( - State(state): State>, - Query(_query): Query, -) -> Result, (StatusCode, String)> { - let workspace = state.workspace.as_ref().ok_or(( - StatusCode::SERVICE_UNAVAILABLE, - "Workspace not available".to_string(), - ))?; - - // Build tree from list_all (flat list of all paths) - let all_paths = workspace - .list_all() - .await - .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; - - // Collect unique directories and files - let mut entries: Vec = Vec::new(); - let mut seen_dirs: std::collections::HashSet = std::collections::HashSet::new(); - - for path in &all_paths { - // Add parent directories - let parts: Vec<&str> = path.split('/').collect(); - for i in 0..parts.len().saturating_sub(1) { - let dir_path = parts[..=i].join("/"); - if seen_dirs.insert(dir_path.clone()) { - entries.push(TreeEntry { - path: dir_path, - is_dir: true, - }); - } - } - // Add the file itself - entries.push(TreeEntry { - path: path.clone(), - is_dir: false, - }); - } - - entries.sort_by(|a, b| a.path.cmp(&b.path)); - - Ok(Json(MemoryTreeResponse { entries })) -} - -#[derive(Deserialize)] -struct ListQuery { - path: Option, -} - -async fn memory_list_handler( - State(state): State>, - Query(query): Query, -) -> Result, (StatusCode, String)> { - let workspace = state.workspace.as_ref().ok_or(( - StatusCode::SERVICE_UNAVAILABLE, - "Workspace not available".to_string(), - ))?; - - let path = query.path.as_deref().unwrap_or(""); - let entries = workspace - .list(path) - .await - .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; - - let list_entries: Vec = entries - .iter() - .map(|e| ListEntry { - name: e.path.rsplit('/').next().unwrap_or(&e.path).to_string(), - path: e.path.clone(), - is_dir: e.is_directory, - updated_at: e.updated_at.map(|dt| dt.to_rfc3339()), - }) - .collect(); - - Ok(Json(MemoryListResponse { - path: path.to_string(), - entries: list_entries, - })) -} - -#[derive(Deserialize)] -struct ReadQuery { - path: String, -} - -async fn memory_read_handler( - State(state): State>, - Query(query): Query, -) -> Result, (StatusCode, String)> { - let workspace = state.workspace.as_ref().ok_or(( - StatusCode::SERVICE_UNAVAILABLE, - "Workspace not available".to_string(), - ))?; - - let doc = workspace - .read(&query.path) - .await - .map_err(|e| (StatusCode::NOT_FOUND, e.to_string()))?; - - Ok(Json(MemoryReadResponse { - path: query.path, - content: doc.content, - updated_at: Some(doc.updated_at.to_rfc3339()), - })) -} - -async fn memory_write_handler( - State(state): State>, - Json(req): Json, -) -> Result, (StatusCode, String)> { - let workspace = state.workspace.as_ref().ok_or(( - StatusCode::SERVICE_UNAVAILABLE, - "Workspace not available".to_string(), - ))?; - - // Route through layer-aware methods when a layer is specified. - // - // Note: unlike MemoryWriteTool, this endpoint does NOT block writes to - // identity files (IDENTITY.md, SOUL.md, etc.). The HTTP API is an - // authenticated admin interface; the supervisor uses it to seed identity - // files at startup. Identity-file protection is enforced at the tool - // layer (LLM-facing) where the write originates from an untrusted agent. - if let Some(ref layer_name) = req.layer { - let result = if req.append { - workspace - .append_to_layer(layer_name, &req.path, &req.content, req.force) - .await - } else { - workspace - .write_to_layer(layer_name, &req.path, &req.content, req.force) - .await - } - .map_err(|e| { - use crate::error::WorkspaceError; - let status = match &e { - WorkspaceError::LayerNotFound { .. } => StatusCode::BAD_REQUEST, - WorkspaceError::LayerReadOnly { .. } => StatusCode::FORBIDDEN, - WorkspaceError::PrivacyRedirectFailed => StatusCode::UNPROCESSABLE_ENTITY, - _ => StatusCode::INTERNAL_SERVER_ERROR, - }; - (status, e.to_string()) - })?; - return Ok(Json(MemoryWriteResponse { - path: req.path, - status: "written", - redirected: Some(result.redirected), - actual_layer: Some(result.actual_layer), - })); - } - - // Non-layer path: honor the append field - if req.append { - workspace - .append(&req.path, &req.content) - .await - .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; - } else { - workspace - .write(&req.path, &req.content) - .await - .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; - } - - Ok(Json(MemoryWriteResponse { - path: req.path, - status: "written", - redirected: None, - actual_layer: None, - })) -} - -async fn memory_search_handler( - State(state): State>, - Json(req): Json, -) -> Result, (StatusCode, String)> { - let workspace = state.workspace.as_ref().ok_or(( - StatusCode::SERVICE_UNAVAILABLE, - "Workspace not available".to_string(), - ))?; - - let limit = req.limit.unwrap_or(10); - let results = workspace - .search(&req.query, limit) - .await - .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; - - let hits: Vec = results - .iter() - .map(|r| SearchHit { - path: r.document_id.to_string(), - content: r.content.clone(), - score: r.score as f64, - }) - .collect(); - - Ok(Json(MemorySearchResponse { results: hits })) -} - // Job handlers moved to handlers/jobs.rs // --- Logs handlers --- async fn logs_events_handler( State(state): State>, + AuthenticatedUser(_user): AuthenticatedUser, ) -> Result { let broadcaster = state.log_broadcaster.as_ref().ok_or(( StatusCode::SERVICE_UNAVAILABLE, @@ -1947,6 +1972,7 @@ async fn logs_events_handler( async fn logs_level_get_handler( State(state): State>, + AuthenticatedUser(_user): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let handle = state.log_level_handle.as_ref().ok_or(( StatusCode::SERVICE_UNAVAILABLE, @@ -1957,6 +1983,7 @@ async fn logs_level_get_handler( async fn logs_level_set_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Json(body): Json, ) -> Result, (StatusCode, String)> { let handle = state.log_level_handle.as_ref().ok_or(( @@ -1973,7 +2000,7 @@ async fn logs_level_set_handler( .set_level(level) .map_err(|e| (StatusCode::BAD_REQUEST, e))?; - tracing::info!("Log level changed to '{}'", handle.current_level()); + tracing::info!(user_id = %user.user_id, "Log level changed to '{}'", handle.current_level()); Ok(Json(serde_json::json!({ "level": handle.current_level() }))) } @@ -1981,6 +2008,7 @@ async fn logs_level_set_handler( async fn extensions_list_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let ext_mgr = state.extension_manager.as_ref().ok_or(( StatusCode::NOT_IMPLEMENTED, @@ -1988,7 +2016,7 @@ async fn extensions_list_handler( ))?; let installed = ext_mgr - .list(None, false) + .list(None, false, &user.user_id) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; @@ -2048,6 +2076,7 @@ async fn extensions_list_handler( async fn extensions_tools_handler( State(state): State>, + AuthenticatedUser(_user): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let registry = state.tool_registry.as_ref().ok_or(( StatusCode::SERVICE_UNAVAILABLE, @@ -2068,6 +2097,7 @@ async fn extensions_tools_handler( async fn extensions_install_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Json(req): Json, ) -> Result, (StatusCode, String)> { // When extension manager isn't available, check registry entries for a helpful message @@ -2103,7 +2133,7 @@ async fn extensions_install_handler( }); match ext_mgr - .install(&req.name, req.url.as_deref(), kind_hint) + .install(&req.name, req.url.as_deref(), kind_hint, &user.user_id) .await { Ok(result) => { @@ -2111,7 +2141,7 @@ async fn extensions_install_handler( // Auto-activate WASM tools after install (install = active). if result.kind == crate::extensions::ExtensionKind::WasmTool { - if let Err(e) = ext_mgr.activate(&req.name).await { + if let Err(e) = ext_mgr.activate(&req.name, &user.user_id).await { tracing::debug!( extension = %req.name, error = %e, @@ -2123,7 +2153,7 @@ async fn extensions_install_handler( // expansion and for first-time auth when credentials are already // configured (e.g., built-in providers). We only surface an auth_url // when the extension reports it is awaiting authorization. - match ext_mgr.auth(&req.name).await { + match ext_mgr.auth(&req.name, &user.user_id).await { Ok(auth_result) if auth_result.auth_url().is_some() => { // Scope expansion or initial OAuth: user needs to authorize resp.auth_url = auth_result.auth_url().map(String::from); @@ -2140,6 +2170,7 @@ async fn extensions_install_handler( async fn extensions_activate_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(name): Path, ) -> Result, (StatusCode, String)> { let ext_mgr = state.extension_manager.as_ref().ok_or(( @@ -2147,14 +2178,14 @@ async fn extensions_activate_handler( "Extension manager not available (secrets store required)".to_string(), ))?; - match ext_mgr.activate(&name).await { + match ext_mgr.activate(&name, &user.user_id).await { Ok(result) => { // Activation loaded the WASM module. Check if the tool needs // OAuth scope expansion (e.g., adding google-docs when gmail // already has a token but missing the documents scope). // Initial OAuth setup is triggered via configure. let mut resp = ActionResponse::ok(result.message); - if let Ok(auth_result) = ext_mgr.auth(&name).await + if let Ok(auth_result) = ext_mgr.auth(&name, &user.user_id).await && auth_result.auth_url().is_some() { resp.auth_url = auth_result.auth_url().map(String::from); @@ -2172,10 +2203,10 @@ async fn extensions_activate_handler( } // Activation failed due to auth; try authenticating first. - match ext_mgr.auth(&name).await { + match ext_mgr.auth(&name, &user.user_id).await { Ok(auth_result) if auth_result.is_authenticated() => { // Auth succeeded, retry activation. - match ext_mgr.activate(&name).await { + match ext_mgr.activate(&name, &user.user_id).await { Ok(result) => Ok(Json(ActionResponse::ok(result.message))), Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))), } @@ -2206,22 +2237,57 @@ async fn extensions_activate_handler( /// Redirect `/projects/{id}` to `/projects/{id}/` so relative paths in /// the served HTML resolve within the project namespace. -async fn project_redirect_handler(Path(project_id): Path) -> impl IntoResponse { - axum::response::Redirect::permanent(&format!("/projects/{project_id}/")) +async fn project_redirect_handler( + State(state): State>, + super::auth::AuthenticatedUser(user): super::auth::AuthenticatedUser, + Path(project_id): Path, +) -> impl IntoResponse { + if !verify_project_ownership(&state, &project_id, &user.user_id).await { + return (StatusCode::NOT_FOUND, "Not found").into_response(); + } + axum::response::Redirect::permanent(&format!("/projects/{project_id}/")).into_response() } /// Serve `index.html` when hitting `/projects/{project_id}/`. -async fn project_index_handler(Path(project_id): Path) -> impl IntoResponse { +async fn project_index_handler( + State(state): State>, + super::auth::AuthenticatedUser(user): super::auth::AuthenticatedUser, + Path(project_id): Path, +) -> impl IntoResponse { + if !verify_project_ownership(&state, &project_id, &user.user_id).await { + return (StatusCode::NOT_FOUND, "Not found").into_response(); + } serve_project_file(&project_id, "index.html").await } /// Serve any file under `/projects/{project_id}/{path}`. async fn project_file_handler( + State(state): State>, + super::auth::AuthenticatedUser(user): super::auth::AuthenticatedUser, Path((project_id, path)): Path<(String, String)>, ) -> impl IntoResponse { + if !verify_project_ownership(&state, &project_id, &user.user_id).await { + return (StatusCode::NOT_FOUND, "Not found").into_response(); + } serve_project_file(&project_id, &path).await } +/// Check that a project directory belongs to a job owned by the given user. +/// Returns false if the store is unavailable or the project is not found. +async fn verify_project_ownership(state: &GatewayState, project_id: &str, user_id: &str) -> bool { + let Some(ref store) = state.store else { + return false; + }; + // The project_id is a sandbox job UUID used as the directory name. + let Ok(job_id) = project_id.parse::() else { + return false; + }; + match store.get_sandbox_job(job_id).await { + Ok(Some(job)) => job.user_id == user_id, + _ => false, + } +} + /// Shared logic: resolve the file inside `~/.ironclaw/projects/{project_id}/`, /// guard against path traversal, and stream the content with the right MIME type. async fn serve_project_file(project_id: &str, path: &str) -> axum::response::Response { @@ -2264,6 +2330,7 @@ async fn serve_project_file(project_id: &str, path: &str) -> axum::response::Res async fn extensions_remove_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(name): Path, ) -> Result, (StatusCode, String)> { let ext_mgr = state.extension_manager.as_ref().ok_or(( @@ -2271,7 +2338,7 @@ async fn extensions_remove_handler( "Extension manager not available (secrets store required)".to_string(), ))?; - match ext_mgr.remove(&name).await { + match ext_mgr.remove(&name, &user.user_id).await { Ok(message) => Ok(Json(ActionResponse::ok(message))), Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))), } @@ -2279,6 +2346,7 @@ async fn extensions_remove_handler( async fn extensions_registry_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Query(params): Query, ) -> Json { let query = params.query.unwrap_or_default(); @@ -2311,7 +2379,7 @@ async fn extensions_registry_handler( let installed: std::collections::HashSet<(String, String)> = if let Some(ext_mgr) = state.extension_manager.as_ref() { ext_mgr - .list(None, false) + .list(None, false, &user.user_id) .await .unwrap_or_default() .into_iter() @@ -2342,6 +2410,7 @@ async fn extensions_registry_handler( async fn extensions_setup_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(name): Path, ) -> Result, (StatusCode, String)> { let ext_mgr = state.extension_manager.as_ref().ok_or(( @@ -2350,12 +2419,12 @@ async fn extensions_setup_handler( ))?; let setup = ext_mgr - .get_setup_schema(&name) + .get_setup_schema(&name, &user.user_id) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; let kind = ext_mgr - .list(None, false) + .list(None, false, &user.user_id) .await .ok() .and_then(|list| list.into_iter().find(|e| e.name == name)) @@ -2372,6 +2441,7 @@ async fn extensions_setup_handler( async fn extensions_setup_submit_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(name): Path, Json(req): Json, ) -> Result, (StatusCode, String)> { @@ -2382,9 +2452,12 @@ async fn extensions_setup_submit_handler( // Clear auth mode regardless of outcome so the next user message goes // through to the LLM instead of being intercepted as a token. - clear_auth_mode(&state).await; + clear_auth_mode(&state, &user.user_id).await; - match ext_mgr.configure(&name, &req.secrets, &req.fields).await { + match ext_mgr + .configure(&name, &req.secrets, &req.fields, &user.user_id) + .await + { Ok(result) => { let mut resp = if result.verification.is_some() || result.activated { ActionResponse::ok(result.message) @@ -2401,11 +2474,14 @@ async fn extensions_setup_submit_handler( if result.verification.is_none() { // Broadcast auth_completed so the chat UI can dismiss any in-progress // auth card or setup modal that was triggered by tool_auth/tool_activate. - state.sse.broadcast(SseEvent::AuthCompleted { - extension_name: name.clone(), - success: result.activated, - message: resp.message.clone(), - }); + state.sse.broadcast_for_user( + &user.user_id, + SseEvent::AuthCompleted { + extension_name: name.clone(), + success: result.activated, + message: resp.message.clone(), + }, + ); } Ok(Json(resp)) } @@ -2462,6 +2538,7 @@ async fn pairing_approve_handler( async fn routines_runs_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(id): Path, ) -> Result, (StatusCode, String)> { let store = state.store.as_ref().ok_or(( @@ -2472,6 +2549,17 @@ async fn routines_runs_handler( let routine_id = Uuid::parse_str(&id) .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?; + // Verify ownership before listing runs. + let routine = store + .get_routine(routine_id) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? + .ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?; + + if routine.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Routine not found".to_string())); + } + let runs = store .list_routine_runs(routine_id, 50) .await @@ -2501,12 +2589,13 @@ async fn routines_runs_handler( async fn settings_list_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result, StatusCode> { let store = state .store .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; - let rows = store.list_settings(&state.user_id).await.map_err(|e| { + let rows = store.list_settings(&user.user_id).await.map_err(|e| { tracing::error!("Failed to list settings: {}", e); StatusCode::INTERNAL_SERVER_ERROR })?; @@ -2525,6 +2614,7 @@ async fn settings_list_handler( async fn settings_get_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(key): Path, ) -> Result, StatusCode> { let store = state @@ -2532,7 +2622,7 @@ async fn settings_get_handler( .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; let row = store - .get_setting_full(&state.user_id, &key) + .get_setting_full(&user.user_id, &key) .await .map_err(|e| { tracing::error!("Failed to get setting '{}': {}", key, e); @@ -2549,6 +2639,7 @@ async fn settings_get_handler( async fn settings_set_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(key): Path, Json(body): Json, ) -> Result { @@ -2557,7 +2648,7 @@ async fn settings_set_handler( .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; store - .set_setting(&state.user_id, &key, &body.value) + .set_setting(&user.user_id, &key, &body.value) .await .map_err(|e| { tracing::error!("Failed to set setting '{}': {}", key, e); @@ -2569,6 +2660,7 @@ async fn settings_set_handler( async fn settings_delete_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(key): Path, ) -> Result { let store = state @@ -2576,7 +2668,7 @@ async fn settings_delete_handler( .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; store - .delete_setting(&state.user_id, &key) + .delete_setting(&user.user_id, &key) .await .map_err(|e| { tracing::error!("Failed to delete setting '{}': {}", key, e); @@ -2588,12 +2680,13 @@ async fn settings_delete_handler( async fn settings_export_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result, StatusCode> { let store = state .store .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; - let settings = store.get_all_settings(&state.user_id).await.map_err(|e| { + let settings = store.get_all_settings(&user.user_id).await.map_err(|e| { tracing::error!("Failed to export settings: {}", e); StatusCode::INTERNAL_SERVER_ERROR })?; @@ -2603,6 +2696,7 @@ async fn settings_export_handler( async fn settings_import_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Json(body): Json, ) -> Result { let store = state @@ -2610,7 +2704,7 @@ async fn settings_import_handler( .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; store - .set_all_settings(&state.user_id, &body.settings) + .set_all_settings(&user.user_id, &body.settings) .await .map_err(|e| { tracing::error!("Failed to import settings: {}", e); @@ -2624,6 +2718,7 @@ async fn settings_import_handler( async fn gateway_status_handler( State(state): State>, + AuthenticatedUser(_user): AuthenticatedUser, ) -> Json { let sse_connections = state.sse.connection_count(); let ws_connections = state @@ -2870,8 +2965,9 @@ mod tests { fn test_gateway_state(ext_mgr: Option>) -> Arc { Arc::new(GatewayState { msg_tx: tokio::sync::RwLock::new(None), - sse: SseManager::new(), + sse: Arc::new(SseManager::new()), workspace: None, + workspace_pool: None, session_manager: None, log_broadcaster: None, log_level_handle: None, @@ -2880,14 +2976,14 @@ mod tests { store: None, job_manager: None, prompt_queue: None, - user_id: "test".to_string(), + default_user_id: "test".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: None, llm_provider: None, skill_registry: None, skill_catalog: None, scheduler: None, - chat_rate_limiter: RateLimiter::new(30, 60), + chat_rate_limiter: PerUserRateLimiter::new(30, 60), oauth_rate_limiter: RateLimiter::new(10, 60), webhook_rate_limiter: RateLimiter::new(10, 60), registry_entries: vec![], @@ -2951,12 +3047,18 @@ mod tests { "BOT_TOKEN": "dummy-token" } }); - let req = axum::http::Request::builder() + let mut req = axum::http::Request::builder() .method("POST") .uri(format!("/api/extensions/{channel_name}/setup")) .header("content-type", "application/json") .body(Body::from(req_body.to_string())) .expect("request"); + // Inject AuthenticatedUser so the handler's extractor succeeds + // without needing the full auth middleware layer. + req.extensions_mut().insert(UserIdentity { + user_id: "test".to_string(), + workspace_read_scopes: Vec::new(), + }); let resp = ServiceExt::>::oneshot(app, req) .await @@ -3029,12 +3131,18 @@ mod tests { "telegram_bot_token": "123456789:ABCdefGhI" } }); - let req = axum::http::Request::builder() + let mut req = axum::http::Request::builder() .method("POST") .uri("/api/extensions/telegram/setup") .header("content-type", "application/json") .body(Body::from(req_body.to_string())) .expect("request"); + // Inject AuthenticatedUser so the handler's extractor succeeds + // without needing the full auth middleware layer. + req.extensions_mut().insert(UserIdentity { + user_id: "test".to_string(), + workspace_read_scopes: Vec::new(), + }); let resp = ServiceExt::>::oneshot(app, req) .await @@ -3056,7 +3164,12 @@ mod tests { break; } match timeout(remaining, receiver.recv()).await { - Ok(Ok(crate::channels::web::types::SseEvent::AuthRequired { .. })) => { + Ok(Ok(scoped)) + if matches!( + scoped.event, + crate::channels::web::types::SseEvent::AuthRequired { .. } + ) => + { panic!("verification responses should not emit auth_required SSE events") } Ok(Ok(_)) => continue, @@ -3077,7 +3190,8 @@ mod tests { let state = test_gateway_state(None); let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); - let bound = start_server(addr, state.clone(), "test-token".to_string()) + let auth = MultiAuthState::single("test-token".to_string(), "test".to_string()); + let bound = start_server(addr, state.clone(), auth) .await .expect("server should start"); @@ -3239,7 +3353,7 @@ mod tests { scopes: vec![], user_id: "test".to_string(), secrets, - sse_sender: None, + sse_manager: None, gateway_token: None, token_exchange_extra_params: std::collections::HashMap::new(), client_id_secret_name: None, @@ -3287,7 +3401,8 @@ mod tests { ))); let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(secrets.clone()); - let (sender, mut receiver) = tokio::sync::broadcast::channel(4); + let sse_mgr = Arc::new(SseManager::new()); + let mut receiver = sse_mgr.sender().subscribe(); let Some(created_at) = expired_flow_created_at() else { eprintln!("Skipping expired OAuth flow SSE test: monotonic uptime below expiry window"); return; @@ -3307,7 +3422,7 @@ mod tests { scopes: vec![], user_id: "test".to_string(), secrets, - sse_sender: Some(sender), + sse_manager: Some(sse_mgr), gateway_token: None, token_exchange_extra_params: std::collections::HashMap::new(), client_id_secret_name: None, @@ -3333,7 +3448,7 @@ mod tests { .expect("response"); assert_eq!(resp.status(), StatusCode::OK); - match receiver.recv().await.expect("auth_completed event") { + match receiver.recv().await.expect("auth_completed event").event { crate::channels::web::types::SseEvent::AuthCompleted { extension_name, success, @@ -3410,7 +3525,7 @@ mod tests { scopes: vec![], user_id: "test".to_string(), secrets, - sse_sender: None, + sse_manager: None, gateway_token: None, token_exchange_extra_params: std::collections::HashMap::new(), client_id_secret_name: None, @@ -3497,7 +3612,7 @@ mod tests { scopes: vec![], user_id: "test".to_string(), secrets, - sse_sender: None, + sse_manager: None, gateway_token: None, token_exchange_extra_params: std::collections::HashMap::new(), client_id_secret_name: None, @@ -3718,4 +3833,36 @@ mod tests { let exists = secrets.exists("test", &state_key).await.unwrap_or(true); assert!(!exists, "CSRF nonce should be deleted after use"); } + + #[test] + fn test_is_local_origin_localhost() { + assert!(is_local_origin("http://localhost:3001")); + assert!(is_local_origin("http://localhost")); + assert!(is_local_origin("https://localhost:3001")); + } + + #[test] + fn test_is_local_origin_ipv4() { + assert!(is_local_origin("http://127.0.0.1:3001")); + assert!(is_local_origin("http://127.0.0.1")); + } + + #[test] + fn test_is_local_origin_ipv6() { + assert!(is_local_origin("http://[::1]:3001")); + assert!(is_local_origin("http://[::1]")); + } + + #[test] + fn test_is_local_origin_rejects_remote() { + assert!(!is_local_origin("http://evil.com")); + assert!(!is_local_origin("http://localhost.evil.com")); + assert!(!is_local_origin("http://192.168.1.1:3001")); + } + + #[test] + fn test_is_local_origin_rejects_garbage() { + assert!(!is_local_origin("not-a-url")); + assert!(!is_local_origin("")); + } } diff --git a/src/channels/web/sse.rs b/src/channels/web/sse.rs index 7b952346..46841e19 100644 --- a/src/channels/web/sse.rs +++ b/src/channels/web/sse.rs @@ -17,9 +17,25 @@ use crate::channels::web::types::SseEvent; /// Prevents resource exhaustion from connection flooding. const MAX_CONNECTIONS: u64 = 100; +/// Envelope for broadcast events: carries an optional user scope. +/// +/// `user_id = None` means the event is global (e.g. Heartbeat) and delivered +/// to all subscribers. `user_id = Some(id)` means the event is only delivered +/// to subscribers that match that user_id. +#[derive(Debug, Clone)] +pub(crate) struct ScopedEvent { + pub(crate) user_id: Option, + pub(crate) event: SseEvent, +} + /// Manages SSE broadcast to all connected browser tabs. +/// +/// In multi-user mode, events are scoped by user_id so that each subscriber +/// only receives events intended for their user (plus global events like +/// Heartbeat). In single-user mode, all events are delivered to all subscribers +/// (backwards compatible). pub struct SseManager { - tx: broadcast::Sender, + tx: broadcast::Sender, connection_count: Arc, max_connections: u64, } @@ -45,7 +61,7 @@ impl SseManager { /// only be called before the server starts accepting connections (i.e., /// during startup wiring). Calling it after connections are established /// will break connection tracking and allow exceeding `MAX_CONNECTIONS`. - pub fn from_sender(tx: broadcast::Sender) -> Self { + pub(crate) fn from_sender(tx: broadcast::Sender) -> Self { Self { tx, connection_count: Arc::new(AtomicU64::new(0)), @@ -53,15 +69,28 @@ impl SseManager { } } - /// Broadcast an event to all connected clients. - pub fn broadcast(&self, event: SseEvent) { - // Ignore send errors (no receivers is fine) - let _ = self.tx.send(event); + /// Get a clone of the broadcast sender for use by other components. + pub(crate) fn sender(&self) -> broadcast::Sender { + self.tx.clone() } - /// Get a clone of the broadcast sender for use by other components. - pub fn sender(&self) -> broadcast::Sender { - self.tx.clone() + /// Broadcast an event to all connected clients (global/unscoped). + pub fn broadcast(&self, event: SseEvent) { + let _ = self.tx.send(ScopedEvent { + user_id: None, + event, + }); + } + + /// Broadcast an event scoped to a specific user. + /// + /// Only subscribers for this user_id (or unscoped subscribers) will + /// receive the event. + pub fn broadcast_for_user(&self, user_id: &str, event: SseEvent) { + let _ = self.tx.send(ScopedEvent { + user_id: Some(user_id.to_string()), + event, + }); } /// Get current number of active connections. @@ -71,11 +100,15 @@ impl SseManager { /// Create a raw broadcast subscription for non-SSE consumers (e.g. WebSocket). /// - /// Returns a stream of `SseEvent` values and increments/decrements the - /// connection counter on creation/drop, just like `subscribe()` does for SSE. + /// When `user_id` is `Some`, only events scoped to that user (or global + /// events) are delivered. When `None`, all events are delivered (single-user + /// backwards compatibility). /// /// Returns `None` if the maximum connection limit has been reached. - pub fn subscribe_raw(&self) -> Option + Send + 'static + use<>> { + pub fn subscribe_raw( + &self, + user_id: Option, + ) -> 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); @@ -91,7 +124,19 @@ impl SseManager { .ok()?; let rx = self.tx.subscribe(); - let stream = BroadcastStream::new(rx).filter_map(|result| result.ok()); + let stream = BroadcastStream::new(rx).filter_map(move |result| match result { + Ok(scoped) => { + // Global events (user_id=None) always pass through. + // Scoped events only pass if the subscriber matches (or subscriber is unscoped). + match (&user_id, &scoped.user_id) { + (_, None) => Some(scoped.event), // global -> all + (None, _) => Some(scoped.event), // unscoped subscriber -> all + (Some(sub), Some(ev)) if sub == ev => Some(scoped.event), // match + _ => None, // different user -> skip + } + } + Err(_) => None, + }); Some(CountedStream { inner: stream, @@ -101,9 +146,13 @@ impl SseManager { /// Create a new SSE stream for a client connection. /// + /// When `user_id` is `Some`, only events for that user (or global events) + /// are delivered. When `None`, all events are delivered. + /// /// Returns `None` if the maximum connection limit has been reached. pub fn subscribe( &self, + user_id: Option, ) -> Option> + Send + 'static + use<>>> { // Atomically increment only if below the limit. let counter = Arc::clone(&self.connection_count); @@ -120,9 +169,23 @@ impl SseManager { let rx = self.tx.subscribe(); let stream = BroadcastStream::new(rx) - .filter_map(|result| result.ok()) - .map(|event| { - let data = serde_json::to_string(&event).unwrap_or_default(); + .filter_map(move |result| match result { + Ok(scoped) => match (&user_id, &scoped.user_id) { + (_, None) => Some(scoped.event), + (None, _) => Some(scoped.event), + (Some(sub), Some(ev)) if sub == ev => Some(scoped.event), + _ => None, + }, + Err(_) => None, + }) + .filter_map(|event| { + let data = match serde_json::to_string(&event) { + Ok(s) => s, + Err(e) => { + tracing::warn!("Failed to serialize SSE event: {}", e); + return None; + } + }; let event_type = match &event { SseEvent::Response { .. } => "response", SseEvent::Thinking { .. } => "thinking", @@ -147,7 +210,7 @@ impl SseManager { SseEvent::TurnCost { .. } => "turn_cost", SseEvent::ExtensionStatus { .. } => "extension_status", }; - Ok(Event::default().event(event_type).data(data)) + Some(Ok(Event::default().event(event_type).data(data))) }); // Wrap in a stream that decrements on drop @@ -215,16 +278,14 @@ mod tests { #[tokio::test] async fn test_broadcast_to_receiver() { let manager = SseManager::new(); - let mut rx = BroadcastStream::new(manager.tx.subscribe()); + let mut stream = Box::pin(manager.subscribe_raw(None).expect("should subscribe")); manager.broadcast(SseEvent::Status { message: "test".to_string(), thread_id: None, }); - let event = rx.next().await; - assert!(event.is_some()); - let event = event.unwrap().unwrap(); + let event = stream.next().await.unwrap(); match event { SseEvent::Status { message, .. } => assert_eq!(message, "test"), _ => panic!("unexpected event type"), @@ -234,7 +295,7 @@ mod tests { #[tokio::test] async fn test_subscribe_raw_receives_events() { let manager = SseManager::new(); - let mut stream = Box::pin(manager.subscribe_raw().expect("should subscribe")); + let mut stream = Box::pin(manager.subscribe_raw(None).expect("should subscribe")); assert_eq!(manager.connection_count(), 1); @@ -254,7 +315,7 @@ mod tests { async fn test_subscribe_raw_decrements_on_drop() { let manager = SseManager::new(); { - let _stream = Box::pin(manager.subscribe_raw().expect("should subscribe")); + let _stream = Box::pin(manager.subscribe_raw(None).expect("should subscribe")); assert_eq!(manager.connection_count(), 1); } // Stream dropped, counter should decrement @@ -264,8 +325,8 @@ mod tests { #[tokio::test] async fn test_subscribe_raw_multiple_subscribers() { let manager = SseManager::new(); - let mut s1 = Box::pin(manager.subscribe_raw().expect("should subscribe")); - let mut s2 = Box::pin(manager.subscribe_raw().expect("should subscribe")); + let mut s1 = Box::pin(manager.subscribe_raw(None).expect("should subscribe")); + let mut s2 = Box::pin(manager.subscribe_raw(None).expect("should subscribe")); assert_eq!(manager.connection_count(), 2); manager.broadcast(SseEvent::Heartbeat); @@ -286,12 +347,51 @@ mod tests { let mut manager = SseManager::new(); manager.max_connections = 2; // Low limit for testing - let _s1 = Box::pin(manager.subscribe_raw().expect("first should succeed")); - let _s2 = Box::pin(manager.subscribe_raw().expect("second should succeed")); + let _s1 = Box::pin(manager.subscribe_raw(None).expect("first should succeed")); + let _s2 = Box::pin(manager.subscribe_raw(None).expect("second should succeed")); assert_eq!(manager.connection_count(), 2); // Third should be rejected - assert!(manager.subscribe_raw().is_none()); - assert!(manager.subscribe().is_none()); + assert!(manager.subscribe_raw(None).is_none()); + assert!(manager.subscribe(None).is_none()); + } + + #[tokio::test] + async fn test_scoped_events_filtered_by_user() { + let manager = SseManager::new(); + let mut alice = Box::pin( + manager + .subscribe_raw(Some("alice".to_string())) + .expect("subscribe"), + ); + let mut bob = Box::pin( + manager + .subscribe_raw(Some("bob".to_string())) + .expect("subscribe"), + ); + + // Send event scoped to alice + manager.broadcast_for_user( + "alice", + SseEvent::Status { + message: "alice only".to_string(), + thread_id: None, + }, + ); + + // Send global event + manager.broadcast(SseEvent::Heartbeat); + + // Alice gets her scoped event + let e = alice.next().await.unwrap(); + assert!(matches!(e, SseEvent::Status { .. })); + + // Alice also gets the global heartbeat + let e = alice.next().await.unwrap(); + assert!(matches!(e, SseEvent::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 } } diff --git a/src/channels/web/test_helpers.rs b/src/channels/web/test_helpers.rs index 8751be6a..802512a6 100644 --- a/src/channels/web/test_helpers.rs +++ b/src/channels/web/test_helpers.rs @@ -10,7 +10,8 @@ use std::sync::Arc; use tokio::sync::mpsc; use crate::channels::IncomingMessage; -use crate::channels::web::server::{GatewayState, RateLimiter, start_server}; +use crate::channels::web::auth::MultiAuthState; +use crate::channels::web::server::{GatewayState, PerUserRateLimiter, RateLimiter, start_server}; use crate::channels::web::sse::SseManager; use crate::channels::web::ws::WsConnectionTracker; @@ -64,8 +65,9 @@ impl TestGatewayBuilder { pub fn build(self) -> Arc { Arc::new(GatewayState { msg_tx: tokio::sync::RwLock::new(self.msg_tx), - sse: SseManager::new(), + sse: Arc::new(SseManager::new()), workspace: None, + workspace_pool: None, session_manager: None, log_broadcaster: None, log_level_handle: None, @@ -74,14 +76,14 @@ impl TestGatewayBuilder { store: None, job_manager: None, prompt_queue: None, - user_id: self.user_id, + default_user_id: self.user_id, shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: self.llm_provider, skill_registry: None, skill_catalog: None, scheduler: None, - chat_rate_limiter: RateLimiter::new(30, 60), + chat_rate_limiter: PerUserRateLimiter::new(30, 60), oauth_rate_limiter: RateLimiter::new(10, 60), webhook_rate_limiter: RateLimiter::new(10, 60), registry_entries: Vec::new(), @@ -98,11 +100,26 @@ impl TestGatewayBuilder { self, auth_token: &str, ) -> Result<(SocketAddr, Arc), crate::error::ChannelError> { + let auth = MultiAuthState::single(auth_token.to_string(), "test-user".to_string()); let state = self.build(); let addr: SocketAddr = "127.0.0.1:0" .parse() - .expect("hard-coded address must parse"); - let bound = start_server(addr, state.clone(), auth_token.to_string()).await?; + .expect("hard-coded address must parse"); // safety: constant literal + let bound = start_server(addr, state.clone(), auth).await?; + Ok((bound, state)) + } + + /// Build the state and start a gateway server with multi-user auth. + /// Returns the bound address and the shared state. + pub async fn start_multi( + self, + auth: MultiAuthState, + ) -> Result<(SocketAddr, Arc), crate::error::ChannelError> { + let state = self.build(); + let addr: SocketAddr = "127.0.0.1:0" + .parse() + .expect("hard-coded address must parse"); // safety: constant literal + let bound = start_server(addr, state.clone(), auth).await?; Ok((bound, state)) } } diff --git a/src/channels/web/tests/mod.rs b/src/channels/web/tests/mod.rs new file mode 100644 index 00000000..fa6db197 --- /dev/null +++ b/src/channels/web/tests/mod.rs @@ -0,0 +1,3 @@ +//! Integration tests for the web gateway module. + +mod multi_tenant; diff --git a/src/channels/web/tests/multi_tenant.rs b/src/channels/web/tests/multi_tenant.rs new file mode 100644 index 00000000..55010831 --- /dev/null +++ b/src/channels/web/tests/multi_tenant.rs @@ -0,0 +1,796 @@ +//! Multi-tenant isolation tests for the web gateway. +//! +//! Tests cover workspace pool scoping, job handler isolation, and auth +//! enforcement on protected endpoints. Uses `LibSqlBackend::new_local()` +//! with a temporary directory for a real (but ephemeral) database. + +use std::collections::HashMap; +use std::sync::Arc; +use std::time::Duration; + +use axum::Router; +use axum::body::Body; +use axum::http::{Method, Request, StatusCode}; +use axum::middleware; +use axum::routing::{delete, get, post}; +use tower::ServiceExt; +use uuid::Uuid; + +use crate::channels::web::auth::{ + AuthenticatedUser, MultiAuthState, UserIdentity, auth_middleware, +}; +use crate::channels::web::server::{ + ActiveConfigSnapshot, GatewayState, PerUserRateLimiter, PromptQueue, RateLimiter, WorkspacePool, +}; +use crate::channels::web::sse::SseManager; + +// ── Helpers ──────────────────────────────────────────────────────────── + +/// Create a two-user `MultiAuthState` for alice and bob. +fn two_user_auth() -> MultiAuthState { + let mut tokens = HashMap::new(); + tokens.insert( + "tok-alice".to_string(), + UserIdentity { + user_id: "alice".to_string(), + workspace_read_scopes: vec!["shared".to_string()], + }, + ); + tokens.insert( + "tok-bob".to_string(), + UserIdentity { + user_id: "bob".to_string(), + workspace_read_scopes: vec!["shared".to_string(), "alice".to_string()], + }, + ); + MultiAuthState::multi(tokens) +} + +/// Build a `GatewayState` with configurable store and prompt queue. +fn build_state( + store: Option>, + prompt_queue: Option, +) -> Arc { + Arc::new(GatewayState { + msg_tx: tokio::sync::RwLock::new(None), + sse: Arc::new(SseManager::new()), + workspace: None, + workspace_pool: None, + session_manager: None, + log_broadcaster: None, + log_level_handle: None, + extension_manager: None, + tool_registry: None, + store, + job_manager: None, + prompt_queue, + default_user_id: "test".to_string(), + shutdown_tx: tokio::sync::RwLock::new(None), + ws_tracker: None, + llm_provider: None, + skill_registry: None, + skill_catalog: None, + scheduler: None, + chat_rate_limiter: PerUserRateLimiter::new(30, 60), + oauth_rate_limiter: RateLimiter::new(10, 60), + webhook_rate_limiter: RateLimiter::new(10, 60), + registry_entries: Vec::new(), + cost_guard: None, + routine_engine: Arc::new(tokio::sync::RwLock::new(None)), + startup_time: std::time::Instant::now(), + active_config: ActiveConfigSnapshot::default(), + }) +} + +/// Create a libSQL-backed test database in a temporary directory. +/// +/// Returns the database and a `TempDir` guard — the database file is +/// deleted when the guard is dropped. +#[cfg(feature = "libsql")] +async fn test_db() -> (Arc, tempfile::TempDir) { + use crate::db::Database; + let dir = tempfile::tempdir().expect("failed to create temp dir"); // safety: test-only + let path = dir.path().join("test.db"); + let backend = crate::db::libsql::LibSqlBackend::new_local(&path) + .await + .expect("failed to create test LibSqlBackend"); // safety: test-only + backend + .run_migrations() + .await + .expect("failed to run migrations"); // safety: test-only + (Arc::new(backend) as Arc, dir) +} + +/// Build a minimal Routine for testing. +fn make_routine(user_id: &str, name: &str) -> crate::agent::routine::Routine { + let now = chrono::Utc::now(); + crate::agent::routine::Routine { + id: Uuid::new_v4(), + name: name.to_string(), + description: format!("Test routine: {name}"), + user_id: user_id.to_string(), + enabled: true, + trigger: crate::agent::routine::Trigger::Cron { + schedule: "0 9 * * *".to_string(), + timezone: None, + }, + action: crate::agent::routine::RoutineAction::Lightweight { + prompt: "hello".to_string(), + context_paths: vec![], + max_tokens: 1024, + use_tools: false, + max_tool_rounds: 3, + }, + guardrails: crate::agent::routine::RoutineGuardrails { + cooldown: Duration::from_secs(60), + max_concurrent: 1, + dedup_window: None, + }, + notify: crate::agent::routine::NotifyConfig { + channel: None, + user: None, + on_success: false, + on_failure: true, + on_attention: true, + }, + last_run_at: None, + next_fire_at: None, + run_count: 0, + consecutive_failures: 0, + state: serde_json::json!({}), + created_at: now, + updated_at: now, + } +} + +/// Build a minimal SandboxJobRecord for testing. +fn make_sandbox_job(user_id: &str, task: &str) -> crate::history::SandboxJobRecord { + let now = chrono::Utc::now(); + crate::history::SandboxJobRecord { + id: Uuid::new_v4(), + task: task.to_string(), + status: "completed".to_string(), + user_id: user_id.to_string(), + project_dir: format!("/tmp/test-{}", Uuid::new_v4()), + success: Some(true), + failure_reason: None, + created_at: now, + started_at: Some(now), + completed_at: Some(now), + credential_grants_json: "[]".to_string(), + } +} + +// ═══════════════════════════════════════════════════════════════════════ +// WorkspacePool Tests +// ═══════════════════════════════════════════════════════════════════════ + +#[cfg(feature = "libsql")] +mod workspace_pool { + use super::*; + use crate::config::{WorkspaceConfig, WorkspaceSearchConfig}; + use crate::workspace::EmbeddingCacheConfig; + use crate::workspace::layer::MemoryLayer; + + #[tokio::test] + async fn test_workspace_pool_applies_search_config() { + let (db, _dir) = test_db().await; + let search_config = WorkspaceSearchConfig { + rrf_k: 42, + ..Default::default() + }; + let pool = WorkspacePool::new( + db, + None, + EmbeddingCacheConfig::default(), + search_config, + WorkspaceConfig::default(), + ); + let identity = UserIdentity { + user_id: "alice".to_string(), + workspace_read_scopes: vec![], + }; + let ws = pool.get_or_create(&identity).await; + assert_eq!(ws.user_id(), "alice"); + } + + #[tokio::test] + async fn test_workspace_pool_applies_memory_layers() { + let (db, _dir) = test_db().await; + let layers = vec![MemoryLayer { + name: "shared-layer".to_string(), + scope: "shared".to_string(), + writable: false, + sensitivity: Default::default(), + }]; + let ws_config = WorkspaceConfig { + memory_layers: layers, + read_scopes: vec![], + }; + let pool = WorkspacePool::new( + db, + None, + EmbeddingCacheConfig::default(), + WorkspaceSearchConfig::default(), + ws_config, + ); + let identity = UserIdentity { + user_id: "alice".to_string(), + workspace_read_scopes: vec![], + }; + let ws = pool.get_or_create(&identity).await; + // Memory layer scope "shared" should appear in read_user_ids. + assert!( + ws.read_user_ids().contains(&"shared".to_string()), + "expected 'shared' in read_user_ids, got {:?}", + ws.read_user_ids() + ); + } + + #[tokio::test] + async fn test_workspace_pool_applies_identity_read_scopes() { + let (db, _dir) = test_db().await; + let pool = WorkspacePool::new( + db, + None, + EmbeddingCacheConfig::default(), + WorkspaceSearchConfig::default(), + WorkspaceConfig::default(), + ); + let identity = UserIdentity { + user_id: "bob".to_string(), + workspace_read_scopes: vec!["alice".to_string(), "shared".to_string()], + }; + let ws = pool.get_or_create(&identity).await; + assert_eq!(ws.user_id(), "bob"); + assert!( + ws.read_user_ids().contains(&"alice".to_string()), + "expected 'alice' in read_user_ids from identity scopes" + ); + assert!( + ws.read_user_ids().contains(&"shared".to_string()), + "expected 'shared' in read_user_ids from identity scopes" + ); + } + + #[tokio::test] + async fn test_workspace_pool_caches_per_user() { + let (db, _dir) = test_db().await; + let pool = WorkspacePool::new( + db, + None, + EmbeddingCacheConfig::default(), + WorkspaceSearchConfig::default(), + WorkspaceConfig::default(), + ); + let alice_id = UserIdentity { + user_id: "alice".to_string(), + workspace_read_scopes: vec![], + }; + let bob_id = UserIdentity { + user_id: "bob".to_string(), + workspace_read_scopes: vec![], + }; + + let alice_ws1 = pool.get_or_create(&alice_id).await; + let alice_ws2 = pool.get_or_create(&alice_id).await; + let bob_ws = pool.get_or_create(&bob_id).await; + + // Same user gets the same Arc. + assert!(Arc::ptr_eq(&alice_ws1, &alice_ws2)); + // Different users get different instances. + assert!(!Arc::ptr_eq(&alice_ws1, &bob_ws)); + assert_eq!(alice_ws1.user_id(), "alice"); + assert_eq!(bob_ws.user_id(), "bob"); + } + + #[tokio::test] + async fn test_workspace_pool_combines_global_and_identity_scopes() { + let (db, _dir) = test_db().await; + let ws_config = WorkspaceConfig { + memory_layers: vec![], + read_scopes: vec!["global-shared".to_string()], + }; + let pool = WorkspacePool::new( + db, + None, + EmbeddingCacheConfig::default(), + WorkspaceSearchConfig::default(), + ws_config, + ); + let identity = UserIdentity { + user_id: "alice".to_string(), + workspace_read_scopes: vec!["token-scope".to_string()], + }; + let ws = pool.get_or_create(&identity).await; + let scopes = ws.read_user_ids(); + // Primary scope + assert!(scopes.contains(&"alice".to_string())); + // Global config scope + assert!( + scopes.contains(&"global-shared".to_string()), + "expected global scope 'global-shared', got {:?}", + scopes + ); + // Token identity scope + assert!( + scopes.contains(&"token-scope".to_string()), + "expected token scope 'token-scope', got {:?}", + scopes + ); + } +} + +// ═══════════════════════════════════════════════════════════════════════ +// Jobs Handler Isolation Tests +// ═══════════════════════════════════════════════════════════════════════ + +#[cfg(feature = "libsql")] +mod jobs_isolation { + use super::*; + use crate::channels::web::handlers::jobs::{ + jobs_cancel_handler, jobs_prompt_handler, jobs_restart_handler, jobs_summary_handler, + }; + // SandboxStore methods are accessed through the Database supertrait. + + /// Build a router with job endpoints behind multi-user auth. + fn jobs_router(state: Arc, auth: MultiAuthState) -> Router { + Router::new() + .route("/api/jobs/summary", get(jobs_summary_handler)) + .route("/api/jobs/{id}/cancel", post(jobs_cancel_handler)) + .route("/api/jobs/{id}/restart", post(jobs_restart_handler)) + .route("/api/jobs/{id}/prompt", post(jobs_prompt_handler)) + .layer(middleware::from_fn_with_state(auth, auth_middleware)) + .with_state(state) + } + + #[tokio::test] + async fn test_jobs_summary_scoped_to_user() { + let (db, _dir) = test_db().await; + + // Insert sandbox jobs for alice and bob. + let alice_job = make_sandbox_job("alice", "alice task"); + let bob_job = make_sandbox_job("bob", "bob task"); + db.save_sandbox_job(&alice_job).await.unwrap(); + db.save_sandbox_job(&bob_job).await.unwrap(); + + let state = build_state(Some(db), None); + let auth = two_user_auth(); + let app = jobs_router(state, auth); + + // Alice should see 1 job. + let req = Request::builder() + .uri("/api/jobs/summary") + .header("Authorization", "Bearer tok-alice") + .body(Body::empty()) + .unwrap(); + let resp = app.clone().oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + let body: serde_json::Value = + serde_json::from_slice(&axum::body::to_bytes(resp.into_body(), 4096).await.unwrap()) + .unwrap(); + assert_eq!(body["total"], 1, "alice should see only her own jobs"); + + // Bob should see 1 job. + let req = Request::builder() + .uri("/api/jobs/summary") + .header("Authorization", "Bearer tok-bob") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + let body: serde_json::Value = + serde_json::from_slice(&axum::body::to_bytes(resp.into_body(), 4096).await.unwrap()) + .unwrap(); + assert_eq!(body["total"], 1, "bob should see only his own jobs"); + } + + #[tokio::test] + async fn test_jobs_restart_rejects_other_user() { + let (db, _dir) = test_db().await; + + // Insert a failed sandbox job owned by alice. + let mut alice_job = make_sandbox_job("alice", "alice task"); + alice_job.status = "failed".to_string(); + alice_job.success = Some(false); + db.save_sandbox_job(&alice_job).await.unwrap(); + + let state = build_state(Some(db), None); + let auth = two_user_auth(); + let app = jobs_router(state, auth); + + // Bob tries to restart alice's job. + let req = Request::builder() + .method(Method::POST) + .uri(format!("/api/jobs/{}/restart", alice_job.id)) + .header("Authorization", "Bearer tok-bob") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!( + resp.status(), + StatusCode::NOT_FOUND, + "bob should not be able to restart alice's job" + ); + } + + #[tokio::test] + async fn test_jobs_prompt_works_for_agent_jobs() { + let (db, _dir) = test_db().await; + + // Insert a running sandbox job owned by alice in claude_code mode. + let mut alice_job = make_sandbox_job("alice", "prompt test"); + alice_job.status = "running".to_string(); + alice_job.success = None; + alice_job.completed_at = None; + db.save_sandbox_job(&alice_job).await.unwrap(); + db.update_sandbox_job_mode(alice_job.id, "claude_code") + .await + .unwrap(); + + let prompt_queue: PromptQueue = + Arc::new(tokio::sync::Mutex::new(std::collections::HashMap::new())); + let state = build_state(Some(db), Some(prompt_queue.clone())); + let auth = two_user_auth(); + let app = jobs_router(state, auth); + + // Alice prompts her own job. + let req = Request::builder() + .method(Method::POST) + .uri(format!("/api/jobs/{}/prompt", alice_job.id)) + .header("Authorization", "Bearer tok-alice") + .header("Content-Type", "application/json") + .body(Body::from( + serde_json::to_string(&serde_json::json!({"content": "hello"})).unwrap(), + )) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!( + resp.status(), + StatusCode::OK, + "alice should be able to prompt her own job" + ); + + // Verify prompt was enqueued. + let queue = prompt_queue.lock().await; + assert!( + queue.contains_key(&alice_job.id), + "prompt queue should contain alice's job" + ); + } + + #[tokio::test] + async fn test_jobs_prompt_rejects_other_user() { + let (db, _dir) = test_db().await; + + let mut alice_job = make_sandbox_job("alice", "alice task"); + alice_job.status = "running".to_string(); + alice_job.success = None; + alice_job.completed_at = None; + db.save_sandbox_job(&alice_job).await.unwrap(); + db.update_sandbox_job_mode(alice_job.id, "claude_code") + .await + .unwrap(); + + let prompt_queue: PromptQueue = + Arc::new(tokio::sync::Mutex::new(std::collections::HashMap::new())); + let state = build_state(Some(db), Some(prompt_queue)); + let auth = two_user_auth(); + let app = jobs_router(state, auth); + + // Bob tries to prompt alice's job. + let req = Request::builder() + .method(Method::POST) + .uri(format!("/api/jobs/{}/prompt", alice_job.id)) + .header("Authorization", "Bearer tok-bob") + .header("Content-Type", "application/json") + .body(Body::from( + serde_json::to_string(&serde_json::json!({"content": "sneaky"})).unwrap(), + )) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!( + resp.status(), + StatusCode::NOT_FOUND, + "bob should not be able to prompt alice's job" + ); + } + + #[tokio::test] + async fn test_jobs_cancel_rejects_other_user() { + let (db, _dir) = test_db().await; + + let mut alice_job = make_sandbox_job("alice", "alice running"); + alice_job.status = "running".to_string(); + alice_job.success = None; + alice_job.completed_at = None; + db.save_sandbox_job(&alice_job).await.unwrap(); + + let state = build_state(Some(db), None); + let auth = two_user_auth(); + let app = jobs_router(state, auth); + + // Bob tries to cancel alice's job. + let req = Request::builder() + .method(Method::POST) + .uri(format!("/api/jobs/{}/cancel", alice_job.id)) + .header("Authorization", "Bearer tok-bob") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!( + resp.status(), + StatusCode::NOT_FOUND, + "bob should not be able to cancel alice's job" + ); + } +} + +// ═══════════════════════════════════════════════════════════════════════ +// Routines Isolation Tests +// ═══════════════════════════════════════════════════════════════════════ + +#[cfg(feature = "libsql")] +mod routines_isolation { + use super::*; + use crate::channels::web::handlers::routines::{ + routines_delete_handler, routines_detail_handler, routines_list_handler, + routines_summary_handler, routines_toggle_handler, + }; + // RoutineStore methods are accessed through the Database supertrait. + + fn routines_router(state: Arc, auth: MultiAuthState) -> Router { + Router::new() + .route("/api/routines", get(routines_list_handler)) + .route("/api/routines/summary", get(routines_summary_handler)) + .route("/api/routines/{id}", get(routines_detail_handler)) + .route("/api/routines/{id}/toggle", post(routines_toggle_handler)) + .route("/api/routines/{id}", delete(routines_delete_handler)) + .layer(middleware::from_fn_with_state(auth, auth_middleware)) + .with_state(state) + } + + #[tokio::test] + async fn test_routines_isolation() { + let (db, _dir) = test_db().await; + + // Create routines for alice and bob. + let alice_routine = make_routine("alice", "alice-daily"); + let bob_routine = make_routine("bob", "bob-daily"); + db.create_routine(&alice_routine).await.unwrap(); + db.create_routine(&bob_routine).await.unwrap(); + + let state = build_state(Some(db), None); + let auth = two_user_auth(); + let app = routines_router(state, auth); + + // Alice sees only her routine in the list. + let req = Request::builder() + .uri("/api/routines") + .header("Authorization", "Bearer tok-alice") + .body(Body::empty()) + .unwrap(); + let resp = app.clone().oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + let body: serde_json::Value = + serde_json::from_slice(&axum::body::to_bytes(resp.into_body(), 8192).await.unwrap()) + .unwrap(); + let routines = body["routines"].as_array().unwrap(); + assert_eq!(routines.len(), 1, "alice should see only her routines"); + assert_eq!(routines[0]["name"], "alice-daily"); + + // Bob sees only his routine. + let req = Request::builder() + .uri("/api/routines") + .header("Authorization", "Bearer tok-bob") + .body(Body::empty()) + .unwrap(); + let resp = app.clone().oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + let body: serde_json::Value = + serde_json::from_slice(&axum::body::to_bytes(resp.into_body(), 8192).await.unwrap()) + .unwrap(); + let routines = body["routines"].as_array().unwrap(); + assert_eq!(routines.len(), 1, "bob should see only his routines"); + assert_eq!(routines[0]["name"], "bob-daily"); + + // Bob cannot view alice's routine detail. + let req = Request::builder() + .uri(format!("/api/routines/{}", alice_routine.id)) + .header("Authorization", "Bearer tok-bob") + .body(Body::empty()) + .unwrap(); + let resp = app.clone().oneshot(req).await.unwrap(); + assert_eq!( + resp.status(), + StatusCode::NOT_FOUND, + "bob should not see alice's routine detail" + ); + + // Bob cannot toggle alice's routine. + let req = Request::builder() + .method(Method::POST) + .uri(format!("/api/routines/{}/toggle", alice_routine.id)) + .header("Authorization", "Bearer tok-bob") + .body(Body::empty()) + .unwrap(); + let resp = app.clone().oneshot(req).await.unwrap(); + assert_eq!( + resp.status(), + StatusCode::NOT_FOUND, + "bob should not toggle alice's routine" + ); + + // Bob cannot delete alice's routine. + let req = Request::builder() + .method(Method::DELETE) + .uri(format!("/api/routines/{}", alice_routine.id)) + .header("Authorization", "Bearer tok-bob") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!( + resp.status(), + StatusCode::NOT_FOUND, + "bob should not delete alice's routine" + ); + } +} + +// ═══════════════════════════════════════════════════════════════════════ +// Handler Auth Enforcement Tests +// ═══════════════════════════════════════════════════════════════════════ + +mod auth_enforcement { + use super::*; + + /// Dummy handler that extracts `AuthenticatedUser` — if the auth middleware + /// rejects the request, this handler is never reached. + async fn authed_handler(AuthenticatedUser(_user): AuthenticatedUser) -> &'static str { + "ok" + } + + /// Build a router with the real auth middleware and dummy handlers at all + /// the paths we want to verify require authentication. + fn auth_test_router(auth: MultiAuthState) -> Router { + let state = build_state(None, None); + Router::new() + // Routines + .route("/api/routines", get(authed_handler)) + .route("/api/routines/summary", get(authed_handler)) + .route("/api/routines/{id}", get(authed_handler)) + .route("/api/routines/{id}/toggle", post(authed_handler)) + .route("/api/routines/{id}", delete(authed_handler)) + // Skills + .route("/api/skills", get(authed_handler)) + .route("/api/skills/search", post(authed_handler)) + .route("/api/skills/install", post(authed_handler)) + .route("/api/skills/{name}", delete(authed_handler)) + // Logs + .route("/api/logs/events", get(authed_handler)) + .route("/api/logs/level", get(authed_handler).put(authed_handler)) + // Gateway status + .route("/api/gateway/status", get(authed_handler)) + .layer(middleware::from_fn_with_state(auth, auth_middleware)) + .with_state(state) + } + + /// Send a request without auth and assert it returns UNAUTHORIZED. + async fn assert_requires_auth(app: &Router, method: Method, uri: &str) { + let req = Request::builder() + .method(method.clone()) + .uri(uri) + .body(Body::empty()) + .unwrap(); + let resp = app.clone().oneshot(req).await.unwrap(); + assert_eq!( + resp.status(), + StatusCode::UNAUTHORIZED, + "{} {} should require auth", + method, + uri + ); + } + + /// Send a request with a valid token and assert it succeeds. + async fn assert_passes_with_token(app: &Router, method: Method, uri: &str, token: &str) { + let req = Request::builder() + .method(method.clone()) + .uri(uri) + .header("Authorization", format!("Bearer {token}")) + .body(Body::empty()) + .unwrap(); + let resp = app.clone().oneshot(req).await.unwrap(); + assert_eq!( + resp.status(), + StatusCode::OK, + "{} {} should pass with valid token", + method, + uri + ); + } + + #[tokio::test] + async fn test_routines_handlers_require_auth() { + let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string()); + let app = auth_test_router(auth); + let id = Uuid::new_v4(); + + assert_requires_auth(&app, Method::GET, "/api/routines").await; + assert_requires_auth(&app, Method::GET, "/api/routines/summary").await; + assert_requires_auth(&app, Method::GET, &format!("/api/routines/{id}")).await; + assert_requires_auth(&app, Method::POST, &format!("/api/routines/{id}/toggle")).await; + assert_requires_auth(&app, Method::DELETE, &format!("/api/routines/{id}")).await; + } + + #[tokio::test] + async fn test_skills_handlers_require_auth() { + let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string()); + let app = auth_test_router(auth); + + assert_requires_auth(&app, Method::GET, "/api/skills").await; + assert_requires_auth(&app, Method::POST, "/api/skills/search").await; + assert_requires_auth(&app, Method::POST, "/api/skills/install").await; + assert_requires_auth(&app, Method::DELETE, "/api/skills/test-skill").await; + } + + #[tokio::test] + async fn test_logs_handlers_require_auth() { + let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string()); + let app = auth_test_router(auth); + + assert_requires_auth(&app, Method::GET, "/api/logs/events").await; + assert_requires_auth(&app, Method::GET, "/api/logs/level").await; + assert_requires_auth(&app, Method::PUT, "/api/logs/level").await; + } + + #[tokio::test] + async fn test_gateway_status_requires_auth() { + let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string()); + let app = auth_test_router(auth); + + assert_requires_auth(&app, Method::GET, "/api/gateway/status").await; + } + + #[tokio::test] + async fn test_valid_token_passes_all_endpoints() { + let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string()); + let app = auth_test_router(auth); + let id = Uuid::new_v4(); + + assert_passes_with_token(&app, Method::GET, "/api/routines", "secret-tok").await; + assert_passes_with_token(&app, Method::GET, "/api/skills", "secret-tok").await; + assert_passes_with_token(&app, Method::GET, "/api/logs/events", "secret-tok").await; + assert_passes_with_token(&app, Method::GET, "/api/gateway/status", "secret-tok").await; + assert_passes_with_token( + &app, + Method::GET, + &format!("/api/routines/{id}"), + "secret-tok", + ) + .await; + } + + #[tokio::test] + async fn test_wrong_token_rejected_on_all_endpoints() { + let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string()); + let app = auth_test_router(auth); + + // Wrong token should be rejected. + let req = Request::builder() + .uri("/api/routines") + .header("Authorization", "Bearer wrong-tok") + .body(Body::empty()) + .unwrap(); + let resp = app.clone().oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + + let req = Request::builder() + .uri("/api/gateway/status") + .header("Authorization", "Bearer wrong-tok") + .body(Body::empty()) + .unwrap(); + let resp = app.oneshot(req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); + } +} diff --git a/src/channels/web/ws.rs b/src/channels/web/ws.rs index 470c3422..3a601679 100644 --- a/src/channels/web/ws.rs +++ b/src/channels/web/ws.rs @@ -62,7 +62,11 @@ impl Default for WsConnectionTracker { /// /// When either task ends (client disconnect or broadcast closed), both are /// cleaned up. -pub async fn handle_ws_connection(socket: WebSocket, state: Arc) { +pub async fn handle_ws_connection( + socket: WebSocket, + state: Arc, + user: crate::channels::web::auth::UserIdentity, +) { let (mut ws_sink, mut ws_stream) = socket.split(); // Track connection @@ -71,9 +75,9 @@ pub async fn handle_ws_connection(socket: WebSocket, state: Arc) { } let tracker_for_drop = state.ws_tracker.clone(); - // Subscribe to broadcast events (same source as SSE). + // Subscribe to broadcast events (same source as SSE), scoped to this user. // Reject if we've hit the connection limit. - let Some(raw_stream) = state.sse.subscribe_raw() else { + let Some(raw_stream) = state.sse.subscribe_raw(Some(user.user_id.clone())) else { tracing::warn!("WebSocket rejected: too many connections"); // Decrement the WS tracker we already incremented above. if let Some(ref tracker) = tracker_for_drop { @@ -117,7 +121,7 @@ pub async fn handle_ws_connection(socket: WebSocket, state: Arc) { }); // Receiver task: read client frames and route to agent - let user_id = state.user_id.clone(); + let user_id = user.user_id; while let Some(Ok(frame)) = ws_stream.next().await { match frame { Message::Text(text) => { @@ -263,10 +267,14 @@ async fn handle_client_message( token, } => { if let Some(ref ext_mgr) = state.extension_manager { - match ext_mgr.configure_token(&extension_name, &token).await { + match ext_mgr + .configure_token(&extension_name, &token, user_id) + .await + { Ok(result) => { if result.verification.is_some() { - state.sse.broadcast( + state.sse.broadcast_for_user( + user_id, crate::channels::web::types::SseEvent::AuthRequired { extension_name: extension_name.clone(), instructions: Some(result.message), @@ -275,8 +283,9 @@ async fn handle_client_message( }, ); } else { - crate::channels::web::server::clear_auth_mode(state).await; - state.sse.broadcast( + crate::channels::web::server::clear_auth_mode(state, user_id).await; + state.sse.broadcast_for_user( + user_id, crate::channels::web::types::SseEvent::AuthCompleted { extension_name, success: true, @@ -288,7 +297,8 @@ async fn handle_client_message( Err(e) => { let msg = format!("Auth failed: {}", e); if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) { - state.sse.broadcast( + state.sse.broadcast_for_user( + user_id, crate::channels::web::types::SseEvent::AuthRequired { extension_name: extension_name.clone(), instructions: Some(msg.clone()), @@ -311,7 +321,7 @@ async fn handle_client_message( } } WsClientMessage::AuthCancel { .. } => { - crate::channels::web::server::clear_auth_mode(state).await; + crate::channels::web::server::clear_auth_mode(state, user_id).await; } WsClientMessage::Ping => { let _ = direct_tx.send(WsServerMessage::Pong).await; @@ -498,8 +508,9 @@ mod tests { GatewayState { msg_tx: tokio::sync::RwLock::new(msg_tx), - sse: SseManager::new(), + sse: Arc::new(SseManager::new()), workspace: None, + workspace_pool: None, session_manager: None, log_broadcaster: None, log_level_handle: None, @@ -509,13 +520,13 @@ mod tests { job_manager: None, prompt_queue: None, scheduler: None, - user_id: "test".to_string(), + default_user_id: "test".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: None, skill_registry: None, skill_catalog: None, - chat_rate_limiter: crate::channels::web::server::RateLimiter::new(30, 60), + chat_rate_limiter: crate::channels::web::server::PerUserRateLimiter::new(30, 60), oauth_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60), webhook_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60), registry_entries: Vec::new(), diff --git a/src/cli/oauth_defaults.rs b/src/cli/oauth_defaults.rs index 531d474e..3b57872f 100644 --- a/src/cli/oauth_defaults.rs +++ b/src/cli/oauth_defaults.rs @@ -447,8 +447,8 @@ pub struct PendingOAuthFlow { pub user_id: String, /// Secrets store reference for token persistence. pub secrets: Arc, - /// SSE broadcast sender for notifying the web UI. - pub sse_sender: Option>, + /// SSE broadcast manager for notifying the web UI. + pub sse_manager: Option>, /// Gateway auth token for authenticating with the platform token exchange proxy. pub gateway_token: Option, /// Additional form params for the token exchange request. diff --git a/src/config/channels.rs b/src/config/channels.rs index d249dd18..d9c2c0a9 100644 --- a/src/config/channels.rs +++ b/src/config/channels.rs @@ -2,6 +2,7 @@ use std::collections::HashMap; use std::path::PathBuf; use secrecy::SecretString; +use serde::Deserialize; use crate::bootstrap::ironclaw_base_dir; use crate::config::helpers::{optional_env, parse_bool_env, parse_optional_env}; @@ -45,6 +46,26 @@ pub struct GatewayConfig { /// Bearer token for authentication. Random hex generated at startup if unset. pub auth_token: Option, pub user_id: String, + /// Additional user scopes for workspace reads. + /// + /// When set, the workspace will be able to read (search, read, list) from + /// these additional user scopes while writes remain isolated to `user_id`. + /// Parsed from `WORKSPACE_READ_SCOPES` (comma-separated). + pub workspace_read_scopes: Vec, + /// Memory layer definitions (JSON in env var, or from external config). + pub memory_layers: Vec, + /// Multi-user token map. When set, each token maps to a user identity. + /// Parsed from `GATEWAY_USER_TOKENS` (JSON string). When absent, falls back + /// to single-user mode via `auth_token` + `user_id`. + pub user_tokens: Option>, +} + +/// Per-user token configuration for multi-user mode. +#[derive(Debug, Clone, Deserialize)] +pub struct UserTokenConfig { + pub user_id: String, + #[serde(default)] + pub workspace_read_scopes: Vec, } /// Signal channel configuration (signal-cli daemon HTTP/JSON-RPC). @@ -115,6 +136,118 @@ impl ChannelsConfig { .or_else(|| cs.gateway_user_id.clone()) .unwrap_or_else(|| owner_id.to_string()); + let memory_layers: Vec = + match optional_env("MEMORY_LAYERS")? { + Some(json_str) => { + serde_json::from_str(&json_str).map_err(|e| ConfigError::InvalidValue { + key: "MEMORY_LAYERS".to_string(), + message: format!("must be valid JSON array of layer objects: {e}"), + })? + } + None => crate::workspace::layer::MemoryLayer::default_for_user(&user_id), + }; + + // Validate layer names and scopes + for layer in &memory_layers { + if layer.name.trim().is_empty() { + return Err(ConfigError::InvalidValue { + key: "MEMORY_LAYERS".to_string(), + message: "layer name must not be empty".to_string(), + }); + } + if layer.name.len() > 64 { + return Err(ConfigError::InvalidValue { + key: "MEMORY_LAYERS".to_string(), + message: format!("layer name '{}' exceeds 64 characters", layer.name), + }); + } + if !layer + .name + .chars() + .all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '-') + { + return Err(ConfigError::InvalidValue { + key: "MEMORY_LAYERS".to_string(), + message: format!( + "layer name '{}' contains invalid characters \ + (allowed: a-z, A-Z, 0-9, _, -)", + layer.name + ), + }); + } + if layer.scope.trim().is_empty() { + return Err(ConfigError::InvalidValue { + key: "MEMORY_LAYERS".to_string(), + message: format!("layer '{}' has an empty scope", layer.name), + }); + } + } + + // Check for duplicate layer names + { + let mut seen = std::collections::HashSet::new(); + for layer in &memory_layers { + if !seen.insert(&layer.name) { + return Err(ConfigError::InvalidValue { + key: "MEMORY_LAYERS".to_string(), + message: format!("duplicate layer name '{}'", layer.name), + }); + } + } + } + + let user_tokens: Option> = + match optional_env("GATEWAY_USER_TOKENS")? { + Some(json_str) => { + let tokens: HashMap = serde_json::from_str( + &json_str, + ) + .map_err(|e| ConfigError::InvalidValue { + key: "GATEWAY_USER_TOKENS".to_string(), + message: format!( + "must be valid JSON object mapping tokens to user configs: {e}" + ), + })?; + if tokens.is_empty() { + return Err(ConfigError::InvalidValue { + key: "GATEWAY_USER_TOKENS".to_string(), + message: + "token map is empty — remove the variable to use single-user mode" + .to_string(), + }); + } + for (tok, cfg) in &tokens { + if cfg.user_id.trim().is_empty() { + return Err(ConfigError::InvalidValue { + key: "GATEWAY_USER_TOKENS".to_string(), + message: format!( + "token '{}...' has an empty user_id", + &tok[..tok.len().min(8)] + ), + }); + } + } + Some(tokens) + } + None => None, + }; + let workspace_read_scopes: Vec = optional_env("WORKSPACE_READ_SCOPES")? + .map(|s| { + s.split(',') + .map(|s| s.trim().to_string()) + .filter(|s| !s.is_empty()) + .collect() + }) + .unwrap_or_default(); + + for scope in &workspace_read_scopes { + if scope.len() > 128 { + return Err(ConfigError::InvalidValue { + key: "WORKSPACE_READ_SCOPES".to_string(), + message: format!("scope '{}...' exceeds 128 characters", &scope[..32]), + }); + } + } Some(GatewayConfig { host: optional_env("GATEWAY_HOST")? .or_else(|| cs.gateway_host.clone()) @@ -126,6 +259,9 @@ impl ChannelsConfig { auth_token: optional_env("GATEWAY_AUTH_TOKEN")? .or_else(|| cs.gateway_auth_token.clone()), user_id, + workspace_read_scopes, + memory_layers, + user_tokens, }) } else { None @@ -281,6 +417,9 @@ mod tests { port: 3000, auth_token: Some("tok-abc".to_string()), user_id: "default".to_string(), + workspace_read_scopes: vec![], + memory_layers: vec![], + user_tokens: None, }; assert_eq!(cfg.host, "127.0.0.1"); assert_eq!(cfg.port, 3000); @@ -295,6 +434,9 @@ mod tests { port: 3001, auth_token: None, user_id: "anon".to_string(), + workspace_read_scopes: vec![], + memory_layers: vec![], + user_tokens: None, }; assert!(cfg.auth_token.is_none()); } diff --git a/src/db/libsql/jobs.rs b/src/db/libsql/jobs.rs index 208d348b..297a9282 100644 --- a/src/db/libsql/jobs.rs +++ b/src/db/libsql/jobs.rs @@ -230,6 +230,49 @@ impl JobStore for LibSqlBackend { Ok(jobs) } + async fn list_agent_jobs_for_user( + &self, + user_id: &str, + ) -> Result, DatabaseError> { + let conn = self.connect().await?; + let mut rows = conn + .query( + r#" + SELECT id, title, status, user_id, failure_reason, + created_at, started_at, completed_at + FROM agent_jobs WHERE source = 'direct' AND user_id = ?1 + ORDER BY created_at DESC + "#, + params![user_id], + ) + .await + .map_err(|e| DatabaseError::Query(e.to_string()))?; + + let mut jobs = Vec::new(); + while let Some(row) = rows + .next() + .await + .map_err(|e| DatabaseError::Query(e.to_string()))? + { + let id_str = get_text(&row, 0); + let Ok(id) = id_str.parse() else { + tracing::warn!("Skipping agent job with invalid UUID: {}", id_str); + continue; + }; + jobs.push(AgentJobRecord { + id, + title: get_text(&row, 1), + status: get_text(&row, 2), + user_id: get_text(&row, 3), + failure_reason: get_opt_text(&row, 4), + created_at: get_ts(&row, 5), + started_at: get_opt_ts(&row, 6), + completed_at: get_opt_ts(&row, 7), + }); + } + Ok(jobs) + } + async fn get_agent_job_failure_reason( &self, id: Uuid, @@ -277,6 +320,32 @@ impl JobStore for LibSqlBackend { Ok(summary) } + async fn agent_job_summary_for_user( + &self, + user_id: &str, + ) -> Result { + let conn = self.connect().await?; + let mut rows = conn + .query( + "SELECT status, COUNT(*) as cnt FROM agent_jobs WHERE source = 'direct' AND user_id = ?1 GROUP BY status", + params![user_id], + ) + .await + .map_err(|e| DatabaseError::Query(e.to_string()))?; + + let mut summary = AgentJobSummary::default(); + while let Some(row) = rows + .next() + .await + .map_err(|e| DatabaseError::Query(e.to_string()))? + { + let status = get_text(&row, 0); + let count = get_i64(&row, 1) as usize; + summary.add_count(&status, count); + } + Ok(summary) + } + async fn save_action(&self, job_id: Uuid, action: &ActionRecord) -> Result<(), DatabaseError> { let conn = self.connect().await?; let duration_ms = action.duration.as_millis() as i64; diff --git a/src/db/mod.rs b/src/db/mod.rs index 0c84d35d..c0594bda 100644 --- a/src/db/mod.rs +++ b/src/db/mod.rs @@ -409,7 +409,15 @@ pub trait JobStore: Send + Sync { async fn mark_job_stuck(&self, id: Uuid) -> Result<(), DatabaseError>; async fn get_stuck_jobs(&self) -> Result, DatabaseError>; async fn list_agent_jobs(&self) -> Result, DatabaseError>; + async fn list_agent_jobs_for_user( + &self, + user_id: &str, + ) -> Result, DatabaseError>; async fn agent_job_summary(&self) -> Result; + async fn agent_job_summary_for_user( + &self, + user_id: &str, + ) -> Result; /// Get the failure reason for a single agent job (O(1) lookup). async fn get_agent_job_failure_reason(&self, id: Uuid) -> Result, DatabaseError>; diff --git a/src/db/postgres.rs b/src/db/postgres.rs index cfa10997..a2c686d3 100644 --- a/src/db/postgres.rs +++ b/src/db/postgres.rs @@ -249,10 +249,24 @@ impl JobStore for PgBackend { self.store.list_agent_jobs().await } + async fn list_agent_jobs_for_user( + &self, + user_id: &str, + ) -> Result, DatabaseError> { + self.store.list_agent_jobs_for_user(user_id).await + } + async fn agent_job_summary(&self) -> Result { self.store.agent_job_summary().await } + async fn agent_job_summary_for_user( + &self, + user_id: &str, + ) -> Result { + self.store.agent_job_summary_for_user(user_id).await + } + async fn get_agent_job_failure_reason( &self, id: Uuid, diff --git a/src/extensions/manager.rs b/src/extensions/manager.rs index df5de72d..7da9e980 100644 --- a/src/extensions/manager.rs +++ b/src/extensions/manager.rs @@ -411,9 +411,8 @@ pub struct ExtensionManager { installed_relay_extensions: RwLock>, /// Last activation error for each WASM channel (ephemeral, cleared on success). activation_errors: RwLock>, - /// SSE broadcast sender (set post-construction via `set_sse_sender()`). - sse_sender: - RwLock>>, + /// SSE broadcast manager (set post-construction via `set_sse_sender()`). + sse_manager: RwLock>>, /// Shared registry of pending OAuth flows for gateway-routed callbacks. /// /// Keyed by CSRF `state` parameter. Populated in `start_wasm_oauth()` @@ -484,7 +483,7 @@ impl ExtensionManager { pub async fn active_tool_names(&self) -> HashSet { let mut names = HashSet::new(); - match self.list(None, false).await { + match self.list(None, false, &self.user_id).await { Ok(extensions) => { for extension in extensions { match extension.kind { @@ -550,7 +549,7 @@ impl ExtensionManager { active_channel_names: RwLock::new(HashSet::new()), installed_relay_extensions: RwLock::new(HashSet::new()), activation_errors: RwLock::new(HashMap::new()), - sse_sender: RwLock::new(None), + sse_manager: RwLock::new(None), pending_oauth_flows: crate::cli::oauth_defaults::new_pending_oauth_registry(), gateway_token: std::env::var("GATEWAY_AUTH_TOKEN").ok(), relay_config: crate::config::RelayConfig::from_env(), @@ -892,25 +891,18 @@ impl ExtensionManager { *self.relay_channel_manager.write().await = Some(channel_manager); } - /// Check if a channel name corresponds to a relay extension (has stored team_id + /// Check if a channel name corresponds to a relay extension (has stored stream token /// or is tracked in the installed relay extensions set). - pub async fn is_relay_channel(&self, name: &str) -> bool { + pub async fn is_relay_channel(&self, name: &str, user_id: &str) -> bool { // Check in-memory installed set first (supports no-store mode) if self.installed_relay_extensions.read().await.contains(name) { return true; } - // Then check persistent settings - if let Some(ref store) = self.store { - let team_id_key = format!("relay:{}:team_id", name); - store - .get_setting(&self.user_id, &team_id_key) - .await - .ok() - .flatten() - .is_some() - } else { - false - } + // Then check for stored stream token + self.secrets + .exists(user_id, &format!("relay:{}:stream_token", name)) + .await + .unwrap_or(false) } /// Restore persisted relay channels after startup. @@ -921,18 +913,18 @@ impl ExtensionManager { /// /// Call this only after `set_relay_channel_manager()` or `set_channel_runtime()`. /// Otherwise, each activation attempt fails with "Channel manager not initialized". - pub async fn restore_relay_channels(&self) { - let persisted = self.load_persisted_active_channels().await; + pub async fn restore_relay_channels(&self, user_id: &str) { + let persisted = self.load_persisted_active_channels(user_id).await; let already_active = self.active_channel_names.read().await.clone(); for name in &persisted { if already_active.contains(name) { continue; } - if !self.is_relay_channel(name).await { + if !self.is_relay_channel(name, user_id).await { continue; } - match self.activate_stored_relay(name).await { + match self.activate_stored_relay(name, user_id).await { Ok(_) => { tracing::debug!(channel = %name, "Restored persisted relay channel"); } @@ -987,7 +979,7 @@ impl ExtensionManager { /// Persist the set of active channel names to the settings store. /// /// Saved under key `activated_channels` so channels auto-activate on restart. - async fn persist_active_channels(&self) { + async fn persist_active_channels(&self, user_id: &str) { let Some(ref store) = self.store else { return; }; @@ -1000,7 +992,7 @@ impl ExtensionManager { .collect(); let value = serde_json::json!(names); if let Err(e) = store - .set_setting(&self.user_id, "activated_channels", &value) + .set_setting(user_id, "activated_channels", &value) .await { tracing::warn!(error = %e, "Failed to persist activated_channels setting"); @@ -1011,11 +1003,11 @@ impl ExtensionManager { /// /// Returns channel names that were activated in a prior session so they can /// be auto-activated at startup. - pub async fn load_persisted_active_channels(&self) -> Vec { + pub async fn load_persisted_active_channels(&self, user_id: &str) -> Vec { let Some(ref store) = self.store else { return Vec::new(); }; - match store.get_setting(&self.user_id, "activated_channels").await { + match store.get_setting(user_id, "activated_channels").await { Ok(Some(value)) => match serde_json::from_value(value) { Ok(names) => names, Err(e) => { @@ -1032,11 +1024,8 @@ impl ExtensionManager { } /// Set the SSE broadcast sender for pushing extension status events to the web UI. - pub async fn set_sse_sender( - &self, - sender: tokio::sync::broadcast::Sender, - ) { - *self.sse_sender.write().await = Some(sender); + pub async fn set_sse_sender(&self, sse: Arc) { + *self.sse_manager.write().await = Some(sse); } /// Returns the pending OAuth flow registry for sharing with the web gateway. @@ -1141,8 +1130,8 @@ 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 sender) = *self.sse_sender.read().await { - let _ = sender.send(crate::channels::web::types::SseEvent::ExtensionStatus { + if let Some(ref sse) = *self.sse_manager.read().await { + sse.broadcast(crate::channels::web::types::SseEvent::ExtensionStatus { extension_name: name.to_string(), status: status.to_string(), message: message.map(|m| m.to_string()), @@ -1186,6 +1175,7 @@ impl ExtensionManager { name: &str, url: Option<&str>, kind_hint: Option, + user_id: &str, ) -> Result { let sanitized_url = url.map(sanitize_url_for_logging); tracing::info!(extension = %name, url = ?sanitized_url, kind = ?kind_hint, "Installing extension"); @@ -1193,7 +1183,7 @@ impl ExtensionManager { // If we have a registry entry, use it (prefer kind_hint to resolve collisions) if let Some(entry) = self.registry.get_with_kind(name, kind_hint).await { - return self.install_from_entry(&entry).await.map_err(|e| { + return self.install_from_entry(&entry, user_id).await.map_err(|e| { tracing::error!(extension = %name, error = %e, "Extension install failed"); e }); @@ -1203,7 +1193,7 @@ impl ExtensionManager { if let Some(url) = url { let kind = kind_hint.unwrap_or_else(|| infer_kind_from_url(url)); return match kind { - ExtensionKind::McpServer => self.install_mcp_from_url(name, url).await, + ExtensionKind::McpServer => self.install_mcp_from_url(name, url, user_id).await, ExtensionKind::WasmTool => self.install_wasm_tool_from_url(name, url).await, ExtensionKind::WasmChannel => { self.install_wasm_channel_from_url(name, url, None).await @@ -1234,31 +1224,35 @@ impl ExtensionManager { /// /// Read-only for WASM extensions; may initiate OAuth for MCP servers. /// To provide secrets, use [`configure()`] instead. - pub async fn auth(&self, name: &str) -> Result { + pub async fn auth(&self, name: &str, user_id: &str) -> Result { // Clean up expired pending auths self.cleanup_expired_auths().await; // Determine what kind of extension this is - let kind = self.determine_installed_kind(name).await?; + let kind = self.determine_installed_kind(name, user_id).await?; match kind { - ExtensionKind::McpServer => self.auth_mcp(name).await, - ExtensionKind::WasmTool => self.auth_wasm_tool(name).await, - ExtensionKind::WasmChannel => self.auth_wasm_channel_status(name).await, - ExtensionKind::ChannelRelay => self.auth_channel_relay(name).await, + ExtensionKind::McpServer => self.auth_mcp(name, user_id).await, + ExtensionKind::WasmTool => self.auth_wasm_tool(name, user_id).await, + ExtensionKind::WasmChannel => self.auth_wasm_channel_status(name, user_id).await, + ExtensionKind::ChannelRelay => self.auth_channel_relay(name, user_id).await, } } /// Activate an installed (and optionally authenticated) extension. - pub async fn activate(&self, name: &str) -> Result { + pub async fn activate( + &self, + name: &str, + user_id: &str, + ) -> Result { Self::validate_extension_name(name)?; - let kind = self.determine_installed_kind(name).await?; + let kind = self.determine_installed_kind(name, user_id).await?; match kind { - ExtensionKind::McpServer => self.activate_mcp(name).await, - ExtensionKind::WasmTool => self.activate_wasm_tool(name).await, - ExtensionKind::WasmChannel => self.activate_wasm_channel(name).await, - ExtensionKind::ChannelRelay => self.activate_channel_relay(name).await, + ExtensionKind::McpServer => self.activate_mcp(name, user_id).await, + ExtensionKind::WasmTool => self.activate_wasm_tool(name, user_id).await, + ExtensionKind::WasmChannel => self.activate_wasm_channel(name, user_id).await, + ExtensionKind::ChannelRelay => self.activate_channel_relay(name, user_id).await, } } @@ -1270,16 +1264,16 @@ impl ExtensionManager { &self, kind_filter: Option, include_available: bool, + user_id: &str, ) -> Result, ExtensionError> { let mut extensions = Vec::new(); // List MCP servers if kind_filter.is_none() || kind_filter == Some(ExtensionKind::McpServer) { - match self.load_mcp_servers().await { + match self.load_mcp_servers(user_id).await { Ok(servers) => { for server in &servers.servers { - let authenticated = - is_authenticated(server, &self.secrets, &self.user_id).await; + let authenticated = is_authenticated(server, &self.secrets, user_id).await; let clients = self.mcp_clients.read().await; let active = clients.contains_key(&server.name); @@ -1337,7 +1331,7 @@ impl ExtensionManager { .get_with_kind(&name, Some(ExtensionKind::WasmTool)) .await; let display_name = registry_entry.as_ref().map(|e| e.display_name.clone()); - let auth_state = self.check_tool_auth_status(&name).await; + let auth_state = self.check_tool_auth_status(&name, user_id).await; let version = if let Some(ref cap_path) = discovered.capabilities_path { tokio::fs::read(cap_path) .await @@ -1384,7 +1378,7 @@ impl ExtensionManager { let errors = self.activation_errors.read().await; for (name, discovered) in channels { let active = active_names.contains(&name); - let auth_state = self.check_channel_auth_status(&name).await; + let auth_state = self.check_channel_auth_status(&name, user_id).await; let activation_error = errors.get(&name).cloned(); let registry_entry = self .registry @@ -1436,7 +1430,7 @@ impl ExtensionManager { let active_names = self.active_channel_names.read().await; for name in installed.iter() { let active = active_names.contains(name); - let has_token = self.is_relay_channel(name).await; + let has_token = self.is_relay_channel(name, user_id).await; let registry_entry = self .registry .get_with_kind(name, Some(ExtensionKind::ChannelRelay)) @@ -1499,9 +1493,9 @@ impl ExtensionManager { } /// Remove an installed extension. - pub async fn remove(&self, name: &str) -> Result { + pub async fn remove(&self, name: &str, user_id: &str) -> Result { Self::validate_extension_name(name)?; - let kind = self.determine_installed_kind(name).await?; + let kind = self.determine_installed_kind(name, user_id).await?; // Clean up any in-progress OAuth flows for this extension. // TCP mode: abort the listener task so port 9876 is freed immediately. @@ -1535,7 +1529,7 @@ impl ExtensionManager { self.mcp_clients.write().await.remove(name); // Remove from config - self.remove_mcp_server(name) + self.remove_mcp_server(name, user_id) .await .map_err(|e| ExtensionError::Config(e.to_string()))?; @@ -1595,7 +1589,7 @@ impl ExtensionManager { ExtensionKind::WasmChannel => { // Remove from active set and persist self.active_channel_names.write().await.remove(name); - self.persist_active_channels().await; + self.persist_active_channels(user_id).await; // Clear stale activation errors so reinstall starts clean self.activation_errors.write().await.remove(name); @@ -1629,15 +1623,14 @@ impl ExtensionManager { // Remove from active channels self.active_channel_names.write().await.remove(name); - self.persist_active_channels().await; + self.persist_active_channels(user_id).await; self.activation_errors.write().await.remove(name); - // Remove stored team_id - if let Some(ref store) = self.store { - let _ = store - .delete_setting(&self.user_id, &format!("relay:{}:team_id", name)) - .await; - } + // Remove stored stream token + let _ = self + .secrets + .delete(user_id, &format!("relay:{}:stream_token", name)) + .await; // Stop webhook traffic before removing the channel from the managers. self.clear_relay_webhook_state().await; @@ -1672,13 +1665,17 @@ impl ExtensionManager { /// /// The upgrade preserves authentication secrets — only the `.wasm` binary /// (and `.capabilities.json`) are replaced. - pub async fn upgrade(&self, name: Option<&str>) -> Result { + pub async fn upgrade( + &self, + name: Option<&str>, + user_id: &str, + ) -> Result { // Collect extensions to check let mut candidates: Vec<(String, ExtensionKind)> = Vec::new(); if let Some(name) = name { Self::validate_extension_name(name)?; - let kind = self.determine_installed_kind(name).await?; + let kind = self.determine_installed_kind(name, user_id).await?; if kind == ExtensionKind::McpServer { return Err(ExtensionError::Other( "MCP servers don't have WIT versions and cannot be upgraded this way" @@ -1716,7 +1713,7 @@ impl ExtensionManager { let mut outcomes = Vec::new(); for (ext_name, kind) in &candidates { - let outcome = self.upgrade_one(ext_name, *kind).await; + let outcome = self.upgrade_one(ext_name, *kind, user_id).await; outcomes.push(outcome); } @@ -1742,7 +1739,7 @@ impl ExtensionManager { } /// Upgrade a single WASM extension if its WIT version is outdated. - async fn upgrade_one(&self, name: &str, kind: ExtensionKind) -> UpgradeOutcome { + async fn upgrade_one(&self, name: &str, kind: ExtensionKind, user_id: &str) -> UpgradeOutcome { let (cap_dir, host_wit) = match kind { ExtensionKind::WasmTool => (&self.wasm_tools_dir, crate::tools::wasm::WIT_TOOL_VERSION), ExtensionKind::WasmChannel => ( @@ -1838,7 +1835,7 @@ impl ExtensionManager { } // Reinstall from registry - match self.install_from_entry(&entry).await { + match self.install_from_entry(&entry, user_id).await { Ok(_) => { tracing::info!( extension = %name, @@ -1867,9 +1864,13 @@ impl ExtensionManager { } /// Get detailed info about an installed extension (version, wit_version, host compatibility). - pub async fn extension_info(&self, name: &str) -> Result { + pub async fn extension_info( + &self, + name: &str, + user_id: &str, + ) -> Result { Self::validate_extension_name(name)?; - let kind = self.determine_installed_kind(name).await?; + let kind = self.determine_installed_kind(name, user_id).await?; match kind { ExtensionKind::WasmTool => { @@ -1950,10 +1951,11 @@ impl ExtensionManager { async fn load_mcp_servers( &self, + user_id: &str, ) -> Result { if let Some(ref store) = self.store { - crate::tools::mcp::config::load_mcp_servers_from_db(store.as_ref(), &self.user_id).await + crate::tools::mcp::config::load_mcp_servers_from_db(store.as_ref(), user_id).await } else { crate::tools::mcp::config::load_mcp_servers().await } @@ -1962,8 +1964,9 @@ impl ExtensionManager { async fn get_mcp_server( &self, name: &str, + user_id: &str, ) -> Result { - let servers = self.load_mcp_servers().await?; + let servers = self.load_mcp_servers(user_id).await?; servers.get(name).cloned().ok_or_else(|| { crate::tools::mcp::config::ConfigError::ServerNotFound { name: name.to_string(), @@ -1974,11 +1977,11 @@ impl ExtensionManager { async fn add_mcp_server( &self, config: McpServerConfig, + user_id: &str, ) -> Result<(), crate::tools::mcp::config::ConfigError> { config.validate()?; if let Some(ref store) = self.store { - crate::tools::mcp::config::add_mcp_server_db(store.as_ref(), &self.user_id, config) - .await + crate::tools::mcp::config::add_mcp_server_db(store.as_ref(), user_id, config).await } else { crate::tools::mcp::config::add_mcp_server(config).await } @@ -1987,10 +1990,10 @@ impl ExtensionManager { async fn remove_mcp_server( &self, name: &str, + user_id: &str, ) -> Result<(), crate::tools::mcp::config::ConfigError> { if let Some(ref store) = self.store { - crate::tools::mcp::config::remove_mcp_server_db(store.as_ref(), &self.user_id, name) - .await + crate::tools::mcp::config::remove_mcp_server_db(store.as_ref(), user_id, name).await } else { crate::tools::mcp::config::remove_mcp_server(name).await } @@ -2001,8 +2004,11 @@ impl ExtensionManager { async fn install_from_entry( &self, entry: &RegistryEntry, + user_id: &str, ) -> Result { - let primary_result = self.try_install_from_source(entry, &entry.source).await; + let primary_result = self + .try_install_from_source(entry, &entry.source, user_id) + .await; match fallback_decision(&primary_result, &entry.fallback_source) { FallbackDecision::Return => primary_result, FallbackDecision::TryFallback => { @@ -2017,7 +2023,7 @@ impl ExtensionManager { primary_error = %primary_err, "Primary install failed, trying fallback source" ); - match self.try_install_from_source(entry, fallback).await { + match self.try_install_from_source(entry, fallback, user_id).await { Ok(result) => Ok(result), Err(fallback_err) => { tracing::error!( @@ -2037,6 +2043,7 @@ impl ExtensionManager { &self, entry: &RegistryEntry, source: &ExtensionSource, + user_id: &str, ) -> Result { match entry.kind { ExtensionKind::McpServer => { @@ -2049,7 +2056,7 @@ impl ExtensionManager { )); } }; - self.install_mcp_from_url(&entry.name, &url).await + self.install_mcp_from_url(&entry.name, &url, user_id).await } ExtensionKind::WasmTool => match source { ExtensionSource::WasmDownload { @@ -2133,9 +2140,10 @@ impl ExtensionManager { &self, name: &str, url: &str, + user_id: &str, ) -> Result { // Check if already installed - if self.get_mcp_server(name).await.is_ok() { + if self.get_mcp_server(name, user_id).await.is_ok() { return Err(ExtensionError::AlreadyInstalled(name.to_string())); } @@ -2144,7 +2152,7 @@ impl ExtensionManager { .validate() .map_err(|e| ExtensionError::InvalidUrl(e.to_string()))?; - self.add_mcp_server(config) + self.add_mcp_server(config, user_id) .await .map_err(|e| ExtensionError::Config(e.to_string()))?; @@ -2505,14 +2513,14 @@ impl ExtensionManager { }) } - async fn auth_mcp(&self, name: &str) -> Result { + async fn auth_mcp(&self, name: &str, user_id: &str) -> Result { let server = self - .get_mcp_server(name) + .get_mcp_server(name, user_id) .await .map_err(|e| ExtensionError::NotInstalled(e.to_string()))?; // Check if already authenticated - if is_authenticated(&server, &self.secrets, &self.user_id).await { + if is_authenticated(&server, &self.secrets, user_id).await { return Ok(AuthResult::authenticated(name, ExtensionKind::McpServer)); } @@ -2520,7 +2528,7 @@ impl ExtensionManager { // open in the same browser. The gateway's /oauth/callback handler will // complete the token exchange. if self.should_use_gateway_mode() { - return match self.auth_mcp_build_url(name, &server).await { + return match self.auth_mcp_build_url(name, &server, user_id).await { Ok(result) => Ok(result), Err(ExtensionError::AuthNotSupported(_)) => Ok(AuthResult::awaiting_token( name, @@ -2537,14 +2545,14 @@ impl ExtensionManager { } // CLI/local mode: run the full blocking OAuth flow (opens browser, waits for callback) - match authorize_mcp_server(&server, &self.secrets, &self.user_id).await { + match authorize_mcp_server(&server, &self.secrets, user_id).await { Ok(_token) => { tracing::info!("MCP server '{}' authenticated via OAuth", name); Ok(AuthResult::authenticated(name, ExtensionKind::McpServer)) } Err(crate::tools::mcp::auth::AuthError::NotSupported) => { // Server doesn't support OAuth, try building a URL - match self.auth_mcp_build_url(name, &server).await { + match self.auth_mcp_build_url(name, &server, user_id).await { Ok(result) => Ok(result), Err(_) => Ok(AuthResult::awaiting_token( name, @@ -2584,6 +2592,7 @@ impl ExtensionManager { &self, name: &str, server: &McpServerConfig, + user_id: &str, ) -> Result { // Try to discover OAuth metadata and build a URL the user can open manually let metadata = discover_full_oauth_metadata(&server.url) @@ -2672,9 +2681,9 @@ impl ExtensionManager { provider: Some(format!("mcp:{}", name)), validation_endpoint: None, scopes, - user_id: self.user_id.clone(), + user_id: user_id.to_string(), secrets: Arc::clone(&self.secrets), - sse_sender: self.sse_sender.read().await.clone(), + sse_manager: self.sse_manager.read().await.clone(), gateway_token: self.gateway_token.clone(), token_exchange_extra_params, client_id_secret_name: if server.oauth.is_none() { @@ -2715,7 +2724,11 @@ impl ExtensionManager { } } - async fn auth_wasm_tool(&self, name: &str) -> Result { + async fn auth_wasm_tool( + &self, + name: &str, + user_id: &str, + ) -> Result { // Read the capabilities file to get auth config let cap_path = self .wasm_tools_dir @@ -2747,7 +2760,7 @@ impl ExtensionManager { let params = CreateSecretParams::new(&auth.secret_name, &value).with_provider(name.to_string()); self.secrets - .create(&self.user_id, params) + .create(user_id, params) .await .map_err(|e| ExtensionError::AuthFailed(e.to_string()))?; @@ -2757,7 +2770,7 @@ impl ExtensionManager { // Check if already authenticated (with scope expansion detection) let token_exists = self .secrets - .exists(&self.user_id, &auth.secret_name) + .exists(user_id, &auth.secret_name) .await .unwrap_or(false); @@ -2765,9 +2778,11 @@ impl ExtensionManager { // If this tool has OAuth config, check whether new scopes are needed let needs_reauth = if let Some(ref oauth) = auth.oauth { let merged = self - .collect_shared_scopes(&auth.secret_name, &oauth.scopes) + .collect_shared_scopes(&auth.secret_name, &oauth.scopes, user_id) + .await; + let needs = self + .needs_scope_expansion(&auth.secret_name, &merged, user_id) .await; - let needs = self.needs_scope_expansion(&auth.secret_name, &merged).await; tracing::debug!( tool = name, secret_name = %auth.secret_name, @@ -2790,7 +2805,10 @@ impl ExtensionManager { // But only if credentials are available — if the tool has setup secrets // for client_id/secret that aren't configured yet, return needs_setup. if let Some(ref oauth) = auth.oauth { - if self.needs_setup_credentials(name, &auth, oauth).await { + if self + .needs_setup_credentials(name, &auth, oauth, user_id) + .await + { let display = auth.display_name.as_deref().unwrap_or(name); return Ok(AuthResult::needs_setup( name, @@ -2804,7 +2822,7 @@ impl ExtensionManager { } return self - .start_wasm_oauth(name, &auth, oauth) + .start_wasm_oauth(name, &auth, oauth, user_id) .await .map_err(|e| ExtensionError::AuthFailed(e.to_string())); } @@ -2824,7 +2842,7 @@ impl ExtensionManager { } /// Determine the auth readiness of a WASM channel. - async fn check_channel_auth_status(&self, name: &str) -> ToolAuthState { + async fn check_channel_auth_status(&self, name: &str, user_id: &str) -> ToolAuthState { let cap_path = self .wasm_channels_dir .join(format!("{}.capabilities.json", name)); @@ -2849,7 +2867,7 @@ impl ExtensionManager { let all_provided = futures::future::join_all( required .iter() - .map(|s| self.secrets.exists(&self.user_id, &s.name)), + .map(|s| self.secrets.exists(user_id, &s.name)), ) .await .into_iter() @@ -2885,6 +2903,7 @@ impl ExtensionManager { &self, secret_name: &str, base_scopes: &[String], + _user_id: &str, ) -> Vec { let mut all_scopes: std::collections::BTreeSet = base_scopes.iter().cloned().collect(); @@ -2905,14 +2924,19 @@ impl ExtensionManager { } /// Check whether the stored scopes are insufficient for the merged scopes. - async fn needs_scope_expansion(&self, secret_name: &str, merged_scopes: &[String]) -> bool { + async fn needs_scope_expansion( + &self, + secret_name: &str, + merged_scopes: &[String], + user_id: &str, + ) -> bool { if merged_scopes.is_empty() { return false; } let scopes_key = format!("{}_scopes", secret_name); let stored_scopes: std::collections::HashSet = - match self.secrets.get_decrypted(&self.user_id, &scopes_key).await { + match self.secrets.get_decrypted(user_id, &scopes_key).await { Ok(secret) => { let scopes: std::collections::HashSet = secret .expose() @@ -2980,6 +3004,7 @@ impl ExtensionManager { name: &str, auth: &crate::tools::wasm::AuthCapabilitySchema, oauth: &crate::tools::wasm::OAuthConfigSchema, + user_id: &str, ) -> bool { let builtin = crate::cli::oauth_defaults::builtin_credentials(&auth.secret_name); let (id_entry, secret_entry) = self.find_setup_credential_names(name).await; @@ -3005,7 +3030,7 @@ impl ExtensionManager { continue; } let resolved = self - .resolve_oauth_credential(inline, env, fallback, Some(setup_name)) + .resolve_oauth_credential(inline, env, fallback, Some(setup_name), user_id) .await .is_some(); if !resolved { @@ -3025,10 +3050,11 @@ impl ExtensionManager { env_var_name: &Option, builtin_value: Option<&str>, setup_secret_name: Option<&str>, + user_id: &str, ) -> Option { // 1. Check secrets store (entered via Setup tab) if let Some(secret_name) = setup_secret_name - && let Ok(secret) = self.secrets.get_decrypted(&self.user_id, secret_name).await + && let Ok(secret) = self.secrets.get_decrypted(user_id, secret_name).await { let val = secret.expose(); if !val.is_empty() { @@ -3062,6 +3088,7 @@ impl ExtensionManager { name: &str, auth: &crate::tools::wasm::AuthCapabilitySchema, oauth: &crate::tools::wasm::OAuthConfigSchema, + user_id: &str, ) -> Result { use crate::cli::oauth_defaults; @@ -3082,6 +3109,7 @@ impl ExtensionManager { &oauth.client_id_env, builtin.as_ref().map(|c| c.client_id), setup_client_id_name.as_deref(), + user_id, ) .await .ok_or_else(|| { @@ -3110,6 +3138,7 @@ impl ExtensionManager { &oauth.client_secret_env, builtin.as_ref().map(|c| c.client_secret), setup_client_secret_name.as_deref(), + user_id, ) .await; @@ -3122,7 +3151,7 @@ impl ExtensionManager { // Merge scopes from all tools sharing this provider let merged_scopes = self - .collect_shared_scopes(&auth.secret_name, &oauth.scopes) + .collect_shared_scopes(&auth.secret_name, &oauth.scopes, user_id) .await; // Build authorization URL with CSRF state @@ -3169,9 +3198,9 @@ impl ExtensionManager { provider: auth.provider.clone(), validation_endpoint: auth.validation_endpoint.clone(), scopes: merged_scopes, - user_id: self.user_id.clone(), + user_id: user_id.to_string(), secrets: Arc::clone(&self.secrets), - sse_sender: self.sse_sender.read().await.clone(), + sse_manager: self.sse_manager.read().await.clone(), gateway_token: self.gateway_token.clone(), token_exchange_extra_params: std::collections::HashMap::new(), client_id_secret_name: None, @@ -3199,9 +3228,9 @@ impl ExtensionManager { let secret_name = auth.secret_name.clone(); let provider = auth.provider.clone(); let validation_endpoint = auth.validation_endpoint.clone(); - let user_id = self.user_id.clone(); + let user_id = user_id.to_string(); let secrets = Arc::clone(&self.secrets); - let sse_sender = self.sse_sender.read().await.clone(); + let sse_manager = self.sse_manager.read().await.clone(); let ext_name = name.to_string(); let task_handle = tokio::spawn(async move { @@ -3280,8 +3309,8 @@ impl ExtensionManager { } } - if let Some(ref sender) = sse_sender { - let _ = sender.send(crate::channels::web::types::SseEvent::AuthCompleted { + if let Some(ref sse) = sse_manager { + sse.broadcast(crate::channels::web::types::SseEvent::AuthCompleted { extension_name: ext_name, success, message, @@ -3351,7 +3380,7 @@ impl ExtensionManager { } /// Determine the auth readiness of a WASM tool. - async fn check_tool_auth_status(&self, name: &str) -> ToolAuthState { + async fn check_tool_auth_status(&self, name: &str, user_id: &str) -> ToolAuthState { let Some(cap_file) = self.load_tool_capabilities(name).await else { return ToolAuthState::NoAuth; }; @@ -3402,7 +3431,7 @@ impl ExtensionManager { if let Some(ref auth) = cap_file.auth { let has_token = self .secrets - .exists(&self.user_id, &auth.secret_name) + .exists(user_id, &auth.secret_name) .await .unwrap_or(false) || auth @@ -3420,15 +3449,36 @@ impl ExtensionManager { // No auth section — setup_is_complete was already checked above, // so if we reach here the setup requirements are satisfied. - if cap_file.setup.is_none() { - return ToolAuthState::NoAuth; - } + let setup = match &cap_file.setup { + Some(s) => s, + None => return ToolAuthState::NoAuth, + }; - ToolAuthState::Ready + let all_provided = futures::future::join_all( + setup + .required_secrets + .iter() + .filter(|s| !s.optional) + .filter(|s| !Self::is_auto_resolved_oauth_field(&s.name, &cap_file)) + .map(|s| self.secrets.exists(user_id, &s.name)), + ) + .await + .into_iter() + .all(|r| r.unwrap_or(false)); + + if all_provided { + ToolAuthState::Ready + } else { + ToolAuthState::NeedsSetup + } } /// Check auth status for a WASM channel (read-only). - async fn auth_wasm_channel_status(&self, name: &str) -> Result { + async fn auth_wasm_channel_status( + &self, + name: &str, + user_id: &str, + ) -> Result { let cap_path = self .wasm_channels_dir .join(format!("{}.capabilities.json", name)); @@ -3463,7 +3513,7 @@ impl ExtensionManager { } if !self .secrets - .exists(&self.user_id, &secret.name) + .exists(user_id, &secret.name) .await .unwrap_or(false) { @@ -3485,7 +3535,11 @@ impl ExtensionManager { )) } - async fn activate_mcp(&self, name: &str) -> Result { + async fn activate_mcp( + &self, + name: &str, + user_id: &str, + ) -> Result { // Check if already activated { let clients = self.mcp_clients.read().await; @@ -3509,7 +3563,7 @@ impl ExtensionManager { } let server = self - .get_mcp_server(name) + .get_mcp_server(name, user_id) .await .map_err(|e| ExtensionError::NotInstalled(e.to_string()))?; @@ -3518,7 +3572,7 @@ impl ExtensionManager { &self.mcp_session_manager, &self.mcp_process_manager, Some(Arc::clone(&self.secrets)), - &self.user_id, + user_id, ) .await .map_err(|e| ExtensionError::ActivationFailed(e.to_string()))?; @@ -3576,7 +3630,11 @@ impl ExtensionManager { }) } - async fn activate_wasm_tool(&self, name: &str) -> Result { + async fn activate_wasm_tool( + &self, + name: &str, + user_id: &str, + ) -> Result { // Check if already active if self.tool_registry.has(name).await { return Ok(ActivateResult { @@ -3590,7 +3648,7 @@ impl ExtensionManager { // Check auth status — block activation if required secrets are missing. // NeedsAuth (OAuth not yet completed) is allowed because configure() loads // the tool first, then starts the OAuth flow to obtain the token. - let auth_state = self.check_tool_auth_status(name).await; + let auth_state = self.check_tool_auth_status(name, user_id).await; if auth_state == ToolAuthState::NeedsSetup { return Err(ExtensionError::ActivationFailed(format!( "Tool '{}' requires configuration. Use the setup form to provide credentials.", @@ -3670,14 +3728,18 @@ impl ExtensionManager { /// Loads the channel from its WASM file, injects credentials and config, /// registers it with the webhook router, and hot-adds it to the channel manager /// so its stream feeds into the agent loop. - async fn activate_wasm_channel(&self, name: &str) -> Result { + async fn activate_wasm_channel( + &self, + name: &str, + user_id: &str, + ) -> Result { // If already active, re-inject credentials and refresh webhook secret. // Handles the case where a channel was loaded at startup before the // user saved secrets via the web UI. { let active = self.active_channel_names.read().await; if active.contains(name) { - return self.refresh_active_channel(name).await; + return self.refresh_active_channel(name, user_id).await; } } @@ -3704,7 +3766,7 @@ impl ExtensionManager { }; // Check auth status first - let auth_state = self.check_channel_auth_status(name).await; + let auth_state = self.check_channel_auth_status(name, user_id).await; if auth_state != ToolAuthState::Ready && auth_state != ToolAuthState::NoAuth { return Err(ExtensionError::ActivationFailed(format!( "Channel '{}' requires configuration. Use the setup form to provide credentials.", @@ -3914,7 +3976,7 @@ impl ExtensionManager { .insert(channel_name.clone()); // Persist activation state so the channel auto-activates on restart - self.persist_active_channels().await; + self.persist_active_channels(&self.user_id).await; tracing::info!(channel = %channel_name, "Hot-activated WASM channel"); @@ -3930,7 +3992,11 @@ impl ExtensionManager { /// /// Called when the user saves new secrets via the setup form for a channel /// that was loaded at startup (possibly without credentials). - async fn refresh_active_channel(&self, name: &str) -> Result { + async fn refresh_active_channel( + &self, + name: &str, + user_id: &str, + ) -> Result { let router = { let rt_guard = self.channel_runtime.read().await; match rt_guard.as_ref() { @@ -3964,7 +4030,7 @@ impl ExtensionManager { &existing_channel, Some(self.secrets.as_ref()), name, - &self.user_id, + user_id, ) .await { @@ -4013,7 +4079,7 @@ impl ExtensionManager { // Refresh webhook secret if let Ok(secret) = self .secrets - .get_decrypted(&self.user_id, &webhook_secret_name) + .get_decrypted(user_id, &webhook_secret_name) .await { router @@ -4028,10 +4094,7 @@ impl ExtensionManager { // Refresh signature key if let Some(ref sig_key_name) = sig_key_secret_name - && let Ok(key_secret) = self - .secrets - .get_decrypted(&self.user_id, sig_key_name) - .await + && let Ok(key_secret) = self.secrets.get_decrypted(user_id, sig_key_name).await { match router .register_signature_key(name, key_secret.expose()) @@ -4050,7 +4113,7 @@ impl ExtensionManager { if let Some(ref hmac_secret_name_ref) = hmac_secret_name { match self .secrets - .get_decrypted(&self.user_id, hmac_secret_name_ref) + .get_decrypted(user_id, hmac_secret_name_ref) .await { Ok(secret) => { @@ -4108,9 +4171,9 @@ impl ExtensionManager { // ── Channel-relay extension methods ────────────────────────────────── /// Derive a stable instance ID from the relay config and user_id. - fn relay_instance_id(&self, config: &crate::config::RelayConfig) -> String { + fn relay_instance_id(&self, config: &crate::config::RelayConfig, user_id: &str) -> String { config.instance_id.clone().unwrap_or_else(|| { - uuid::Uuid::new_v5(&uuid::Uuid::NAMESPACE_DNS, self.user_id.as_bytes()).to_string() + uuid::Uuid::new_v5(&uuid::Uuid::NAMESPACE_DNS, user_id.as_bytes()).to_string() }) } @@ -4119,9 +4182,13 @@ impl ExtensionManager { /// For Slack: initiates OAuth flow (redirect-based). /// For Telegram: accepts a bot token, registers it with channel-relay, /// and stores the returned stream token. - async fn auth_channel_relay(&self, name: &str) -> Result { - // Check if already authenticated (has stored team_id) - if self.is_relay_channel(name).await { + async fn auth_channel_relay( + &self, + name: &str, + user_id: &str, + ) -> Result { + // Check if already authenticated (stream token exists) + if self.is_relay_channel(name, user_id).await { return Ok(AuthResult::authenticated(name, ExtensionKind::ChannelRelay)); } @@ -4140,12 +4207,10 @@ impl ExtensionManager { // state and appends it to the post-OAuth redirect URL. let state_nonce = uuid::Uuid::new_v4().to_string(); let state_key = format!("relay:{}:oauth_state", name); - let _ = self.secrets.delete(&self.user_id, &state_key).await; + // Delete any stale nonce before storing the new one + let _ = self.secrets.delete(user_id, &state_key).await; self.secrets - .create( - &self.user_id, - CreateSecretParams::new(&state_key, &state_nonce), - ) + .create(user_id, CreateSecretParams::new(&state_key, &state_nonce)) .await .map_err(|e| ExtensionError::AuthFailed(format!("Failed to store OAuth state: {e}")))?; @@ -4163,23 +4228,40 @@ impl ExtensionManager { } /// Activate a channel-relay extension. - async fn activate_channel_relay(&self, name: &str) -> Result { + async fn activate_channel_relay( + &self, + name: &str, + user_id: &str, + ) -> Result { + let token_key = format!("relay:{}:stream_token", name); let team_id_key = format!("relay:{}:team_id", name); - let store = self.store.as_ref().ok_or(ExtensionError::AuthRequired)?; - let team_id = store - .get_setting(&self.user_id, &team_id_key) - .await - .ok() - .flatten() - .and_then(|v| v.as_str().map(|s| s.to_string())) - .filter(|s| !s.is_empty()) - .ok_or(ExtensionError::AuthRequired)?; + // Check if we have a stream token + // Verify auth: stream token must exist (even though we don't use it in this constructor path) + let _stream_token = match self.secrets.get_decrypted(user_id, &token_key).await { + Ok(secret) => secret.expose().to_string(), + Err(_) => { + return Err(ExtensionError::AuthRequired); + } + }; + + // Get team_id from settings + let team_id = if let Some(ref store) = self.store { + store + .get_setting(user_id, &team_id_key) + .await + .ok() + .flatten() + .and_then(|v| v.as_str().map(|s| s.to_string())) + .unwrap_or_default() + } else { + String::new() + }; // Use relay config captured at startup let relay_config = self.relay_config()?; - let instance_id = self.relay_instance_id(relay_config); + let instance_id = self.relay_instance_id(relay_config, user_id); let client = crate::channels::relay::RelayClient::new( relay_config.url.clone(), @@ -4206,13 +4288,6 @@ impl ExtensionManager { event_rx, ); - // Callback URL is now set during OAuth flow, not via PUT /callbacks. - // The relay webhook endpoint path is still needed for the web gateway. - tracing::info!( - webhook_path = %relay_config.webhook_path, - "Relay channel activated (callback URL set during OAuth)" - ); - // Hot-add to channel manager let cm_guard = self.relay_channel_manager.read().await; let channel_mgr = cm_guard.as_ref().ok_or_else(|| { @@ -4236,7 +4311,7 @@ impl ExtensionManager { .write() .await .insert(name.to_string()); - self.persist_active_channels().await; + self.persist_active_channels(user_id).await; // Broadcast status let status_msg = "Slack connected via channel relay".to_string(); @@ -4252,12 +4327,16 @@ impl ExtensionManager { } /// Activate a channel-relay extension from stored credentials (for startup reconnect). - pub async fn activate_stored_relay(&self, name: &str) -> Result<(), ExtensionError> { - self.activate_channel_relay(name).await?; + pub async fn activate_stored_relay( + &self, + name: &str, + user_id: &str, + ) -> Result<(), ExtensionError> { self.installed_relay_extensions .write() .await .insert(name.to_string()); + self.activate_channel_relay(name, user_id).await?; Ok(()) } @@ -4266,9 +4345,13 @@ impl ExtensionManager { /// This is a read-only check — it never modifies `installed_relay_extensions`. /// To mark a relay extension as installed, use `activate_stored_relay()` or /// the explicit install flow. - async fn determine_installed_kind(&self, name: &str) -> Result { + async fn determine_installed_kind( + &self, + name: &str, + user_id: &str, + ) -> Result { // Check MCP servers first - if self.get_mcp_server(name).await.is_ok() { + if self.get_mcp_server(name, user_id).await.is_ok() { return Ok(ExtensionKind::McpServer); } @@ -4288,8 +4371,8 @@ impl ExtensionManager { if self.installed_relay_extensions.read().await.contains(name) { return Ok(ExtensionKind::ChannelRelay); } - // Also check if there's a stored team_id (persisted across restarts) - if self.is_relay_channel(name).await { + // Also check if there's a stored stream token (persisted across restarts) + if self.is_relay_channel(name, user_id).await { return Ok(ExtensionKind::ChannelRelay); } @@ -4424,9 +4507,10 @@ impl ExtensionManager { pub async fn get_setup_schema( &self, name: &str, + user_id: &str, ) -> Result { Self::validate_extension_name(name)?; - let kind = self.determine_installed_kind(name).await?; + let kind = self.determine_installed_kind(name, user_id).await?; match kind { ExtensionKind::WasmChannel => { let cap_path = self @@ -4449,7 +4533,7 @@ impl ExtensionManager { for secret in &cap_file.setup.required_secrets { let provided = self .secrets - .exists(&self.user_id, &secret.name) + .exists(user_id, &secret.name) .await .unwrap_or(false); secrets.push(crate::channels::web::types::SecretFieldInfo { @@ -4486,7 +4570,7 @@ impl ExtensionManager { } let provided = self .secrets - .exists(&self.user_id, &secret.name) + .exists(user_id, &secret.name) .await .unwrap_or(false); secrets.push(crate::channels::web::types::SecretFieldInfo { @@ -4849,9 +4933,10 @@ impl ExtensionManager { name: &str, secrets: &std::collections::HashMap, fields: &std::collections::HashMap, + user_id: &str, ) -> Result { Self::validate_extension_name(name)?; - let kind = self.determine_installed_kind(name).await?; + let kind = self.determine_installed_kind(name, user_id).await?; // Load allowed secret names and tool setup field definitions from capabilities. let mut channel_cap_file: Option = None; @@ -4907,7 +4992,7 @@ impl ExtensionManager { } ExtensionKind::McpServer => { let server = self - .get_mcp_server(name) + .get_mcp_server(name, user_id) .await .map_err(|e| ExtensionError::NotInstalled(e.to_string()))?; let mut names = std::collections::HashSet::new(); @@ -4993,7 +5078,7 @@ impl ExtensionManager { let params = CreateSecretParams::new(secret_name, trimmed_value).with_provider(name.to_string()); self.secrets - .create(&self.user_id, params) + .create(user_id, params) .await .map_err(|e| ExtensionError::AuthFailed(e.to_string()))?; } @@ -5071,7 +5156,7 @@ impl ExtensionManager { .is_some_and(|v| !v.trim().is_empty()); let already_stored = self .secrets - .exists(&self.user_id, &secret_def.name) + .exists(user_id, &secret_def.name) .await .unwrap_or(false); if !already_provided && !already_stored { @@ -5083,7 +5168,7 @@ impl ExtensionManager { let params = CreateSecretParams::new(&secret_def.name, &hex_value) .with_provider(name.to_string()); self.secrets - .create(&self.user_id, params) + .create(user_id, params) .await .map_err(|e| ExtensionError::AuthFailed(e.to_string()))?; tracing::info!( @@ -5119,7 +5204,7 @@ impl ExtensionManager { // For tools, save and attempt auto-activation, then check auth. if kind == ExtensionKind::WasmTool { - match self.activate_wasm_tool(name).await { + match self.activate_wasm_tool(name, user_id).await { Ok(result) => { // Delete existing OAuth token so auth() starts a fresh flow. // Done AFTER activation succeeds to avoid losing tokens on failure. @@ -5128,20 +5213,14 @@ impl ExtensionManager { && let Some(ref auth_cfg) = cap.auth && auth_cfg.oauth.is_some() { + let _ = self.secrets.delete(user_id, &auth_cfg.secret_name).await; let _ = self .secrets - .delete(&self.user_id, &auth_cfg.secret_name) + .delete(user_id, &format!("{}_scopes", auth_cfg.secret_name)) .await; let _ = self .secrets - .delete(&self.user_id, &format!("{}_scopes", auth_cfg.secret_name)) - .await; - let _ = self - .secrets - .delete( - &self.user_id, - &format!("{}_refresh_token", auth_cfg.secret_name), - ) + .delete(user_id, &format!("{}_refresh_token", auth_cfg.secret_name)) .await; } @@ -5150,7 +5229,7 @@ impl ExtensionManager { let mut auth_url = None; // Box::pin breaks the async recursion cycle: // auth() → auth_wasm_tool() → (OAuth) → configure() → auth() - if let Ok(auth_result) = Box::pin(self.auth(name)).await { + if let Ok(auth_result) = Box::pin(self.auth(name, user_id)).await { auth_url = auth_result.auth_url().map(String::from); } let message = if auth_url.is_some() { @@ -5192,9 +5271,9 @@ impl ExtensionManager { // Activate the extension now that secrets are saved. // Dispatch by kind — WasmTool was already handled above with an early return. let activate_result = match kind { - ExtensionKind::WasmChannel => self.activate_wasm_channel(name).await, - ExtensionKind::McpServer => self.activate_mcp(name).await, - ExtensionKind::ChannelRelay => self.activate_channel_relay(name).await, + ExtensionKind::WasmChannel => self.activate_wasm_channel(name, user_id).await, + ExtensionKind::McpServer => self.activate_mcp(name, user_id).await, + ExtensionKind::ChannelRelay => self.activate_channel_relay(name, user_id).await, ExtensionKind::WasmTool => { return Ok(ConfigureResult { message: format!("Configuration saved for '{}'.", name), @@ -5269,8 +5348,9 @@ impl ExtensionManager { &self, name: &str, token: &str, + user_id: &str, ) -> Result { - let kind = self.determine_installed_kind(name).await?; + let kind = self.determine_installed_kind(name, user_id).await?; let secret_name = match kind { ExtensionKind::WasmChannel => { let cap_path = self @@ -5289,12 +5369,7 @@ impl ExtensionManager { if s.optional { continue; } - if !self - .secrets - .exists(&self.user_id, &s.name) - .await - .unwrap_or(false) - { + if !self.secrets.exists(user_id, &s.name).await.unwrap_or(false) { target = Some(s.name.clone()); break; } @@ -5321,7 +5396,7 @@ impl ExtensionManager { if let Some(ref auth) = cap.auth { if !self .secrets - .exists(&self.user_id, &auth.secret_name) + .exists(user_id, &auth.secret_name) .await .unwrap_or(false) { @@ -5330,12 +5405,7 @@ impl ExtensionManager { // Auth secret exists, find first missing setup secret let mut found = None; for s in &setup.required_secrets { - if !self - .secrets - .exists(&self.user_id, &s.name) - .await - .unwrap_or(false) - { + if !self.secrets.exists(user_id, &s.name).await.unwrap_or(false) { found = Some(s.name.clone()); break; } @@ -5359,7 +5429,7 @@ impl ExtensionManager { } ExtensionKind::McpServer => { let server = self - .get_mcp_server(name) + .get_mcp_server(name, user_id) .await .map_err(|e| ExtensionError::NotInstalled(e.to_string()))?; server.token_secret_name() @@ -5369,7 +5439,7 @@ impl ExtensionManager { let mut secrets = std::collections::HashMap::new(); secrets.insert(secret_name, token.to_string()); - self.configure(name, &secrets, &std::collections::HashMap::new()) + self.configure(name, &secrets, &std::collections::HashMap::new(), user_id) .await } @@ -5930,8 +6000,8 @@ mod tests { wasm_runtime, tools_dir, channels_dir, - None, // tunnel_url - "test".to_string(), + None, // tunnel_url + "test".to_string(), // user_id store, vec![], ) @@ -6049,7 +6119,12 @@ mod tests { fields.insert("llm_backend".to_string(), "openai".to_string()); let result = mgr - .configure("switch-llm", &std::collections::HashMap::new(), &fields) + .configure( + "switch-llm", + &std::collections::HashMap::new(), + &fields, + "test-user", + ) .await .expect("save configuration"); @@ -6097,7 +6172,12 @@ mod tests { fields.insert("session".to_string(), "overwrite".to_string()); let err = match mgr - .configure("evil-tool", &std::collections::HashMap::new(), &fields) + .configure( + "evil-tool", + &std::collections::HashMap::new(), + &fields, + "test-user", + ) .await { Ok(_) => panic!("disallowed setting_path should fail"), @@ -6128,7 +6208,7 @@ mod tests { let runtime = Arc::new(crate::tools::wasm::WasmToolRuntime::new(config).expect("runtime")); let mgr = make_test_manager(Some(runtime), dir.path().to_path_buf()); - let err = mgr.activate("nonexistent").await.unwrap_err(); + let err = mgr.activate("nonexistent", "test").await.unwrap_err(); let msg = err.to_string(); assert!( !msg.contains("WASM runtime not available"), @@ -6152,7 +6232,7 @@ mod tests { let mgr = make_test_manager(None, dir.path().to_path_buf()); - let err = mgr.activate("fake").await.unwrap_err(); + let err = mgr.activate("fake", "test").await.unwrap_err(); let msg = err.to_string(); assert!( msg.contains("WASM runtime not available"), @@ -6187,7 +6267,7 @@ mod tests { #[tokio::test] async fn test_upgrade_no_installed_extensions() { let manager = make_manager_with_temp_dirs(); - let result = manager.upgrade(None).await.unwrap(); + let result = manager.upgrade(None, "test").await.unwrap(); assert!(result.results.is_empty()); assert!(result.message.contains("No WASM extensions installed")); } @@ -6196,7 +6276,7 @@ mod tests { async fn test_upgrade_mcp_server_rejected() { let manager = make_manager_with_temp_dirs(); // MCP servers can't be upgraded via tool_upgrade - let err = manager.upgrade(Some("some-mcp")).await; + let err = manager.upgrade(Some("some-mcp"), "test").await; // It will fail with NotInstalled because there's no MCP server named "some-mcp", // but if it were installed, the MCP code path would be rejected. assert!(err.is_err()); @@ -6222,7 +6302,7 @@ mod tests { let manager = make_manager_custom_dirs(dir.path().join("tools"), channels_dir); - let result = manager.upgrade(Some("test-channel")).await.unwrap(); + let result = manager.upgrade(Some("test-channel"), "test").await.unwrap(); assert_eq!(result.results.len(), 1); assert_eq!(result.results[0].status, "already_up_to_date"); } @@ -6247,7 +6327,10 @@ mod tests { let manager = make_manager_custom_dirs(dir.path().join("tools"), channels_dir); - let result = manager.upgrade(Some("custom-channel")).await.unwrap(); + let result = manager + .upgrade(Some("custom-channel"), "test") + .await + .unwrap(); assert_eq!(result.results.len(), 1); assert_eq!(result.results[0].status, "not_in_registry"); } @@ -6502,6 +6585,7 @@ mod tests { "123456789:ABCdefGhI".to_string(), )]), &std::collections::HashMap::new(), + "test", ) .await .map_err(|err| format!("configure succeeds: {err}"))?; @@ -6532,7 +6616,7 @@ mod tests { "telegram should be hot-added to the running channel manager", )?; require_eq( - manager.load_persisted_active_channels().await, + manager.load_persisted_active_channels("test").await, vec!["telegram".to_string()], "persisted active channels", )?; @@ -6630,6 +6714,7 @@ mod tests { "123456789:ABCdefGhI".to_string(), )]), &std::collections::HashMap::new(), + "test", ) .await .map_err(|err| format!("configure returned challenge: {err}"))?; @@ -6943,7 +7028,7 @@ mod tests { ); // Calling determine_installed_kind for a non-installed name returns NotInstalled - let result = mgr.determine_installed_kind("slack-relay").await; + let result = mgr.determine_installed_kind("slack-relay", "test").await; assert!(result.is_err(), "Should return NotInstalled"); // Crucially: installed_relay_extensions must still be empty @@ -6958,8 +7043,8 @@ mod tests { let dir = tempfile::tempdir().expect("temp dir"); let mgr = make_test_manager(None, dir.path().to_path_buf()); - // With no DB store, is_relay_channel always returns false - assert!(!mgr.is_relay_channel("slack-relay").await); + // No token stored → not a relay channel + assert!(!mgr.is_relay_channel("slack-relay", "test").await); } #[tokio::test] @@ -6967,7 +7052,10 @@ mod tests { let dir = tempfile::tempdir().expect("temp dir"); let mgr = make_test_manager(None, dir.path().to_path_buf()); - let err = mgr.activate_channel_relay("slack-relay").await.unwrap_err(); + let err = mgr + .activate_channel_relay("slack-relay", "test") + .await + .unwrap_err(); assert!( matches!(err, ExtensionError::AuthRequired), "expected AuthRequired, got: {err:?}" @@ -7011,7 +7099,7 @@ mod tests { assert!(cm.get_channel("slack-relay").await.is_some()); // Remove should succeed and shut down the channel - let result = mgr.remove("slack-relay").await; + let result = mgr.remove("slack-relay", "test").await; assert!(result.is_ok(), "remove should succeed: {:?}", result.err()); // installed_relay_extensions should be cleared @@ -7080,7 +7168,7 @@ mod tests { scopes: vec![], user_id: "test".to_string(), secrets: Arc::clone(&secrets), - sse_sender: None, + sse_manager: None, gateway_token: None, token_exchange_extra_params: std::collections::HashMap::new(), client_id_secret_name: None, @@ -7104,7 +7192,7 @@ mod tests { scopes: vec![], user_id: "test".to_string(), secrets, - sse_sender: None, + sse_manager: None, gateway_token: None, token_exchange_extra_params: std::collections::HashMap::new(), client_id_secret_name: None, @@ -7112,7 +7200,7 @@ mod tests { }, ); - let result = mgr.remove("gmail").await; + let result = mgr.remove("gmail", "test").await; assert!(result.is_ok(), "remove should succeed: {:?}", result.err()); tokio::task::yield_now().await; @@ -7158,7 +7246,7 @@ mod tests { .await .insert("telegram".to_string(), "channel failed".to_string()); - let result = mgr.remove("telegram").await; + let result = mgr.remove("telegram", "test").await; assert!(result.is_ok(), "remove should succeed: {:?}", result.err()); assert!( @@ -7568,7 +7656,7 @@ mod tests { .expect("store SECRET_A"); // configure_token should target SECRET_B (the first missing one) - let _result = mgr.configure_token("multi", "value-b").await; + let _result = mgr.configure_token("multi", "value-b", "test").await; // configure will fail at activation (no real WASM runtime), but the // secret should still have been stored before activation was attempted. // Check that SECRET_B was stored. @@ -7608,7 +7696,7 @@ mod tests { let mgr = make_manager_custom_dirs(dir.path().join("tools"), channels_dir); // auth() should return a result without storing anything - let result = mgr.auth("test-ch").await; + let result = mgr.auth("test-ch", "test").await; assert!(result.is_ok(), "auth should succeed: {:?}", result.err()); // No secrets should have been created @@ -7651,7 +7739,7 @@ mod tests { let mgr = make_manager_custom_dirs(dir.path().join("tools"), channels_dir); let result = mgr - .auth("telegram") + .auth("telegram", "test") .await .map_err(|err| format!("telegram auth status: {err}"))?; let instructions = result @@ -7784,7 +7872,12 @@ mod tests { ); let result = mgr - .configure("test-relay", &secrets, &std::collections::HashMap::new()) + .configure( + "test-relay", + &secrets, + &std::collections::HashMap::new(), + "test", + ) .await; assert!( result.is_ok(), diff --git a/src/history/store.rs b/src/history/store.rs index f0b593c2..d6570b3c 100644 --- a/src/history/store.rs +++ b/src/history/store.rs @@ -842,6 +842,38 @@ impl Store { .collect()) } + pub async fn list_agent_jobs_for_user( + &self, + user_id: &str, + ) -> Result, DatabaseError> { + let conn = self.conn().await?; + let rows = conn + .query( + r#" + SELECT id, title, status, user_id, failure_reason, + created_at, started_at, completed_at + FROM agent_jobs WHERE source = 'direct' AND user_id = $1 + ORDER BY created_at DESC + "#, + &[&user_id], + ) + .await?; + + Ok(rows + .iter() + .map(|r| AgentJobRecord { + id: r.get("id"), + title: r.get("title"), + status: r.get("status"), + user_id: r.get::<_, Option>("user_id").unwrap_or_default(), + created_at: r.get("created_at"), + started_at: r.get("started_at"), + completed_at: r.get("completed_at"), + failure_reason: r.get("failure_reason"), + }) + .collect()) + } + /// Get the failure reason for a single agent job. pub async fn get_agent_job_failure_reason( &self, @@ -875,6 +907,27 @@ impl Store { } Ok(summary) } + + pub async fn agent_job_summary_for_user( + &self, + user_id: &str, + ) -> Result { + let conn = self.conn().await?; + let rows = conn + .query( + "SELECT status, COUNT(*) as cnt FROM agent_jobs WHERE source = 'direct' AND user_id = $1 GROUP BY status", + &[&user_id], + ) + .await?; + + let mut summary = AgentJobSummary::default(); + for row in &rows { + let status: String = row.get("status"); + let count: i64 = row.get("cnt"); + summary.add_count(&status, count as usize); + } + Ok(summary) + } } // ==================== Job Events ==================== diff --git a/src/main.rs b/src/main.rs index 2cf8fd53..dd224f47 100644 --- a/src/main.rs +++ b/src/main.rs @@ -589,15 +589,46 @@ async fn async_main() -> anyhow::Result<()> { // ── Gateway channel ──────────────────────────────────────────────── let mut gateway_url: Option = None; - let mut sse_sender: Option< - tokio::sync::broadcast::Sender, - > = None; + let mut sse_manager: Option> = None; if let Some(ref gw_config) = config.channels.gateway { - let mut gw = - GatewayChannel::new(gw_config.clone()).with_llm_provider(Arc::clone(&components.llm)); + // Build multi-user auth state if user_tokens is configured, else single-user. + let mut gw = if let Some(ref user_tokens) = gw_config.user_tokens { + use ironclaw::channels::web::auth::{MultiAuthState, UserIdentity}; + let tokens = user_tokens + .iter() + .map(|(token, cfg)| { + ( + token.clone(), + UserIdentity { + user_id: cfg.user_id.clone(), + workspace_read_scopes: cfg.workspace_read_scopes.clone(), + }, + ) + }) + .collect(); + let auth = MultiAuthState::multi(tokens); + GatewayChannel::new_multi_auth(gw_config.clone(), auth) + } else { + GatewayChannel::new(gw_config.clone()) + }; + gw = gw.with_llm_provider(Arc::clone(&components.llm)); if let Some(ref ws) = components.workspace { gw = gw.with_workspace(Arc::clone(ws)); } + // Create per-user workspace pool for multi-user mode. + if let Some(ref db) = components.db { + let emb_cache_config = ironclaw::workspace::EmbeddingCacheConfig { + max_entries: config.embeddings.cache_size, + }; + let pool = Arc::new(ironclaw::channels::web::server::WorkspacePool::new( + Arc::clone(db), + components.embeddings.clone(), + emb_cache_config, + config.search.clone(), + config.workspace.clone(), + )); + gw = gw.with_workspace_pool(pool); + } gw = gw.with_session_manager(Arc::clone(&session_manager)); gw = gw.with_log_broadcaster(Arc::clone(&log_broadcaster)); gw = gw.with_log_level_handle(Arc::clone(&log_level_handle)); @@ -648,8 +679,12 @@ async fn async_main() -> anyhow::Result<()> { let mut rx = tx.subscribe(); let gw_state = Arc::clone(gw.state()); tokio::spawn(async move { - while let Ok((_job_id, event)) = rx.recv().await { - gw_state.sse.broadcast(event); + while let Ok((_job_id, user_id, event)) = rx.recv().await { + if user_id.is_empty() { + gw_state.sse.broadcast(event); + } else { + gw_state.sse.broadcast_for_user(&user_id, event); + } } }); } @@ -691,7 +726,7 @@ async fn async_main() -> anyhow::Result<()> { // Capture SSE sender and routine engine slot before moving gw into channels. // IMPORTANT: This must come after all `with_*` calls since `rebuild_state` // creates a new SseManager, which would orphan this sender. - sse_sender = Some(gw.state().sse.sender()); + sse_manager = Some(Arc::clone(&gw.state().sse)); channel_names.push("gateway".to_string()); channels.add(Box::new(gw)).await; } @@ -754,6 +789,14 @@ async fn async_main() -> anyhow::Result<()> { .register_message_tools(Arc::clone(&channels), components.extension_manager.clone()) .await; + // Default user ID for extension operations (single-user mode). + let ext_user_id = config + .channels + .gateway + .as_ref() + .map(|g| g.user_id.clone()) + .unwrap_or_else(|| "default".to_string()); + // Wire up channel runtime for hot-activation of WASM channels. if let Some(ref ext_mgr) = components.extension_manager && let Some((rt, ps, router)) = wasm_channel_runtime_state.take() @@ -774,12 +817,14 @@ async fn async_main() -> anyhow::Result<()> { // Auto-activate WASM channels that were active in a previous session. // Relay channels are handled separately below via restore_relay_channels(). - let persisted = ext_mgr.load_persisted_active_channels().await; + let persisted = ext_mgr.load_persisted_active_channels(&ext_user_id).await; for name in &persisted { - if active_at_startup.contains(name) || ext_mgr.is_relay_channel(name).await { + if active_at_startup.contains(name) + || ext_mgr.is_relay_channel(name, &ext_user_id).await + { continue; } - match ext_mgr.activate(name).await { + match ext_mgr.activate(name, &ext_user_id).await { Ok(result) => { tracing::debug!( channel = %name, @@ -804,14 +849,14 @@ async fn async_main() -> anyhow::Result<()> { ext_mgr .set_relay_channel_manager(Arc::clone(&channels)) .await; - ext_mgr.restore_relay_channels().await; + ext_mgr.restore_relay_channels(&ext_user_id).await; } // Wire SSE sender into extension manager for broadcasting status events. if let Some(ref ext_mgr) = components.extension_manager - && let Some(ref sender) = sse_sender + && let Some(ref sse) = sse_manager { - ext_mgr.set_sse_sender(sender.clone()).await; + ext_mgr.set_sse_sender(Arc::clone(sse)).await; } // Snapshot memory for trace recording before the agent starts @@ -849,7 +894,7 @@ async fn async_main() -> anyhow::Result<()> { skills_config: config.skills.clone(), hooks: components.hooks, cost_guard: components.cost_guard, - sse_tx: sse_sender, + sse_tx: sse_manager, http_interceptor, transcription: config.transcription.create_provider().map(|p| { Arc::new(ironclaw::llm::transcription::TranscriptionMiddleware::new( diff --git a/src/orchestrator/api.rs b/src/orchestrator/api.rs index 8d77c581..00f8a4da 100644 --- a/src/orchestrator/api.rs +++ b/src/orchestrator/api.rs @@ -40,7 +40,8 @@ pub struct OrchestratorState { pub job_manager: Arc, pub token_store: TokenStore, /// Broadcast channel for job events (consumed by the web gateway SSE). - pub job_event_tx: Option>, + /// Tuple: (job_id, user_id, event). + 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. @@ -49,6 +50,9 @@ pub struct OrchestratorState { pub secrets_store: Option>, /// User ID for secret lookups (single-tenant, typically "default"). pub user_id: String, + /// In-memory cache of job_id → user_id for SSE scoping. Populated when + /// sandbox jobs are created, avoiding a DB round-trip on every job event. + pub job_owner_cache: Arc>>, } /// The orchestrator's internal API server. @@ -351,9 +355,45 @@ async fn job_event_handler( }, }; - // Broadcast via the channel (if configured) + // Broadcast via the channel (if configured). + // Look up the job owner from the in-memory cache (populated at job creation). if let Some(ref tx) = state.job_event_tx { - let _ = tx.send((job_id, sse_event)); + let cached_uid = state + .job_owner_cache + .read() + .unwrap_or_else(|e| e.into_inner()) + .get(&job_id) + .cloned(); + + let user_id = match cached_uid { + Some(uid) => uid, + None => { + // Cache miss: fall back to DB lookup and populate cache. + let uid = match state.store.as_ref() { + Some(store) => store + .get_sandbox_job(job_id) + .await + .ok() + .flatten() + .map(|j| j.user_id), + None => None, + }; + if let Some(ref uid) = uid { + state + .job_owner_cache + .write() + .unwrap_or_else(|e| e.into_inner()) + .insert(job_id, uid.clone()); + } + uid.unwrap_or_default() + } + }; + + if user_id.is_empty() { + let _ = tx.send((job_id, String::new(), sse_event)); + } else { + let _ = tx.send((job_id, user_id, sse_event)); + } } Ok(StatusCode::OK) @@ -480,6 +520,7 @@ mod tests { store: None, secrets_store: None, user_id: "default".to_string(), + job_owner_cache: Arc::new(std::sync::RwLock::new(HashMap::new())), } } @@ -709,6 +750,7 @@ mod tests { store: None, secrets_store: Some(secrets_store), user_id: "default".to_string(), + job_owner_cache: Arc::new(std::sync::RwLock::new(HashMap::new())), }; let router = OrchestratorApi::router(state); @@ -744,6 +786,7 @@ mod tests { store: None, secrets_store: None, user_id: "default".to_string(), + job_owner_cache: Arc::new(std::sync::RwLock::new(HashMap::new())), }; let job_id = Uuid::new_v4(); @@ -769,8 +812,10 @@ mod tests { let resp = router.oneshot(req).await.unwrap(); assert_eq!(resp.status(), StatusCode::OK); - let (recv_id, event) = rx.recv().await.unwrap(); + let (recv_id, recv_uid, event) = rx.recv().await.unwrap(); assert_eq!(recv_id, job_id); + // No store configured, so user_id falls back to empty string. + assert_eq!(recv_uid, ""); match event { SseEvent::JobMessage { job_id: jid, @@ -799,6 +844,7 @@ mod tests { store: None, secrets_store: None, user_id: "default".to_string(), + job_owner_cache: Arc::new(std::sync::RwLock::new(HashMap::new())), }; let job_id = Uuid::new_v4(); @@ -824,7 +870,7 @@ mod tests { let resp = router.oneshot(req).await.unwrap(); assert_eq!(resp.status(), StatusCode::OK); - let (_recv_id, event) = rx.recv().await.unwrap(); + let (_recv_id, _recv_uid, event) = rx.recv().await.unwrap(); match event { SseEvent::JobToolUse { tool_name, .. } => { assert_eq!(tool_name, "shell"); @@ -847,6 +893,7 @@ mod tests { store: None, secrets_store: None, user_id: "default".to_string(), + job_owner_cache: Arc::new(std::sync::RwLock::new(HashMap::new())), }; let job_id = Uuid::new_v4(); @@ -869,7 +916,7 @@ mod tests { let resp = router.oneshot(req).await.unwrap(); assert_eq!(resp.status(), StatusCode::OK); - let (_recv_id, event) = rx.recv().await.unwrap(); + let (_recv_id, _recv_uid, event) = rx.recv().await.unwrap(); // Unknown event types fall through to JobStatus assert!(matches!(event, SseEvent::JobStatus { .. })); } diff --git a/src/orchestrator/mod.rs b/src/orchestrator/mod.rs index d6e028a5..896b5648 100644 --- a/src/orchestrator/mod.rs +++ b/src/orchestrator/mod.rs @@ -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, } @@ -134,6 +134,7 @@ pub async fn setup_orchestrator( store: db.cloned(), secrets_store: secrets_store.cloned(), user_id: "default".to_string(), + job_owner_cache: Arc::new(std::sync::RwLock::new(std::collections::HashMap::new())), }; tokio::spawn(async move { diff --git a/src/tools/builtin/extension_tools.rs b/src/tools/builtin/extension_tools.rs index cb0f71dd..fba61613 100644 --- a/src/tools/builtin/extension_tools.rs +++ b/src/tools/builtin/extension_tools.rs @@ -130,7 +130,7 @@ impl Tool for ToolInstallTool { async fn execute( &self, params: serde_json::Value, - _ctx: &JobContext, + ctx: &JobContext, ) -> Result { let start = std::time::Instant::now(); @@ -150,7 +150,7 @@ impl Tool for ToolInstallTool { let result = self .manager - .install(name, url, kind_hint) + .install(name, url, kind_hint, &ctx.user_id) .await .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?; @@ -205,7 +205,7 @@ impl Tool for ToolAuthTool { async fn execute( &self, params: serde_json::Value, - _ctx: &JobContext, + ctx: &JobContext, ) -> Result { let start = std::time::Instant::now(); @@ -213,13 +213,13 @@ impl Tool for ToolAuthTool { let result = self .manager - .auth(name) + .auth(name, &ctx.user_id) .await .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?; // Auto-activate after successful auth so tools are available immediately if result.is_authenticated() { - match self.manager.activate(name).await { + match self.manager.activate(name, &ctx.user_id).await { Ok(activate_result) => { let output = serde_json::json!({ "status": "authenticated_and_activated", @@ -304,13 +304,13 @@ impl Tool for ToolActivateTool { async fn execute( &self, params: serde_json::Value, - _ctx: &JobContext, + ctx: &JobContext, ) -> Result { let start = std::time::Instant::now(); let name = require_str(¶ms, "name")?; - match self.manager.activate(name).await { + match self.manager.activate(name, &ctx.user_id).await { Ok(result) => { let output = serde_json::to_value(&result) .unwrap_or_else(|_| serde_json::json!({"error": "serialization failed"})); @@ -329,12 +329,12 @@ impl Tool for ToolActivateTool { // Activation failed due to missing auth; initiate auth flow // so the agent loop can show the auth card. - match self.manager.auth(name).await { + match self.manager.auth(name, &ctx.user_id).await { Ok(auth_result) if auth_result.is_authenticated() => { // Auth succeeded (e.g. env var was set); retry activation. let result = self .manager - .activate(name) + .activate(name, &ctx.user_id) .await .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?; let output = serde_json::to_value(&result).unwrap_or_else( @@ -404,7 +404,7 @@ impl Tool for ToolListTool { async fn execute( &self, params: serde_json::Value, - _ctx: &JobContext, + ctx: &JobContext, ) -> Result { let start = std::time::Instant::now(); @@ -425,7 +425,7 @@ impl Tool for ToolListTool { let extensions = self .manager - .list(kind_filter, include_available) + .list(kind_filter, include_available, &ctx.user_id) .await .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?; @@ -477,7 +477,7 @@ impl Tool for ToolRemoveTool { async fn execute( &self, params: serde_json::Value, - _ctx: &JobContext, + ctx: &JobContext, ) -> Result { let start = std::time::Instant::now(); @@ -485,7 +485,7 @@ impl Tool for ToolRemoveTool { let message = self .manager - .remove(name) + .remove(name, &ctx.user_id) .await .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?; @@ -541,7 +541,7 @@ impl Tool for ToolUpgradeTool { async fn execute( &self, params: serde_json::Value, - _ctx: &JobContext, + ctx: &JobContext, ) -> Result { let start = std::time::Instant::now(); @@ -549,7 +549,7 @@ impl Tool for ToolUpgradeTool { let result = self .manager - .upgrade(name) + .upgrade(name, &ctx.user_id) .await .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?; @@ -603,7 +603,7 @@ impl Tool for ExtensionInfoTool { async fn execute( &self, params: serde_json::Value, - _ctx: &JobContext, + ctx: &JobContext, ) -> Result { let start = std::time::Instant::now(); @@ -611,7 +611,7 @@ impl Tool for ExtensionInfoTool { let info = self .manager - .extension_info(name) + .extension_info(name, &ctx.user_id) .await .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?; diff --git a/src/tools/builtin/job.rs b/src/tools/builtin/job.rs index 0933ee40..86d7e44d 100644 --- a/src/tools/builtin/job.rs +++ b/src/tools/builtin/job.rs @@ -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, SseEvent)>, + event_tx: tokio::sync::broadcast::Sender<(Uuid, String, SseEvent)>, inject_tx: tokio::sync::mpsc::Sender, ) -> Self { self.event_tx = Some(event_tx); diff --git a/src/tools/builtin/memory.rs b/src/tools/builtin/memory.rs index edbc4f1c..501ccf46 100644 --- a/src/tools/builtin/memory.rs +++ b/src/tools/builtin/memory.rs @@ -21,6 +21,35 @@ use crate::context::JobContext; use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str}; use crate::workspace::{Workspace, paths}; +// ── WorkspaceResolver ────────────────────────────────────────────── + +/// Resolves a workspace for a given user ID. +/// +/// In single-user mode, always returns the same workspace. +/// In multi-tenant mode, creates per-user workspaces on demand. +#[async_trait] +pub trait WorkspaceResolver: Send + Sync { + async fn resolve(&self, user_id: &str) -> Arc; +} + +/// Returns a fixed workspace regardless of user ID (single-user mode). +pub struct FixedWorkspaceResolver { + workspace: Arc, +} + +impl FixedWorkspaceResolver { + pub fn new(workspace: Arc) -> Self { + Self { workspace } + } +} + +#[async_trait] +impl WorkspaceResolver for FixedWorkspaceResolver { + async fn resolve(&self, _user_id: &str) -> Arc { + Arc::clone(&self.workspace) + } +} + /// Detect paths that are clearly local filesystem references, not workspace-memory docs. /// /// Examples: @@ -62,13 +91,20 @@ fn map_write_err(e: crate::error::WorkspaceError) -> ToolError { /// The agent should call this tool before answering questions about /// prior work, decisions, preferences, or any historical context. pub struct MemorySearchTool { - workspace: Arc, + resolver: Arc, } impl MemorySearchTool { - /// Create a new memory search tool. - pub fn new(workspace: Arc) -> Self { - Self { workspace } + /// Create a new memory search tool with a workspace resolver. + pub fn new(resolver: Arc) -> Self { + Self { resolver } + } + + /// Create from a fixed workspace (backward compatibility). + pub fn from_workspace(workspace: Arc) -> Self { + Self { + resolver: Arc::new(FixedWorkspaceResolver::new(workspace)), + } } } @@ -107,7 +143,7 @@ impl Tool for MemorySearchTool { async fn execute( &self, params: serde_json::Value, - _ctx: &JobContext, + ctx: &JobContext, ) -> Result { let start = std::time::Instant::now(); @@ -119,8 +155,8 @@ impl Tool for MemorySearchTool { .unwrap_or(5) .min(20) as usize; - let results = self - .workspace + let workspace = self.resolver.resolve(&ctx.user_id).await; + let results = workspace .search(query, limit) .await .map_err(|e| ToolError::ExecutionFailed(format!("Search failed: {}", e)))?; @@ -151,13 +187,20 @@ impl Tool for MemorySearchTool { /// Use this to persist important information that should be remembered /// across sessions: decisions, preferences, facts, lessons learned. pub struct MemoryWriteTool { - workspace: Arc, + resolver: Arc, } impl MemoryWriteTool { - /// Create a new memory write tool. - pub fn new(workspace: Arc) -> Self { - Self { workspace } + /// Create a new memory write tool with a workspace resolver. + pub fn new(resolver: Arc) -> Self { + Self { resolver } + } + + /// Create from a fixed workspace (backward compatibility). + pub fn from_workspace(workspace: Arc) -> Self { + Self { + resolver: Arc::new(FixedWorkspaceResolver::new(workspace)), + } } } @@ -231,19 +274,21 @@ impl Tool for MemoryWriteTool { ))); } + let workspace = self.resolver.resolve(&ctx.user_id).await; + // Bootstrap target: clear BOOTSTRAP.md to mark first-run ritual complete. // Handled early because it accepts empty content (unlike other targets). if target == "bootstrap" { // Write empty content to effectively disable the bootstrap injection. // system_prompt_for_context() skips empty files. - self.workspace + workspace .write(paths::BOOTSTRAP, "") .await .map_err(map_write_err)?; // Also set the in-memory flag so BOOTSTRAP.md injection stops // immediately without waiting for a restart. - self.workspace.mark_bootstrap_completed(); + workspace.mark_bootstrap_completed(); let output = serde_json::json!({ "status": "cleared", @@ -289,12 +334,12 @@ impl Tool for MemoryWriteTool { // Otherwise, use default workspace methods (which include injection scanning). let layer_result = if let Some(layer_name) = layer { let result = if append { - self.workspace + workspace .append_to_layer(layer_name, &resolved_path, content, force) .await .map_err(map_write_err)? } else { - self.workspace + workspace .write_to_layer(layer_name, &resolved_path, content, force) .await .map_err(map_write_err)? @@ -307,31 +352,33 @@ impl Tool for MemoryWriteTool { match target { "memory" => { if append { - self.workspace + workspace .append_memory(content) .await .map_err(map_write_err)?; } else { - self.workspace + workspace .write(paths::MEMORY, content) .await .map_err(map_write_err)?; } } "daily_log" => { - self.workspace + let tz = crate::timezone::parse_timezone(&ctx.user_timezone) + .unwrap_or(chrono_tz::Tz::UTC); + workspace .append_daily_log_tz(content, tz) .await .map_err(map_write_err)?; } _ => { if append { - self.workspace + workspace .append(&resolved_path, content) .await .map_err(map_write_err)?; } else { - self.workspace + workspace .write(&resolved_path, content) .await .map_err(map_write_err)?; @@ -361,12 +408,12 @@ impl Tool for MemoryWriteTool { }; let mut synced_docs: Vec<&str> = Vec::new(); if normalized_path == paths::PROFILE { - match self.workspace.sync_profile_documents().await { + match workspace.sync_profile_documents().await { Ok(true) => { tracing::info!("profile write: synced USER.md + assistant-directives.md"); synced_docs.extend_from_slice(&[paths::USER, paths::ASSISTANT_DIRECTIVES]); - self.workspace.mark_bootstrap_completed(); + workspace.mark_bootstrap_completed(); let toml_path = crate::settings::Settings::default_toml_path(); if let Ok(Some(mut settings)) = crate::settings::Settings::load_toml(&toml_path) && !settings.profile_onboarding_completed @@ -416,13 +463,20 @@ impl Tool for MemoryWriteTool { /// /// Use this to read the full content of any file in the workspace. pub struct MemoryReadTool { - workspace: Arc, + resolver: Arc, } impl MemoryReadTool { - /// Create a new memory read tool. - pub fn new(workspace: Arc) -> Self { - Self { workspace } + /// Create a new memory read tool with a workspace resolver. + pub fn new(resolver: Arc) -> Self { + Self { resolver } + } + + /// Create from a fixed workspace (backward compatibility). + pub fn from_workspace(workspace: Arc) -> Self { + Self { + resolver: Arc::new(FixedWorkspaceResolver::new(workspace)), + } } } @@ -456,7 +510,7 @@ impl Tool for MemoryReadTool { async fn execute( &self, params: serde_json::Value, - _ctx: &JobContext, + ctx: &JobContext, ) -> Result { let start = std::time::Instant::now(); @@ -470,8 +524,8 @@ impl Tool for MemoryReadTool { ))); } - let doc = self - .workspace + let workspace = self.resolver.resolve(&ctx.user_id).await; + let doc = workspace .read(path) .await .map_err(|e| ToolError::ExecutionFailed(format!("Read failed: {}", e)))?; @@ -495,20 +549,27 @@ impl Tool for MemoryReadTool { /// /// Returns a hierarchical view of files and directories with configurable depth. pub struct MemoryTreeTool { - workspace: Arc, + resolver: Arc, } impl MemoryTreeTool { - /// Create a new memory tree tool. - pub fn new(workspace: Arc) -> Self { - Self { workspace } + /// Create a new memory tree tool with a workspace resolver. + pub fn new(resolver: Arc) -> Self { + Self { resolver } + } + + /// Create from a fixed workspace (backward compatibility). + pub fn from_workspace(workspace: Arc) -> Self { + Self { + resolver: Arc::new(FixedWorkspaceResolver::new(workspace)), + } } /// Recursively build tree structure. /// /// Returns a compact format where directories end with `/` and may have children. async fn build_tree( - &self, + workspace: &Arc, path: &str, current_depth: usize, max_depth: usize, @@ -517,8 +578,7 @@ impl MemoryTreeTool { return Ok(Vec::new()); } - let entries = self - .workspace + let entries = workspace .list(path) .await .map_err(|e| ToolError::ExecutionFailed(format!("Tree failed: {}", e)))?; @@ -533,8 +593,13 @@ impl MemoryTreeTool { }; if entry.is_directory && current_depth < max_depth { - let children = - Box::pin(self.build_tree(&entry.path, current_depth + 1, max_depth)).await?; + let children = Box::pin(Self::build_tree( + workspace, + &entry.path, + current_depth + 1, + max_depth, + )) + .await?; if children.is_empty() { result.push(serde_json::Value::String(display_path)); } else { @@ -584,7 +649,7 @@ impl Tool for MemoryTreeTool { async fn execute( &self, params: serde_json::Value, - _ctx: &JobContext, + ctx: &JobContext, ) -> Result { let start = std::time::Instant::now(); @@ -596,7 +661,8 @@ impl Tool for MemoryTreeTool { .unwrap_or(1) .clamp(1, 10) as usize; - let tree = self.build_tree(path, 1, depth).await?; + let workspace = self.resolver.resolve(&ctx.user_id).await; + let tree = Self::build_tree(&workspace, path, 1, depth).await?; // Compact output: just the tree array Ok(ToolOutput::success( @@ -650,7 +716,7 @@ mod tests { #[test] fn test_memory_search_schema() { let workspace = make_test_workspace(); - let tool = MemorySearchTool::new(workspace); + let tool = MemorySearchTool::from_workspace(workspace); assert_eq!(tool.name(), "memory_search"); assert!(!tool.requires_sanitization()); @@ -668,7 +734,7 @@ mod tests { #[test] fn test_memory_write_schema() { let workspace = make_test_workspace(); - let tool = MemoryWriteTool::new(workspace); + let tool = MemoryWriteTool::from_workspace(workspace); assert_eq!(tool.name(), "memory_write"); @@ -681,7 +747,7 @@ mod tests { #[test] fn test_memory_read_schema() { let workspace = make_test_workspace(); - let tool = MemoryReadTool::new(workspace); + let tool = MemoryReadTool::from_workspace(workspace); assert_eq!(tool.name(), "memory_read"); @@ -698,7 +764,7 @@ mod tests { #[test] fn test_memory_tree_schema() { let workspace = make_test_workspace(); - let tool = MemoryTreeTool::new(workspace); + let tool = MemoryTreeTool::from_workspace(workspace); assert_eq!(tool.name(), "memory_tree"); @@ -711,7 +777,7 @@ mod tests { #[tokio::test] async fn test_memory_write_rejects_injection_to_identity_file() { let workspace = make_test_workspace(); - let tool = MemoryWriteTool::new(workspace); + let tool = MemoryWriteTool::from_workspace(workspace); let ctx = JobContext::default(); let params = serde_json::json!({ @@ -733,4 +799,176 @@ mod tests { } } } + + // Regression tests for per-user workspace scoping (multi-tenant mode). + // See: https://github.com/nearai/ironclaw/pull/1118 + // Bug: memory tools used a single startup workspace regardless of which + // user was chatting. Fix: resolve workspace per-request via JobContext.user_id. + + #[cfg(feature = "postgres")] + mod resolver_tests { + use super::*; + + fn make_test_workspace_for_user(user_id: &str) -> Arc { + Arc::new(Workspace::new( + user_id, + deadpool_postgres::Pool::builder(deadpool_postgres::Manager::new( + tokio_postgres::Config::new(), + tokio_postgres::NoTls, + )) + .build() + .unwrap(), + )) + } + + #[tokio::test] + async fn test_fixed_workspace_resolver_ignores_user_id() { + let ws = make_test_workspace_for_user("alice"); + let resolver = FixedWorkspaceResolver::new(Arc::clone(&ws)); + + let ws_alice = resolver.resolve("alice").await; + let ws_bob = resolver.resolve("bob").await; + + // Both should return the exact same Arc (pointer equality) + assert!(Arc::ptr_eq(&ws_alice, &ws_bob)); + assert_eq!(ws_alice.user_id(), "alice"); + } + + /// Tracking resolver that records which user_ids were requested. + struct TrackingWorkspaceResolver { + inner: FixedWorkspaceResolver, + resolved_users: std::sync::Mutex>, + } + + impl TrackingWorkspaceResolver { + fn new(workspace: Arc) -> Self { + Self { + inner: FixedWorkspaceResolver::new(workspace), + resolved_users: std::sync::Mutex::new(Vec::new()), + } + } + + fn resolved_users(&self) -> Vec { + self.resolved_users.lock().unwrap().clone() + } + } + + #[async_trait] + impl WorkspaceResolver for TrackingWorkspaceResolver { + async fn resolve(&self, user_id: &str) -> Arc { + self.resolved_users + .lock() + .unwrap() + .push(user_id.to_string()); + self.inner.resolve(user_id).await + } + } + + #[tokio::test] + async fn test_memory_search_uses_job_context_user_id() { + let ws = make_test_workspace_for_user("default"); + let tracker = Arc::new(TrackingWorkspaceResolver::new(ws)); + let tool = MemorySearchTool::new(tracker.clone() as Arc); + + // Execute with user_id "alice" + let ctx_alice = JobContext::with_user("alice", "test", "test"); + let params = serde_json::json!({"query": "test"}); + // The search will fail (no real DB) but we only care about resolver call + let _ = tool.execute(params, &ctx_alice).await; + + // Execute with user_id "bob" + let ctx_bob = JobContext::with_user("bob", "test", "test"); + let params = serde_json::json!({"query": "test"}); + let _ = tool.execute(params, &ctx_bob).await; + + let resolved = tracker.resolved_users(); + assert_eq!(resolved, vec!["alice", "bob"]); + } + + #[tokio::test] + async fn test_memory_write_uses_job_context_user_id() { + let ws = make_test_workspace_for_user("default"); + let tracker = Arc::new(TrackingWorkspaceResolver::new(ws)); + let tool = MemoryWriteTool::new(tracker.clone() as Arc); + + // Execute with user_id "alice" + let ctx_alice = JobContext::with_user("alice", "test", "test"); + let params = serde_json::json!({ + "content": "remember this", + "target": "daily_log", + }); + let _ = tool.execute(params, &ctx_alice).await; + + // Execute with user_id "bob" + let ctx_bob = JobContext::with_user("bob", "test", "test"); + let params = serde_json::json!({ + "content": "remember that", + "target": "daily_log", + }); + let _ = tool.execute(params, &ctx_bob).await; + + let resolved = tracker.resolved_users(); + assert_eq!(resolved, vec!["alice", "bob"]); + } + } + + #[cfg(feature = "libsql")] + mod per_user_resolver_tests { + use super::*; + + async fn make_test_db() -> Arc { + use crate::db::libsql::LibSqlBackend; + let temp_dir = tempfile::tempdir().expect("tempdir"); + let db_path = temp_dir.path().join("resolver_test.db"); + let backend = LibSqlBackend::new_local(&db_path) + .await + .expect("LibSqlBackend"); + ::run_migrations(&backend) + .await + .expect("migrations"); + // Leak the tempdir so it outlives the test (cleaned up on process exit). + std::mem::forget(temp_dir); + Arc::new(backend) + } + + #[tokio::test] + async fn test_workspace_pool_resolver_returns_different_workspaces() { + let db = make_test_db().await; + + let pool = crate::channels::web::server::WorkspacePool::new( + db, + None, + crate::workspace::EmbeddingCacheConfig::default(), + crate::config::WorkspaceSearchConfig::default(), + crate::config::WorkspaceConfig::default(), + ); + + let ws_alice = pool.resolve("alice").await; + let ws_bob = pool.resolve("bob").await; + + // Different user IDs should get different workspaces + assert_eq!(ws_alice.user_id(), "alice"); + assert_eq!(ws_bob.user_id(), "bob"); + assert!(!Arc::ptr_eq(&ws_alice, &ws_bob)); + } + + #[tokio::test] + async fn test_workspace_pool_resolver_caches_workspace() { + let db = make_test_db().await; + + let pool = crate::channels::web::server::WorkspacePool::new( + db, + None, + crate::workspace::EmbeddingCacheConfig::default(), + crate::config::WorkspaceSearchConfig::default(), + crate::config::WorkspaceConfig::default(), + ); + + let ws1 = pool.resolve("alice").await; + let ws2 = pool.resolve("alice").await; + + // Same user_id should return the same cached Arc (pointer equality) + assert!(Arc::ptr_eq(&ws1, &ws2)); + } + } } diff --git a/src/tools/builtin/mod.rs b/src/tools/builtin/mod.rs index 8ba8e57b..d196b12c 100644 --- a/src/tools/builtin/mod.rs +++ b/src/tools/builtin/mod.rs @@ -6,7 +6,7 @@ mod file; mod http; mod job; mod json; -mod memory; +pub mod memory; mod message; pub mod path_utils; mod restart; diff --git a/src/tools/registry.rs b/src/tools/registry.rs index dff09a5c..bc3be144 100644 --- a/src/tools/registry.rs +++ b/src/tools/registry.rs @@ -334,15 +334,37 @@ impl ToolRegistry { tracing::debug!("Registered 5 development tools"); } - /// Register memory tools with a workspace. + /// Register memory tools with a workspace resolver. + /// + /// Memory tools require a workspace resolver for persistence. Call this after + /// `register_builtin_tools()` if you have a workspace available. + pub fn register_memory_tools_with_resolver( + &self, + resolver: Arc, + ) { + self.register_sync(Arc::new(MemorySearchTool::new(Arc::clone(&resolver)))); + self.register_sync(Arc::new(MemoryWriteTool::new(Arc::clone(&resolver)))); + self.register_sync(Arc::new(MemoryReadTool::new(Arc::clone(&resolver)))); + self.register_sync(Arc::new(MemoryTreeTool::new(resolver))); + + tracing::debug!("Registered 4 memory tools"); + } + + /// Register memory tools with a fixed workspace (backward compatibility). /// /// Memory tools require a workspace for persistence. Call this after /// `register_builtin_tools()` if you have a workspace available. pub fn register_memory_tools(&self, workspace: Arc) { - self.register_sync(Arc::new(MemorySearchTool::new(Arc::clone(&workspace)))); - self.register_sync(Arc::new(MemoryWriteTool::new(Arc::clone(&workspace)))); - self.register_sync(Arc::new(MemoryReadTool::new(Arc::clone(&workspace)))); - self.register_sync(Arc::new(MemoryTreeTool::new(workspace))); + self.register_sync(Arc::new(MemorySearchTool::from_workspace(Arc::clone( + &workspace, + )))); + self.register_sync(Arc::new(MemoryWriteTool::from_workspace(Arc::clone( + &workspace, + )))); + self.register_sync(Arc::new(MemoryReadTool::from_workspace(Arc::clone( + &workspace, + )))); + self.register_sync(Arc::new(MemoryTreeTool::from_workspace(workspace))); tracing::debug!("Registered 4 memory tools"); } @@ -361,7 +383,11 @@ impl ToolRegistry { job_manager: Option>, store: Option>, job_event_tx: Option< - tokio::sync::broadcast::Sender<(uuid::Uuid, crate::channels::web::types::SseEvent)>, + tokio::sync::broadcast::Sender<( + uuid::Uuid, + String, + crate::channels::web::types::SseEvent, + )>, >, inject_tx: Option>, prompt_queue: Option, diff --git a/src/worker/job.rs b/src/worker/job.rs index ba5d47b9..b2e3f7e6 100644 --- a/src/worker/job.rs +++ b/src/worker/job.rs @@ -48,8 +48,8 @@ pub struct WorkerDeps { pub hooks: Arc, pub timeout: Duration, pub use_planning: bool, - /// SSE broadcast sender for live job event streaming to the web gateway. - pub sse_tx: Option>, + /// SSE manager 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 /// are pre-approved for autonomous execution. @@ -138,7 +138,7 @@ impl Worker { } // Broadcast SSE for live web UI updates - if let Some(ref tx) = self.deps.sse_tx { + 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 { @@ -203,7 +203,7 @@ impl Worker { _ => None, }; if let Some(event) = event { - let _ = tx.send(event); + sse.broadcast(event); } } } diff --git a/tests/e2e_advanced_traces.rs b/tests/e2e_advanced_traces.rs index 2b9fac29..b3efc8d9 100644 --- a/tests/e2e_advanced_traces.rs +++ b/tests/e2e_advanced_traces.rs @@ -661,7 +661,7 @@ mod advanced { .await .expect("failed to inject test token"); - let activate_result = ext_mgr.activate("mock-notion").await; + let activate_result = ext_mgr.activate("mock-notion", "default").await; assert!( activate_result.is_ok(), "activation failed: {:?}", diff --git a/tests/module_init_integration.rs b/tests/module_init_integration.rs index c75ccc6f..3aea7984 100644 --- a/tests/module_init_integration.rs +++ b/tests/module_init_integration.rs @@ -216,7 +216,7 @@ async fn extension_manager_with_process_manager_constructs() { ); // Verify the manager is functional — list returns Ok. - let result = manager.list(None, false).await; + let result = manager.list(None, false, "test").await; assert!(result.is_ok(), "list should succeed on empty manager"); assert!(result.unwrap().is_empty()); } diff --git a/tests/multi_tenant_integration.rs b/tests/multi_tenant_integration.rs new file mode 100644 index 00000000..02eb60e8 --- /dev/null +++ b/tests/multi_tenant_integration.rs @@ -0,0 +1,1059 @@ +//! Integration tests for multi-tenant auth, isolation, and per-user scoping. +//! +//! These tests verify that multi-tenant infrastructure works correctly: +//! - Token-to-identity mapping via MultiAuthState +//! - Per-user SSE event scoping (user A doesn't see user B's events) +//! - Per-user rate limiting (user A exhausting limit doesn't block user B) +//! - Auth middleware inserts correct UserIdentity into request extensions +//! - WebSocket connections are scoped to the authenticated user + +use std::collections::HashMap; +use std::net::SocketAddr; +use std::sync::Arc; +use std::time::Duration; + +use axum::Router; +use axum::body::Body; +use axum::http::{Request, StatusCode}; +use axum::middleware; +use axum::routing::{get, post}; +use tower::ServiceExt; + +use ironclaw::channels::web::auth::{ + AuthenticatedUser, MultiAuthState, UserIdentity, auth_middleware, +}; +use ironclaw::channels::web::server::{GatewayState, PerUserRateLimiter, RateLimiter}; +use ironclaw::channels::web::sse::SseManager; +use ironclaw::channels::web::test_helpers::TestGatewayBuilder; +use ironclaw::channels::web::ws::WsConnectionTracker; +use ironclaw::context::JobContext; +use ironclaw::db::Database; + +// --------------------------------------------------------------------------- +// Helpers +// --------------------------------------------------------------------------- + +const ALICE_TOKEN: &str = "tok-alice-secret"; +const BOB_TOKEN: &str = "tok-bob-secret"; +const ALICE_USER_ID: &str = "alice"; +const BOB_USER_ID: &str = "bob"; + +/// Build a MultiAuthState with two users. +fn two_user_auth() -> MultiAuthState { + let mut tokens = HashMap::new(); + tokens.insert( + ALICE_TOKEN.to_string(), + UserIdentity { + user_id: ALICE_USER_ID.to_string(), + workspace_read_scopes: Vec::new(), + }, + ); + tokens.insert( + BOB_TOKEN.to_string(), + UserIdentity { + user_id: BOB_USER_ID.to_string(), + workspace_read_scopes: vec!["shared".to_string()], + }, + ); + MultiAuthState::multi(tokens) +} + +/// Build a test Router that echoes the authenticated user_id back. +fn user_echo_app(auth: MultiAuthState) -> Router { + async fn echo_user(AuthenticatedUser(user): AuthenticatedUser) -> String { + user.user_id + } + + async fn echo_user_with_scopes(AuthenticatedUser(user): AuthenticatedUser) -> String { + format!("{}:{}", user.user_id, user.workspace_read_scopes.join(",")) + } + + Router::new() + .route("/api/whoami", get(echo_user)) + .route("/api/whoami/scopes", get(echo_user_with_scopes)) + .route("/api/action", post(echo_user)) + .route("/api/chat/events", get(echo_user)) // SSE endpoint (allows query token) + .layer(middleware::from_fn_with_state(auth, auth_middleware)) +} + +// =========================================================================== +// Auth: token-to-identity mapping +// =========================================================================== + +#[tokio::test] +async fn alice_token_resolves_to_alice_identity() { + let app = user_echo_app(two_user_auth()); + let resp = app + .oneshot( + Request::builder() + .uri("/api/whoami") + .header("Authorization", format!("Bearer {ALICE_TOKEN}")) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::OK); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + assert_eq!(std::str::from_utf8(&body).unwrap(), ALICE_USER_ID); +} + +#[tokio::test] +async fn bob_token_resolves_to_bob_identity() { + let app = user_echo_app(two_user_auth()); + let resp = app + .oneshot( + Request::builder() + .uri("/api/whoami") + .header("Authorization", format!("Bearer {BOB_TOKEN}")) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::OK); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + assert_eq!(std::str::from_utf8(&body).unwrap(), BOB_USER_ID); +} + +#[tokio::test] +async fn bob_identity_carries_workspace_read_scopes() { + let app = user_echo_app(two_user_auth()); + let resp = app + .oneshot( + Request::builder() + .uri("/api/whoami/scopes") + .header("Authorization", format!("Bearer {BOB_TOKEN}")) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::OK); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + assert_eq!(std::str::from_utf8(&body).unwrap(), "bob:shared"); +} + +#[tokio::test] +async fn unknown_token_rejected() { + let app = user_echo_app(two_user_auth()); + let resp = app + .oneshot( + Request::builder() + .uri("/api/whoami") + .header("Authorization", "Bearer unknown-token") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); +} + +#[tokio::test] +async fn no_token_rejected() { + let app = user_echo_app(two_user_auth()); + let resp = app + .oneshot( + Request::builder() + .uri("/api/whoami") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); +} + +#[tokio::test] +async fn alice_token_does_not_authenticate_as_bob() { + let app = user_echo_app(two_user_auth()); + let resp = app + .oneshot( + Request::builder() + .uri("/api/whoami") + .header("Authorization", format!("Bearer {ALICE_TOKEN}")) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + let user_id = std::str::from_utf8(&body).unwrap(); + assert_eq!(user_id, ALICE_USER_ID); + assert_ne!(user_id, BOB_USER_ID); +} + +// =========================================================================== +// Auth: query token on SSE/WS endpoints +// =========================================================================== + +#[tokio::test] +async fn query_token_works_for_sse_endpoint_multi_user() { + let app = user_echo_app(two_user_auth()); + let resp = app + .oneshot( + Request::builder() + .uri(format!("/api/chat/events?token={ALICE_TOKEN}")) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::OK); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + assert_eq!(std::str::from_utf8(&body).unwrap(), ALICE_USER_ID); +} + +#[tokio::test] +async fn query_token_rejected_for_non_sse_endpoint_multi_user() { + let app = user_echo_app(two_user_auth()); + let resp = app + .oneshot( + Request::builder() + .uri(format!("/api/whoami?token={ALICE_TOKEN}")) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); +} + +#[tokio::test] +async fn query_token_rejected_for_post_multi_user() { + let app = user_echo_app(two_user_auth()); + let resp = app + .oneshot( + Request::builder() + .method("POST") + .uri(format!("/api/action?token={ALICE_TOKEN}")) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); +} + +// =========================================================================== +// Per-user rate limiting +// =========================================================================== + +#[test] +fn per_user_rate_limiter_isolates_users() { + let limiter = PerUserRateLimiter::new(3, 60); + + // Alice uses all 3 requests + assert!(limiter.check("alice")); + assert!(limiter.check("alice")); + assert!(limiter.check("alice")); + // Alice is now rate-limited + assert!(!limiter.check("alice")); + + // Bob is unaffected — gets his own 3 requests + assert!(limiter.check("bob")); + assert!(limiter.check("bob")); + assert!(limiter.check("bob")); + assert!(!limiter.check("bob")); +} + +#[test] +fn per_user_rate_limiter_different_users_independent() { + let limiter = PerUserRateLimiter::new(2, 60); + + // Interleave requests from different users + assert!(limiter.check("alice")); + assert!(limiter.check("bob")); + assert!(limiter.check("alice")); + assert!(limiter.check("bob")); + + // Both exhausted independently + assert!(!limiter.check("alice")); + assert!(!limiter.check("bob")); + + // Charlie is fresh + assert!(limiter.check("charlie")); +} + +#[test] +fn per_user_rate_limiter_single_user_mode() { + // In single-user mode, only one user_id is used + let limiter = PerUserRateLimiter::new(5, 60); + for _ in 0..5 { + assert!(limiter.check("default")); + } + assert!(!limiter.check("default")); +} + +// =========================================================================== +// SSE event scoping +// =========================================================================== + +#[tokio::test] +async fn sse_scoped_event_only_delivered_to_target_user() { + use ironclaw::channels::web::types::SseEvent; + use tokio_stream::StreamExt; + + let manager = SseManager::new(); + let mut alice_stream = Box::pin( + manager + .subscribe_raw(Some(ALICE_USER_ID.to_string())) + .expect("subscribe"), + ); + let mut bob_stream = Box::pin( + manager + .subscribe_raw(Some(BOB_USER_ID.to_string())) + .expect("subscribe"), + ); + + // Send event scoped to alice + manager.broadcast_for_user( + ALICE_USER_ID, + SseEvent::Status { + message: "alice's event".to_string(), + thread_id: None, + }, + ); + + // Send global heartbeat (both should get it) + manager.broadcast(SseEvent::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"), + _ => panic!("Expected Status, got {:?}", e), + } + + // Alice also gets heartbeat + let e = alice_stream.next().await.unwrap(); + assert!(matches!(e, SseEvent::Heartbeat)); + + // Bob only gets the heartbeat (alice's event was filtered) + let e = bob_stream.next().await.unwrap(); + assert!(matches!(e, SseEvent::Heartbeat)); +} + +#[tokio::test] +async fn sse_global_event_delivered_to_all_users() { + use ironclaw::channels::web::types::SseEvent; + use tokio_stream::StreamExt; + + let manager = SseManager::new(); + let mut alice = Box::pin( + manager + .subscribe_raw(Some(ALICE_USER_ID.to_string())) + .expect("subscribe"), + ); + let mut bob = Box::pin( + manager + .subscribe_raw(Some(BOB_USER_ID.to_string())) + .expect("subscribe"), + ); + + manager.broadcast(SseEvent::Status { + message: "global announcement".to_string(), + thread_id: None, + }); + + let ea = alice.next().await.unwrap(); + let eb = bob.next().await.unwrap(); + match (&ea, &eb) { + (SseEvent::Status { message: a, .. }, SseEvent::Status { message: b, .. }) => { + assert_eq!(a, "global announcement"); + assert_eq!(b, "global announcement"); + } + _ => panic!("Expected Status events"), + } +} + +#[tokio::test] +async fn sse_user_b_event_not_visible_to_user_a() { + use ironclaw::channels::web::types::SseEvent; + use tokio_stream::StreamExt; + + let manager = SseManager::new(); + let mut alice = Box::pin( + manager + .subscribe_raw(Some(ALICE_USER_ID.to_string())) + .expect("subscribe"), + ); + + // Send event for bob only + manager.broadcast_for_user( + BOB_USER_ID, + SseEvent::Response { + content: "bob's secret".to_string(), + thread_id: "t1".to_string(), + }, + ); + + // Send heartbeat so alice has something to receive + manager.broadcast(SseEvent::Heartbeat); + + // Alice should only get heartbeat, not bob's response + let e = alice.next().await.unwrap(); + assert!( + matches!(e, SseEvent::Heartbeat), + "Expected Heartbeat, got {:?}", + e + ); +} + +#[tokio::test] +async fn sse_unscoped_subscriber_receives_all_events() { + use ironclaw::channels::web::types::SseEvent; + use tokio_stream::StreamExt; + + let manager = SseManager::new(); + // Unscoped subscriber (None user_id) — backwards-compatible single-user mode + let mut stream = Box::pin(manager.subscribe_raw(None).expect("subscribe")); + + manager.broadcast_for_user( + ALICE_USER_ID, + SseEvent::Status { + message: "alice only".to_string(), + thread_id: None, + }, + ); + manager.broadcast_for_user( + BOB_USER_ID, + SseEvent::Status { + message: "bob only".to_string(), + thread_id: None, + }, + ); + manager.broadcast(SseEvent::Heartbeat); + + // Unscoped subscriber gets ALL three events + let e1 = stream.next().await.unwrap(); + let e2 = stream.next().await.unwrap(); + let e3 = stream.next().await.unwrap(); + + match &e1 { + SseEvent::Status { message, .. } => assert_eq!(message, "alice only"), + _ => panic!("Expected alice's Status"), + } + match &e2 { + SseEvent::Status { message, .. } => assert_eq!(message, "bob only"), + _ => panic!("Expected bob's Status"), + } + assert!(matches!(e3, SseEvent::Heartbeat)); +} + +// =========================================================================== +// MultiAuthState: edge cases +// =========================================================================== + +#[test] +fn multi_auth_state_empty_token_not_valid() { + let state = MultiAuthState::single("real-token".to_string(), "user1".to_string()); + assert!(state.authenticate("").is_none()); +} + +#[test] +fn multi_auth_state_first_token_is_none_in_multi_user_mode() { + let auth = two_user_auth(); + // first_token() returns None in multi-user mode to avoid exposing tokens. + assert!(auth.first_token().is_none()); +} + +#[test] +fn multi_auth_state_first_identity_returns_valid_user() { + let auth = two_user_auth(); + let identity = auth.first_identity().unwrap(); + assert!(identity.user_id == ALICE_USER_ID || identity.user_id == BOB_USER_ID); +} + +#[test] +fn multi_auth_state_token_prefix_not_valid() { + // Ensure partial token matches don't authenticate + let state = MultiAuthState::single("secret-token-123".to_string(), "user1".to_string()); + assert!(state.authenticate("secret-token").is_none()); + assert!(state.authenticate("secret-token-1234").is_none()); + assert!(state.authenticate("secret-token-123").is_some()); +} + +// =========================================================================== +// Connection counting with user scoping +// =========================================================================== + +#[tokio::test] +async fn sse_connection_count_tracks_scoped_subscribers() { + let manager = SseManager::new(); + assert_eq!(manager.connection_count(), 0); + + let _alice = Box::pin( + manager + .subscribe_raw(Some(ALICE_USER_ID.to_string())) + .expect("subscribe"), + ); + assert_eq!(manager.connection_count(), 1); + + let _bob = Box::pin( + manager + .subscribe_raw(Some(BOB_USER_ID.to_string())) + .expect("subscribe"), + ); + assert_eq!(manager.connection_count(), 2); + + drop(_alice); + assert_eq!(manager.connection_count(), 1); + + drop(_bob); + assert_eq!(manager.connection_count(), 0); +} + +// =========================================================================== +// GatewayState construction: multi-user fields +// =========================================================================== + +#[test] +fn gateway_state_has_multi_tenant_fields() { + // Verify the GatewayState struct accepts all multi-tenant fields. + // This is a compile-time check that the conflict resolution didn't + // drop any fields. + let state = GatewayState { + msg_tx: tokio::sync::RwLock::new(None), + sse: Arc::new(SseManager::new()), + workspace: None, + workspace_pool: None, // Multi-tenant: per-user workspace pool + session_manager: None, + log_broadcaster: None, + log_level_handle: None, + extension_manager: None, + tool_registry: None, + store: None, + job_manager: None, + prompt_queue: None, + scheduler: None, + default_user_id: "fallback".to_string(), // Multi-tenant: renamed from user_id + shutdown_tx: tokio::sync::RwLock::new(None), + ws_tracker: Some(Arc::new(WsConnectionTracker::new())), + llm_provider: None, + skill_registry: None, + skill_catalog: None, + chat_rate_limiter: PerUserRateLimiter::new(30, 60), // Multi-tenant: per-user + oauth_rate_limiter: RateLimiter::new(10, 60), + registry_entries: Vec::new(), + cost_guard: None, + routine_engine: Arc::new(tokio::sync::RwLock::new(None)), + startup_time: std::time::Instant::now(), + webhook_rate_limiter: RateLimiter::new(10, 60), + active_config: Default::default(), + }; + + assert_eq!(state.default_user_id, "fallback"); + assert!(state.workspace_pool.is_none()); +} + +// =========================================================================== +// Full-server handler-level tests (real HTTP through auth middleware) +// =========================================================================== + +/// Build a MultiAuthState with two users and start a real server. +async fn start_multi_user_server() -> (SocketAddr, Arc) { + let (agent_tx, _agent_rx) = tokio::sync::mpsc::channel(64); + let auth = two_user_auth(); + TestGatewayBuilder::new() + .msg_tx(agent_tx) + .start_multi(auth) + .await + .expect("Failed to start multi-user test server") +} + +#[tokio::test] +async fn full_server_alice_can_access_protected_endpoint() { + let (addr, _state) = start_multi_user_server().await; + + let client = reqwest::Client::new(); + let resp = client + .get(format!("http://{}/api/gateway/status", addr)) + .header("Authorization", format!("Bearer {}", ALICE_TOKEN)) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); +} + +#[tokio::test] +async fn full_server_bob_can_access_protected_endpoint() { + let (addr, _state) = start_multi_user_server().await; + + let client = reqwest::Client::new(); + let resp = client + .get(format!("http://{}/api/gateway/status", addr)) + .header("Authorization", format!("Bearer {}", BOB_TOKEN)) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); +} + +#[tokio::test] +async fn full_server_unknown_token_returns_401() { + let (addr, _state) = start_multi_user_server().await; + + let client = reqwest::Client::new(); + let resp = client + .get(format!("http://{}/api/gateway/status", addr)) + .header("Authorization", "Bearer wrong-token") + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 401); +} + +#[tokio::test] +async fn full_server_no_auth_header_returns_401() { + let (addr, _state) = start_multi_user_server().await; + + let client = reqwest::Client::new(); + let resp = client + .get(format!("http://{}/api/gateway/status", addr)) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 401); +} + +#[tokio::test] +async fn full_server_health_is_public() { + let (addr, _state) = start_multi_user_server().await; + + let client = reqwest::Client::new(); + let resp = client + .get(format!("http://{}/api/health", addr)) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); +} + +#[tokio::test] +async fn full_server_chat_send_accepted_for_alice() { + let (agent_tx, mut agent_rx) = tokio::sync::mpsc::channel(64); + let auth = two_user_auth(); + let (addr, _state) = TestGatewayBuilder::new() + .msg_tx(agent_tx) + .start_multi(auth) + .await + .expect("Failed to start server"); + + let client = reqwest::Client::new(); + let resp = client + .post(format!("http://{}/api/chat/send", addr)) + .header("Authorization", format!("Bearer {}", ALICE_TOKEN)) + .header("Content-Type", "application/json") + .body(r#"{"content":"hello from alice"}"#) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 202); // ACCEPTED + + // Verify the message reached the agent channel + let msg = tokio::time::timeout(Duration::from_secs(2), agent_rx.recv()) + .await + .expect("Timed out waiting for agent message") + .expect("Agent channel closed"); + + assert_eq!(msg.content, "hello from alice"); + assert_eq!(msg.channel, "gateway"); +} + +#[tokio::test] +async fn full_server_chat_send_rejected_without_auth() { + let (addr, _state) = start_multi_user_server().await; + + let client = reqwest::Client::new(); + let resp = client + .post(format!("http://{}/api/chat/send", addr)) + .header("Content-Type", "application/json") + .body(r#"{"content":"unauthorized message"}"#) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 401); +} + +#[tokio::test] +async fn full_server_query_token_works_for_sse() { + let (addr, _state) = start_multi_user_server().await; + + let client = reqwest::Client::new(); + // SSE endpoint should accept query token + let resp = client + .get(format!( + "http://{}/api/chat/events?token={}", + addr, ALICE_TOKEN + )) + .send() + .await + .unwrap(); + + // Should get 200 (SSE stream starts) + assert_eq!(resp.status(), 200); +} + +#[tokio::test] +async fn full_server_query_token_rejected_for_non_sse() { + let (addr, _state) = start_multi_user_server().await; + + let client = reqwest::Client::new(); + // Non-SSE endpoint should NOT accept query token + let resp = client + .get(format!( + "http://{}/api/gateway/status?token={}", + addr, ALICE_TOKEN + )) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 401); +} + +#[tokio::test] +async fn full_server_jobs_endpoint_returns_503_without_db() { + let (addr, _state) = start_multi_user_server().await; + + let client = reqwest::Client::new(); + // Jobs endpoint requires database — should return 503 (no DB configured) + // but NOT 401 (auth should pass) + let resp = client + .get(format!("http://{}/api/jobs", addr)) + .header("Authorization", format!("Bearer {}", ALICE_TOKEN)) + .send() + .await + .unwrap(); + + // Without a database, this should return a server error, not an auth error + let status = resp.status().as_u16(); + assert_ne!(status, 401, "Should not be auth error — token is valid"); + assert_ne!(status, 403, "Should not be forbidden — token is valid"); +} + +#[tokio::test] +async fn full_server_jobs_endpoint_rejected_without_auth() { + let (addr, _state) = start_multi_user_server().await; + + let client = reqwest::Client::new(); + let resp = client + .get(format!("http://{}/api/jobs", addr)) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 401); +} + +#[tokio::test] +async fn full_server_ws_multi_user_event_isolation() { + use futures::StreamExt; + use ironclaw::channels::web::types::SseEvent; + use tokio_tungstenite::tungstenite::Message; + use tokio_tungstenite::tungstenite::client::IntoClientRequest; + + let (addr, state) = start_multi_user_server().await; + + // Connect Alice's WS + let alice_url = format!("ws://{}/api/chat/ws?token={}", addr, ALICE_TOKEN); + let mut alice_req = alice_url.into_client_request().unwrap(); + alice_req.headers_mut().insert( + "Origin", + format!("http://127.0.0.1:{}", addr.port()).parse().unwrap(), + ); + let (mut alice_ws, _) = tokio_tungstenite::connect_async(alice_req) + .await + .expect("Alice WS connect failed"); + + // Connect Bob's WS + let bob_url = format!("ws://{}/api/chat/ws?token={}", addr, BOB_TOKEN); + let mut bob_req = bob_url.into_client_request().unwrap(); + bob_req.headers_mut().insert( + "Origin", + format!("http://127.0.0.1:{}", addr.port()).parse().unwrap(), + ); + let (mut bob_ws, _) = tokio_tungstenite::connect_async(bob_req) + .await + .expect("Bob WS connect failed"); + + tokio::time::sleep(Duration::from_millis(100)).await; + + // Broadcast an event scoped to Alice only + state.sse.broadcast_for_user( + ALICE_USER_ID, + SseEvent::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); + + // Alice should get her scoped event + let alice_msg = tokio::time::timeout(Duration::from_secs(2), alice_ws.next()) + .await + .expect("Alice WS timed out") + .expect("Alice stream ended") + .expect("Alice WS error"); + + if let Message::Text(text) = alice_msg { + let parsed: serde_json::Value = serde_json::from_str(&text).unwrap(); + assert_eq!(parsed["type"], "event"); + assert_eq!(parsed["event_type"], "status"); + assert_eq!(parsed["data"]["message"], "alice-only-event"); + } else { + panic!("Expected Text frame from Alice WS, got {:?}", alice_msg); + } + + // Bob should only get the heartbeat, NOT alice's event + let bob_msg = tokio::time::timeout(Duration::from_secs(2), bob_ws.next()) + .await + .expect("Bob WS timed out") + .expect("Bob stream ended") + .expect("Bob WS error"); + + if let Message::Text(text) = bob_msg { + let parsed: serde_json::Value = serde_json::from_str(&text).unwrap(); + assert_eq!(parsed["type"], "event"); + assert_eq!( + parsed["event_type"], "heartbeat", + "Bob should only see heartbeat, not alice's event. Got: {}", + text + ); + } else { + panic!("Expected Text frame from Bob WS, got {:?}", bob_msg); + } + + alice_ws.close(None).await.ok(); + bob_ws.close(None).await.ok(); +} + +// =========================================================================== +// DB-backed job ownership tests (libSQL in-memory) +// =========================================================================== + +/// Start a multi-user server with a real (in-memory) database. +#[cfg(feature = "libsql")] +async fn start_multi_user_server_with_db() -> ( + SocketAddr, + Arc, + Arc, + tempfile::TempDir, +) { + let temp_dir = tempfile::tempdir().expect("failed to create temp dir"); + let path = temp_dir.path().join("test.db"); + let backend = ironclaw::db::libsql::LibSqlBackend::new_local(&path) + .await + .expect("failed to create test DB"); + backend + .run_migrations() + .await + .expect("failed to run migrations"); + let db: Arc = Arc::new(backend); + let (agent_tx, _agent_rx) = tokio::sync::mpsc::channel(64); + let auth = two_user_auth(); + + // Build state manually so we can inject the DB + let state = Arc::new(GatewayState { + msg_tx: tokio::sync::RwLock::new(Some(agent_tx)), + sse: Arc::new(SseManager::new()), + workspace: None, + workspace_pool: None, + session_manager: None, + log_broadcaster: None, + log_level_handle: None, + extension_manager: None, + tool_registry: None, + store: Some(Arc::clone(&db)), + job_manager: None, + prompt_queue: None, + scheduler: None, + default_user_id: ALICE_USER_ID.to_string(), + shutdown_tx: tokio::sync::RwLock::new(None), + ws_tracker: Some(Arc::new(WsConnectionTracker::new())), + llm_provider: None, + skill_registry: None, + skill_catalog: None, + chat_rate_limiter: PerUserRateLimiter::new(30, 60), + oauth_rate_limiter: RateLimiter::new(10, 60), + registry_entries: Vec::new(), + cost_guard: None, + routine_engine: Arc::new(tokio::sync::RwLock::new(None)), + startup_time: std::time::Instant::now(), + webhook_rate_limiter: RateLimiter::new(10, 60), + active_config: Default::default(), + }); + + let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); + let bound = ironclaw::channels::web::server::start_server(addr, state.clone(), auth) + .await + .expect("Failed to start server with DB"); + + (bound, state, db, temp_dir) +} + +#[cfg(feature = "libsql")] +#[tokio::test] +async fn full_server_alice_sees_own_jobs_only() { + let (addr, _state, db, _tmp) = start_multi_user_server_with_db().await; + + // Create jobs owned by Alice and Bob + let alice_job = JobContext::with_user(ALICE_USER_ID, "Alice's job", "Alice's work"); + let bob_job = JobContext::with_user(BOB_USER_ID, "Bob's job", "Bob's work"); + let alice_job_id = alice_job.job_id; + + db.save_job(&alice_job).await.unwrap(); + db.save_job(&bob_job).await.unwrap(); + + let client = reqwest::Client::new(); + + // Alice lists jobs — should only see her own + let resp = client + .get(format!("http://{}/api/jobs", addr)) + .header("Authorization", format!("Bearer {}", ALICE_TOKEN)) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); + let body: serde_json::Value = resp.json().await.unwrap(); + let jobs = body["jobs"].as_array().unwrap(); + + // Alice should see exactly 1 job + assert_eq!(jobs.len(), 1, "Alice should see only her own job"); + assert_eq!(jobs[0]["id"], alice_job_id.to_string()); + assert_eq!(jobs[0]["title"], "Alice's job"); +} + +#[cfg(feature = "libsql")] +#[tokio::test] +async fn full_server_bob_cannot_see_alice_job_detail() { + let (addr, _state, db, _tmp) = start_multi_user_server_with_db().await; + + // Create a job owned by Alice + let alice_job = JobContext::with_user(ALICE_USER_ID, "Alice's secret job", "Private"); + let alice_job_id = alice_job.job_id; + db.save_job(&alice_job).await.unwrap(); + + let client = reqwest::Client::new(); + + // Bob tries to access Alice's job by ID — should get 404 (not 403, to prevent enumeration) + let resp = client + .get(format!("http://{}/api/jobs/{}", addr, alice_job_id)) + .header("Authorization", format!("Bearer {}", BOB_TOKEN)) + .send() + .await + .unwrap(); + + assert_eq!( + resp.status(), + 404, + "Bob should not be able to see Alice's job" + ); +} + +#[cfg(feature = "libsql")] +#[tokio::test] +async fn full_server_alice_can_see_own_job_detail() { + let (addr, _state, db, _tmp) = start_multi_user_server_with_db().await; + + let alice_job = JobContext::with_user(ALICE_USER_ID, "Alice's visible job", "Details here"); + let alice_job_id = alice_job.job_id; + db.save_job(&alice_job).await.unwrap(); + + let client = reqwest::Client::new(); + + let resp = client + .get(format!("http://{}/api/jobs/{}", addr, alice_job_id)) + .header("Authorization", format!("Bearer {}", ALICE_TOKEN)) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); + let body: serde_json::Value = resp.json().await.unwrap(); + assert_eq!(body["id"], alice_job_id.to_string()); + assert_eq!(body["title"], "Alice's visible job"); +} + +#[cfg(feature = "libsql")] +#[tokio::test] +async fn full_server_bob_sees_own_jobs_only() { + let (addr, _state, db, _tmp) = start_multi_user_server_with_db().await; + + // Create multiple jobs for each user + for i in 0..3 { + let aj = JobContext::with_user(ALICE_USER_ID, format!("Alice job {}", i), ""); + db.save_job(&aj).await.unwrap(); + } + for i in 0..2 { + let bj = JobContext::with_user(BOB_USER_ID, format!("Bob job {}", i), ""); + db.save_job(&bj).await.unwrap(); + } + + let client = reqwest::Client::new(); + + // Bob lists jobs + let resp = client + .get(format!("http://{}/api/jobs", addr)) + .header("Authorization", format!("Bearer {}", BOB_TOKEN)) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); + let body: serde_json::Value = resp.json().await.unwrap(); + let jobs = body["jobs"].as_array().unwrap(); + + assert_eq!( + jobs.len(), + 2, + "Bob should see only his 2 jobs, not Alice's 3" + ); + for job in jobs { + let title = job["title"].as_str().unwrap(); + assert!( + title.starts_with("Bob job"), + "Bob should only see his own jobs, got: {}", + title + ); + } +} + +#[cfg(feature = "libsql")] +#[tokio::test] +async fn full_server_nonexistent_job_returns_404() { + let (addr, _state, _db, _tmp) = start_multi_user_server_with_db().await; + + let client = reqwest::Client::new(); + let fake_id = uuid::Uuid::new_v4(); + + let resp = client + .get(format!("http://{}/api/jobs/{}", addr, fake_id)) + .header("Authorization", format!("Bearer {}", ALICE_TOKEN)) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 404); +} diff --git a/tests/multi_tenant_system_prompt.rs b/tests/multi_tenant_system_prompt.rs new file mode 100644 index 00000000..ece794bf --- /dev/null +++ b/tests/multi_tenant_system_prompt.rs @@ -0,0 +1,240 @@ +//! Tests proving that multi-tenant system prompts are broken. +//! +//! 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. +//! +//! These tests: +//! 1. Seed identity files for two users (alice, bob) in the database +//! 2. Send messages as each user +//! 3. Verify the system prompt in captured LLM requests contains the +//! 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. + +#[cfg(feature = "libsql")] +mod support; + +#[cfg(feature = "libsql")] +mod tests { + use std::sync::Arc; + use std::time::Duration; + + use ironclaw::channels::IncomingMessage; + use ironclaw::llm::Role; + use ironclaw::workspace::Workspace; + + use crate::support::test_rig::TestRigBuilder; + use crate::support::trace_llm::{LlmTrace, TraceResponse, TraceStep}; + + const TIMEOUT: Duration = Duration::from_secs(15); + + const ALICE_USER_ID: &str = "alice"; + const BOB_USER_ID: &str = "bob"; + + const ALICE_IDENTITY: &str = "You are Alice's personal assistant. \ + Alice is a software engineer who lives in Seattle."; + const BOB_IDENTITY: &str = "You are Bob's personal assistant. \ + Bob is a marine biologist who lives in Miami."; + + /// Create a simple trace that returns a canned text response. + /// We need one step per message we plan to send. + fn simple_trace(num_steps: usize) -> LlmTrace { + let steps: Vec = (0..num_steps) + .map(|i| TraceStep { + request_hint: None, + response: TraceResponse::Text { + content: format!("Response {}", i), + input_tokens: 100, + output_tokens: 10, + }, + expected_tool_results: Vec::new(), + }) + .collect(); + + // Create separate turns for each step so the trace replays correctly. + let turns: Vec = steps + .into_iter() + .enumerate() + .map(|(i, step)| crate::support::trace_llm::TraceTurn { + user_input: format!("message {}", i), + steps: vec![step], + expects: Default::default(), + }) + .collect(); + + LlmTrace::new("test-model", turns) + } + + /// Seed identity files for a user by creating a workspace scoped to that + /// user and writing IDENTITY.md. + async fn seed_identity(db: &Arc, user_id: &str, content: &str) { + let ws = Workspace::new_with_db(user_id, db.clone()); + ws.write("IDENTITY.md", content) + .await + .unwrap_or_else(|e| panic!("Failed to seed IDENTITY.md for {user_id}: {e}")); + } + + /// Extract the system prompt from captured LLM requests. + /// + /// The system prompt is the first message with role=System in the first + /// LLM request for a given turn. + fn extract_system_prompt(requests: &[Vec]) -> Option { + requests.last().and_then(|msgs| { + msgs.iter() + .find(|m| matches!(m.role, Role::System)) + .map(|m| m.content.clone()) + }) + } + + // ----------------------------------------------------------------------- + // Test 1: Alice's identity should appear in system prompt when messaging + // as Alice. + // ----------------------------------------------------------------------- + + #[tokio::test] + async fn alice_system_prompt_contains_alice_identity() { + let trace = simple_trace(1); + let rig = TestRigBuilder::new().with_trace(trace).build().await; + + // Seed alice's identity into the database + let db = rig.database(); + seed_identity(db, ALICE_USER_ID, ALICE_IDENTITY).await; + + // Send a message AS alice (using her user_id) + let msg = IncomingMessage::new("test", ALICE_USER_ID, "Hello, who am I?"); + rig.send_incoming(msg).await; + let _responses = rig.wait_for_responses(1, TIMEOUT).await; + + // The system prompt sent to the LLM should contain Alice's identity + let requests = rig.captured_llm_requests(); + let system_prompt = + extract_system_prompt(&requests).expect("Expected a system prompt in the LLM request"); + + assert!( + system_prompt.contains("Alice is a software engineer"), + "System prompt should contain Alice's identity when messaging as Alice.\n\ + Actual system prompt:\n{system_prompt}" + ); + + rig.shutdown(); + } + + // ----------------------------------------------------------------------- + // Test 2: Bob's identity should appear in system prompt when messaging + // as Bob. + // ----------------------------------------------------------------------- + + #[tokio::test] + async fn bob_system_prompt_contains_bob_identity() { + let trace = simple_trace(1); + let rig = TestRigBuilder::new().with_trace(trace).build().await; + + // Seed bob's identity into the database + let db = rig.database(); + seed_identity(db, BOB_USER_ID, BOB_IDENTITY).await; + + // Send a message AS bob + let msg = IncomingMessage::new("test", BOB_USER_ID, "Hello, who am I?"); + rig.send_incoming(msg).await; + let _responses = rig.wait_for_responses(1, TIMEOUT).await; + + // The system prompt should contain Bob's identity + let requests = rig.captured_llm_requests(); + let system_prompt = + extract_system_prompt(&requests).expect("Expected a system prompt in the LLM request"); + + assert!( + system_prompt.contains("Bob is a marine biologist"), + "System prompt should contain Bob's identity when messaging as Bob.\n\ + Actual system prompt:\n{system_prompt}" + ); + + rig.shutdown(); + } + + // ----------------------------------------------------------------------- + // Test 3: Alice's identity must NOT appear in Bob's system prompt. + // ----------------------------------------------------------------------- + + #[tokio::test] + async fn alice_identity_does_not_leak_into_bob_prompt() { + let trace = simple_trace(1); + let rig = TestRigBuilder::new().with_trace(trace).build().await; + + // Seed BOTH users' identities + let db = rig.database(); + seed_identity(db, ALICE_USER_ID, ALICE_IDENTITY).await; + seed_identity(db, BOB_USER_ID, BOB_IDENTITY).await; + + // Send a message AS bob + let msg = IncomingMessage::new("test", BOB_USER_ID, "Tell me about myself"); + rig.send_incoming(msg).await; + let _responses = rig.wait_for_responses(1, TIMEOUT).await; + + // Bob's prompt must NOT contain Alice's identity + let requests = rig.captured_llm_requests(); + let system_prompt = extract_system_prompt(&requests); + + if let Some(ref prompt) = system_prompt { + assert!( + !prompt.contains("Alice is a software engineer"), + "Alice's identity LEAKED into Bob's system prompt!\n\ + System prompt:\n{prompt}" + ); + } + // Also verify Bob's identity IS present (compound check) + let prompt = system_prompt.expect("Expected a system prompt in the LLM request"); + assert!( + prompt.contains("Bob is a marine biologist"), + "Bob's own identity should be in his system prompt.\n\ + Actual system prompt:\n{prompt}" + ); + + rig.shutdown(); + } + + // ----------------------------------------------------------------------- + // Test 4: Bob's identity must NOT appear in Alice's system prompt. + // ----------------------------------------------------------------------- + + #[tokio::test] + async fn bob_identity_does_not_leak_into_alice_prompt() { + let trace = simple_trace(1); + let rig = TestRigBuilder::new().with_trace(trace).build().await; + + // Seed BOTH users' identities + let db = rig.database(); + seed_identity(db, ALICE_USER_ID, ALICE_IDENTITY).await; + seed_identity(db, BOB_USER_ID, BOB_IDENTITY).await; + + // Send a message AS alice + let msg = IncomingMessage::new("test", ALICE_USER_ID, "Tell me about myself"); + rig.send_incoming(msg).await; + let _responses = rig.wait_for_responses(1, TIMEOUT).await; + + // Alice's prompt must NOT contain Bob's identity + let requests = rig.captured_llm_requests(); + let system_prompt = extract_system_prompt(&requests); + + if let Some(ref prompt) = system_prompt { + assert!( + !prompt.contains("Bob is a marine biologist"), + "Bob's identity LEAKED into Alice's system prompt!\n\ + System prompt:\n{prompt}" + ); + } + // Also verify Alice's identity IS present + let prompt = system_prompt.expect("Expected a system prompt in the LLM request"); + assert!( + prompt.contains("Alice is a software engineer"), + "Alice's own identity should be in her system prompt.\n\ + Actual system prompt:\n{prompt}" + ); + + rig.shutdown(); + } +} diff --git a/tests/openai_compat_integration.rs b/tests/openai_compat_integration.rs index 2a472d00..16568246 100644 --- a/tests/openai_compat_integration.rs +++ b/tests/openai_compat_integration.rs @@ -191,8 +191,9 @@ async fn start_test_server_with_provider( ) -> (SocketAddr, Arc) { let state = Arc::new(GatewayState { msg_tx: tokio::sync::RwLock::new(None), - sse: SseManager::new(), + sse: Arc::new(SseManager::new()), workspace: None, + workspace_pool: None, session_manager: None, log_broadcaster: None, log_level_handle: None, @@ -202,13 +203,13 @@ async fn start_test_server_with_provider( job_manager: None, prompt_queue: None, scheduler: None, - user_id: "test-user".to_string(), + default_user_id: "test-user".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: Some(llm_provider), skill_registry: None, skill_catalog: None, - chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60), + chat_rate_limiter: ironclaw::channels::web::server::PerUserRateLimiter::new(30, 60), oauth_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60), webhook_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60), registry_entries: Vec::new(), @@ -218,8 +219,12 @@ async fn start_test_server_with_provider( active_config: ironclaw::channels::web::server::ActiveConfigSnapshot::default(), }); + let auth = ironclaw::channels::web::auth::MultiAuthState::single( + AUTH_TOKEN.to_string(), + "test-user".to_string(), + ); let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); - let bound_addr = start_server(addr, state.clone(), AUTH_TOKEN.to_string()) + let bound_addr = start_server(addr, state.clone(), auth) .await .expect("Failed to start test server"); @@ -684,8 +689,9 @@ async fn test_no_llm_provider_returns_503() { // Create state WITHOUT llm_provider let state = Arc::new(GatewayState { msg_tx: tokio::sync::RwLock::new(None), - sse: SseManager::new(), + sse: Arc::new(SseManager::new()), workspace: None, + workspace_pool: None, session_manager: None, log_broadcaster: None, log_level_handle: None, @@ -695,13 +701,13 @@ async fn test_no_llm_provider_returns_503() { job_manager: None, prompt_queue: None, scheduler: None, - user_id: "test-user".to_string(), + default_user_id: "test-user".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: None, // No LLM! skill_registry: None, skill_catalog: None, - chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60), + chat_rate_limiter: ironclaw::channels::web::server::PerUserRateLimiter::new(30, 60), oauth_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60), webhook_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60), registry_entries: Vec::new(), @@ -711,10 +717,12 @@ async fn test_no_llm_provider_returns_503() { active_config: ironclaw::channels::web::server::ActiveConfigSnapshot::default(), }); + let auth = ironclaw::channels::web::auth::MultiAuthState::single( + AUTH_TOKEN.to_string(), + "test-user".to_string(), + ); let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); - let bound_addr = start_server(addr, state, AUTH_TOKEN.to_string()) - .await - .unwrap(); + let bound_addr = start_server(addr, state, auth).await.unwrap(); let url = format!("http://{}/v1/chat/completions", bound_addr); let resp = client() @@ -741,9 +749,10 @@ async fn test_chat_completions_body_too_large() { let state = ironclaw::channels::web::test_helpers::TestGatewayBuilder::new() .llm_provider(llm_provider) .build(); - let auth_state = ironclaw::channels::web::auth::AuthState { - token: AUTH_TOKEN.to_string(), - }; + let auth_state = ironclaw::channels::web::auth::MultiAuthState::single( + AUTH_TOKEN.to_string(), + "test-user".to_string(), + ); let app = Router::new() .route( diff --git a/tests/support/gateway_workflow_harness.rs b/tests/support/gateway_workflow_harness.rs index d33c6fe0..7f9d3dff 100644 --- a/tests/support/gateway_workflow_harness.rs +++ b/tests/support/gateway_workflow_harness.rs @@ -13,8 +13,11 @@ use ironclaw::agent::routine_engine::RoutineEngine; use ironclaw::agent::{Agent, AgentDeps, SessionManager as AgentSessionManager}; use ironclaw::app::{AppBuilder, AppBuilderFlags}; use ironclaw::channels::IncomingMessage; +use ironclaw::channels::web::auth::MultiAuthState; use ironclaw::channels::web::log_layer::LogBroadcaster; -use ironclaw::channels::web::server::{GatewayState, RateLimiter, start_server}; +use ironclaw::channels::web::server::{ + GatewayState, PerUserRateLimiter, RateLimiter, start_server, +}; use ironclaw::channels::web::sse::SseManager; use ironclaw::channels::web::ws::WsConnectionTracker; use ironclaw::config::{Config, RegistryProviderConfig, RoutineConfig}; @@ -211,8 +214,9 @@ impl GatewayWorkflowHarness { let gateway_state = Arc::new(GatewayState { msg_tx: tokio::sync::RwLock::new(Some(gw_tx)), - sse: SseManager::new(), + sse: Arc::new(SseManager::new()), workspace: components.workspace.clone(), + workspace_pool: None, session_manager: Some(Arc::clone(&agent_session_manager)), log_broadcaster: None, log_level_handle: None, @@ -222,13 +226,13 @@ impl GatewayWorkflowHarness { job_manager: None, prompt_queue: None, scheduler: Some(scheduler_slot.clone()), - user_id: user_id.clone(), + default_user_id: user_id.clone(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: Some(Arc::clone(&components.llm)), skill_registry: components.skill_registry.clone(), skill_catalog: components.skill_catalog.clone(), - chat_rate_limiter: RateLimiter::new(120, 60), + chat_rate_limiter: PerUserRateLimiter::new(120, 60), oauth_rate_limiter: RateLimiter::new(10, 60), webhook_rate_limiter: RateLimiter::new(10, 60), registry_entries: Vec::new(), @@ -254,7 +258,7 @@ impl GatewayWorkflowHarness { skills_config: components.config.skills.clone(), hooks: components.hooks, cost_guard: components.cost_guard, - sse_tx: Some(gateway_state.sse.sender()), + sse_tx: None, http_interceptor: None, transcription: None, document_extraction: None, @@ -288,10 +292,11 @@ impl GatewayWorkflowHarness { } let auth_token = "gateway-test-token".to_string(); + let auth = MultiAuthState::single(auth_token.clone(), user_id.clone()); let addr = start_server( "127.0.0.1:0".parse().expect("valid localhost addr"), Arc::clone(&gateway_state), - auth_token.clone(), + auth, ) .await .expect("failed to start gateway server"); diff --git a/tests/ws_gateway_integration.rs b/tests/ws_gateway_integration.rs index 556c5dcc..43277389 100644 --- a/tests/ws_gateway_integration.rs +++ b/tests/ws_gateway_integration.rs @@ -39,8 +39,9 @@ async fn start_test_server() -> ( let state = Arc::new(GatewayState { msg_tx: tokio::sync::RwLock::new(Some(agent_tx)), - sse: SseManager::new(), + sse: Arc::new(SseManager::new()), workspace: None, + workspace_pool: None, session_manager: None, log_broadcaster: None, log_level_handle: None, @@ -50,13 +51,13 @@ async fn start_test_server() -> ( job_manager: None, prompt_queue: None, scheduler: None, - user_id: "test-user".to_string(), + default_user_id: "test-user".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: None, skill_registry: None, skill_catalog: None, - chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60), + chat_rate_limiter: ironclaw::channels::web::server::PerUserRateLimiter::new(30, 60), oauth_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60), webhook_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60), registry_entries: Vec::new(), @@ -66,8 +67,12 @@ async fn start_test_server() -> ( active_config: ironclaw::channels::web::server::ActiveConfigSnapshot::default(), }); + let auth = ironclaw::channels::web::auth::MultiAuthState::single( + AUTH_TOKEN.to_string(), + "test-user".to_string(), + ); let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); - let bound_addr = start_server(addr, state.clone(), AUTH_TOKEN.to_string()) + let bound_addr = start_server(addr, state.clone(), auth) .await .expect("Failed to start test server"); From 3fdb18779699b68a7d429048a0b232e7afffff3c Mon Sep 17 00:00:00 2001 From: Illia Polosukhin Date: Mon, 23 Mar 2026 21:59:14 -0700 Subject: [PATCH 04/20] refactor(tools): auto-compact WASM tool schemas, add descriptions, improve credential prompts (#1525) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix(tools): add missing description, parameters, and improve credential prompts Silence three categories of startup warnings emitted by CapabilitiesFile::validate() and WasmToolLoader: 1. "description" field missing → add tool descriptions to all manifests 2. "parameters" field missing → add action-enum parameter schemas 3. Short credential prompts (<30 chars) → append source URLs Affects: github, gmail, google-calendar, google-docs, google-drive, google-sheets, google-slides, slack, telegram, llm-context, feishu. [skip-regression-check] Co-Authored-By: Claude Opus 4.6 (1M context) * refactor(tools): auto-compact WASM tool schemas from module exports Replace the manual `parameters` field in capabilities JSON with automatic schema compaction. WasmToolSchemas::compact_schema() derives a compact advertised schema from the WASM module's schema() export by keeping only required and enum-constrained properties. The full schema remains available via tool_info(detail: "schema"). This eliminates: - The `parameters` field from CapabilitiesFile and all 11 sidecar JSONs - The "missing parameters" startup warning from the loader - Manual maintenance of duplicate schema data The `description` field in capabilities JSON is retained. [skip-regression-check] Co-Authored-By: Claude Opus 4.6 (1M context) * fix(tests): remove cap_file.parameters reference in test_rig The parameters field was removed from CapabilitiesFile in the previous commit. Update test_rig.rs to match — schema is now auto-compacted from the WASM module export, no sidecar override needed. [skip-regression-check] Co-Authored-By: Claude Opus 4.6 (1M context) * fix(tools): handle oneOf schemas in compact_schema, add tool name to warning Address PR review feedback: - compact_schema now collects properties from oneOf/anyOf/allOf variants, fixing GitHub-style schemas that have no top-level properties - Use HashSet for required lookup instead of Vec::contains - Add tool name to "Capabilities file not found" warning for consistency [skip-regression-check] Co-Authored-By: Claude Opus 4.6 (1M context) * fix(tools): merge oneOf const values into enum, cap property collection Address review feedback from @serrrfirat: 1. Merge const values across oneOf variants into a single enum array, so the LLM sees all valid actions (not just the first variant's const). 2. Cap property collection at 100 to bound allocations. 3. Also keep properties with const constraint (single-variant case). 4. Update doc comment to describe variant collection and design choices around variant-level required fields. [skip-regression-check] Co-Authored-By: Claude Opus 4.6 (1M context) --------- Co-authored-by: Claude Opus 4.6 (1M context) --- channels-src/feishu/feishu.capabilities.json | 2 +- src/tools/wasm/capabilities_schema.rs | 122 ++------ src/tools/wasm/loader.rs | 94 +++--- src/tools/wasm/wrapper.rs | 285 +++++++++++++++--- tests/support/test_rig.rs | 15 +- .../github/github-tool.capabilities.json | 1 + tools-src/gmail/gmail-tool.capabilities.json | 3 +- .../google-calendar-tool.capabilities.json | 3 +- .../google-docs-tool.capabilities.json | 3 +- .../google-drive-tool.capabilities.json | 3 +- .../google-sheets-tool.capabilities.json | 3 +- .../google-slides-tool.capabilities.json | 3 +- .../llm-context-tool.capabilities.json | 1 + tools-src/slack/slack-tool.capabilities.json | 3 +- .../telegram/telegram-tool.capabilities.json | 3 +- .../web-search-tool.capabilities.json | 34 --- 16 files changed, 322 insertions(+), 256 deletions(-) diff --git a/channels-src/feishu/feishu.capabilities.json b/channels-src/feishu/feishu.capabilities.json index 877a293a..a228cc4e 100644 --- a/channels-src/feishu/feishu.capabilities.json +++ b/channels-src/feishu/feishu.capabilities.json @@ -21,7 +21,7 @@ }, { "name": "feishu_app_secret", - "prompt": "Enter your Feishu/Lark App Secret", + "prompt": "Enter your Feishu/Lark App Secret (from your app settings at open.feishu.cn)", "optional": false }, { diff --git a/src/tools/wasm/capabilities_schema.rs b/src/tools/wasm/capabilities_schema.rs index 482aca83..b2758329 100644 --- a/src/tools/wasm/capabilities_schema.rs +++ b/src/tools/wasm/capabilities_schema.rs @@ -47,12 +47,6 @@ pub struct CapabilitiesFile { #[serde(default)] pub description: Option, - /// JSON Schema for the tool's input parameters. - /// Used as the `Tool::parameters_schema()` return value. - /// If omitted, a permissive fallback is used (with a warning). - #[serde(default)] - pub parameters: Option, - /// Extension version (semver). #[serde(default)] pub version: Option, @@ -103,9 +97,6 @@ pub struct CapabilitiesFile { /// Maximum length for the description field to prevent memory abuse. const MAX_DESCRIPTION_CHARS: usize = 4096; -/// Maximum serialized size of the parameters schema JSON. -const MAX_PARAMETERS_SCHEMA_BYTES: usize = 64 * 1024; - impl CapabilitiesFile { /// Parse from JSON string. pub fn from_json(json: &str) -> Result { @@ -135,18 +126,6 @@ impl CapabilitiesFile { ); self.description = Some(truncated.to_string()); } - // Drop oversized parameters schema (issue #977) - if let Some(ref params) = self.parameters { - let size = params.to_string().len(); - if size > MAX_PARAMETERS_SCHEMA_BYTES { - tracing::warn!( - "Capabilities parameters schema dropped ({} bytes exceeds {} limit)", - size, - MAX_PARAMETERS_SCHEMA_BYTES, - ); - self.parameters = None; - } - } } /// Merge nested `capabilities` wrapper into top-level fields. @@ -171,7 +150,6 @@ impl CapabilitiesFile { if let Some(inner) = self.capabilities.take() { let inner = inner.resolve_nested_inner(depth + 1); self.description = self.description.or(inner.description); - self.parameters = self.parameters.or(inner.parameters); self.http = self.http.or(inner.http); self.secrets = self.secrets.or(inner.secrets); self.tool_invoke = self.tool_invoke.or(inner.tool_invoke); @@ -1424,26 +1402,12 @@ mod tests { ); } - // ── Tool description and parameters schema ────────────────────────── + // ── Tool description ──────────────────────────────────────────────── #[test] - fn test_parse_description_and_parameters() { + fn test_parse_description() { let json = r#"{ - "description": "Search the web using Brave Search API", - "parameters": { - "type": "object", - "properties": { - "query": { - "type": "string", - "description": "Search query" - }, - "count": { - "type": "integer", - "description": "Number of results" - } - }, - "required": ["query"] - } + "description": "Search the web using Brave Search API" }"#; let caps = CapabilitiesFile::from_json(json).unwrap(); @@ -1451,28 +1415,10 @@ mod tests { caps.description.as_deref(), Some("Search the web using Brave Search API") ); - let params = caps.parameters.unwrap(); - assert_eq!(params["type"], "object"); - assert!(params["properties"]["query"].is_object()); - assert_eq!(params["required"][0], "query"); } #[test] - fn test_parse_description_only() { - let json = r#"{ - "description": "A tool without explicit parameters schema" - }"#; - - let caps = CapabilitiesFile::from_json(json).unwrap(); - assert_eq!( - caps.description.as_deref(), - Some("A tool without explicit parameters schema") - ); - assert!(caps.parameters.is_none()); - } - - #[test] - fn test_parse_without_description_or_parameters() { + fn test_parse_without_description() { let json = r#"{ "http": { "allowlist": [{ "host": "api.example.com" }] @@ -1484,24 +1430,28 @@ mod tests { caps.description.is_none(), "description should be None when not provided" ); - assert!( - caps.parameters.is_none(), - "parameters should be None when not provided" - ); + } + + #[test] + fn test_parameters_field_silently_ignored() { + // Backward compat: old capabilities files with "parameters" still parse. + let json = r#"{ + "description": "A tool", + "parameters": { + "type": "object", + "properties": { "action": { "type": "string" } } + } + }"#; + + let caps = CapabilitiesFile::from_json(json).unwrap(); + assert_eq!(caps.description.as_deref(), Some("A tool")); } #[test] fn test_resolve_nested_description_promoted() { let json = r#"{ "capabilities": { - "description": "Inner tool description", - "parameters": { - "type": "object", - "properties": { - "input": { "type": "string" } - }, - "required": ["input"] - } + "description": "Inner tool description" } }"#; @@ -1511,10 +1461,6 @@ mod tests { Some("Inner tool description"), "description should be promoted from inner capabilities" ); - assert!( - caps.parameters.is_some(), - "parameters should be promoted from inner capabilities" - ); } #[test] @@ -1564,32 +1510,4 @@ mod tests { desc.len() ); } - - /// Regression test for issue #977: oversized parameters schema is dropped. - #[test] - fn test_oversized_parameters_schema_dropped() { - // Build a parameters schema larger than MAX_PARAMETERS_SCHEMA_BYTES - let mut properties = serde_json::Map::new(); - for i in 0..2000 { - properties.insert( - format!("field_{i}"), - serde_json::json!({ - "type": "string", - "description": "x".repeat(50) - }), - ); - } - let schema = serde_json::json!({ - "type": "object", - "properties": properties, - }); - let json = serde_json::json!({ - "parameters": schema, - }); - let caps = CapabilitiesFile::from_json(&json.to_string()).unwrap(); - assert!( - caps.parameters.is_none(), - "oversized parameters schema should be dropped" - ); - } } diff --git a/src/tools/wasm/loader.rs b/src/tools/wasm/loader.rs index 3b5f7a0c..b50fc717 100644 --- a/src/tools/wasm/loader.rs +++ b/src/tools/wasm/loader.rs @@ -123,73 +123,51 @@ impl WasmToolLoader { } let wasm_bytes = fs::read(wasm_path).await?; - // Read capabilities (optional) and extract OAuth refresh config, - // tool description, and parameter schema. - let (capabilities, oauth_refresh, description, schema) = - if let Some(cap_path) = capabilities_path { - if cap_path.exists() { - let cap_bytes = fs::read(cap_path).await?; - let cap_file = CapabilitiesFile::from_bytes(&cap_bytes) - .map_err(|e| WasmLoadError::InvalidCapabilities(e.to_string()))?; - cap_file.validate(name); + // Read capabilities (optional) and extract OAuth refresh config + // and tool description. Parameter schema is auto-derived from the + // WASM module's schema() export (see WasmToolSchemas::compact_schema). + let (capabilities, oauth_refresh, description) = if let Some(cap_path) = capabilities_path { + if cap_path.exists() { + let cap_bytes = fs::read(cap_path).await?; + let cap_file = CapabilitiesFile::from_bytes(&cap_bytes) + .map_err(|e| WasmLoadError::InvalidCapabilities(e.to_string()))?; + cap_file.validate(name); - // Check WIT version compatibility - check_wit_version_compat( - name, - cap_file.wit_version.as_deref(), - crate::tools::wasm::WIT_TOOL_VERSION, - )?; + // Check WIT version compatibility + check_wit_version_compat( + name, + cap_file.wit_version.as_deref(), + crate::tools::wasm::WIT_TOOL_VERSION, + )?; - let caps = cap_file.to_capabilities(); - let oauth = resolve_oauth_refresh_config(&cap_file); - let desc = cap_file.description.clone(); - // Validate parameters schema before accepting it. - let params = cap_file.parameters.clone().and_then(|p| { - let errors = crate::tools::validate_tool_schema(&p, name); - if errors.is_empty() { - Some(p) - } else { - tracing::warn!( - tool = name, - ?errors, - "Invalid parameters schema in capabilities.json, \ - using permissive fallback" - ); - None - } - }); - if desc.is_none() { - tracing::warn!( - tool = name, - path = %cap_path.display(), - "Capabilities file missing \"description\" field; \ - tool will use generic fallback description" - ); - } - if params.is_none() && cap_file.parameters.is_none() { - tracing::warn!( - tool = name, - path = %cap_path.display(), - "Capabilities file missing \"parameters\" field; \ - tool will accept any JSON object (permissive fallback)" - ); - } - (caps, oauth, desc, params) - } else { + let caps = cap_file.to_capabilities(); + let oauth = resolve_oauth_refresh_config(&cap_file); + let desc = cap_file.description.clone(); + if desc.is_none() { tracing::warn!( + tool = name, path = %cap_path.display(), - "Capabilities file not found, using default (no permissions)" + "Capabilities file missing \"description\" field; \ + tool will use generic fallback description" ); - (Capabilities::default(), None, None, None) } + (caps, oauth, desc) } else { tracing::warn!( tool = name, - "No capabilities file for WASM tool; \ - tool will use generic fallback description and accept any JSON object" + path = %cap_path.display(), + "Capabilities file not found, using default (no permissions)" ); - (Capabilities::default(), None, None, None) - }; + (Capabilities::default(), None, None) + } + } else { + tracing::warn!( + tool = name, + "No capabilities file for WASM tool; \ + tool will use generic fallback description" + ); + (Capabilities::default(), None, None) + }; // Register the tool self.registry @@ -200,7 +178,7 @@ impl WasmToolLoader { capabilities, limits: None, description: description.as_deref(), - schema, + schema: None, secrets_store: self.secrets_store.clone(), oauth_refresh, }) diff --git a/src/tools/wasm/wrapper.rs b/src/tools/wasm/wrapper.rs index 679f33ab..33fcedb9 100644 --- a/src/tools/wasm/wrapper.rs +++ b/src/tools/wasm/wrapper.rs @@ -656,12 +656,125 @@ impl WasmToolSchemas { } fn new(discovery: serde_json::Value) -> Self { + let advertised = Self::compact_schema(&discovery); Self { - advertised: Self::permissive_schema(), + advertised, discovery, } } + /// Derive a compact advertised schema from the full discovery schema. + /// + /// Collects properties from top-level `properties` and from + /// `oneOf`/`anyOf`/`allOf` variants. Keeps only properties that are in + /// the top-level `required` array or carry an `enum`/`const` constraint. + /// For properties defined via `const` across multiple variants (e.g. + /// `"action": {"const": "get_repo"}` in each `oneOf` branch), the `const` + /// values are merged into a single `enum` array. + /// + /// Variant-level `required` fields (e.g. `owner`, `repo` required within + /// each `oneOf` variant but not top-level) are intentionally omitted from + /// the compact schema — the LLM can discover them via + /// `tool_info(detail: "schema")`. + /// + /// At most `MAX_COMPACT_PROPERTIES` properties are collected to bound + /// allocations from adversarial schemas. + fn compact_schema(discovery: &serde_json::Value) -> serde_json::Value { + const MAX_COMPACT_PROPERTIES: usize = 100; + + let required: std::collections::HashSet = discovery + .get("required") + .and_then(|r| r.as_array()) + .map(|arr| { + arr.iter() + .filter_map(|v| v.as_str().map(String::from)) + .collect() + }) + .unwrap_or_default(); + + // Collect properties from top-level and oneOf/anyOf/allOf variants. + // For properties with `const` across variants, merge into an `enum`. + let mut all_properties = serde_json::Map::new(); + // Track const values per property to merge into enum. + let mut const_values: std::collections::HashMap> = + std::collections::HashMap::new(); + + if let Some(props) = discovery.get("properties").and_then(|p| p.as_object()) { + for (k, v) in props { + if all_properties.len() >= MAX_COMPACT_PROPERTIES { + break; + } + all_properties.insert(k.clone(), v.clone()); + } + } + for key in ["oneOf", "anyOf", "allOf"] { + if let Some(variants) = discovery.get(key).and_then(|v| v.as_array()) { + for variant in variants { + if let Some(props) = variant.get("properties").and_then(|p| p.as_object()) { + for (k, v) in props { + if all_properties.len() >= MAX_COMPACT_PROPERTIES + && !all_properties.contains_key(k) + { + continue; + } + // Track const values for merging into enum. + if let Some(c) = v.get("const") { + const_values.entry(k.clone()).or_default().push(c.clone()); + } + all_properties.entry(k.clone()).or_insert_with(|| v.clone()); + } + } + } + } + } + + // Merge collected const values into enum arrays. + for (name, values) in &const_values { + if values.len() > 1 + && let Some(prop) = all_properties.get_mut(name) + { + let mut merged = prop.clone(); + if let Some(obj) = merged.as_object_mut() { + obj.remove("const"); + obj.insert("enum".to_string(), serde_json::Value::Array(values.clone())); + } + *prop = merged; + } + } + + if all_properties.is_empty() { + return Self::permissive_schema(); + } + + let kept: serde_json::Map = all_properties + .into_iter() + .filter(|(name, prop)| { + required.contains(name) || prop.get("enum").is_some() || prop.get("const").is_some() + }) + .collect(); + + if kept.is_empty() { + return Self::permissive_schema(); + } + + let kept_required: Vec = required + .iter() + .filter(|name| kept.contains_key(name.as_str())) + .map(|name| serde_json::Value::String(name.clone())) + .collect(); + + let mut result = serde_json::json!({ + "type": "object", + "properties": kept, + "additionalProperties": true, + }); + if !kept_required.is_empty() { + result["required"] = serde_json::Value::Array(kept_required); + } + + result + } + fn with_override(&self, schema: serde_json::Value) -> Self { Self { advertised: schema.clone(), @@ -1655,7 +1768,7 @@ mod tests { } #[tokio::test] - async fn test_advertised_schema_stays_permissive_until_sidecar_override() { + async fn test_advertised_schema_auto_compacted_from_discovery() { let discovery_schema = serde_json::json!({ "type": "object", "properties": { @@ -1675,42 +1788,7 @@ mod tests { wrapper.schemas = super::WasmToolSchemas::new(discovery_schema.clone()); wrapper.description = "Search documents".to_string(); - // Advertised schema stays permissive; discovery holds the typed schema - assert_eq!( - wrapper.parameters_schema(), - serde_json::json!({ - "type": "object", - "properties": {}, - "additionalProperties": true - }) - ); - assert_eq!(wrapper.discovery_schema(), discovery_schema); - - // Raw description is clean — no tool_info hint baked in - assert!(!wrapper.description().contains("tool_info")); - - // But schema() composes the hint at display time when advertised is permissive - let schema = wrapper.schema(); - assert!( - schema.description.contains("tool_info"), - "schema().description should contain tool_info hint: {}", - schema.description - ); - assert!( - schema.description.contains("include_schema: true"), - "hint should mention include_schema: true: {}", - schema.description - ); - - // After sidecar override, both schemas match and hint disappears - let wrapper = wrapper.with_schema(serde_json::json!({ - "type": "object", - "properties": { - "query": { "type": "string" } - }, - "required": ["query"] - })); - + // Advertised schema is auto-compacted: keeps required props, drops optional assert_eq!( wrapper.parameters_schema(), serde_json::json!({ @@ -1718,20 +1796,143 @@ mod tests { "properties": { "query": { "type": "string" } }, - "required": ["query"] + "required": ["query"], + "additionalProperties": true }) ); - assert_eq!(wrapper.discovery_schema(), wrapper.parameters_schema()); + // Discovery retains the full schema + assert_eq!(wrapper.discovery_schema(), discovery_schema); - // With typed schema, schema() should NOT include tool_info hint + // Compacted schema has typed properties, so no tool_info hint needed let schema = wrapper.schema(); assert!( !schema.description.contains("tool_info"), - "schema().description should not contain tool_info hint when typed: {}", + "schema().description should not contain tool_info hint when auto-compacted: {}", schema.description ); } + #[test] + fn test_compact_schema_keeps_required_and_enum_properties() { + let schema = serde_json::json!({ + "type": "object", + "properties": { + "action": { + "type": "string", + "enum": ["list", "get", "create"], + "description": "The operation" + }, + "query": { "type": "string" }, + "limit": { "type": "integer" }, + "format": { + "type": "string", + "enum": ["json", "csv"] + } + }, + "required": ["action"] + }); + + let compacted = super::WasmToolSchemas::compact_schema(&schema); + let props = compacted["properties"].as_object().unwrap(); + + // action: required + enum → kept + assert!(props.contains_key("action")); + // format: has enum → kept + assert!(props.contains_key("format")); + // query: not required, no enum → dropped + assert!(!props.contains_key("query")); + // limit: not required, no enum → dropped + assert!(!props.contains_key("limit")); + // additionalProperties lets the LLM still pass dropped props + assert_eq!(compacted["additionalProperties"], true); + assert_eq!(compacted["required"], serde_json::json!(["action"])); + } + + #[test] + fn test_compact_schema_falls_back_to_permissive_when_empty() { + // No required, no enum → permissive fallback + let schema = serde_json::json!({ + "type": "object", + "properties": { + "query": { "type": "string" }, + "limit": { "type": "integer" } + } + }); + + let compacted = super::WasmToolSchemas::compact_schema(&schema); + assert!(compacted["properties"].as_object().unwrap().is_empty()); + } + + #[test] + fn test_compact_schema_handles_no_properties() { + let schema = serde_json::json!({ "type": "object" }); + let compacted = super::WasmToolSchemas::compact_schema(&schema); + assert!(compacted["properties"].as_object().unwrap().is_empty()); + } + + #[test] + fn test_compact_schema_handles_oneof_variants() { + // GitHub-style schema: oneOf with no top-level properties, const per variant + let schema = serde_json::json!({ + "type": "object", + "required": ["action"], + "oneOf": [ + { + "properties": { + "action": { "const": "get_repo" }, + "owner": { "type": "string" }, + "repo": { "type": "string" } + }, + "required": ["action", "owner", "repo"] + }, + { + "properties": { + "action": { "const": "list_issues" }, + "owner": { "type": "string" }, + "repo": { "type": "string" }, + "state": { "type": "string", "enum": ["open", "closed", "all"] } + }, + "required": ["action", "owner", "repo"] + } + ] + }); + + let compacted = super::WasmToolSchemas::compact_schema(&schema); + let props = compacted["properties"].as_object().unwrap(); + + // action: required + const values merged into enum → kept + let action = &props["action"]; + assert!( + action.get("enum").is_some(), + "action const values should be merged into enum: {action}" + ); + let action_enum = action["enum"].as_array().unwrap(); + assert!( + action_enum.contains(&serde_json::json!("get_repo")), + "enum should contain get_repo" + ); + assert!( + action_enum.contains(&serde_json::json!("list_issues")), + "enum should contain list_issues" + ); + assert!( + action.get("const").is_none(), + "const should be removed after merging into enum" + ); + + // state: has enum → kept + assert!( + props.contains_key("state"), + "state should be kept (has enum)" + ); + // owner/repo: not in top-level required, no enum → intentionally dropped + // (variant-level required is omitted; discoverable via tool_info) + assert!(!props.contains_key("owner"), "owner should be dropped"); + assert!(!props.contains_key("repo"), "repo should be dropped"); + assert_eq!(compacted["additionalProperties"], true); + assert_eq!(compacted["required"], serde_json::json!(["action"])); + } + #[test] fn test_capabilities_default() { let caps = Capabilities::default(); diff --git a/tests/support/test_rig.rs b/tests/support/test_rig.rs index be2b3bb2..19ce5aa0 100644 --- a/tests/support/test_rig.rs +++ b/tests/support/test_rig.rs @@ -701,7 +701,7 @@ impl TestRigBuilder { let wasm_bytes = tokio::fs::read(&spec.wasm_path) .await .unwrap_or_else(|e| panic!("read {}: {e}", spec.wasm_path.display())); - let (capabilities, description, schema) = + let (capabilities, description) = if let Some(cap_path) = &spec.capabilities_path { if cap_path.exists() { let cap_bytes = tokio::fs::read(cap_path) @@ -709,16 +709,12 @@ impl TestRigBuilder { .unwrap_or_else(|e| panic!("read {}: {e}", cap_path.display())); let cap_file = CapabilitiesFile::from_bytes(&cap_bytes) .expect("parse capabilities.json"); - ( - cap_file.to_capabilities(), - cap_file.description.clone(), - cap_file.parameters.clone(), - ) + (cap_file.to_capabilities(), cap_file.description.clone()) } else { - (Capabilities::default(), None, None) + (Capabilities::default(), None) } } else { - (Capabilities::default(), None, None) + (Capabilities::default(), None) }; let prepared = runtime @@ -730,9 +726,6 @@ impl TestRigBuilder { if let Some(desc) = description { wrapper = wrapper.with_description(desc); } - if let Some(s) = schema { - wrapper = wrapper.with_schema(s); - } if let Some(interceptor) = &http_interceptor { wrapper = wrapper.with_http_interceptor(Arc::clone(interceptor)); } diff --git a/tools-src/github/github-tool.capabilities.json b/tools-src/github/github-tool.capabilities.json index 61bbd55f..77370510 100644 --- a/tools-src/github/github-tool.capabilities.json +++ b/tools-src/github/github-tool.capabilities.json @@ -1,6 +1,7 @@ { "version": "0.2.1", "wit_version": "0.3.0", + "description": "Manage GitHub repositories, issues, pull requests, reviews, and workflows. Supports listing, creating, commenting, merging PRs, and triggering GitHub Actions.", "capabilities": { "webhook": { "hmac_secret_name": "github_webhook_secret", diff --git a/tools-src/gmail/gmail-tool.capabilities.json b/tools-src/gmail/gmail-tool.capabilities.json index 2e11d32b..fab6dcb3 100644 --- a/tools-src/gmail/gmail-tool.capabilities.json +++ b/tools-src/gmail/gmail-tool.capabilities.json @@ -1,6 +1,7 @@ { "version": "0.2.0", "wit_version": "0.3.0", + "description": "Read, search, send, draft, and reply to emails via Gmail. Supports Gmail search query syntax (is:unread, from:, subject:, after:, etc.).", "http": { "allowlist": [ { @@ -53,7 +54,7 @@ }, { "name": "google_oauth_client_secret", - "prompt": "Google OAuth Client Secret" + "prompt": "Google OAuth Client Secret (from console.cloud.google.com/apis/credentials)" } ] } diff --git a/tools-src/google-calendar/google-calendar-tool.capabilities.json b/tools-src/google-calendar/google-calendar-tool.capabilities.json index 15e756ae..f9869288 100644 --- a/tools-src/google-calendar/google-calendar-tool.capabilities.json +++ b/tools-src/google-calendar/google-calendar-tool.capabilities.json @@ -1,6 +1,7 @@ { "version": "0.2.0", "wit_version": "0.3.0", + "description": "View, create, update, and delete Google Calendar events. Supports timed events, all-day events, attendees, locations, and free text search.", "http": { "allowlist": [ { @@ -52,7 +53,7 @@ }, { "name": "google_oauth_client_secret", - "prompt": "Google OAuth Client Secret" + "prompt": "Google OAuth Client Secret (from console.cloud.google.com/apis/credentials)" } ] } diff --git a/tools-src/google-docs/google-docs-tool.capabilities.json b/tools-src/google-docs/google-docs-tool.capabilities.json index 7a365c1d..2a34ce94 100644 --- a/tools-src/google-docs/google-docs-tool.capabilities.json +++ b/tools-src/google-docs/google-docs-tool.capabilities.json @@ -1,6 +1,7 @@ { "version": "0.2.0", "wit_version": "0.3.0", + "description": "Create, read, edit, and format Google Docs documents. Supports text insert/delete/replace, formatting (bold, italic, font, color, size), paragraph styling, tables, and lists.", "http": { "allowlist": [ { @@ -52,7 +53,7 @@ }, { "name": "google_oauth_client_secret", - "prompt": "Google OAuth Client Secret" + "prompt": "Google OAuth Client Secret (from console.cloud.google.com/apis/credentials)" } ] } diff --git a/tools-src/google-drive/google-drive-tool.capabilities.json b/tools-src/google-drive/google-drive-tool.capabilities.json index 53667933..a5e60125 100644 --- a/tools-src/google-drive/google-drive-tool.capabilities.json +++ b/tools-src/google-drive/google-drive-tool.capabilities.json @@ -1,6 +1,7 @@ { "version": "0.2.0", "wit_version": "0.3.0", + "description": "Search, access, upload, share, and organize files and folders in Google Drive. Supports personal drives and shared (organizational) drives.", "http": { "allowlist": [ { @@ -57,7 +58,7 @@ }, { "name": "google_oauth_client_secret", - "prompt": "Google OAuth Client Secret" + "prompt": "Google OAuth Client Secret (from console.cloud.google.com/apis/credentials)" } ] } diff --git a/tools-src/google-sheets/google-sheets-tool.capabilities.json b/tools-src/google-sheets/google-sheets-tool.capabilities.json index 624c4381..ceadb8f1 100644 --- a/tools-src/google-sheets/google-sheets-tool.capabilities.json +++ b/tools-src/google-sheets/google-sheets-tool.capabilities.json @@ -1,6 +1,7 @@ { "version": "0.2.0", "wit_version": "0.3.0", + "description": "Create, read, write, and format Google Sheets spreadsheets. Supports cell operations using A1 notation, sheet (tab) management, and cell formatting.", "http": { "allowlist": [ { @@ -52,7 +53,7 @@ }, { "name": "google_oauth_client_secret", - "prompt": "Google OAuth Client Secret" + "prompt": "Google OAuth Client Secret (from console.cloud.google.com/apis/credentials)" } ] } diff --git a/tools-src/google-slides/google-slides-tool.capabilities.json b/tools-src/google-slides/google-slides-tool.capabilities.json index 17334bc0..2d3c378e 100644 --- a/tools-src/google-slides/google-slides-tool.capabilities.json +++ b/tools-src/google-slides/google-slides-tool.capabilities.json @@ -1,6 +1,7 @@ { "version": "0.2.0", "wit_version": "0.3.0", + "description": "Create, read, edit, and format Google Slides presentations. Supports slide management, text operations, shapes, images, text formatting, and paragraph alignment.", "http": { "allowlist": [ { @@ -52,7 +53,7 @@ }, { "name": "google_oauth_client_secret", - "prompt": "Google OAuth Client Secret" + "prompt": "Google OAuth Client Secret (from console.cloud.google.com/apis/credentials)" } ] } diff --git a/tools-src/llm-context/llm-context-tool.capabilities.json b/tools-src/llm-context/llm-context-tool.capabilities.json index 72061eaa..5ea3fe7d 100644 --- a/tools-src/llm-context/llm-context-tool.capabilities.json +++ b/tools-src/llm-context/llm-context-tool.capabilities.json @@ -1,6 +1,7 @@ { "version": "0.1.0", "wit_version": "0.3.0", + "description": "Fetch pre-extracted web content from Brave Search for grounding LLM answers. Returns actual page content (text chunks, tables, code) relevant to the query, ready for RAG or fact-checking.", "capabilities": { "http": { "allowlist": [ diff --git a/tools-src/slack/slack-tool.capabilities.json b/tools-src/slack/slack-tool.capabilities.json index 8b9060d7..5ac9f49c 100644 --- a/tools-src/slack/slack-tool.capabilities.json +++ b/tools-src/slack/slack-tool.capabilities.json @@ -1,6 +1,7 @@ { "version": "0.2.0", "wit_version": "0.3.0", + "description": "Send messages, list channels, read history, add reactions, and get user information in Slack.", "http": { "allowlist": [ { @@ -57,7 +58,7 @@ }, { "name": "slack_oauth_client_secret", - "prompt": "Slack OAuth Client Secret" + "prompt": "Slack OAuth Client Secret (from api.slack.com/apps > Basic Information)" } ] } diff --git a/tools-src/telegram/telegram-tool.capabilities.json b/tools-src/telegram/telegram-tool.capabilities.json index 665baedd..02b451ee 100644 --- a/tools-src/telegram/telegram-tool.capabilities.json +++ b/tools-src/telegram/telegram-tool.capabilities.json @@ -1,6 +1,7 @@ { "version": "0.2.0", "wit_version": "0.3.0", + "description": "Read and send messages from a Telegram user account. Supports contacts, chat history, message search, sending, forwarding, and deletion via encrypted MTProto.", "http": { "allowlist": [ { @@ -35,7 +36,7 @@ }, { "name": "telegram_api_hash", - "prompt": "Telegram API Hash" + "prompt": "Telegram API Hash (from my.telegram.org/apps — alphanumeric string)" } ] } diff --git a/tools-src/web-search/web-search-tool.capabilities.json b/tools-src/web-search/web-search-tool.capabilities.json index 9c2559ab..26c48b53 100644 --- a/tools-src/web-search/web-search-tool.capabilities.json +++ b/tools-src/web-search/web-search-tool.capabilities.json @@ -2,40 +2,6 @@ "version": "0.2.0", "wit_version": "0.3.0", "description": "Search the web using Brave Search. Returns titles, URLs, descriptions, and publication dates for matching web pages. Supports filtering by country, language, and freshness. Authentication is handled via the 'brave_api_key' secret injected by the host.", - "parameters": { - "type": "object", - "properties": { - "query": { - "type": "string", - "description": "The search query to look up on the web" - }, - "count": { - "type": "integer", - "description": "Number of results to return (1-20, default 5)", - "minimum": 1, - "maximum": 20, - "default": 5 - }, - "country": { - "type": "string", - "description": "2-letter uppercase country code to bias results (e.g. 'US', 'DE', 'JP')" - }, - "search_lang": { - "type": "string", - "description": "2-letter lowercase language code for search results (e.g. 'en', 'de', 'fr')" - }, - "ui_lang": { - "type": "string", - "description": "Locale in language-region format (e.g. 'en-US', 'de-DE')" - }, - "freshness": { - "type": "string", - "description": "Filter by discovery time: 'pd' (past day), 'pw' (past week), 'pm' (past month), 'py' (past year), or date range 'YYYY-MM-DDtoYYYY-MM-DD'" - } - }, - "required": ["query"], - "additionalProperties": false - }, "capabilities": { "http": { "allowlist": [ From 5847479fd851726e7e1e848b45bcf48a195f9aa9 Mon Sep 17 00:00:00 2001 From: Illia Polosukhin Date: Mon, 23 Mar 2026 22:24:26 -0700 Subject: [PATCH 05/20] fix(agent): persist /model selection to .env, TOML, and DB (#1581) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix(agent): persist /model selection to .env, TOML, and DB The /model command only wrote selected_model to the DB and config.toml, but env vars from ~/.ironclaw/.env (e.g. NEARAI_MODEL) have the highest priority in LlmConfig::resolve_model(). The .env value was never updated, so it always shadowed the new model on restart. Now persist_selected_model updates all three persistence layers: 1. The backend-specific model env var in ~/.ironclaw/.env (only if the var already exists, to avoid injecting new vars) 2. The config.toml file (created if absent, since TOML > DB priority) 3. The DB settings table (for completeness) Also adds diagnostic logging when the DB store is unavailable. Co-Authored-By: Claude Opus 4.6 (1M context) * fix(agent): address PR review — backend from deps, exact .env match Review feedback: - Use resolved llm_backend from AgentDeps instead of re-reading from disk/env (fixes DB-only backend detection, eliminates redundant I/O) - Match .env var with exact "KEY=" prefix and skip commented lines (prevents false matches on NEARAI_MODEL_VERSION etc.) - TOML is now loaded once (no double-read for backend + model update) Co-Authored-By: Claude Opus 4.6 (1M context) --------- Co-authored-by: Claude Opus 4.6 (1M context) --- src/agent/agent_loop.rs | 3 + src/agent/commands.rs | 52 ++++++++++- src/agent/dispatcher.rs | 3 + src/main.rs | 1 + src/settings.rs | 108 ++++++++++++++++++++++ src/testing/mod.rs | 1 + tests/e2e_telegram_message_routing.rs | 1 + tests/support/gateway_workflow_harness.rs | 1 + tests/support/test_rig.rs | 1 + 9 files changed, 168 insertions(+), 3 deletions(-) diff --git a/src/agent/agent_loop.rs b/src/agent/agent_loop.rs index ee91ea9a..3ab369b1 100644 --- a/src/agent/agent_loop.rs +++ b/src/agent/agent_loop.rs @@ -169,6 +169,9 @@ pub struct AgentDeps { pub sandbox_readiness: crate::agent::routine_engine::SandboxReadiness, /// Software builder for self-repair tool rebuilding. pub builder: Option>, + /// Resolved LLM backend identifier (e.g., "nearai", "openai", "groq"). + /// Used by `/model` persistence to determine which env var to update. + pub llm_backend: String, } /// The main agent that coordinates all components. diff --git a/src/agent/commands.rs b/src/agent/commands.rs index 75c99359..b6aff3c0 100644 --- a/src/agent/commands.rs +++ b/src/agent/commands.rs @@ -841,12 +841,50 @@ impl Agent { .await { tracing::warn!("Failed to persist model to DB: {}", e); + } else { + tracing::debug!("Persisted selected_model to DB: {}", model); } + } else { + tracing::warn!("No database store available — model choice will not persist to DB"); } - // 2. Update TOML config file if it exists (sync I/O in spawn_blocking). + // 2. 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 || { + // 2a. Update the backend-specific model env var in ~/.ironclaw/.env. + // + // Env vars have the HIGHEST priority in LlmConfig::resolve_model() + // (env var > TOML > DB > default). If the .env file has e.g. + // NEARAI_MODEL=old-model, it shadows everything else. We must + // update this var or the /model change is invisible on restart. + let registry = crate::llm::ProviderRegistry::load(); + let model_env = registry.model_env_var(&backend); + let env_var_prefix = format!("{}=", model_env); + + // Only update the .env file if the var is actually set there + // (avoid injecting new vars the user never configured). + let env_path = crate::bootstrap::ironclaw_env_path(); + let env_has_var = std::fs::read_to_string(&env_path) + .ok() + .is_some_and(|content| { + content.lines().any(|line| { + let trimmed = line.trim_start(); + !trimmed.starts_with('#') && trimmed.starts_with(&env_var_prefix) + }) + }); + if env_has_var { + if let Err(e) = crate::bootstrap::upsert_bootstrap_var(model_env, &model_owned) { + tracing::warn!("Failed to update {} in .env: {}", model_env, e); + } else { + tracing::debug!("Updated {} in .env to {}", model_env, model_owned); + } + } + + // 2b. Update (or create) the TOML config file. + // + // The TOML overlay has higher priority than DB settings on + // startup, so it MUST stay in sync with the DB. let toml_path = crate::settings::Settings::default_toml_path(); match crate::settings::Settings::load_toml(&toml_path) { Ok(Some(mut settings)) => { @@ -856,7 +894,15 @@ impl Agent { } } Ok(None) => { - // No config file on disk; nothing to update. + // No config file yet — create one so the model choice + // survives restarts even when the DB is unavailable. + let settings = crate::settings::Settings { + selected_model: Some(model_owned), + ..Default::default() + }; + if let Err(e) = settings.save_toml(&toml_path) { + tracing::warn!("Failed to create config.toml for model persistence: {}", e); + } } Err(e) => { tracing::warn!("Failed to load config.toml for model persistence: {}", e); @@ -865,7 +911,7 @@ impl Agent { }) .await { - tracing::warn!("Model TOML persistence task failed: {}", e); + tracing::warn!("Model persistence task failed: {}", e); } } } diff --git a/src/agent/dispatcher.rs b/src/agent/dispatcher.rs index 5d39866b..a195458d 100644 --- a/src/agent/dispatcher.rs +++ b/src/agent/dispatcher.rs @@ -1233,6 +1233,7 @@ mod tests { document_extraction: None, sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, builder: None, + llm_backend: "nearai".to_string(), }; Agent::new( @@ -2100,6 +2101,7 @@ mod tests { document_extraction: None, sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, builder: None, + llm_backend: "nearai".to_string(), }; Agent::new( @@ -2220,6 +2222,7 @@ mod tests { document_extraction: None, sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, builder: None, + llm_backend: "nearai".to_string(), }; Agent::new( diff --git a/src/main.rs b/src/main.rs index dd224f47..eab01264 100644 --- a/src/main.rs +++ b/src/main.rs @@ -912,6 +912,7 @@ async fn async_main() -> anyhow::Result<()> { ironclaw::agent::routine_engine::SandboxReadiness::DockerUnavailable }, builder: components.builder, + llm_backend: config.llm.backend.clone(), }; let channels_for_warnings = Arc::clone(&channels); diff --git a/src/settings.rs b/src/settings.rs index 2340f0d2..1bb1a8f7 100644 --- a/src/settings.rs +++ b/src/settings.rs @@ -1297,6 +1297,92 @@ mod tests { assert_eq!(loaded.heartbeat.interval_secs, 900); } + /// Regression: /model writes a single key ("selected_model") to the DB via + /// set_setting(). On restart, get_all_settings() returns ALL keys including + /// wizard-written defaults. The single-key update must survive the full + /// from_db_map() round trip. + #[test] + fn db_single_key_model_update_survives_roundtrip() { + // Step 1: Wizard writes full settings to DB (including selected_model + // from initial setup). + let wizard_settings = Settings { + llm_backend: Some("nearai".to_string()), + selected_model: Some("old-wizard-model".to_string()), + ..Default::default() + }; + let mut db: std::collections::HashMap = + wizard_settings.to_db_map(); + + // Step 2: User runs /model new-model — persist_selected_model writes + // a single key, overwriting the wizard value. + db.insert( + "selected_model".to_string(), + serde_json::Value::String("new-model".to_string()), + ); + + // Step 3: On restart, from_db_map() rebuilds Settings from the full + // DB map. + let restored = Settings::from_db_map(&db); + assert_eq!( + restored.selected_model, + Some("new-model".to_string()), + "/model change must survive DB round trip" + ); + } + + /// Regression: TOML overlay must not clobber a DB-persisted selected_model + /// when the TOML file matches the DB. This is the normal case after /model + /// successfully writes to both DB and TOML. + #[test] + fn toml_overlay_preserves_matching_model() { + // DB settings with new model from /model command. + let mut db_settings = Settings { + llm_backend: Some("nearai".to_string()), + selected_model: Some("new-model".to_string()), + ..Default::default() + }; + + // TOML also updated by /model command to the same value. + let toml_settings = Settings { + selected_model: Some("new-model".to_string()), + ..Default::default() + }; + + db_settings.merge_from(&toml_settings); + assert_eq!( + db_settings.selected_model, + Some("new-model".to_string()), + "TOML overlay must not clobber matching model" + ); + } + + /// Regression: when /model updates DB but TOML write fails, a stale TOML + /// file would overwrite the DB value. This test documents the priority: + /// TOML > DB (by design). persist_selected_model MUST update the TOML. + #[test] + fn stale_toml_overwrites_db_model() { + // DB has the new model from /model. + let mut db_settings = Settings { + selected_model: Some("new-model".to_string()), + ..Default::default() + }; + + // TOML still has the old model (write failed or was not attempted). + let stale_toml = Settings { + selected_model: Some("old-model".to_string()), + ..Default::default() + }; + + db_settings.merge_from(&stale_toml); + // This documents the current priority: TOML wins over DB. + // The fix in persist_selected_model ensures TOML is always updated. + assert_eq!( + db_settings.selected_model, + Some("old-model".to_string()), + "TOML overlay has higher priority than DB (by design)" + ); + } + /// Regression test: /model command must persist selected_model to TOML config. /// Prior to the fix, `set_model()` only changed the in-memory provider and the /// choice was lost on restart. @@ -1322,6 +1408,28 @@ mod tests { assert_eq!(reloaded.selected_model, Some("new-model".to_string())); } + /// Regression: /model must create config.toml when it doesn't exist, so the + /// model survives restarts. Previously the Ok(None) case was a no-op. + #[test] + fn toml_created_when_missing_for_model_persist() { + let dir = tempfile::tempdir().unwrap(); + let path = dir.path().join("config.toml"); + + // No config.toml yet (fresh install, no wizard). + assert!(Settings::load_toml(&path).unwrap().is_none()); + + // Simulate what persist_selected_model now does for the Ok(None) case. + let settings = Settings { + selected_model: Some("new-model".to_string()), + ..Default::default() + }; + settings.save_toml(&path).unwrap(); + + // Verify the model survived. + let loaded = Settings::load_toml(&path).unwrap().unwrap(); + assert_eq!(loaded.selected_model, Some("new-model".to_string())); + } + #[test] fn toml_missing_file_returns_none() { let result = Settings::load_toml(std::path::Path::new("/tmp/nonexistent_config.toml")); diff --git a/src/testing/mod.rs b/src/testing/mod.rs index a633e91c..e580b169 100644 --- a/src/testing/mod.rs +++ b/src/testing/mod.rs @@ -563,6 +563,7 @@ impl TestHarnessBuilder { document_extraction: None, sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, builder: None, + llm_backend: "nearai".to_string(), }; TestHarness { diff --git a/tests/e2e_telegram_message_routing.rs b/tests/e2e_telegram_message_routing.rs index fe9a9b04..ead164eb 100644 --- a/tests/e2e_telegram_message_routing.rs +++ b/tests/e2e_telegram_message_routing.rs @@ -200,6 +200,7 @@ mod tests { document_extraction: None, sandbox_readiness: ironclaw::agent::SandboxReadiness::DisabledByConfig, builder: None, + llm_backend: "nearai".to_string(), }; let gateway = Arc::new(TestChannel::new()); diff --git a/tests/support/gateway_workflow_harness.rs b/tests/support/gateway_workflow_harness.rs index 7f9d3dff..e4620f70 100644 --- a/tests/support/gateway_workflow_harness.rs +++ b/tests/support/gateway_workflow_harness.rs @@ -264,6 +264,7 @@ impl GatewayWorkflowHarness { document_extraction: None, sandbox_readiness: ironclaw::agent::SandboxReadiness::DisabledByConfig, builder: None, + llm_backend: "nearai".to_string(), }, channels, None, diff --git a/tests/support/test_rig.rs b/tests/support/test_rig.rs index 19ce5aa0..624bb054 100644 --- a/tests/support/test_rig.rs +++ b/tests/support/test_rig.rs @@ -761,6 +761,7 @@ impl TestRigBuilder { document_extraction: None, sandbox_readiness: ironclaw::agent::SandboxReadiness::Available, // tests don't use real Docker builder: None, + llm_backend: "nearai".to_string(), }; // 7. Create TestChannel and ChannelManager. From fb3548956bf6b1cc4fb31cb753b4fa24a7cfec68 Mon Sep 17 00:00:00 2001 From: nearfamiliarcow Date: Tue, 24 Mar 2026 03:46:22 -0400 Subject: [PATCH 06/20] fix(tunnel): managed tunnels target wrong port and die from SIGPIPE (#1093) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix(tunnel): target webhook server port instead of gateway port start_managed_tunnel() always used the gateway port (3000) for the tunnel target. Webhook routes live on the webhook server (HTTP_PORT, default 8080), not the gateway. The old code never read config.channels.http — no configuration could work around this. Extracts resolve_tunnel_target() with regression tests. * fix(tunnel): prevent SIGPIPE and fix default port fallback Two fixes for managed tunnel subprocess lifetime: 1. After extracting the public URL from stdout/stderr, the pipe reader was dropped (Rust ownership). The tunnel binary's next log write hit the closed pipe and got SIGPIPE — killing it silently. Fix: drain pipes in background tasks stored in TunnelProcess. Storing without reading isn't enough — the OS pipe buffer fills up and the process blocks instead. 2. When neither HTTP_PORT nor gateway is configured, the tunnel fell back to 127.0.0.1:3000. But the webhook server defaults to 0.0.0.0:8080 in this case. Now the tunnel matches that fallback. Affects ngrok (stdout), cloudflare (stderr), and custom (stdout). Tailscale uses a daemon and is not affected by SIGPIPE. * fix(tunnel): simplify drain loops and suppress CI false positives Simplify `while let Ok(Ok(Some(line)))` drain pattern to `while let Ok(Some(line))` — the extra Ok wrapper was unnecessary. Add `// safety: test-only` to assert_eq! lines in test module to suppress the "No panics in production code" CI check which greps the diff without understanding Rust's #[cfg(test)] module boundaries. --------- Co-authored-by: firat.sertgoz --- src/tunnel/cloudflare.rs | 28 +++++--- src/tunnel/custom.rs | 23 +++++-- src/tunnel/mod.rs | 136 ++++++++++++++++++++++++++++++++++----- src/tunnel/ngrok.rs | 31 ++++++--- 4 files changed, 179 insertions(+), 39 deletions(-) diff --git a/src/tunnel/cloudflare.rs b/src/tunnel/cloudflare.rs index 2c0ceb2a..9cc51bd4 100644 --- a/src/tunnel/cloudflare.rs +++ b/src/tunnel/cloudflare.rs @@ -111,10 +111,23 @@ impl Tunnel for CloudflareTunnel { } } - // Drain stderr in the background to prevent SIGPIPE/buffer stalls. - tokio::spawn(async move { while let Ok(Some(_)) = reader.next_line().await {} }); + if let Ok(mut guard) = self.url.write() { + *guard = Some(public_url.clone()); + } - // Drain stdout silently. + // We took ownership of cloudflared's stderr pipe above to parse the URL. + // cloudflared continues writing logs for its entire lifetime. If we drop + // the reader, the pipe closes and cloudflared gets SIGPIPE on its next + // write. We can't just store the reader without reading — the OS pipe + // buffer fills up and cloudflared blocks. So we drain it in a background + // task. The task exits naturally when cloudflared is killed (EOF). + let drain_handle = tokio::spawn(async move { + while let Ok(Some(line)) = reader.next_line().await { + tracing::trace!("cloudflared: {line}"); + } + }); + + // Drain stdout silently to prevent SIGPIPE/buffer stalls. if let Some(stdout) = stdout { tokio::spawn(async move { let mut out_reader = tokio::io::BufReader::new(stdout).lines(); @@ -122,12 +135,11 @@ impl Tunnel for CloudflareTunnel { }); } - if let Ok(mut guard) = self.url.write() { - *guard = Some(public_url.clone()); - } - let mut guard = self.proc.lock().await; - *guard = Some(TunnelProcess { child }); + *guard = Some(TunnelProcess { + child, + _pipe_drain: Some(drain_handle), + }); Ok(public_url) } diff --git a/src/tunnel/custom.rs b/src/tunnel/custom.rs index 9a2be403..2fffa264 100644 --- a/src/tunnel/custom.rs +++ b/src/tunnel/custom.rs @@ -73,6 +73,7 @@ impl Tunnel for CustomTunnel { let stderr = child.stderr.take(); let mut public_url = format!("http://{local_host}:{local_port}"); + let mut drain_handle: Option> = None; if self.url_pattern.is_some() && let Some(stdout) = stdout @@ -103,17 +104,26 @@ impl Tunnel for CustomTunnel { Err(_) => {} } } - // Drain remaining stdout to prevent SIGPIPE/buffer stalls. - tokio::spawn(async move { while let Ok(Some(_)) = reader.next_line().await {} }); + // We took ownership of the process's stdout pipe above to parse the + // URL. The process may continue writing to stdout for its lifetime. + // If we drop the reader, the pipe closes and the process gets SIGPIPE. + // We can't just store the reader without reading — the OS pipe buffer + // fills up and the process blocks. So we drain it in a background task. + // The task exits naturally when the process is killed (EOF). + drain_handle = Some(tokio::spawn(async move { + while let Ok(Some(line)) = reader.next_line().await { + tracing::trace!("custom-tunnel: {line}"); + } + })); } else if let Some(stdout) = stdout { - // No url_pattern: still drain stdout to prevent pipe stalls. + // No url_pattern: still drain stdout to prevent SIGPIPE/buffer stalls. tokio::spawn(async move { let mut reader = tokio::io::BufReader::new(stdout).lines(); while let Ok(Some(_)) = reader.next_line().await {} }); } - // Drain stderr silently. + // Drain stderr to prevent SIGPIPE/buffer stalls. if let Some(stderr) = stderr { tokio::spawn(async move { let mut reader = tokio::io::BufReader::new(stderr).lines(); @@ -126,7 +136,10 @@ impl Tunnel for CustomTunnel { } let mut guard = self.proc.lock().await; - *guard = Some(TunnelProcess { child }); + *guard = Some(TunnelProcess { + child, + _pipe_drain: drain_handle, + }); Ok(public_url) } diff --git a/src/tunnel/mod.rs b/src/tunnel/mod.rs index fa028834..a6869eda 100644 --- a/src/tunnel/mod.rs +++ b/src/tunnel/mod.rs @@ -66,6 +66,10 @@ pub trait Tunnel: Send + Sync { /// Wraps a spawned tunnel child process. pub(crate) struct TunnelProcess { pub child: tokio::process::Child, + /// Background task that drains the process's output pipe (stdout or stderr). + /// Must stay alive or the process dies (SIGPIPE from closed pipe) or hangs + /// (OS pipe buffer fills up, blocking the process's writes). + pub _pipe_drain: Option>, } pub(crate) type SharedProcess = Arc>>; @@ -182,6 +186,22 @@ pub fn create_tunnel(config: &TunnelProviderConfig) -> Result (&str, u16) { + if let Some(ref http) = channels.http { + return (http.host.as_str(), http.port); + } + if let Some(ref gw) = channels.gateway { + return (gw.host.as_str(), gw.port); + } + ("0.0.0.0", 8080) +} + /// Start a managed tunnel if configured and no static URL is already set. /// /// Returns the (potentially mutated) config with `tunnel.public_url` set, @@ -201,28 +221,17 @@ pub async fn start_managed_tunnel( return (config, None); }; - let gateway_port = config - .channels - .gateway - .as_ref() - .map(|g| g.port) - .unwrap_or(3000); - let gateway_host = config - .channels - .gateway - .as_ref() - .map(|g| g.host.as_str()) - .unwrap_or("127.0.0.1"); + let (tunnel_host, tunnel_port) = resolve_tunnel_target(&config.channels); match create_tunnel(provider_config) { Ok(Some(tunnel)) => { tracing::debug!( "Starting {} tunnel on {}:{}...", tunnel.name(), - gateway_host, - gateway_port + tunnel_host, + tunnel_port ); - match tunnel.start(gateway_host, gateway_port).await { + match tunnel.start(tunnel_host, tunnel_port).await { Ok(url) => { tracing::debug!("Tunnel started: {}", url); config.tunnel.public_url = Some(url); @@ -383,10 +392,105 @@ mod tests { { let mut guard = proc.lock().await; - *guard = Some(TunnelProcess { child }); + *guard = Some(TunnelProcess { + child, + _pipe_drain: None, + }); } kill_shared(&proc).await.unwrap(); assert!(proc.lock().await.is_none()); } + + // ── Port selection regression tests ────────────────────────────── + + fn base_channels() -> crate::config::ChannelsConfig { + crate::config::ChannelsConfig { + cli: crate::config::CliConfig { enabled: false }, + http: None, + gateway: None, + signal: None, + wasm_channels_dir: std::env::temp_dir().join("ironclaw-test-channels"), + wasm_channels_enabled: false, + wasm_channel_owner_ids: std::collections::HashMap::new(), + } + } + + fn channels_with_http(host: &str, port: u16) -> crate::config::ChannelsConfig { + let mut c = base_channels(); + c.http = Some(crate::config::HttpConfig { + host: host.to_string(), + port, + webhook_secret: None, + user_id: "test".to_string(), + }); + c.gateway = Some(crate::config::GatewayConfig { + host: "127.0.0.1".to_string(), + port: 3000, + auth_token: None, + user_id: "test".to_string(), + }); + c + } + + fn channels_gateway_only(host: &str, port: u16) -> crate::config::ChannelsConfig { + let mut c = base_channels(); + c.gateway = Some(crate::config::GatewayConfig { + host: host.to_string(), + port, + auth_token: None, + user_id: "test".to_string(), + }); + c + } + + fn channels_neither() -> crate::config::ChannelsConfig { + base_channels() + } + + #[test] + fn tunnel_target_prefers_http_port() { + let channels = channels_with_http("0.0.0.0", 8080); + let (host, port) = resolve_tunnel_target(&channels); + assert_eq!(host, "0.0.0.0"); // safety: test-only + assert_eq!(port, 8080); // safety: test-only + } + + #[test] + fn tunnel_target_falls_back_to_gateway() { + let channels = channels_gateway_only("10.0.0.1", 4000); + let (host, port) = resolve_tunnel_target(&channels); + assert_eq!(host, "10.0.0.1"); // safety: test-only + assert_eq!(port, 4000); // safety: test-only + } + + #[test] + fn tunnel_target_defaults_to_webhook_fallback() { + let channels = channels_neither(); + let (host, port) = resolve_tunnel_target(&channels); + // Matches the webhook server's hardcoded fallback in main.rs + assert_eq!(host, "0.0.0.0"); // safety: test-only + assert_eq!(port, 8080); // safety: test-only + } + + #[test] + fn tunnel_target_http_takes_priority_over_gateway() { + let channels = channels_with_http("192.168.1.1", 9090); + let (host, port) = resolve_tunnel_target(&channels); + // Should use HTTP config, not gateway's 127.0.0.1:3000 + assert_eq!(host, "192.168.1.1"); // safety: test-only + assert_eq!(port, 9090); // safety: test-only + } + + #[test] + fn tunnel_target_no_http_no_gateway_matches_webhook_fallback() { + // When HTTP_PORT is not set and gateway is not configured (e.g. WASM + // channels exist but no explicit HTTP config), the webhook server in + // main.rs binds to 0.0.0.0:8080 as a hardcoded fallback. The tunnel + // must target the same address so webhook traffic reaches the right + // server. + let channels = channels_neither(); + let (host, port) = resolve_tunnel_target(&channels); + assert_eq!((host, port), ("0.0.0.0", 8080)); // safety: test-only + } } diff --git a/src/tunnel/ngrok.rs b/src/tunnel/ngrok.rs index 80a5cc46..66642e3b 100644 --- a/src/tunnel/ngrok.rs +++ b/src/tunnel/ngrok.rs @@ -110,12 +110,24 @@ impl Tunnel for NgrokTunnel { } } - // Drain stdout silently — ngrok only emits low-level connection events - // to stdout; the pipe must be consumed to prevent SIGPIPE/buffer stalls. - tokio::spawn(async move { while let Ok(Some(_)) = reader.next_line().await {} }); + if let Ok(mut guard) = self.url.write() { + *guard = Some(public_url.clone()); + } - // Drain stderr silently — with --log stdout all meaningful output goes - // to stdout; stderr only needs to be consumed to prevent pipe stalls. + // We took ownership of ngrok's stdout pipe above to parse the URL. + // ngrok continues writing logs to stdout for its entire lifetime. + // If we drop the reader, the pipe closes and ngrok gets SIGPIPE on + // its next write → process dies. We can't just store the reader + // without reading — the OS pipe buffer (~64KB) fills up and ngrok + // blocks. So we drain it in a background task. The task exits + // naturally when ngrok is killed (EOF on the pipe). + let drain_handle = tokio::spawn(async move { + while let Ok(Some(line)) = reader.next_line().await { + tracing::trace!("ngrok: {line}"); + } + }); + + // Drain stderr silently to prevent SIGPIPE/buffer stalls. if let Some(stderr) = stderr { tokio::spawn(async move { let mut err_reader = tokio::io::BufReader::new(stderr).lines(); @@ -123,12 +135,11 @@ impl Tunnel for NgrokTunnel { }); } - if let Ok(mut guard) = self.url.write() { - *guard = Some(public_url.clone()); - } - let mut guard = self.proc.lock().await; - *guard = Some(TunnelProcess { child }); + *guard = Some(TunnelProcess { + child, + _pipe_drain: Some(drain_handle), + }); Ok(public_url) } From 01678be61d6a95ed3051772f6fe128b63c187b1e Mon Sep 17 00:00:00 2001 From: Zaki Manian Date: Tue, 24 Mar 2026 02:41:33 -0700 Subject: [PATCH 07/20] fix(routines): normalize status display across web and CLI (#1469) * fix(routines): normalize status display across web and CLI surfaces (#1319) - Use Display (lowercase) instead of Debug (PascalCase) for RunStatus serialization in web handler - Update JavaScript status class mapping to match lowercase values from the API - Enrich CLI `routines list` to show running/attention states by querying last run status [skip-regression-check] Co-Authored-By: Claude Opus 4.6 (1M context) * fix(routines): address review -- batch last-run query, consistent status, simplify ternary (#1319) - Parallelize last-run lookups with join_all to avoid N+1 sequential queries - Normalize status in /api/routines/{id}/runs handler to match lowercase convention - Remove redundant 'running' check in app.js runStatusClass logic Co-Authored-By: Claude Opus 4.6 (1M context) * fix(db): replace N+1 last-run-status queries with batch method The CLI routines list was firing a separate list_routine_runs query per routine to determine each one's last run status. For large routine sets this overwhelms the connection pool. Add batch_get_last_run_status to the Database trait with implementations for both PostgreSQL (DISTINCT ON + ORDER BY) and libSQL (correlated subquery + in-memory filter). Update the CLI to call the batch method once instead of N times. Co-Authored-By: Claude Opus 4.6 (1M context) * style: cargo fmt https://claude.ai/code/session_01Va9wwvATNWFAx35GG7Zek7 --------- Co-authored-by: Claude Opus 4.6 (1M context) --- src/channels/web/handlers/routines.rs | 4 +- src/channels/web/server.rs | 2 +- src/channels/web/static/app.js | 6 +- src/cli/routines.rs | 27 ++-- src/db/libsql/routines.rs | 50 +++++++ src/db/mod.rs | 9 ++ src/db/postgres.rs | 8 ++ src/history/store.rs | 34 +++++ tests/batch_last_run_status_tests.rs | 191 ++++++++++++++++++++++++++ 9 files changed, 317 insertions(+), 14 deletions(-) create mode 100644 tests/batch_last_run_status_tests.rs diff --git a/src/channels/web/handlers/routines.rs b/src/channels/web/handlers/routines.rs index d27adca2..fc56b187 100644 --- a/src/channels/web/handlers/routines.rs +++ b/src/channels/web/handlers/routines.rs @@ -114,7 +114,7 @@ pub async fn routines_detail_handler( trigger_type: run.trigger_type.clone(), started_at: run.started_at.to_rfc3339(), completed_at: run.completed_at.map(|dt| dt.to_rfc3339()), - status: format!("{:?}", run.status), + status: run.status.to_string(), result_summary: run.result_summary.clone(), tokens_used: run.tokens_used, job_id: run.job_id, @@ -324,7 +324,7 @@ pub async fn routines_runs_handler( trigger_type: run.trigger_type.clone(), started_at: run.started_at.to_rfc3339(), completed_at: run.completed_at.map(|dt| dt.to_rfc3339()), - status: format!("{:?}", run.status), + status: run.status.to_string(), result_summary: run.result_summary.clone(), tokens_used: run.tokens_used, job_id: run.job_id, diff --git a/src/channels/web/server.rs b/src/channels/web/server.rs index aaa479fa..fa29040e 100644 --- a/src/channels/web/server.rs +++ b/src/channels/web/server.rs @@ -2572,7 +2572,7 @@ async fn routines_runs_handler( trigger_type: run.trigger_type.clone(), started_at: run.started_at.to_rfc3339(), completed_at: run.completed_at.map(|dt| dt.to_rfc3339()), - status: format!("{:?}", run.status), + status: run.status.to_string(), result_summary: run.result_summary.clone(), tokens_used: run.tokens_used, job_id: run.job_id, diff --git a/src/channels/web/static/app.js b/src/channels/web/static/app.js index ddcfc828..6b366482 100644 --- a/src/channels/web/static/app.js +++ b/src/channels/web/static/app.js @@ -4265,9 +4265,9 @@ function renderRoutineDetail(routine) { + 'TriggerStartedCompletedStatusSummaryTokens' + ''; for (const run of routine.recent_runs) { - const runStatusClass = run.status === 'Ok' ? 'completed' - : run.status === 'Failed' ? 'failed' - : run.status === 'Attention' ? 'stuck' + const runStatusClass = run.status === 'ok' ? 'completed' + : run.status === 'failed' ? 'failed' + : run.status === 'attention' ? 'stuck' : 'in_progress'; html += '' + '' + escapeHtml(run.trigger_type) + '' diff --git a/src/cli/routines.rs b/src/cli/routines.rs index ebef8839..287663f6 100644 --- a/src/cli/routines.rs +++ b/src/cli/routines.rs @@ -10,7 +10,7 @@ use clap::Subcommand; use uuid::Uuid; use crate::agent::routine::{ - NotifyConfig, Routine, RoutineAction, RoutineGuardrails, Trigger, next_cron_fire, + NotifyConfig, Routine, RoutineAction, RoutineGuardrails, RunStatus, Trigger, next_cron_fire, }; use crate::db::Database; @@ -251,15 +251,26 @@ async fn list( ); println!("{}", "-".repeat(130)); + // Fetch last-run status for all routines in a single batch query + let routine_ids: Vec = filtered.iter().map(|r| r.id).collect(); + let last_run_results = db + .batch_get_last_run_status(&routine_ids) + .await + .unwrap_or_default(); + for r in &filtered { - let status = if r.enabled { - if r.consecutive_failures > 0 { - format!("err({})", r.consecutive_failures) - } else { - "active".to_string() - } - } else { + let last_run_status = last_run_results.get(&r.id).copied(); + + let status = if !r.enabled { "disabled".to_string() + } else if last_run_status == Some(RunStatus::Running) { + "running".to_string() + } else if r.consecutive_failures > 0 { + format!("err({})", r.consecutive_failures) + } else if last_run_status == Some(RunStatus::Attention) { + "attention".to_string() + } else { + "active".to_string() }; let next_fire = r diff --git a/src/db/libsql/routines.rs b/src/db/libsql/routines.rs index 6702cc1b..69c9f5c0 100644 --- a/src/db/libsql/routines.rs +++ b/src/db/libsql/routines.rs @@ -462,6 +462,56 @@ impl RoutineStore for LibSqlBackend { Ok(counts) } + async fn batch_get_last_run_status( + &self, + routine_ids: &[Uuid], + ) -> Result, DatabaseError> { + if routine_ids.is_empty() { + return Ok(HashMap::new()); + } + + let conn = self.connect().await?; + + // SQLite doesn't support ANY($1), so we query all latest runs and filter in memory. + // Uses a subquery to pick only the most recent run per routine. + let mut rows = conn + .query( + "SELECT routine_id, status FROM routine_runs r1 + WHERE started_at = ( + SELECT MAX(started_at) FROM routine_runs r2 + WHERE r2.routine_id = r1.routine_id + ) + GROUP BY routine_id", + params![], + ) + .await + .map_err(|e| { + DatabaseError::Query(format!("Failed to batch get last run status: {}", e)) + })?; + + let routine_id_set: HashSet = routine_ids.iter().copied().collect(); + let mut statuses = HashMap::new(); + + while let Some(row) = rows + .next() + .await + .map_err(|e| DatabaseError::Query(e.to_string()))? + { + let id_str: String = get_text(&row, 0); + let id = Uuid::parse_str(&id_str) + .map_err(|e| DatabaseError::Query(format!("Invalid routine UUID: {}", e)))?; + + if routine_id_set.contains(&id) { + let status_str: String = get_text(&row, 1); + if let std::result::Result::Ok(status) = status_str.parse::() { + statuses.insert(id, status); + } + } + } + + Ok(statuses) + } + async fn link_routine_run_to_job( &self, run_id: Uuid, diff --git a/src/db/mod.rs b/src/db/mod.rs index c0594bda..6d984fed 100644 --- a/src/db/mod.rs +++ b/src/db/mod.rs @@ -528,6 +528,15 @@ pub trait RoutineStore: Send + Sync { &self, routine_ids: &[Uuid], ) -> Result, DatabaseError>; + + /// Fetch the last run status for multiple routines in a single query. + /// Returns a map from routine_id to its most recent RunStatus. + /// Routines with no runs are omitted from the result. + async fn batch_get_last_run_status( + &self, + routine_ids: &[Uuid], + ) -> Result, DatabaseError>; + async fn link_routine_run_to_job( &self, run_id: Uuid, diff --git a/src/db/postgres.rs b/src/db/postgres.rs index a2c686d3..7bf76001 100644 --- a/src/db/postgres.rs +++ b/src/db/postgres.rs @@ -510,6 +510,14 @@ impl RoutineStore for PgBackend { .await } + async fn batch_get_last_run_status( + &self, + routine_ids: &[Uuid], + ) -> Result, DatabaseError> + { + self.store.batch_get_last_run_status(routine_ids).await + } + async fn link_routine_run_to_job( &self, run_id: Uuid, diff --git a/src/history/store.rs b/src/history/store.rs index d6570b3c..1e4cdd82 100644 --- a/src/history/store.rs +++ b/src/history/store.rs @@ -1403,6 +1403,40 @@ impl Store { Ok(counts) } + /// Batch-load the most recent run status for multiple routines in a single query. + /// Uses a window function to pick only the latest run per routine. + #[cfg(feature = "postgres")] + pub async fn batch_get_last_run_status( + &self, + routine_ids: &[Uuid], + ) -> Result, DatabaseError> { + if routine_ids.is_empty() { + return Ok(HashMap::new()); + } + + let conn = self.conn().await?; + let rows = conn + .query( + "SELECT DISTINCT ON (routine_id) routine_id, status + FROM routine_runs + WHERE routine_id = ANY($1) + ORDER BY routine_id, started_at DESC", + &[&routine_ids], + ) + .await?; + + let mut statuses = HashMap::new(); + for row in rows { + let id: Uuid = row.get("routine_id"); + let status_str: String = row.get("status"); + if let std::result::Result::Ok(status) = status_str.parse::() { + statuses.insert(id, status); + } + } + + Ok(statuses) + } + /// Link a routine run to a dispatched job. pub async fn link_routine_run_to_job( &self, diff --git a/tests/batch_last_run_status_tests.rs b/tests/batch_last_run_status_tests.rs new file mode 100644 index 00000000..4bd476ec --- /dev/null +++ b/tests/batch_last_run_status_tests.rs @@ -0,0 +1,191 @@ +//! Tests for batch_get_last_run_status (#1469 N+1 fix). +//! +//! Verifies: +//! 1. Empty input returns empty map +//! 2. Returns the most recent run status per routine +//! 3. Routines with no runs are omitted from result +//! 4. Multiple routines with different statuses are correctly returned + +#[cfg(feature = "libsql")] +mod tests { + use std::sync::Arc; + + use chrono::{Duration, Utc}; + use uuid::Uuid; + + use ironclaw::agent::routine::{ + Routine, RoutineAction, RoutineGuardrails, RoutineRun, RunStatus, Trigger, + }; + use ironclaw::db::Database; + + async fn create_test_db() -> (Arc, tempfile::TempDir) { + use ironclaw::db::libsql::LibSqlBackend; + + let temp_dir = tempfile::tempdir().expect("tempdir"); + let db_path = temp_dir.path().join("test.db"); + let backend = LibSqlBackend::new_local(&db_path) + .await + .expect("LibSqlBackend"); + backend.run_migrations().await.expect("migrations"); + let db: Arc = Arc::new(backend); + (db, temp_dir) + } + + fn make_routine(id: Uuid) -> Routine { + Routine { + id, + name: format!("test-routine-{}", id), + description: "Test routine".to_string(), + user_id: "default".to_string(), + enabled: true, + trigger: Trigger::Manual, + action: RoutineAction::FullJob { + title: "Test job".to_string(), + description: "Test description".to_string(), + max_iterations: 5, + }, + guardrails: RoutineGuardrails { + cooldown: std::time::Duration::from_secs(0), + max_concurrent: 1, + dedup_window: None, + }, + notify: Default::default(), + last_run_at: None, + next_fire_at: None, + run_count: 0, + consecutive_failures: 0, + state: serde_json::json!({}), + created_at: Utc::now(), + updated_at: Utc::now(), + } + } + + fn make_run( + routine_id: Uuid, + status: RunStatus, + started_at: chrono::DateTime, + ) -> RoutineRun { + RoutineRun { + id: Uuid::new_v4(), + routine_id, + trigger_type: "manual".to_string(), + trigger_detail: None, + started_at, + completed_at: if status == RunStatus::Running { + None + } else { + Some(Utc::now()) + }, + status, + result_summary: None, + tokens_used: None, + job_id: None, + created_at: Utc::now(), + } + } + + #[tokio::test] + async fn test_batch_get_last_run_status_empty_input() { + let (db, _tmp) = create_test_db().await; + let result = db + .batch_get_last_run_status(&[]) + .await + .expect("batch query"); + assert!(result.is_empty()); + } + + #[tokio::test] + async fn test_batch_get_last_run_status_returns_latest() { + let (db, _tmp) = create_test_db().await; + + let routine_id = Uuid::new_v4(); + db.create_routine(&make_routine(routine_id)) + .await + .expect("create routine"); + + // Create an older run with Ok status + let older_run = make_run(routine_id, RunStatus::Ok, Utc::now() - Duration::hours(2)); + db.create_routine_run(&older_run) + .await + .expect("create older run"); + db.complete_routine_run(older_run.id, RunStatus::Ok, None, None) + .await + .expect("complete older run"); + + // Create a newer run with Attention status + let newer_run = make_run( + routine_id, + RunStatus::Attention, + Utc::now() - Duration::hours(1), + ); + db.create_routine_run(&newer_run) + .await + .expect("create newer run"); + db.complete_routine_run(newer_run.id, RunStatus::Attention, None, None) + .await + .expect("complete newer run"); + + let result = db + .batch_get_last_run_status(&[routine_id]) + .await + .expect("batch query"); + assert_eq!(result.get(&routine_id), Some(&RunStatus::Attention)); + } + + #[tokio::test] + async fn test_batch_get_last_run_status_omits_routines_without_runs() { + let (db, _tmp) = create_test_db().await; + + let with_runs = Uuid::new_v4(); + let without_runs = Uuid::new_v4(); + db.create_routine(&make_routine(with_runs)) + .await + .expect("create routine"); + db.create_routine(&make_routine(without_runs)) + .await + .expect("create routine"); + + let run = make_run(with_runs, RunStatus::Ok, Utc::now()); + db.create_routine_run(&run).await.expect("create run"); + db.complete_routine_run(run.id, RunStatus::Ok, None, None) + .await + .expect("complete run"); + + let result = db + .batch_get_last_run_status(&[with_runs, without_runs]) + .await + .expect("batch query"); + assert_eq!(result.get(&with_runs), Some(&RunStatus::Ok)); + assert_eq!(result.get(&without_runs), None); + } + + #[tokio::test] + async fn test_batch_get_last_run_status_multiple_routines() { + let (db, _tmp) = create_test_db().await; + + let r1 = Uuid::new_v4(); + let r2 = Uuid::new_v4(); + db.create_routine(&make_routine(r1)) + .await + .expect("create r1"); + db.create_routine(&make_routine(r2)) + .await + .expect("create r2"); + + let run1 = make_run(r1, RunStatus::Running, Utc::now()); + db.create_routine_run(&run1).await.expect("create run1"); + + let run2 = make_run(r2, RunStatus::Failed, Utc::now()); + db.create_routine_run(&run2).await.expect("create run2"); + db.complete_routine_run(run2.id, RunStatus::Failed, None, None) + .await + .expect("complete run2"); + + let result = db + .batch_get_last_run_status(&[r1, r2]) + .await + .expect("batch query"); + assert_eq!(result.get(&r1), Some(&RunStatus::Running)); + assert_eq!(result.get(&r2), Some(&RunStatus::Failed)); + } +} From d3d517fd677f3f1f32f7351df8b310229fb5fba9 Mon Sep 17 00:00:00 2001 From: Zaki Manian Date: Tue, 24 Mar 2026 02:44:25 -0700 Subject: [PATCH 08/20] fix(agent): case-insensitive channel match and user_id filter for event triggers (#1211) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix(agent): case-insensitive channel match and user_id filter for event triggers (#1051, #1076) Event-triggered routines had two bugs preventing them from firing: 1. Channel comparison was case-sensitive (e.g., "Telegram" != "telegram"), while emit_system_event already used eq_ignore_ascii_case. Fixed to match. 2. No user_id scoping — routines from any user were evaluated against every message. Added ownership check so routines only fire for their owner's messages. Also adds periodic event cache refresh (every ~60s) in the cron ticker so web/CLI mutations are picked up without requiring the tool path. Upgrades skip-reason logging from trace to debug for debuggability. Closes #1051 Refs #1076 Co-Authored-By: Claude Opus 4.6 (1M context) * fix: correct refresh_every from 6 to 4 to match 15s default interval The default cron_check_interval_secs is 15s, not 10s. With refresh_every=6, the cache would refresh every 90s instead of the intended ~60s. Fix to 4 ticks (4 * 15s = 60s). Co-Authored-By: Claude Opus 4.6 (1M context) * fix(agent): address #1211 review -- extract routine_matches_message, fix refresh interval Extract user/channel filter logic from check_event_triggers into a standalone pure function routine_matches_message(). Rewrite tests to call this function directly with controlled Routine and IncomingMessage values, so they exercise the real code path and would catch a revert. Add test_no_channel_filter_matches_any_channel for the None channel case. Co-Authored-By: Claude Opus 4.6 (1M context) * ci: re-trigger CI with latest changes Co-Authored-By: Claude Opus 4.6 * fix: add missing IncomingMessage fields in test helper Co-Authored-By: Claude Opus 4.6 * fix(agent): address review -- time-based refresh, trace-level user mismatch, scope guard (#1211) - Use tokio::time::Instant for cache refresh instead of tick counting - Downgrade user-mismatch log to trace to reduce noise - Add early return false for non-Event triggers in routine_matches_message - Fix doc comment to say 'user scope' instead of 'message sender' Co-Authored-By: Claude Opus 4.6 (1M context) * style: run cargo fmt on agent_loop.rs https://claude.ai/code/session_01ABGWibdKVQ3b6pEKtxPPkM * fix(agent): resolve clippy warnings for unused binding and needless borrow Fix unused `content` variable in event trigger guard (use `content: _`) and remove redundant `&` on `message` which was already a reference. https://claude.ai/code/session_01PzBK21BbUAuZbrfLpoz4Xb * fix(test): update check_event_triggers call sites to new single-arg signature The staging merge brought e2e_routine_heartbeat tests that still used the old 3-argument check_event_triggers(user_id, channel, content) signature. Updated all 11 call sites to pass &IncomingMessage directly. [skip-regression-check] https://claude.ai/code/session_012GrkTDrtDFkpJos2hkgTcE * fix(agent): address review feedback on event trigger handling - Use post-hook content for event trigger matching so BeforeInbound hooks that rewrite input are respected - Set MissedTickBehavior::Skip on cron ticker to avoid burst catch-up after delays Co-Authored-By: Claude Opus 4.6 (1M context) * style: cargo fmt https://claude.ai/code/session_01Va9wwvATNWFAx35GG7Zek7 --------- Co-authored-by: Claude Opus 4.6 (1M context) Co-authored-by: firat.sertgoz --- src/agent/agent_loop.rs | 6 +- src/agent/routine_engine.rs | 205 ++++++++++++++++++++++++++++++--- tests/e2e_routine_heartbeat.rs | 44 ++----- 3 files changed, 201 insertions(+), 54 deletions(-) diff --git a/src/agent/agent_loop.rs b/src/agent/agent_loop.rs index 3ab369b1..7961250d 100644 --- a/src/agent/agent_loop.rs +++ b/src/agent/agent_loop.rs @@ -1139,9 +1139,9 @@ impl Agent { && let Submission::UserInput { ref content } = submission && let Some(engine) = self.routine_engine().await { - let fired = engine - .check_event_triggers(&message.user_id, &message.channel, content) - .await; + // 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; if fired > 0 { tracing::debug!( channel = %message.channel, diff --git a/src/agent/routine_engine.rs b/src/agent/routine_engine.rs index 7c7ef5f3..39acb83d 100644 --- a/src/agent/routine_engine.rs +++ b/src/agent/routine_engine.rs @@ -24,7 +24,7 @@ use crate::agent::Scheduler; use crate::agent::routine::{ NotifyConfig, Routine, RoutineAction, RoutineRun, RunStatus, Trigger, next_cron_fire, }; -use crate::channels::OutgoingResponse; +use crate::channels::{IncomingMessage, OutgoingResponse}; use crate::config::RoutineConfig; use crate::context::{JobContext, JobState}; use crate::db::Database; @@ -56,6 +56,40 @@ pub enum SandboxReadiness { DockerUnavailable, } +/// Check whether an event-triggered routine's user/channel filters match an +/// incoming message. +/// +/// Returns `true` if: +/// - The routine has an `Event` trigger (non-Event routines always return `false`) +/// - The routine's `user_id` matches the message's user scope +/// - The routine's channel filter (if any) matches the message channel +/// case-insensitively +/// +/// This is a pure function extracted from `check_event_triggers` so the +/// filter logic can be unit-tested without async infrastructure. +pub(crate) fn routine_matches_message(routine: &Routine, message: &IncomingMessage) -> bool { + // Only Event-triggered routines can match incoming messages. + if !matches!(routine.trigger, Trigger::Event { .. }) { + return false; + } + + // User ownership filter — only fire routines scoped to this user. + if routine.user_id != message.user_id { + return false; + } + + // Channel filter (case-insensitive, matching emit_system_event behavior) + if let Trigger::Event { + channel: Some(ch), .. + } = &routine.trigger + && !ch.eq_ignore_ascii_case(&message.channel) + { + return false; + } + + true +} + /// The routine execution engine. pub struct RoutineEngine { config: RoutineConfig, @@ -167,10 +201,7 @@ impl RoutineEngine { } /// Check incoming message against event triggers. Returns number of routines fired. - /// - /// Accepts only the three fields needed for matching (user scope, channel, - /// message content) so callers never need to clone a full `IncomingMessage`. - pub async fn check_event_triggers(&self, user_id: &str, channel: &str, content: &str) -> usize { + pub async fn check_event_triggers(&self, message: &IncomingMessage, content: &str) -> usize { let cache = self.event_cache.read().await; // Early return if there are no message matchers at all. @@ -208,16 +239,24 @@ impl RoutineEngine { EventMatcher::System { .. } => continue, }; - if routine.user_id != user_id { - continue; - } - - // Channel filter - if let Trigger::Event { - channel: Some(ch), .. - } = &routine.trigger - && ch != channel - { + // User ownership + channel filter (extracted for testability). + if !routine_matches_message(routine, message) { + // User mismatch is expected for multi-user setups — keep at + // trace to avoid one log per routine per inbound message. + if routine.user_id != message.user_id { + tracing::trace!( + routine = %routine.name, + routine_user = %routine.user_id, + message_user = %message.user_id, + "Skipped: user scope mismatch" + ); + } else { + tracing::debug!( + routine = %routine.name, + channel = %message.channel, + "Skipped: channel mismatch" + ); + } continue; } @@ -228,14 +267,14 @@ impl RoutineEngine { // Cooldown check if !self.check_cooldown(routine) { - tracing::trace!(routine = %routine.name, "Skipped: cooldown active"); + tracing::debug!(routine = %routine.name, "Skipped: cooldown active"); continue; } // Concurrent run check (using batch-loaded counts) let running_count = concurrent_counts.get(&routine.id).copied().unwrap_or(0); if running_count >= routine.guardrails.max_concurrent as i64 { - tracing::trace!(routine = %routine.name, "Skipped: max concurrent reached"); + tracing::debug!(routine = %routine.name, "Skipped: max concurrent reached"); continue; } @@ -1781,6 +1820,13 @@ pub fn spawn_cron_ticker( engine.check_cron_triggers().await; let mut ticker = tokio::time::interval(interval); + ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); + // Periodic event cache refresh so web/CLI mutations are picked up + // without requiring tool-path code to call refresh_event_cache(). + // Uses wall-clock elapsed time so the refresh cadence is stable + // regardless of the cron tick interval configuration. + let refresh_interval = Duration::from_secs(60); + let mut last_refresh = tokio::time::Instant::now(); loop { ticker.tick().await; @@ -1788,7 +1834,11 @@ pub fn spawn_cron_ticker( // never races with FullJobWatcher instances from this process. engine.sync_dispatched_runs().await; engine.check_cron_triggers().await; - engine.sync_dispatched_runs().await; + + if last_refresh.elapsed() >= refresh_interval { + engine.refresh_event_cache().await; + last_refresh = tokio::time::Instant::now(); + } } }) } @@ -1854,7 +1904,13 @@ fn strip_html_tags(s: &str) -> String { #[cfg(test)] mod tests { - use crate::agent::routine::{NotifyConfig, RunStatus}; + use chrono::Utc; + use uuid::Uuid; + + use crate::agent::routine::{ + NotifyConfig, Routine, RoutineAction, RoutineGuardrails, RunStatus, Trigger, + }; + use crate::channels::IncomingMessage; use crate::config::RoutineConfig; #[test] @@ -2052,6 +2108,117 @@ mod tests { } } + /// Helper to build a test routine with the given user_id and trigger. + fn make_routine(user_id: &str, trigger: Trigger) -> Routine { + Routine { + id: Uuid::new_v4(), + name: "test".to_string(), + description: String::new(), + user_id: user_id.to_string(), + enabled: true, + trigger, + action: RoutineAction::Lightweight { + prompt: String::new(), + context_paths: vec![], + max_tokens: 1000, + use_tools: false, + max_tool_rounds: 0, + }, + guardrails: RoutineGuardrails::default(), + notify: Default::default(), + last_run_at: None, + next_fire_at: None, + run_count: 0, + consecutive_failures: 0, + state: serde_json::Value::Null, + created_at: Utc::now(), + updated_at: Utc::now(), + } + } + + /// Helper to build a test IncomingMessage. + fn make_message(user_id: &str, channel: &str, content: &str) -> IncomingMessage { + IncomingMessage { + id: Uuid::new_v4(), + channel: channel.to_string(), + user_id: user_id.to_string(), + owner_id: user_id.to_string(), + sender_id: user_id.to_string(), + user_name: None, + content: content.to_string(), + thread_id: None, + conversation_scope_id: None, + received_at: Utc::now(), + metadata: serde_json::Value::Null, + timezone: None, + attachments: vec![], + is_internal: false, + } + } + + /// Regression test for issue #1051: event triggers used case-sensitive + /// channel comparison, so "Telegram" != "telegram" caused silent mismatch. + /// Tests the actual `routine_matches_message` function used in `check_event_triggers`. + #[test] + fn test_channel_filter_is_case_insensitive() { + let routine = make_routine( + "user1", + Trigger::Event { + pattern: ".*".to_string(), + channel: Some("Telegram".to_string()), + }, + ); + let msg = make_message("user1", "telegram", "hello"); + + // Case-insensitive channel match must succeed + assert!(super::routine_matches_message(&routine, &msg)); + + // Exact case must also work + let msg_exact = make_message("user1", "Telegram", "hello"); + assert!(super::routine_matches_message(&routine, &msg_exact)); + + // Different channel must not match + let msg_wrong = make_message("user1", "discord", "hello"); + assert!(!super::routine_matches_message(&routine, &msg_wrong)); + } + + /// Regression test for issue #1051: event triggers did not filter by + /// user_id, so routines from user A could fire on messages from user B. + /// Tests the actual `routine_matches_message` function used in `check_event_triggers`. + #[test] + fn test_event_trigger_requires_user_match() { + let routine = make_routine( + "alice", + Trigger::Event { + pattern: ".*".to_string(), + channel: None, + }, + ); + + // Different user must not match + let msg_bob = make_message("bob", "telegram", "hello"); + assert!(!super::routine_matches_message(&routine, &msg_bob)); + + // Same user must match + let msg_alice = make_message("alice", "telegram", "hello"); + assert!(super::routine_matches_message(&routine, &msg_alice)); + } + + /// When no channel filter is set, any channel should match (given user matches). + #[test] + fn test_no_channel_filter_matches_any_channel() { + let routine = make_routine( + "user1", + Trigger::Event { + pattern: ".*".to_string(), + channel: None, + }, + ); + + let msg = make_message("user1", "whatever_channel", "hello"); + assert!(super::routine_matches_message(&routine, &msg)); + } + #[test] fn test_routine_tool_denylist_blocks_self_management_tools() { let denylisted = vec![ diff --git a/tests/e2e_routine_heartbeat.rs b/tests/e2e_routine_heartbeat.rs index 12125d43..27d8cfdc 100644 --- a/tests/e2e_routine_heartbeat.rs +++ b/tests/e2e_routine_heartbeat.rs @@ -561,11 +561,7 @@ mod tests { "deploy to production now", ); let fired = engine - .check_event_triggers( - &matching_msg.user_id, - &matching_msg.channel, - &matching_msg.content, - ) + .check_event_triggers(&matching_msg, &matching_msg.content) .await; assert!( fired >= 1, @@ -584,11 +580,7 @@ mod tests { "check the staging environment", ); let fired_neg = engine - .check_event_triggers( - &non_matching_msg.user_id, - &non_matching_msg.channel, - &non_matching_msg.content, - ) + .check_event_triggers(&non_matching_msg, &non_matching_msg.content) .await; assert_eq!(fired_neg, 0, "Expected 0 routines fired on non-match"); } @@ -652,7 +644,7 @@ mod tests { "deploy to production now", ); let guest_fired = engine - .check_event_triggers(&guest_msg.user_id, &guest_msg.channel, &guest_msg.content) + .check_event_triggers(&guest_msg, &guest_msg.content) .await; assert_eq!( guest_fired, 0, @@ -677,7 +669,7 @@ mod tests { "deploy to production now", ); let owner_fired = engine - .check_event_triggers(&owner_msg.user_id, &owner_msg.channel, &owner_msg.content) + .check_event_triggers(&owner_msg, &owner_msg.content) .await; assert!( owner_fired >= 1, @@ -906,9 +898,7 @@ mod tests { "default", "test-cooldown trigger", ); - let fired1 = engine - .check_event_triggers(&msg.user_id, &msg.channel, &msg.content) - .await; + let fired1 = engine.check_event_triggers(&msg, &msg.content).await; assert!(fired1 >= 1, "First fire should work"); // Give spawn time, then update last_run_at to simulate recent execution. @@ -923,9 +913,7 @@ mod tests { engine.refresh_event_cache().await; // Second fire should be blocked by cooldown. - let fired2 = engine - .check_event_triggers(&msg.user_id, &msg.channel, &msg.content) - .await; + let fired2 = engine.check_event_triggers(&msg, &msg.content).await; assert_eq!(fired2, 0, "Second fire should be blocked by cooldown"); } @@ -1095,9 +1083,7 @@ mod tests { engine.refresh_event_cache().await; let msg = IncomingMessage::new("test", "default", "DISABLE_ME"); - let fired_before = engine - .check_event_triggers(&msg.user_id, &msg.channel, &msg.content) - .await; + let fired_before = engine.check_event_triggers(&msg, &msg.content).await; assert!(fired_before >= 1, "Expected routine to fire before disable"); // Simulate what routines_toggle_handler now does: update DB, then refresh. @@ -1106,9 +1092,7 @@ mod tests { db.update_routine(&routine).await.expect("update_routine"); engine.refresh_event_cache().await; - let fired_after = engine - .check_event_triggers(&msg.user_id, &msg.channel, &msg.content) - .await; + let fired_after = engine.check_event_triggers(&msg, &msg.content).await; assert_eq!( fired_after, 0, "Disabled routine must not fire after cache refresh" @@ -1134,10 +1118,7 @@ mod tests { let msg = IncomingMessage::new("test", "default", "DELETE_ME"); assert!( - engine - .check_event_triggers(&msg.user_id, &msg.channel, &msg.content) - .await - >= 1, + engine.check_event_triggers(&msg, &msg.content).await >= 1, "Expected routine to fire before delete" ); @@ -1146,9 +1127,7 @@ mod tests { engine.refresh_event_cache().await; assert_eq!( - engine - .check_event_triggers(&msg.user_id, &msg.channel, &msg.content) - .await, + engine.check_event_triggers(&msg, &msg.content).await, 0, "Deleted routine must not fire after cache refresh" ); @@ -1462,8 +1441,9 @@ mod tests { db.create_routine(&routine).await.expect("create_routine"); engine.refresh_event_cache().await; + let trigger_msg = IncomingMessage::new("test", "default", "owner-gate"); let fired = engine - .check_event_triggers("default", "test", "owner-gate") + .check_event_triggers(&trigger_msg, &trigger_msg.content) .await; assert_eq!(fired, 1, "expected one matching event routine"); From 5901451603d164a0e5814855161e5d5b05cc0cf0 Mon Sep 17 00:00:00 2001 From: Pierre LE GUEN <26087574+PierreLeGuen@users.noreply.github.com> Date: Tue, 24 Mar 2026 17:49:13 +0000 Subject: [PATCH 09/20] fix: remove stale stream_token gate from channel-relay activation (#1623) * fix: remove stale stream_token gate from channel-relay activation The relay architecture now uses instance-scoped bearer auth + webhook callbacks, not streaming. The `relay::stream_token` secret was never written by the current OAuth flow, so activation always failed with AuthRequired. Replace stream_token with the team_id setting (already stored by the OAuth callback) as the persistent "auth completed" marker: - is_relay_channel(): check team_id setting instead of stream_token secret - activate_channel_relay(): gate on team_id emptiness, not stream_token - removal flow: delete team_id setting + oauth_state secret - configure(): return empty allowed-secrets set (relay is OAuth-only) - configure_token(): return AuthRequired (no manual token entry) - list(): surface activation_error for relay channels (was hardcoded None) - Clean up stale comments referencing stream_token / "stored token" - Update test to match OAuth-only model (no secrets to pass) Made-with: Cursor * fix: address CI and review feedback - Fix pre-existing tunnel/mod.rs test compilation (missing GatewayConfig fields: memory_layers, user_tokens, workspace_read_scopes) - Log warnings on failed team_id/oauth_state cleanup during removal instead of silently ignoring errors (gemini review) - Also delete legacy stream_token secret during removal for backward compatibility with pre-webhook installs (codex review) Made-with: Cursor --- src/extensions/manager.rs | 102 ++++++++++++++++++-------------------- src/tunnel/mod.rs | 6 +++ 2 files changed, 55 insertions(+), 53 deletions(-) diff --git a/src/extensions/manager.rs b/src/extensions/manager.rs index 7da9e980..39654305 100644 --- a/src/extensions/manager.rs +++ b/src/extensions/manager.rs @@ -891,24 +891,27 @@ impl ExtensionManager { *self.relay_channel_manager.write().await = Some(channel_manager); } - /// Check if a channel name corresponds to a relay extension (has stored stream token + /// Check if a channel name corresponds to a relay extension (has stored team_id /// or is tracked in the installed relay extensions set). pub async fn is_relay_channel(&self, name: &str, user_id: &str) -> bool { // Check in-memory installed set first (supports no-store mode) if self.installed_relay_extensions.read().await.contains(name) { return true; } - // Then check for stored stream token - self.secrets - .exists(user_id, &format!("relay:{}:stream_token", name)) - .await - .unwrap_or(false) + // Check for stored team_id (persisted across restarts by the OAuth callback) + if let Some(ref store) = self.store { + let key = format!("relay:{}:team_id", name); + if let Ok(Some(v)) = store.get_setting(user_id, &key).await { + return v.as_str().is_some_and(|s| !s.is_empty()); + } + } + false } /// Restore persisted relay channels after startup. /// /// Loads the persisted active channel list, filters to relay types (those with - /// a stored stream token), and activates each via `activate_stored_relay()`. + /// a stored team_id setting), and activates each via `activate_stored_relay()`. /// Skips channels that are already active. /// /// Call this only after `set_relay_channel_manager()` or `set_channel_runtime()`. @@ -1428,9 +1431,11 @@ impl ExtensionManager { if kind_filter.is_none() || kind_filter == Some(ExtensionKind::ChannelRelay) { let installed = self.installed_relay_extensions.read().await; let active_names = self.active_channel_names.read().await; + let errors = self.activation_errors.read().await; for name in installed.iter() { let active = active_names.contains(name); - let has_token = self.is_relay_channel(name, user_id).await; + let authenticated = self.is_relay_channel(name, user_id).await; + let activation_error = errors.get(name).cloned(); let registry_entry = self .registry .get_with_kind(name, Some(ExtensionKind::ChannelRelay)) @@ -1443,13 +1448,13 @@ impl ExtensionManager { display_name, description, url: None, - authenticated: has_token, + authenticated, active, tools: Vec::new(), needs_setup: false, has_auth: true, installed: true, - activation_error: None, + activation_error, version: None, }); } @@ -1626,7 +1631,22 @@ impl ExtensionManager { self.persist_active_channels(user_id).await; self.activation_errors.write().await.remove(name); - // Remove stored stream token + // Remove stored team_id setting and clean up secrets + if let Some(ref store) = self.store + && let Err(e) = store + .delete_setting(user_id, &format!("relay:{}:team_id", name)) + .await + { + tracing::warn!(error = %e, name, "Failed to delete relay team_id setting on removal"); + } + if let Err(e) = self + .secrets + .delete(user_id, &format!("relay:{}:oauth_state", name)) + .await + { + tracing::warn!(error = %e, name, "Failed to delete relay oauth_state secret on removal"); + } + // Clean up legacy stream_token secret from pre-webhook installs let _ = self .secrets .delete(user_id, &format!("relay:{}:stream_token", name)) @@ -4181,13 +4201,13 @@ impl ExtensionManager { /// /// For Slack: initiates OAuth flow (redirect-based). /// For Telegram: accepts a bot token, registers it with channel-relay, - /// and stores the returned stream token. + /// and stores the team_id setting. async fn auth_channel_relay( &self, name: &str, user_id: &str, ) -> Result { - // Check if already authenticated (stream token exists) + // Check if already authenticated (team_id setting exists) if self.is_relay_channel(name, user_id).await { return Ok(AuthResult::authenticated(name, ExtensionKind::ChannelRelay)); } @@ -4233,19 +4253,9 @@ impl ExtensionManager { name: &str, user_id: &str, ) -> Result { - let token_key = format!("relay:{}:stream_token", name); let team_id_key = format!("relay:{}:team_id", name); - // Check if we have a stream token - // Verify auth: stream token must exist (even though we don't use it in this constructor path) - let _stream_token = match self.secrets.get_decrypted(user_id, &token_key).await { - Ok(secret) => secret.expose().to_string(), - Err(_) => { - return Err(ExtensionError::AuthRequired); - } - }; - - // Get team_id from settings + // Get team_id from settings (stored by the OAuth callback) let team_id = if let Some(ref store) = self.store { store .get_setting(user_id, &team_id_key) @@ -4258,6 +4268,10 @@ impl ExtensionManager { String::new() }; + if team_id.is_empty() { + return Err(ExtensionError::AuthRequired); + } + // Use relay config captured at startup let relay_config = self.relay_config()?; @@ -4367,11 +4381,11 @@ impl ExtensionManager { return Ok(ExtensionKind::WasmChannel); } - // Check channel-relay extensions (installed in memory or has stored token) + // Check channel-relay extensions (installed in memory or has stored team_id) if self.installed_relay_extensions.read().await.contains(name) { return Ok(ExtensionKind::ChannelRelay); } - // Also check if there's a stored stream token (persisted across restarts) + // Also check if there's a stored team_id setting (persisted across restarts) if self.is_relay_channel(name, user_id).await { return Ok(ExtensionKind::ChannelRelay); } @@ -4999,11 +5013,7 @@ impl ExtensionManager { names.insert(server.token_secret_name()); (names, Vec::new()) } - ExtensionKind::ChannelRelay => { - let mut names = std::collections::HashSet::new(); - names.insert(format!("relay:{}:stream_token", name)); - (names, Vec::new()) - } + ExtensionKind::ChannelRelay => (std::collections::HashSet::new(), Vec::new()), }; let allowed_fields: std::collections::HashSet = @@ -5434,7 +5444,9 @@ impl ExtensionManager { .map_err(|e| ExtensionError::NotInstalled(e.to_string()))?; server.token_secret_name() } - ExtensionKind::ChannelRelay => format!("relay:{}:stream_token", name), + ExtensionKind::ChannelRelay => { + return Err(ExtensionError::AuthRequired); + } }; let mut secrets = std::collections::HashMap::new(); @@ -7043,7 +7055,7 @@ mod tests { let dir = tempfile::tempdir().expect("temp dir"); let mgr = make_test_manager(None, dir.path().to_path_buf()); - // No token stored → not a relay channel + // No store configured, no team_id → not a relay channel assert!(!mgr.is_relay_channel("slack-relay", "test").await); } @@ -7862,19 +7874,13 @@ mod tests { .await .insert("test-relay".to_string()); - // configure() should dispatch to activate_channel_relay(), not - // activate_wasm_channel(). Both will fail (no runtime configured), - // but the error should be about relay config, not WASM channels. - let mut secrets = std::collections::HashMap::new(); - secrets.insert( - "relay:test-relay:stream_token".to_string(), - "tok".to_string(), - ); - + // configure() with empty secrets should dispatch to + // activate_channel_relay(), not activate_wasm_channel(). Relay auth + // is OAuth-only so there are no manual secrets to pass. let result = mgr .configure( "test-relay", - &secrets, + &std::collections::HashMap::new(), &std::collections::HashMap::new(), "test", ) @@ -7886,7 +7892,6 @@ mod tests { ); let result = result.unwrap(); - // Activation will fail (no relay config), but secrets should still be stored assert!( !result.activated, "activation should fail without relay config" @@ -7896,15 +7901,6 @@ mod tests { "error should not mention WASM — got: {}", result.message ); - - // Verify the secret was stored - assert!( - mgr.secrets - .exists("test", "relay:test-relay:stream_token") - .await - .unwrap_or(false), - "configure should have stored the relay stream token" - ); } #[test] fn test_validation_failed_is_distinct_error_variant() { diff --git a/src/tunnel/mod.rs b/src/tunnel/mod.rs index a6869eda..8719b6e1 100644 --- a/src/tunnel/mod.rs +++ b/src/tunnel/mod.rs @@ -429,6 +429,9 @@ mod tests { port: 3000, auth_token: None, user_id: "test".to_string(), + workspace_read_scopes: Vec::new(), + memory_layers: Vec::new(), + user_tokens: None, }); c } @@ -440,6 +443,9 @@ mod tests { port, auth_token: None, user_id: "test".to_string(), + workspace_read_scopes: Vec::new(), + memory_layers: Vec::new(), + user_tokens: None, }); c } From f3da30a4549947e715891732b966b56b73f56fa0 Mon Sep 17 00:00:00 2001 From: Zaki Manian Date: Tue, 24 Mar 2026 11:48:30 -0700 Subject: [PATCH 10/20] perf(agent): optimize approval thread resolution (UUID parsing + lock contention) (#1592) --- src/agent/agent_loop.rs | 3 +- src/agent/session_manager.rs | 223 +++++++++++++++++++++++++++++------ 2 files changed, 191 insertions(+), 35 deletions(-) diff --git a/src/agent/agent_loop.rs b/src/agent/agent_loop.rs index 7961250d..7e950146 100644 --- a/src/agent/agent_loop.rs +++ b/src/agent/agent_loop.rs @@ -1055,10 +1055,11 @@ impl Agent { } else { drop(sess); self.session_manager - .resolve_thread( + .resolve_thread_with_parsed_uuid( &message.user_id, &message.channel, message.conversation_scope(), + approval_thread_uuid, ) .await } diff --git a/src/agent/session_manager.rs b/src/agent/session_manager.rs index 3bf20697..ae98b0b0 100644 --- a/src/agent/session_manager.rs +++ b/src/agent/session_manager.rs @@ -102,11 +102,30 @@ impl SessionManager { /// Resolve an external thread ID to an internal thread. /// /// Returns the session and thread ID. Creates both if they don't exist. + /// Delegates to [`resolve_thread_with_parsed_uuid`](Self::resolve_thread_with_parsed_uuid) + /// with `parsed_uuid: None`. pub async fn resolve_thread( &self, user_id: &str, channel: &str, external_thread_id: Option<&str>, + ) -> (Arc>, Uuid) { + self.resolve_thread_with_parsed_uuid(user_id, channel, external_thread_id, None) + .await + } + + /// Like [`resolve_thread`](Self::resolve_thread), but accepts a pre-parsed + /// UUID to skip redundant parsing when the caller has already validated + /// the external thread ID as a UUID (e.g. the approval routing path). + /// + /// Uses a single read-lock acquisition for both the key lookup and the UUID + /// adoption check to reduce contention under concurrent approval load. + pub async fn resolve_thread_with_parsed_uuid( + &self, + user_id: &str, + channel: &str, + external_thread_id: Option<&str>, + parsed_uuid: Option, ) -> (Arc>, Uuid) { let session = self.get_or_create_session(user_id).await; @@ -116,51 +135,65 @@ impl SessionManager { external_thread_id: external_thread_id.map(String::from), }; - // Check if we have a mapping - { + // Use pre-parsed UUID if available, otherwise parse from string. + let ext_uuid = parsed_uuid + .or_else(|| external_thread_id.and_then(|ext_tid| Uuid::parse_str(ext_tid).ok())); + + // Validate that parsed_uuid (if provided) is consistent with external_thread_id. + #[cfg(debug_assertions)] + if let (Some(parsed), Some(ext_tid)) = (&parsed_uuid, external_thread_id) { + debug_assert_eq!( + Uuid::parse_str(ext_tid).ok().as_ref(), + Some(parsed), + "parsed_uuid must be the parsed form of external_thread_id" + ); + } + + // Single read lock for both the key lookup and UUID adoption check + let adoptable_uuid = { let thread_map = self.thread_map.read().await; + + // Fast path: exact key match if let Some(&thread_id) = thread_map.get(&key) { - // Verify thread still exists in session let sess = session.lock().await; if sess.threads.contains_key(&thread_id) { return (Arc::clone(&session), thread_id); } } - } - // Check if external_thread_id is itself a known thread UUID that - // exists in the session but was never registered in the thread_map - // (e.g. created by chat_new_thread_handler or hydrated from DB). - // We only adopt it if no thread_map entry maps to this UUID — - // otherwise it belongs to a different channel scope. - if let Some(ext_tid) = external_thread_id - && let Ok(ext_uuid) = Uuid::parse_str(ext_tid) - { - let thread_map = self.thread_map.read().await; - let mapped_elsewhere = thread_map.values().any(|&v| v == ext_uuid); - drop(thread_map); + // UUID adoption check (still under the same read lock). + // If external_thread_id is a valid UUID not mapped elsewhere, + // it may be a thread created by chat_new_thread_handler or + // hydrated from DB that we can adopt. + // Only attempt adoption when external_thread_id is Some, preserving + // the invariant that None external_thread_id never triggers adoption. + if external_thread_id.is_some() { + ext_uuid.filter(|&uuid| !thread_map.values().any(|&v| v == uuid)) + } else { + None + } + }; // Single read lock dropped here - if !mapped_elsewhere { - let sess = session.lock().await; - if sess.threads.contains_key(&ext_uuid) { - drop(sess); + // If we found an adoptable UUID, verify it exists in session and acquire write lock + if let Some(ext_uuid) = adoptable_uuid { + let sess = session.lock().await; + if sess.threads.contains_key(&ext_uuid) { + drop(sess); - let mut thread_map = self.thread_map.write().await; - // Re-check after acquiring write lock to prevent race condition - // where another task mapped this UUID between our read and write. - if !thread_map.values().any(|&v| v == ext_uuid) { - thread_map.insert(key, ext_uuid); - drop(thread_map); - // Ensure undo manager exists - let mut undo_managers = self.undo_managers.write().await; - undo_managers - .entry(ext_uuid) - .or_insert_with(|| Arc::new(Mutex::new(UndoManager::new()))); - return (session, ext_uuid); - } - // If it was mapped elsewhere while we were unlocked, fall through - // to create a new thread, preserving channel isolation. + let mut thread_map = self.thread_map.write().await; + // Re-check after acquiring write lock to prevent race condition + // where another task mapped this UUID between our read and write. + if !thread_map.values().any(|&v| v == ext_uuid) { + thread_map.insert(key, ext_uuid); + drop(thread_map); + // Ensure undo manager exists + let mut undo_managers = self.undo_managers.write().await; + undo_managers + .entry(ext_uuid) + .or_insert_with(|| Arc::new(Mutex::new(UndoManager::new()))); + return (session, ext_uuid); } + // If mapped elsewhere while unlocked, fall through to create new thread } } @@ -909,6 +942,44 @@ mod tests { } } + #[tokio::test] + async fn test_resolve_thread_consolidates_read_path() { + // Verify that resolve_thread still correctly handles: + // 1. Fast path: key exists in thread_map + // 2. UUID adoption: external_thread_id is a UUID in session but not in map + // 3. New thread: neither path matches + use crate::agent::session::Thread; + + let manager = SessionManager::new(); + + // Case 1: Normal resolution creates thread and maps it + let (session1, tid1) = manager + .resolve_thread("user1", "chan1", Some("ext-1")) + .await; + // Resolving again with same key should return same thread (fast path) + let (_, tid1_again) = manager + .resolve_thread("user1", "chan1", Some("ext-1")) + .await; + assert_eq!(tid1, tid1_again); + + // Case 2: UUID adoption - insert a thread directly into session + let adopted_id = Uuid::new_v4(); + { + let mut sess = session1.lock().await; + let thread = Thread::with_id(adopted_id, sess.id); + sess.threads.insert(adopted_id, thread); + } + // Resolve with the UUID as external_thread_id -- should adopt it + let (_, resolved) = manager + .resolve_thread("user1", "chan1", Some(&adopted_id.to_string())) + .await; + assert_eq!(resolved, adopted_id); + + // Case 3: Different channel gets different thread + let (_, tid2) = manager.resolve_thread("user1", "chan2", None).await; + assert_ne!(tid1, tid2); + } + #[tokio::test] async fn test_resolve_thread_finds_existing_session_thread_by_uuid() { use crate::agent::session::{Session, Thread}; @@ -947,4 +1018,88 @@ mod tests { "should have exactly 1 thread, not a duplicate" ); } + + #[tokio::test] + async fn test_resolve_thread_with_pre_parsed_uuid_adopts_thread() { + use crate::agent::session::Thread; + + let manager = SessionManager::new(); + let (session, _) = manager.resolve_thread("user1", "chan1", None).await; + + // Manually insert a thread with a known UUID + let known_id = Uuid::new_v4(); + { + let mut sess = session.lock().await; + let thread = Thread::with_id(known_id, sess.id); + sess.threads.insert(known_id, thread); + } + + // Resolve with pre-parsed UUID -- should adopt it without re-parsing + let (_, resolved) = manager + .resolve_thread_with_parsed_uuid( + "user1", + "chan1", + Some(&known_id.to_string()), + Some(known_id), + ) + .await; + assert_eq!(resolved, known_id); + } + + #[tokio::test] + async fn test_resolve_thread_with_parsed_uuid_none_delegates_to_parse() { + use crate::agent::session::Thread; + + let manager = SessionManager::new(); + let (session, _) = manager.resolve_thread("user2", "chan2", None).await; + + // Insert a thread with a known UUID + let known_id = Uuid::new_v4(); + { + let mut sess = session.lock().await; + let thread = Thread::with_id(known_id, sess.id); + sess.threads.insert(known_id, thread); + } + + // Resolve with parsed_uuid=None but a valid UUID string -- should + // fall back to parsing the string and still adopt the thread + let (_, resolved) = manager + .resolve_thread_with_parsed_uuid("user2", "chan2", Some(&known_id.to_string()), None) + .await; + assert_eq!(resolved, known_id); + } + + #[tokio::test] + async fn test_resolve_thread_with_none_external_thread_id_does_not_adopt() { + use crate::agent::session::Thread; + + let manager = SessionManager::new(); + let (session, default_tid) = manager.resolve_thread("user3", "chan3", None).await; + + // Manually insert a thread with a known UUID (simulating a thread + // created by chat_new_thread_handler) + let known_id = Uuid::new_v4(); + { + let mut sess = session.lock().await; + let thread = Thread::with_id(known_id, sess.id); + sess.threads.insert(known_id, thread); + } + + // Resolve with external_thread_id=None but parsed_uuid=Some. + // This should NOT adopt the UUID — the old code prevented adoption + // when external_thread_id was None, and we preserve that invariant. + let (_, resolved) = manager + .resolve_thread_with_parsed_uuid("user3", "chan3", None, Some(known_id)) + .await; + + // Should return the existing default thread, not the injected UUID + assert_eq!( + resolved, default_tid, + "should return existing default thread when external_thread_id is None" + ); + assert_ne!( + resolved, known_id, + "should NOT adopt UUID when external_thread_id is None" + ); + } } From dcb2d89e3a5ed19b30878557adfe505b66484483 Mon Sep 17 00:00:00 2001 From: Henry Park Date: Tue, 24 Mar 2026 13:51:30 -0700 Subject: [PATCH 11/20] Fix hosted OAuth refresh via proxy (#1602) * Fix hosted OAuth refresh via proxy * Address OAuth refresh review feedback * Address new OAuth refresh review comments * Address additional OAuth refresh review feedback * Harden proxy exchange redirects --- src/cli/oauth_defaults.rs | 424 +++++++++++++++++-- src/extensions/manager.rs | 35 +- src/tools/wasm/loader.rs | 141 +++++++ src/tools/wasm/wrapper.rs | 481 ++++++++++++++++++++-- tests/e2e/CLAUDE.md | 9 + tests/e2e/conftest.py | 128 +++++- tests/e2e/mock_llm.py | 61 +++ tests/e2e/scenarios/test_oauth_refresh.py | 227 ++++++++++ 8 files changed, 1407 insertions(+), 99 deletions(-) create mode 100644 tests/e2e/scenarios/test_oauth_refresh.py diff --git a/src/cli/oauth_defaults.rs b/src/cli/oauth_defaults.rs index 3b57872f..e9001909 100644 --- a/src/cli/oauth_defaults.rs +++ b/src/cli/oauth_defaults.rs @@ -62,6 +62,30 @@ pub fn builtin_client_id_override_env(secret_name: &str) -> Option<&'static str> } } +/// Suppress the baked-in desktop OAuth client secret when a hosted proxy is configured. +/// +/// In hosted deployments, IronClaw may resolve the platform Google client ID from +/// environment variables while still falling back to the baked-in desktop secret. +/// That client_id/client_secret mismatch breaks Google token exchange and refresh. +/// +/// When the proxy is configured, the platform will inject the correct server-side +/// secret for matching platform credentials, so the baked-in secret must be omitted. +pub fn hosted_proxy_client_secret( + client_secret: &Option, + builtin: Option<&OAuthCredentials>, + exchange_proxy_configured: bool, +) -> Option { + if !exchange_proxy_configured { + return client_secret.clone(); + } + + let builtin_secret = builtin.map(|credentials| credentials.client_secret); + match (client_secret, builtin_secret) { + (Some(resolved), Some(baked_in)) if resolved == baked_in => None, + _ => client_secret.clone(), + } +} + // ── Shared callback server ────────────────────────────────────────────── // Core OAuth callback infrastructure is defined in `crate::llm::oauth_helpers` @@ -661,6 +685,48 @@ pub struct ProxyTokenExchangeRequest<'a> { pub extra_token_params: &'a HashMap, } +pub struct ProxyRefreshTokenRequest<'a> { + pub proxy_url: &'a str, + pub gateway_token: &'a str, + pub token_url: &'a str, + pub client_id: &'a str, + pub client_secret: Option<&'a str>, + pub refresh_token: &'a str, + pub provider: Option<&'a str>, +} + +fn oauth_token_response_from_json( + token_data: serde_json::Value, + access_token_field: &str, +) -> Result { + let access_token = token_data + .get(access_token_field) + .and_then(|v| v.as_str()) + .ok_or_else(|| { + let fields: Vec<&str> = token_data + .as_object() + .map(|o| o.keys().map(|k| k.as_str()).collect()) + .unwrap_or_default(); + OAuthCallbackError::Io(format!( + "No '{}' field in proxy response (fields present: {:?})", + access_token_field, fields + )) + })? + .to_string(); + + let refresh_token = token_data + .get("refresh_token") + .and_then(|v| v.as_str()) + .map(String::from); + let expires_in = token_data.get("expires_in").and_then(|v| v.as_u64()); + + Ok(OAuthTokenResponse { + access_token, + refresh_token, + expires_in, + }) +} + /// Exchange an OAuth authorization code via the platform's token exchange proxy. /// /// Authenticated via the gateway auth token (Bearer header). The caller may @@ -682,6 +748,7 @@ pub async fn exchange_via_proxy( let client = reqwest::Client::builder() .timeout(Duration::from_secs(60)) + .redirect(reqwest::redirect::Policy::none()) .build() .map_err(|e| OAuthCallbackError::Io(format!("Failed to build HTTP client: {}", e)))?; let mut params = vec![ @@ -724,41 +791,350 @@ pub async fn exchange_via_proxy( .json() .await .map_err(|e| OAuthCallbackError::Io(format!("Failed to parse proxy response: {}", e)))?; + oauth_token_response_from_json(token_data, request.access_token_field) +} - let access_token = token_data - .get(request.access_token_field) - .and_then(|v| v.as_str()) - .ok_or_else(|| { - let fields: Vec<&str> = token_data - .as_object() - .map(|o| o.keys().map(|k| k.as_str()).collect()) - .unwrap_or_default(); - OAuthCallbackError::Io(format!( - "No '{}' field in proxy response (fields present: {:?})", - request.access_token_field, fields - )) - })? - .to_string(); +/// Refresh an OAuth access token via the platform's token refresh proxy. +/// +/// Authenticated via the gateway auth token (Bearer header). The caller may +/// either rely on proxy-side secret lookup or forward a `client_secret` when +/// the provider requires it. +pub async fn refresh_token_via_proxy( + request: ProxyRefreshTokenRequest<'_>, +) -> Result { + if request.gateway_token.is_empty() { + return Err(OAuthCallbackError::Io( + "Gateway auth token is required for proxy token refresh".to_string(), + )); + } - let refresh_token = token_data - .get("refresh_token") - .and_then(|v| v.as_str()) - .map(String::from); - let expires_in = token_data.get("expires_in").and_then(|v| v.as_u64()); + let refresh_url = format!("{}/oauth/refresh", request.proxy_url.trim_end_matches('/')); + let client = reqwest::Client::builder() + .timeout(Duration::from_secs(15)) + .redirect(reqwest::redirect::Policy::none()) + .build() + .map_err(|e| OAuthCallbackError::Io(format!("Failed to build HTTP client: {}", e)))?; - Ok(OAuthTokenResponse { - access_token, - refresh_token, - expires_in, - }) + let mut params = vec![ + ("refresh_token", request.refresh_token.to_string()), + ("token_url", request.token_url.to_string()), + ("client_id", request.client_id.to_string()), + ]; + if let Some(secret) = request.client_secret { + params.push(("client_secret", secret.to_string())); + } + if let Some(provider) = request.provider { + params.push(("provider", provider.to_string())); + } + + let response = client + .post(&refresh_url) + .bearer_auth(request.gateway_token) + .form(¶ms) + .send() + .await + .map_err(|e| { + OAuthCallbackError::Io(format!("Token refresh proxy request failed: {}", e)) + })?; + + if !response.status().is_success() { + let status = response.status(); + let body = response.text().await.unwrap_or_default(); + return Err(OAuthCallbackError::Io(format!( + "Token refresh proxy failed: {} - {}", + status, body + ))); + } + + let token_data: serde_json::Value = response + .json() + .await + .map_err(|e| OAuthCallbackError::Io(format!("Failed to parse proxy response: {}", e)))?; + + oauth_token_response_from_json(token_data, "access_token") } #[cfg(test)] mod tests { + use std::collections::HashMap; + use std::net::SocketAddr; + use std::sync::Arc; + + use axum::extract::{Form, State}; + use axum::http::HeaderMap; + use axum::response::Redirect; + use axum::routing::post; + use axum::{Json, Router}; + use serde_json::json; + use tokio::net::TcpListener; + use tokio::sync::{Mutex, oneshot}; + use crate::cli::oauth_defaults::{ builtin_credentials, callback_host, callback_url, is_loopback_host, landing_html, }; use crate::config::helpers::lock_env; + use crate::testing::credentials::{TEST_OAUTH_CLIENT_ID, TEST_OAUTH_CLIENT_SECRET}; + + #[derive(Clone, Debug, PartialEq, Eq)] + struct RecordedProxyRequest { + authorization: Option, + form: HashMap, + } + + #[derive(Clone)] + struct MockProxyState { + requests: Arc>>, + exchange_redirect_target: String, + refresh_redirect_target: String, + } + + struct MockProxyServer { + addr: SocketAddr, + requests: Arc>>, + shutdown_tx: Option>, + server_task: Option>, + } + + impl MockProxyServer { + async fn start() -> Self { + async fn exchange_handler( + State(state): State, + headers: HeaderMap, + Form(form): Form>, + ) -> Json { + state.requests.lock().await.push(RecordedProxyRequest { + authorization: headers + .get(axum::http::header::AUTHORIZATION) + .and_then(|value| value.to_str().ok()) + .map(str::to_string), + form, + }); + Json(json!({ + "access_token": "proxy-access-token", + "refresh_token": "proxy-refresh-token", + "expires_in": 7200 + })) + } + + async fn refresh_handler( + State(state): State, + headers: HeaderMap, + Form(form): Form>, + ) -> Json { + state.requests.lock().await.push(RecordedProxyRequest { + authorization: headers + .get(axum::http::header::AUTHORIZATION) + .and_then(|value| value.to_str().ok()) + .map(str::to_string), + form, + }); + Json(json!({ + "access_token": "proxy-access-token", + "refresh_token": "proxy-refresh-token", + "expires_in": 7200 + })) + } + + async fn exchange_redirect_handler(State(state): State) -> Redirect { + Redirect::temporary(&state.exchange_redirect_target) + } + + async fn refresh_redirect_handler(State(state): State) -> Redirect { + Redirect::temporary(&state.refresh_redirect_target) + } + + let requests = Arc::new(Mutex::new(Vec::new())); + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("bind mock proxy"); + let addr = listener.local_addr().expect("read mock proxy addr"); + let exchange_redirect_target = format!("http://{addr}/oauth/exchange"); + let refresh_redirect_target = format!("http://{addr}/oauth/refresh"); + let app = Router::new() + .route("/oauth/exchange", post(exchange_handler)) + .route("/oauth/refresh", post(refresh_handler)) + .route("/redirect/oauth/exchange", post(exchange_redirect_handler)) + .route("/redirect/oauth/refresh", post(refresh_redirect_handler)) + .with_state(MockProxyState { + requests: Arc::clone(&requests), + exchange_redirect_target, + refresh_redirect_target, + }); + let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>(); + let server_task = tokio::spawn(async move { + let _ = axum::serve(listener, app) + .with_graceful_shutdown(async { + let _ = shutdown_rx.await; + }) + .await; + }); + + Self { + addr, + requests, + shutdown_tx: Some(shutdown_tx), + server_task: Some(server_task), + } + } + + fn base_url(&self) -> String { + format!("http://{}", self.addr) + } + + fn redirecting_base_url(&self) -> String { + format!("{}/redirect", self.base_url()) + } + + async fn requests(&self) -> Vec { + self.requests.lock().await.clone() + } + + async fn shutdown(mut self) { + if let Some(tx) = self.shutdown_tx.take() { + let _ = tx.send(()); + } + if let Some(task) = self.server_task.take() { + let _ = task.await; + } + } + } + + impl Drop for MockProxyServer { + fn drop(&mut self) { + if let Some(tx) = self.shutdown_tx.take() { + let _ = tx.send(()); + } + if let Some(task) = self.server_task.take() { + task.abort(); + } + } + } + + #[test] + fn test_hosted_proxy_client_secret_suppresses_builtin_secret() { + let builtin = builtin_credentials("google_oauth_token").expect("google builtin creds"); + let client_secret = Some(builtin.client_secret.to_string()); + + let result = super::hosted_proxy_client_secret(&client_secret, Some(&builtin), true); + + assert_eq!(result, None); + } + + #[test] + fn test_hosted_proxy_client_secret_preserves_explicit_secret() { + let builtin = builtin_credentials("google_oauth_token").expect("google builtin creds"); + let client_secret = Some("hosted-server-secret".to_string()); + + let result = super::hosted_proxy_client_secret(&client_secret, Some(&builtin), true); + + assert_eq!(result, client_secret); + } + + #[tokio::test] + async fn test_refresh_token_via_proxy_sends_auth_and_form() { + let server = MockProxyServer::start().await; + + let response = super::refresh_token_via_proxy(super::ProxyRefreshTokenRequest { + proxy_url: &server.base_url(), + gateway_token: "gateway-test-token", + token_url: "https://oauth2.googleapis.com/token", + client_id: TEST_OAUTH_CLIENT_ID, + client_secret: Some(TEST_OAUTH_CLIENT_SECRET), + refresh_token: "refresh-token-123", + provider: Some("google"), + }) + .await + .expect("proxy refresh succeeds"); + + assert_eq!(response.access_token, "proxy-access-token"); + assert_eq!( + response.refresh_token.as_deref(), + Some("proxy-refresh-token") + ); + assert_eq!(response.expires_in, Some(7200)); + + let requests = server.requests().await; + assert_eq!(requests.len(), 1); + assert_eq!( + requests[0].authorization.as_deref(), + Some("Bearer gateway-test-token") + ); + assert_eq!( + requests[0].form.get("token_url").map(String::as_str), + Some("https://oauth2.googleapis.com/token") + ); + assert_eq!( + requests[0].form.get("client_id").map(String::as_str), + Some(TEST_OAUTH_CLIENT_ID) + ); + assert_eq!( + requests[0].form.get("client_secret").map(String::as_str), + Some(TEST_OAUTH_CLIENT_SECRET) + ); + assert_eq!( + requests[0].form.get("refresh_token").map(String::as_str), + Some("refresh-token-123") + ); + assert_eq!( + requests[0].form.get("provider").map(String::as_str), + Some("google") + ); + + server.shutdown().await; + } + + #[tokio::test] + async fn test_exchange_via_proxy_does_not_follow_redirects() { + let server = MockProxyServer::start().await; + + let error = match super::exchange_via_proxy(super::ProxyTokenExchangeRequest { + proxy_url: &server.redirecting_base_url(), + gateway_token: "gateway-test-token", + code: "auth-code-123", + redirect_uri: "http://localhost:3000/oauth/callback", + token_url: "https://oauth2.googleapis.com/token", + client_id: TEST_OAUTH_CLIENT_ID, + client_secret: Some(TEST_OAUTH_CLIENT_SECRET), + access_token_field: "access_token", + code_verifier: Some("code-verifier-123"), + extra_token_params: &HashMap::new(), + }) + .await + { + Ok(_) => panic!("redirected proxy exchange should fail"), + Err(error) => error, + }; + + assert!(error.to_string().contains("307")); + assert!(server.requests().await.is_empty()); + + server.shutdown().await; + } + + #[tokio::test] + async fn test_refresh_token_via_proxy_does_not_follow_redirects() { + let server = MockProxyServer::start().await; + + let error = match super::refresh_token_via_proxy(super::ProxyRefreshTokenRequest { + proxy_url: &server.redirecting_base_url(), + gateway_token: "gateway-test-token", + token_url: "https://oauth2.googleapis.com/token", + client_id: TEST_OAUTH_CLIENT_ID, + client_secret: Some(TEST_OAUTH_CLIENT_SECRET), + refresh_token: "refresh-token-123", + provider: Some("google"), + }) + .await + { + Ok(_) => panic!("redirected proxy refresh should fail"), + Err(error) => error, + }; + + assert!(error.to_string().contains("307")); + assert!(server.requests().await.is_empty()); + + server.shutdown().await; + } #[test] fn test_is_loopback_host() { diff --git a/src/extensions/manager.rs b/src/extensions/manager.rs index 39654305..0f308352 100644 --- a/src/extensions/manager.rs +++ b/src/extensions/manager.rs @@ -53,22 +53,6 @@ struct HostedOAuthFlowStart { flow: crate::cli::oauth_defaults::PendingOAuthFlow, } -fn hosted_proxy_client_secret( - client_secret: &Option, - builtin: Option<&crate::cli::oauth_defaults::OAuthCredentials>, - exchange_proxy_configured: bool, -) -> Option { - if !exchange_proxy_configured { - return client_secret.clone(); - } - - let builtin_secret = builtin.map(|credentials| credentials.client_secret); - match (client_secret, builtin_secret) { - (Some(resolved), Some(baked_in)) if resolved == baked_in => None, - _ => client_secret.clone(), - } -} - fn normalize_oauth_callback_path(path: &str) -> String { let trimmed_path = path.trim_end_matches('/'); if trimmed_path.is_empty() { @@ -3199,7 +3183,7 @@ impl ExtensionManager { // apps. Sending the desktop secret would cause a client_id/secret // mismatch because the container's GOOGLE_OAUTH_CLIENT_ID is the web // app, not the desktop app. - let proxy_client_secret = hosted_proxy_client_secret( + let proxy_client_secret = oauth_defaults::hosted_proxy_client_secret( &client_secret, builtin.as_ref(), oauth_defaults::exchange_proxy_url().is_some(), @@ -5714,7 +5698,7 @@ mod tests { use crate::extensions::manager::{ ChannelRuntimeState, FallbackDecision, TelegramBindingData, TelegramBindingResult, TelegramOwnerBindingState, build_wasm_channel_runtime_config_updates, - combine_install_errors, fallback_decision, hosted_proxy_client_secret, infer_kind_from_url, + combine_install_errors, fallback_decision, infer_kind_from_url, normalize_hosted_callback_url, send_telegram_text_message, telegram_message_matches_verification_code, }; @@ -7966,7 +7950,8 @@ mod tests { let builtin_ref = builtin.as_ref(); let secret = Some(builtin_ref.unwrap().client_secret.to_string()); - let result = hosted_proxy_client_secret(&secret, builtin_ref, true); + let result = + crate::cli::oauth_defaults::hosted_proxy_client_secret(&secret, builtin_ref, true); assert_eq!( result, None, "built-in desktop secret must be suppressed when the exchange proxy is configured" @@ -7978,7 +7963,8 @@ mod tests { let builtin = crate::cli::oauth_defaults::builtin_credentials("google_oauth_token"); let secret = Some("user-entered-custom-secret".to_string()); - let result = hosted_proxy_client_secret(&secret, builtin.as_ref(), true); + let result = + crate::cli::oauth_defaults::hosted_proxy_client_secret(&secret, builtin.as_ref(), true); assert_eq!( result, Some("user-entered-custom-secret".to_string()), @@ -7992,7 +7978,8 @@ mod tests { let builtin_ref = builtin.as_ref(); let secret = Some(builtin_ref.unwrap().client_secret.to_string()); - let result = hosted_proxy_client_secret(&secret, builtin_ref, false); + let result = + crate::cli::oauth_defaults::hosted_proxy_client_secret(&secret, builtin_ref, false); assert_eq!( result, secret, "built-in secret must be kept when the callback will exchange directly" @@ -8003,7 +7990,8 @@ mod tests { fn test_proxy_client_secret_none_stays_none() { let builtin = crate::cli::oauth_defaults::builtin_credentials("google_oauth_token"); - let result = hosted_proxy_client_secret(&None, builtin.as_ref(), true); + let result = + crate::cli::oauth_defaults::hosted_proxy_client_secret(&None, builtin.as_ref(), true); assert_eq!( result, None, "None secret stays None even when the exchange proxy is configured" @@ -8017,7 +8005,8 @@ mod tests { assert!(builtin.is_none()); let secret = Some("dcr-secret".to_string()); - let result = hosted_proxy_client_secret(&secret, builtin.as_ref(), true); + let result = + crate::cli::oauth_defaults::hosted_proxy_client_secret(&secret, builtin.as_ref(), true); assert_eq!( result, Some("dcr-secret".to_string()), diff --git a/src/tools/wasm/loader.rs b/src/tools/wasm/loader.rs index b50fc717..2a7ed040 100644 --- a/src/tools/wasm/loader.rs +++ b/src/tools/wasm/loader.rs @@ -418,6 +418,7 @@ fn resolve_oauth_refresh_config(cap_file: &CapabilitiesFile) -> Option Option, + } + + impl Drop for EnvVarGuard { + fn drop(&mut self) { + // SAFETY: Tests use lock_env() to serialize environment access. + unsafe { + if let Some(ref value) = self.previous { + std::env::set_var(&self.key, value); + } else { + std::env::remove_var(&self.key); + } + } + } + } + + fn set_env_var(key: &str, value: Option<&str>) -> EnvVarGuard { + let previous = std::env::var(key).ok(); + // SAFETY: Tests use lock_env() to serialize environment access. + unsafe { + match value { + Some(value) => std::env::set_var(key, value), + None => std::env::remove_var(key), + } + } + EnvVarGuard { + key: key.to_string(), + previous, + } + } + #[test] fn wit_version_compat_none_is_ok() { // Pre-versioning extensions (no wit_version declared) should always pass @@ -871,6 +917,8 @@ mod tests { config.client_secret, Some(TEST_OAUTH_CLIENT_SECRET.to_string()) ); + assert_eq!(config.exchange_proxy_url, None); + assert_eq!(config.gateway_token, None); assert_eq!(config.secret_name, "google_oauth_token"); assert_eq!(config.provider, Some("google".to_string())); } @@ -931,6 +979,10 @@ mod tests { AuthCapabilitySchema, CapabilitiesFile, OAuthConfigSchema, }; + let _guard = lock_env(); + let _proxy_guard = set_env_var("IRONCLAW_OAUTH_EXCHANGE_URL", None); + let _gateway_token_guard = set_env_var("GATEWAY_AUTH_TOKEN", None); + // google_oauth_token should fall back to built-in credentials let caps = CapabilitiesFile { auth: Some(AuthCapabilitySchema { @@ -952,6 +1004,95 @@ mod tests { let config = config.unwrap(); assert!(!config.client_id.is_empty()); assert!(config.client_secret.is_some()); + assert_eq!(config.exchange_proxy_url, None); + assert_eq!(config.gateway_token, None); + } + + #[test] + fn test_resolve_oauth_refresh_config_hosted_proxy_populates_env_and_suppresses_builtin_secret() + { + use crate::tools::wasm::capabilities_schema::{ + AuthCapabilitySchema, CapabilitiesFile, OAuthConfigSchema, + }; + + let _guard = lock_env(); + let _proxy_guard = set_env_var( + "IRONCLAW_OAUTH_EXCHANGE_URL", + Some("https://compose-api.example.com"), + ); + let _gateway_token_guard = set_env_var("GATEWAY_AUTH_TOKEN", Some("gateway-test-token")); + let _client_id_guard = + set_env_var("GOOGLE_OAUTH_CLIENT_ID", Some("hosted-google-client-id")); + + let caps = CapabilitiesFile { + auth: Some(AuthCapabilitySchema { + secret_name: "google_oauth_token".to_string(), + provider: Some("google".to_string()), + oauth: Some(OAuthConfigSchema { + authorization_url: "https://accounts.google.com/o/oauth2/v2/auth".to_string(), + token_url: "https://oauth2.googleapis.com/token".to_string(), + client_id_env: Some("GOOGLE_OAUTH_CLIENT_ID".to_string()), + ..Default::default() + }), + ..Default::default() + }), + ..Default::default() + }; + + let config = super::resolve_oauth_refresh_config(&caps).expect("hosted oauth config"); + assert_eq!(config.client_id, "hosted-google-client-id"); + assert_eq!(config.client_secret, None); + assert_eq!( + config.exchange_proxy_url.as_deref(), + Some("https://compose-api.example.com") + ); + assert_eq!(config.gateway_token.as_deref(), Some("gateway-test-token")); + } + + #[test] + fn test_resolve_oauth_refresh_config_hosted_proxy_preserves_explicit_secret() { + use crate::tools::wasm::capabilities_schema::{ + AuthCapabilitySchema, CapabilitiesFile, OAuthConfigSchema, + }; + + let _guard = lock_env(); + let _proxy_guard = set_env_var( + "IRONCLAW_OAUTH_EXCHANGE_URL", + Some("https://compose-api.example.com"), + ); + let _gateway_token_guard = set_env_var("GATEWAY_AUTH_TOKEN", Some("gateway-test-token")); + let _client_id_guard = + set_env_var("GOOGLE_OAUTH_CLIENT_ID", Some("hosted-google-client-id")); + let _client_secret_guard = + set_env_var("GOOGLE_OAUTH_CLIENT_SECRET", Some("hosted-server-secret")); + + let caps = CapabilitiesFile { + auth: Some(AuthCapabilitySchema { + secret_name: "google_oauth_token".to_string(), + provider: Some("google".to_string()), + oauth: Some(OAuthConfigSchema { + authorization_url: "https://accounts.google.com/o/oauth2/v2/auth".to_string(), + token_url: "https://oauth2.googleapis.com/token".to_string(), + client_id_env: Some("GOOGLE_OAUTH_CLIENT_ID".to_string()), + client_secret_env: Some("GOOGLE_OAUTH_CLIENT_SECRET".to_string()), + ..Default::default() + }), + ..Default::default() + }), + ..Default::default() + }; + + let config = super::resolve_oauth_refresh_config(&caps).expect("hosted oauth config"); + assert_eq!(config.client_id, "hosted-google-client-id"); + assert_eq!( + config.client_secret.as_deref(), + Some("hosted-server-secret") + ); + assert_eq!( + config.exchange_proxy_url.as_deref(), + Some("https://compose-api.example.com") + ); + assert_eq!(config.gateway_token.as_deref(), Some("gateway-test-token")); } // --------------------------------------------------------------- diff --git a/src/tools/wasm/wrapper.rs b/src/tools/wasm/wrapper.rs index 33fcedb9..05508e97 100644 --- a/src/tools/wasm/wrapper.rs +++ b/src/tools/wasm/wrapper.rs @@ -19,7 +19,7 @@ use wasmtime_wasi::{ResourceTable, WasiCtx, WasiCtxBuilder, WasiView}; use crate::context::JobContext; use crate::llm::recording::{HttpExchangeRequest, HttpExchangeResponse, HttpInterceptor}; use crate::safety::LeakDetector; -use crate::secrets::SecretsStore; +use crate::secrets::{DecryptedSecret, SecretsStore}; use crate::tools::tool::{Tool, ToolError, ToolOutput}; use crate::tools::wasm::capabilities::Capabilities; use crate::tools::wasm::credential_injector::{ @@ -44,6 +44,7 @@ wasmtime::component::bindgen!({ }); // Alias the export interface types for convenience. +use crate::cli::oauth_defaults; use exports::near::agent::tool as wit_tool; /// Configuration needed to refresh an expired OAuth access token. @@ -59,6 +60,10 @@ pub struct OAuthRefreshConfig { pub client_id: String, /// OAuth client_secret (optional, some providers use PKCE without a secret). pub client_secret: Option, + /// Hosted OAuth proxy base URL (e.g., "http://host.docker.internal:8080"). + pub exchange_proxy_url: Option, + /// Gateway auth token for authenticating with the hosted OAuth proxy. + pub gateway_token: Option, /// Secret name of the access token (e.g., "google_oauth_token"). /// The refresh token lives at `{secret_name}_refresh_token`. pub secret_name: String, @@ -1210,6 +1215,53 @@ async fn refresh_oauth_token( user_id: &str, config: &OAuthRefreshConfig, ) -> bool { + let refresh_name = format!("{}_refresh_token", config.secret_name); + + if let Some(proxy_url) = config.exchange_proxy_url.as_deref() { + let Some(gateway_token) = config.gateway_token.as_deref() else { + tracing::warn!( + "OAuth refresh proxy is configured, but no gateway auth token is available" + ); + return false; + }; + + // In hosted mode, the configured exchange proxy owns the outbound token + // refresh and validation policy for the provider token_url. Direct-mode + // HTTPS/private-IP checks remain in place for self-hosted refreshes below. + let refresh_secret = match load_oauth_refresh_secret(store, user_id, &refresh_name).await { + Some(secret) => secret, + None => return false, + }; + let token_response = match oauth_defaults::refresh_token_via_proxy( + oauth_defaults::ProxyRefreshTokenRequest { + proxy_url, + gateway_token, + token_url: &config.token_url, + client_id: &config.client_id, + client_secret: config.client_secret.as_deref(), + refresh_token: refresh_secret.expose(), + provider: config.provider.as_deref(), + }, + ) + .await + { + Ok(response) => response, + Err(error) => { + tracing::warn!(error = %error, "OAuth token refresh via proxy failed"); + return false; + } + }; + + return persist_refreshed_oauth_tokens( + store, + user_id, + config, + &refresh_name, + token_response, + ) + .await; + } + // SSRF defense: token_url comes from the tool's capabilities file. if !config.token_url.starts_with("https://") { tracing::warn!( @@ -1227,19 +1279,6 @@ async fn refresh_oauth_token( return false; } - let refresh_name = format!("{}_refresh_token", config.secret_name); - let refresh_secret = match store.get_decrypted(user_id, &refresh_name).await { - Ok(s) => s, - Err(e) => { - tracing::debug!( - secret_name = %refresh_name, - error = %e, - "No refresh token available, skipping token refresh" - ); - return false; - } - }; - let client = match reqwest::Client::builder() .timeout(Duration::from_secs(15)) .redirect(reqwest::redirect::Policy::none()) @@ -1252,6 +1291,10 @@ async fn refresh_oauth_token( } }; + let refresh_secret = match load_oauth_refresh_secret(store, user_id, &refresh_name).await { + Some(secret) => secret, + None => return false, + }; let mut params = vec![ ("grant_type", "refresh_token".to_string()), ("refresh_token", refresh_secret.expose().to_string()), @@ -1287,22 +1330,55 @@ async fn refresh_oauth_token( return false; } }; - - let new_access_token = match token_data.get("access_token").and_then(|v| v.as_str()) { - Some(t) => t, + let token_response = match token_data.get("access_token").and_then(|v| v.as_str()) { + Some(access_token) => oauth_defaults::OAuthTokenResponse { + access_token: access_token.to_string(), + refresh_token: token_data + .get("refresh_token") + .and_then(|v| v.as_str()) + .map(str::to_string), + expires_in: token_data.get("expires_in").and_then(|v| v.as_u64()), + }, None => { tracing::warn!("Token refresh response missing access_token field"); return false; } }; - // Store the new access token with expiry + persist_refreshed_oauth_tokens(store, user_id, config, &refresh_name, token_response).await +} + +async fn load_oauth_refresh_secret( + store: &(dyn SecretsStore + Send + Sync), + user_id: &str, + refresh_name: &str, +) -> Option { + match store.get_decrypted(user_id, refresh_name).await { + Ok(secret) => Some(secret), + Err(error) => { + tracing::debug!( + secret_name = %refresh_name, + error = %error, + "No refresh token available, skipping token refresh" + ); + None + } + } +} + +async fn persist_refreshed_oauth_tokens( + store: &(dyn SecretsStore + Send + Sync), + user_id: &str, + config: &OAuthRefreshConfig, + refresh_name: &str, + token_response: oauth_defaults::OAuthTokenResponse, +) -> bool { let mut access_params = - crate::secrets::CreateSecretParams::new(&config.secret_name, new_access_token); + crate::secrets::CreateSecretParams::new(&config.secret_name, &token_response.access_token); if let Some(ref provider) = config.provider { access_params = access_params.with_provider(provider); } - if let Some(expires_in) = token_data.get("expires_in").and_then(|v| v.as_u64()) { + if let Some(expires_in) = token_response.expires_in { let expires_at = chrono::Utc::now() + chrono::Duration::seconds(expires_in as i64); access_params = access_params.with_expiry(expires_at); } @@ -1312,10 +1388,8 @@ async fn refresh_oauth_token( return false; } - // Store rotated refresh token if the provider sent a new one - if let Some(new_refresh) = token_data.get("refresh_token").and_then(|v| v.as_str()) { - let mut refresh_params = - crate::secrets::CreateSecretParams::new(&refresh_name, new_refresh); + if let Some(new_refresh) = token_response.refresh_token.as_deref() { + let mut refresh_params = crate::secrets::CreateSecretParams::new(refresh_name, new_refresh); if let Some(ref provider) = config.provider { refresh_params = refresh_params.with_provider(provider); } @@ -1664,9 +1738,18 @@ fn build_tool_usage_hint(tool_name: &str, schema: &serde_json::Value) -> String #[cfg(test)] mod tests { + use std::collections::HashMap; + use std::net::SocketAddr; use std::sync::{Arc, Mutex}; use async_trait::async_trait; + use axum::extract::{Form, State}; + use axum::http::HeaderMap; + use axum::routing::post; + use axum::{Json, Router}; + use serde_json::json; + use tokio::net::TcpListener; + use tokio::sync::{Mutex as AsyncMutex, oneshot}; use uuid::Uuid; use crate::context::JobContext; @@ -1756,6 +1839,95 @@ mod tests { } } + #[derive(Clone, Debug, PartialEq, Eq)] + struct RecordedProxyRequest { + authorization: Option, + form: HashMap, + } + + struct MockProxyServer { + addr: SocketAddr, + requests: Arc>>, + shutdown_tx: Option>, + server_task: Option>, + } + + impl MockProxyServer { + async fn start() -> Self { + async fn refresh_handler( + State(requests): State>>>, + headers: HeaderMap, + Form(form): Form>, + ) -> Json { + requests.lock().await.push(RecordedProxyRequest { + authorization: headers + .get(axum::http::header::AUTHORIZATION) + .and_then(|value| value.to_str().ok()) + .map(str::to_string), + form, + }); + Json(json!({ + "access_token": "mock-refreshed-access-token", + "refresh_token": "mock-rotated-refresh-token", + "expires_in": 3600 + })) + } + + let requests = Arc::new(AsyncMutex::new(Vec::new())); + let app = Router::new() + .route("/oauth/refresh", post(refresh_handler)) + .with_state(Arc::clone(&requests)); + + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("bind mock proxy"); + let addr = listener.local_addr().expect("read mock proxy addr"); + let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>(); + let server_task = tokio::spawn(async move { + let _ = axum::serve(listener, app) + .with_graceful_shutdown(async { + let _ = shutdown_rx.await; + }) + .await; + }); + + Self { + addr, + requests, + shutdown_tx: Some(shutdown_tx), + server_task: Some(server_task), + } + } + + fn base_url(&self) -> String { + format!("http://{}", self.addr) + } + + async fn requests(&self) -> Vec { + self.requests.lock().await.clone() + } + + async fn shutdown(mut self) { + if let Some(tx) = self.shutdown_tx.take() { + let _ = tx.send(()); + } + if let Some(task) = self.server_task.take() { + let _ = task.await; + } + } + } + + impl Drop for MockProxyServer { + fn drop(&mut self) { + if let Some(tx) = self.shutdown_tx.take() { + let _ = tx.send(()); + } + if let Some(task) = self.server_task.take() { + task.abort(); + } + } + } + #[test] fn test_wrapper_creation() { // This test verifies the runtime can be created @@ -2094,8 +2266,6 @@ mod tests { #[tokio::test] async fn test_resolve_host_credentials_bearer() { - use std::collections::HashMap; - use crate::secrets::{ CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore, }; @@ -2141,8 +2311,6 @@ mod tests { #[tokio::test] async fn test_resolve_host_credentials_owner_scope_bearer() { - use std::collections::HashMap; - use crate::secrets::{ CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore, }; @@ -2188,8 +2356,6 @@ mod tests { #[tokio::test] async fn test_execute_resolves_host_credentials_from_owner_scope_context() { - use std::collections::HashMap; - use crate::secrets::{CredentialLocation, CredentialMapping}; use crate::tools::wasm::capabilities::HttpCapability; @@ -2239,8 +2405,6 @@ mod tests { #[tokio::test] async fn test_resolve_host_credentials_missing_secret() { - use std::collections::HashMap; - use crate::secrets::{CredentialLocation, CredentialMapping}; use crate::tools::wasm::capabilities::HttpCapability; use crate::tools::wasm::wrapper::resolve_host_credentials; @@ -2272,8 +2436,6 @@ mod tests { #[tokio::test] async fn test_resolve_host_credentials_skips_refresh_when_not_expired() { - use std::collections::HashMap; - use crate::secrets::{ CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore, }; @@ -2315,6 +2477,8 @@ mod tests { token_url: "https://oauth2.googleapis.com/token".to_string(), client_id: TEST_OAUTH_CLIENT_ID.to_string(), client_secret: Some(TEST_OAUTH_CLIENT_SECRET.to_string()), + exchange_proxy_url: None, + gateway_token: None, secret_name: "google_oauth_token".to_string(), provider: Some("google".to_string()), }; @@ -2331,8 +2495,6 @@ mod tests { #[tokio::test] async fn test_resolve_host_credentials_skips_refresh_no_config() { - use std::collections::HashMap; - use crate::secrets::{ CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore, }; @@ -2376,8 +2538,6 @@ mod tests { #[tokio::test] async fn test_resolve_host_credentials_skips_refresh_no_expires_at() { - use std::collections::HashMap; - use crate::secrets::{ CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore, }; @@ -2417,6 +2577,8 @@ mod tests { token_url: "https://oauth2.googleapis.com/token".to_string(), client_id: TEST_OAUTH_CLIENT_ID.to_string(), client_secret: Some(TEST_OAUTH_CLIENT_SECRET.to_string()), + exchange_proxy_url: None, + gateway_token: None, secret_name: "google_oauth_token".to_string(), provider: Some("google".to_string()), }; @@ -2431,6 +2593,249 @@ mod tests { ); } + #[tokio::test] + async fn test_resolve_host_credentials_refreshes_via_proxy_without_direct_token_url_validation() + { + use crate::secrets::{ + CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore, + }; + use crate::tools::wasm::capabilities::HttpCapability; + use crate::tools::wasm::wrapper::{OAuthRefreshConfig, resolve_host_credentials}; + + let proxy = MockProxyServer::start().await; + let store = test_secrets_store(); + + store + .create( + "user1", + CreateSecretParams::new("google_oauth_token", "expired-access-token") + .with_expiry(chrono::Utc::now() - chrono::Duration::hours(1)), + ) + .await + .unwrap(); + store + .create( + "user1", + CreateSecretParams::new("google_oauth_token_refresh_token", "stored-refresh-token"), + ) + .await + .unwrap(); + + let mut credentials = HashMap::new(); + credentials.insert( + "google_oauth_token".to_string(), + CredentialMapping { + secret_name: "google_oauth_token".to_string(), + location: CredentialLocation::AuthorizationBearer, + host_patterns: vec!["www.googleapis.com".to_string()], + }, + ); + + let caps = Capabilities { + http: Some(HttpCapability { + credentials, + ..Default::default() + }), + ..Default::default() + }; + + let oauth_config = OAuthRefreshConfig { + token_url: "http://127.0.0.1:9/provider-token-endpoint".to_string(), + client_id: "hosted-google-client-id".to_string(), + client_secret: None, + exchange_proxy_url: Some(proxy.base_url()), + gateway_token: Some("gateway-test-token".to_string()), + secret_name: "google_oauth_token".to_string(), + provider: Some("google".to_string()), + }; + + let resolved = + resolve_host_credentials(&caps, Some(&store), "user1", Some(&oauth_config)).await; + assert_eq!(resolved.len(), 1); + assert_eq!( + resolved[0].headers.get("Authorization"), + Some(&"Bearer mock-refreshed-access-token".to_string()) + ); + + let access_secret = store.get("user1", "google_oauth_token").await.unwrap(); + assert!( + access_secret + .expires_at + .expect("refreshed access token expiry") + > chrono::Utc::now() + ); + let access_value = store + .get_decrypted("user1", "google_oauth_token") + .await + .unwrap(); + assert_eq!(access_value.expose(), "mock-refreshed-access-token"); + + let refresh_value = store + .get_decrypted("user1", "google_oauth_token_refresh_token") + .await + .unwrap(); + assert_eq!(refresh_value.expose(), "mock-rotated-refresh-token"); + + let requests = proxy.requests().await; + assert_eq!(requests.len(), 1); + assert_eq!( + requests[0].authorization.as_deref(), + Some("Bearer gateway-test-token") + ); + assert_eq!( + requests[0].form.get("client_id").map(String::as_str), + Some("hosted-google-client-id") + ); + assert_eq!( + requests[0].form.get("token_url").map(String::as_str), + Some("http://127.0.0.1:9/provider-token-endpoint") + ); + assert_eq!( + requests[0].form.get("refresh_token").map(String::as_str), + Some("stored-refresh-token") + ); + assert_eq!( + requests[0].form.get("provider").map(String::as_str), + Some("google") + ); + assert!(!requests[0].form.contains_key("client_secret")); + + proxy.shutdown().await; + } + + #[tokio::test] + async fn test_resolve_host_credentials_skips_refresh_token_lookup_without_gateway_token() { + use crate::secrets::{ + CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore, + }; + use crate::tools::wasm::capabilities::HttpCapability; + use crate::tools::wasm::wrapper::{OAuthRefreshConfig, resolve_host_credentials}; + + let store = RecordingSecretsStore::new(); + + store + .create( + "user1", + CreateSecretParams::new("google_oauth_token", "expired-access-token") + .with_expiry(chrono::Utc::now() - chrono::Duration::hours(1)), + ) + .await + .unwrap(); + store + .create( + "user1", + CreateSecretParams::new("google_oauth_token_refresh_token", "stored-refresh-token"), + ) + .await + .unwrap(); + + let mut credentials = HashMap::new(); + credentials.insert( + "google_oauth_token".to_string(), + CredentialMapping { + secret_name: "google_oauth_token".to_string(), + location: CredentialLocation::AuthorizationBearer, + host_patterns: vec!["www.googleapis.com".to_string()], + }, + ); + + let caps = Capabilities { + http: Some(HttpCapability { + credentials, + ..Default::default() + }), + ..Default::default() + }; + + let oauth_config = OAuthRefreshConfig { + token_url: "https://oauth2.googleapis.com/token".to_string(), + client_id: "hosted-google-client-id".to_string(), + client_secret: None, + exchange_proxy_url: Some("https://compose-api.example.com".to_string()), + gateway_token: None, + secret_name: "google_oauth_token".to_string(), + provider: Some("google".to_string()), + }; + + let resolved = + resolve_host_credentials(&caps, Some(&store), "user1", Some(&oauth_config)).await; + assert!(resolved.is_empty()); + + let lookups = store.decrypted_lookups(); + assert!(lookups.contains(&("user1".to_string(), "google_oauth_token".to_string()))); + assert!(!lookups.contains(&( + "user1".to_string(), + "google_oauth_token_refresh_token".to_string(), + ))); + } + + #[tokio::test] + async fn test_resolve_host_credentials_skips_refresh_token_lookup_for_invalid_direct_token_url() + { + use crate::secrets::{ + CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore, + }; + use crate::tools::wasm::capabilities::HttpCapability; + use crate::tools::wasm::wrapper::{OAuthRefreshConfig, resolve_host_credentials}; + + let store = RecordingSecretsStore::new(); + + store + .create( + "user1", + CreateSecretParams::new("google_oauth_token", "expired-access-token") + .with_expiry(chrono::Utc::now() - chrono::Duration::hours(1)), + ) + .await + .unwrap(); + store + .create( + "user1", + CreateSecretParams::new("google_oauth_token_refresh_token", "stored-refresh-token"), + ) + .await + .unwrap(); + + let mut credentials = HashMap::new(); + credentials.insert( + "google_oauth_token".to_string(), + CredentialMapping { + secret_name: "google_oauth_token".to_string(), + location: CredentialLocation::AuthorizationBearer, + host_patterns: vec!["www.googleapis.com".to_string()], + }, + ); + + let caps = Capabilities { + http: Some(HttpCapability { + credentials, + ..Default::default() + }), + ..Default::default() + }; + + let oauth_config = OAuthRefreshConfig { + token_url: "http://127.0.0.1:9/provider-token-endpoint".to_string(), + client_id: TEST_OAUTH_CLIENT_ID.to_string(), + client_secret: Some(TEST_OAUTH_CLIENT_SECRET.to_string()), + exchange_proxy_url: None, + gateway_token: None, + secret_name: "google_oauth_token".to_string(), + provider: Some("google".to_string()), + }; + + let resolved = + resolve_host_credentials(&caps, Some(&store), "user1", Some(&oauth_config)).await; + assert!(resolved.is_empty()); + + let lookups = store.decrypted_lookups(); + assert!(lookups.contains(&("user1".to_string(), "google_oauth_token".to_string()))); + assert!(!lookups.contains(&( + "user1".to_string(), + "google_oauth_token_refresh_token".to_string(), + ))); + } + #[test] fn test_is_private_ip_v4() { use std::net::IpAddr; diff --git a/tests/e2e/CLAUDE.md b/tests/e2e/CLAUDE.md index 0cf5e6dc..46b7b752 100644 --- a/tests/e2e/CLAUDE.md +++ b/tests/e2e/CLAUDE.md @@ -53,6 +53,7 @@ HEADED=1 pytest scenarios/ | `test_skills.py` | Skills tab UI visibility, ClawHub search (skipped if registry unreachable), install + remove lifecycle | | `test_sse_reconnect.py` | SSE reconnects after programmatic `eventSource.close()` + `connectSSE()`; history is reloaded after reconnect | | `test_tool_approval.py` | Approval card appears, buttons disable on approve/deny, parameters toggle via `page.evaluate("showApproval(...)")`; the waiting-approval regression uses a real HTTP tool call | +| `test_oauth_refresh.py` | Hosted Gmail OAuth regression: complete setup via `/oauth/callback`, expire the stored access token in libSQL, trigger a real `gmail` tool call through `/api/chat/send`, and verify refresh goes through the mock `/oauth/refresh` proxy without forwarding `client_secret` | ## `helpers.py` @@ -75,6 +76,7 @@ All fixtures are defined in `tests/e2e/conftest.py`. Running `pytest scenarios/` | `ironclaw_binary` | Checks `target/debug/ironclaw`; if absent, runs `cargo build --no-default-features --features libsql` (timeout 600s). | | `mock_llm_server` | Starts `mock_llm.py --port 0`, reads the assigned port from stdout, waits for `/v1/models` to return 200. Yields the base URL. | | `ironclaw_server` | Starts the ironclaw binary with a minimal env (see below), waits for `/api/health` (timeout 60s). Yields the base URL. On teardown sends **SIGINT** (not SIGTERM) so the tokio ctrl_c handler triggers a graceful shutdown and LLVM coverage data is flushed. | +| `hosted_oauth_refresh_server` | Starts a second ironclaw instance with a dedicated libSQL DB and `GOOGLE_OAUTH_CLIENT_ID=hosted-google-client-id`, while still pointing `IRONCLAW_OAUTH_EXCHANGE_URL` at `mock_llm.py`. Yields a dict with `base_url`, `db_path`, `gateway_user_id`, and `mock_llm_url` for the hosted refresh regression scenario. | | `browser` | Launches a single Chromium instance (headless by default; set `HEADED=1` for headed). Shared across all tests. | ### Function-scoped fixtures @@ -100,6 +102,8 @@ EMBEDDING_ENABLED=false, SKILLS_ENABLED=true ONBOARD_COMPLETED=true # prevents setup wizard ``` +The `hosted_oauth_refresh_server` fixture uses the same baseline, but with its own DB/home tempdirs and `GOOGLE_OAUTH_CLIENT_ID=hosted-google-client-id` so hosted OAuth flows exercise proxy credential injection instead of the baked-in desktop Google app. + The binary is also started with `--no-onboard`. Coverage env vars (`CARGO_LLVM_COV*`, `LLVM_*`, `CARGO_ENCODED_RUSTFLAGS`, `CARGO_INCREMENTAL`) are forwarded from the outer environment when present. ## Mock LLM (`mock_llm.py`) @@ -113,6 +117,11 @@ python mock_llm.py --port 0 It serves `POST /v1/chat/completions` (streaming + non-streaming) and `GET /v1/models`. Responses are pattern-matched from `CANNED_RESPONSES` against the last user message. Unmatched messages return `"I understand your request."`. The model name reported is always `"mock-model"`. +It also hosts OAuth test endpoints: +- `POST /oauth/exchange` for hosted auth-code exchange +- `POST /oauth/refresh` for hosted refresh-token exchange +- `GET /__mock/oauth/state` and `POST /__mock/oauth/reset` so HTTP E2E scenarios can assert exact proxy payloads and reset counters between setup and refresh assertions + To add a new canned response: ```python # In mock_llm.py diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py index 06c7da03..1496f93f 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -113,6 +113,15 @@ def _reserve_loopback_sockets(count: int) -> list[socket.socket]: raise +def _forward_coverage_env(env: dict[str, str]) -> None: + """Forward cargo-llvm-cov env vars into child processes when present.""" + cov_env_prefixes = ("CARGO_LLVM_COV", "LLVM_") + cov_env_extras = ("CARGO_ENCODED_RUSTFLAGS", "CARGO_INCREMENTAL") + for key, val in os.environ.items(): + if key.startswith(cov_env_prefixes) or key in cov_env_extras: + env[key] = val + + @pytest.fixture(scope="session") def ironclaw_binary(): """Ensure ironclaw binary is built. Returns the binary path.""" @@ -264,14 +273,7 @@ async def ironclaw_server( "IRONCLAW_OAUTH_CALLBACK_URL": "https://oauth.test.example/oauth/callback", "IRONCLAW_OAUTH_EXCHANGE_URL": mock_llm_server, } - # Forward LLVM coverage instrumentation env vars when present - # (allows cargo-llvm-cov to collect profraw data from E2E runs). - # Use prefix matching to stay resilient to cargo-llvm-cov changes. - COV_ENV_PREFIXES = ("CARGO_LLVM_COV", "LLVM_") - COV_ENV_EXTRAS = ("CARGO_ENCODED_RUSTFLAGS", "CARGO_INCREMENTAL") - for key, val in os.environ.items(): - if key.startswith(COV_ENV_PREFIXES) or key in COV_ENV_EXTRAS: - env[key] = val + _forward_coverage_env(env) proc = await asyncio.create_subprocess_exec( ironclaw_binary, "--no-onboard", stdin=asyncio.subprocess.DEVNULL, @@ -310,6 +312,109 @@ async def ironclaw_server( proc.kill() +@pytest.fixture(scope="session") +async def hosted_oauth_refresh_server( + ironclaw_binary, + mock_llm_server, + wasm_tools_dir, +): + """Start a hosted-mode ironclaw instance for OAuth refresh regression tests.""" + reserved = _reserve_loopback_sockets(2) + db_tmpdir = tempfile.TemporaryDirectory(prefix="ironclaw-e2e-hosted-oauth-db-") + home_tmpdir = tempfile.TemporaryDirectory(prefix="ironclaw-e2e-hosted-oauth-home-") + + try: + gateway_port = reserved[0].getsockname()[1] + http_port = reserved[1].getsockname()[1] + for sock in reserved: + if sock.fileno() != -1: + sock.close() + + db_path = os.path.join(db_tmpdir.name, "hosted-oauth-refresh.db") + home_dir = home_tmpdir.name + env = { + "PATH": os.environ.get("PATH", "/usr/bin:/bin"), + "HOME": home_dir, + "IRONCLAW_BASE_DIR": os.path.join(home_dir, ".ironclaw"), + "RUST_LOG": "ironclaw=info", + "RUST_BACKTRACE": "1", + "IRONCLAW_OWNER_ID": OWNER_SCOPE_ID, + "GATEWAY_ENABLED": "true", + "GATEWAY_HOST": "127.0.0.1", + "GATEWAY_PORT": str(gateway_port), + "GATEWAY_AUTH_TOKEN": AUTH_TOKEN, + "GATEWAY_USER_ID": OWNER_SCOPE_ID, + "HTTP_HOST": "127.0.0.1", + "HTTP_PORT": str(http_port), + "HTTP_WEBHOOK_SECRET": HTTP_WEBHOOK_SECRET, + "CLI_ENABLED": "false", + "LLM_BACKEND": "openai_compatible", + "LLM_BASE_URL": mock_llm_server, + "LLM_MODEL": "mock-model", + "DATABASE_BACKEND": "libsql", + "LIBSQL_PATH": db_path, + "SECRETS_MASTER_KEY": "0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef", + "SANDBOX_ENABLED": "false", + "SKILLS_ENABLED": "true", + "ROUTINES_ENABLED": "true", + "HEARTBEAT_ENABLED": "false", + "EMBEDDING_ENABLED": "false", + "WASM_ENABLED": "true", + "WASM_TOOLS_DIR": wasm_tools_dir, + "WASM_CHANNELS_DIR": _WASM_CHANNELS_TMPDIR.name, + "ONBOARD_COMPLETED": "true", + "IRONCLAW_OAUTH_CALLBACK_URL": "https://oauth.test.example/oauth/callback", + "IRONCLAW_OAUTH_EXCHANGE_URL": mock_llm_server, + "GOOGLE_OAUTH_CLIENT_ID": "hosted-google-client-id", + } + _forward_coverage_env(env) + + proc = await asyncio.create_subprocess_exec( + ironclaw_binary, "--no-onboard", + stdin=asyncio.subprocess.DEVNULL, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + env=env, + ) + base_url = f"http://127.0.0.1:{gateway_port}" + try: + await wait_for_ready(f"{base_url}/api/health", timeout=60) + yield { + "base_url": base_url, + "db_path": db_path, + "gateway_user_id": OWNER_SCOPE_ID, + "mock_llm_url": mock_llm_server, + } + except TimeoutError: + returncode = proc.returncode + stderr_bytes = b"" + if proc.stderr: + try: + stderr_bytes = await asyncio.wait_for(proc.stderr.read(8192), timeout=2) + except (asyncio.TimeoutError, Exception): + pass + stderr_text = stderr_bytes.decode("utf-8", errors="replace") + if proc.returncode is None: + proc.kill() + pytest.fail( + f"hosted oauth refresh server failed to start on port {gateway_port} " + f"(returncode={returncode}).\nstderr:\n{stderr_text}" + ) + finally: + if proc.returncode is None: + proc.send_signal(signal.SIGINT) + try: + await asyncio.wait_for(proc.wait(), timeout=10) + except asyncio.TimeoutError: + proc.kill() + finally: + for sock in reserved: + if sock.fileno() != -1: + sock.close() + db_tmpdir.cleanup() + home_tmpdir.cleanup() + + @pytest.fixture(scope="session") async def http_channel_server(ironclaw_server, server_ports): """HTTP webhook channel base URL.""" @@ -362,12 +467,7 @@ async def http_channel_server_without_secret( "IRONCLAW_OAUTH_CALLBACK_URL": "https://oauth.test.example/oauth/callback", "IRONCLAW_OAUTH_EXCHANGE_URL": mock_llm_server, } - # Forward LLVM coverage instrumentation env vars when present - COV_ENV_PREFIXES = ("CARGO_LLVM_COV", "LLVM_") - COV_ENV_EXTRAS = ("CARGO_ENCODED_RUSTFLAGS", "CARGO_INCREMENTAL") - for key, val in os.environ.items(): - if key.startswith(COV_ENV_PREFIXES) or key in COV_ENV_EXTRAS: - env[key] = val + _forward_coverage_env(env) proc = await asyncio.create_subprocess_exec( ironclaw_binary, "--no-onboard", stdin=asyncio.subprocess.DEVNULL, diff --git a/tests/e2e/mock_llm.py b/tests/e2e/mock_llm.py index 359c22d5..1147662c 100644 --- a/tests/e2e/mock_llm.py +++ b/tests/e2e/mock_llm.py @@ -34,6 +34,15 @@ TOOL_CALL_PATTERNS = [ "body": {"label": m.group("label")}, }, ), + ( + re.compile(r"check gmail unread|gmail unread", re.IGNORECASE), + "gmail", + lambda _: { + "action": "list_messages", + "query": "is:unread", + "max_results": 1, + }, + ), (re.compile(r"what time|current time", re.IGNORECASE), "time", lambda _: {"operation": "now"}), ( re.compile( @@ -91,6 +100,15 @@ TOOL_CALL_PATTERNS = [ ] +def _new_oauth_state() -> dict: + return { + "exchange_count": 0, + "refresh_count": 0, + "last_exchange": None, + "last_refresh": None, + } + + def _last_user_content(messages: list[dict]) -> str: for msg in reversed(messages): if msg.get("role") == "user": @@ -272,6 +290,12 @@ async def oauth_exchange(request: web.Request) -> web.Response: specific token params such as RFC 8707 `resource` are forwarded here. """ data = await request.post() + oauth_state = request.app["oauth_state"] + oauth_state["exchange_count"] += 1 + oauth_state["last_exchange"] = { + "authorization": request.headers.get("Authorization"), + "form": dict(data), + } code = data.get("code", "") access_token_field = data.get("access_token_field", "access_token") @@ -290,6 +314,39 @@ async def oauth_exchange(request: web.Request) -> web.Response: }) +async def oauth_refresh(request: web.Request) -> web.Response: + """Mock OAuth token refresh proxy for hosted refresh E2E tests.""" + data = await request.post() + oauth_state = request.app["oauth_state"] + oauth_state["refresh_count"] += 1 + oauth_state["last_refresh"] = { + "authorization": request.headers.get("Authorization"), + "form": dict(data), + } + + if request.headers.get("Authorization") != "Bearer e2e-test-token": + return web.json_response({"error": "invalid_gateway_auth"}, status=401) + if data.get("client_id") != "hosted-google-client-id": + return web.json_response({"error": "invalid_client_id"}, status=400) + if "client_secret" in data: + return web.json_response({"error": "unexpected_client_secret"}, status=400) + + return web.json_response({ + "access_token": "mock-refreshed-access-token", + "refresh_token": "mock-rotated-refresh-token", + "expires_in": 3600, + }) + + +async def oauth_state_handler(request: web.Request) -> web.Response: + return web.json_response(request.app["oauth_state"]) + + +async def oauth_reset(request: web.Request) -> web.Response: + request.app["oauth_state"] = _new_oauth_state() + return web.json_response({"ok": True}) + + async def models(_request: web.Request) -> web.Response: return web.json_response({ "object": "list", @@ -424,12 +481,16 @@ def main(): parser.add_argument("--port", type=int, default=0) args = parser.parse_args() app = web.Application() + app["oauth_state"] = _new_oauth_state() # Register both /v1/ and non-/v1/ paths (rig-core omits the /v1/ prefix) app.router.add_post("/v1/chat/completions", chat_completions) app.router.add_post("/chat/completions", chat_completions) app.router.add_get("/v1/models", models) app.router.add_get("/models", models) app.router.add_post("/oauth/exchange", oauth_exchange) + app.router.add_post("/oauth/refresh", oauth_refresh) + app.router.add_get("/__mock/oauth/state", oauth_state_handler) + app.router.add_post("/__mock/oauth/reset", oauth_reset) # Mock MCP server endpoints app.router.add_post("/mcp", mcp_endpoint) app.router.add_post("/mcp-400", mcp_endpoint_400) diff --git a/tests/e2e/scenarios/test_oauth_refresh.py b/tests/e2e/scenarios/test_oauth_refresh.py new file mode 100644 index 00000000..50871f7f --- /dev/null +++ b/tests/e2e/scenarios/test_oauth_refresh.py @@ -0,0 +1,227 @@ +"""Hosted OAuth refresh HTTP regression test. + +Runs a real ironclaw binary in hosted mode, expires a stored Gmail access +token in the libSQL database, triggers a real gmail tool call through the +chat API, and verifies that refresh uses the hosted proxy endpoint. +""" + +import asyncio +import sqlite3 +from datetime import datetime, timezone +from urllib.parse import parse_qs, urlparse + +import httpx + +from helpers import api_get, api_post + + +def _extract_state(auth_url: str) -> str: + parsed = urlparse(auth_url) + state = parse_qs(parsed.query).get("state", [None])[0] + assert state, f"auth_url should include state: {auth_url}" + return state + + +def _parse_timestamp(value: str | None) -> datetime | None: + if value is None: + return None + return datetime.fromisoformat(value.replace("Z", "+00:00")) + + +def _expire_access_token(db_path: str, user_id: str, secret_name: str) -> None: + with sqlite3.connect(db_path) as conn: + cursor = conn.execute( + """ + UPDATE secrets + SET expires_at = strftime('%Y-%m-%dT%H:%M:%fZ', 'now', '-1 hour') + WHERE user_id = ?1 AND name = ?2 + """, + (user_id, secret_name), + ) + conn.commit() + assert cursor.rowcount == 1, f"Expected one secret row for {user_id}/{secret_name}" + + +def _find_secret_row( + db_path: str, + secret_name: str, +) -> tuple[str, str | None, str | None]: + with sqlite3.connect(db_path) as conn: + row = conn.execute( + """ + SELECT user_id, expires_at, updated_at + FROM secrets + WHERE name = ?1 + ORDER BY updated_at DESC + LIMIT 1 + """, + (secret_name,), + ).fetchone() + assert row is not None, f"Missing secret row for {secret_name}" + return row[0], row[1], row[2] + + +async def _get_extension(base_url: str, name: str) -> dict | None: + response = await api_get(base_url, "/api/extensions", timeout=15) + response.raise_for_status() + for extension in response.json().get("extensions", []): + if extension["name"] == name: + return extension + return None + + +async def _reset_mock_oauth_state(mock_base_url: str) -> None: + async with httpx.AsyncClient() as client: + response = await client.post(f"{mock_base_url}/__mock/oauth/reset", timeout=10) + response.raise_for_status() + + +async def _get_mock_oauth_state(mock_base_url: str) -> dict: + async with httpx.AsyncClient() as client: + response = await client.get(f"{mock_base_url}/__mock/oauth/state", timeout=10) + response.raise_for_status() + return response.json() + + +async def _approve_pending_request(base_url: str, thread_id: str, request_id: str) -> None: + response = await api_post( + base_url, + "/api/chat/approval", + json={"request_id": request_id, "action": "approve", "thread_id": thread_id}, + timeout=15, + ) + assert response.status_code == 202, ( + f"Approval submission failed: {response.status_code} {response.text[:400]}" + ) + + +async def _wait_for_gmail_tool_call(base_url: str, thread_id: str, timeout: float = 30.0) -> dict: + approved_request_ids = set() + for _ in range(int(timeout * 2)): + response = await api_get( + base_url, + f"/api/chat/history?thread_id={thread_id}", + timeout=15, + ) + response.raise_for_status() + history = response.json() + + pending = history.get("pending_approval") + if pending and pending["request_id"] not in approved_request_ids: + await _approve_pending_request(base_url, thread_id, pending["request_id"]) + approved_request_ids.add(pending["request_id"]) + + for turn in history.get("turns", []): + for tool_call in turn.get("tool_calls", []): + if tool_call.get("name") == "gmail": + return history + + await asyncio.sleep(0.5) + + raise AssertionError(f"Timed out waiting for gmail tool call in thread {thread_id}") + + +async def _wait_for_refresh_request(mock_base_url: str, timeout: float = 20.0) -> dict: + for _ in range(int(timeout * 2)): + state = await _get_mock_oauth_state(mock_base_url) + if state.get("refresh_count") == 1: + return state + await asyncio.sleep(0.5) + raise AssertionError("Timed out waiting for exactly one OAuth refresh request") + + +async def test_hosted_gmail_oauth_refresh_uses_proxy(hosted_oauth_refresh_server): + server = hosted_oauth_refresh_server["base_url"] + db_path = hosted_oauth_refresh_server["db_path"] + mock_base_url = hosted_oauth_refresh_server["mock_llm_url"] + + install_response = await api_post( + server, + "/api/extensions/install", + json={"name": "gmail"}, + timeout=180, + ) + assert install_response.status_code == 200, install_response.text + assert install_response.json().get("success") is True + + setup_response = await api_post( + server, + "/api/extensions/gmail/setup", + json={"secrets": {}}, + timeout=30, + ) + assert setup_response.status_code == 200, setup_response.text + setup_data = setup_response.json() + assert setup_data.get("success") is True, setup_data + auth_url = setup_data.get("auth_url") + assert auth_url, setup_data + auth_params = parse_qs(urlparse(auth_url).query) + assert auth_params.get("client_id") == ["hosted-google-client-id"] + + async with httpx.AsyncClient() as client: + callback_response = await client.get( + f"{server}/oauth/callback", + params={"code": "mock_auth_code", "state": _extract_state(auth_url)}, + timeout=30, + follow_redirects=True, + ) + + assert callback_response.status_code == 200, callback_response.text[:400] + callback_body = callback_response.text.lower() + assert "connected" in callback_body or "success" in callback_body + + gmail = await _get_extension(server, "gmail") + assert gmail is not None, "gmail should be installed" + assert gmail["authenticated"] is True, gmail + assert "gmail" in gmail.get("tools", []), gmail + + await _reset_mock_oauth_state(mock_base_url) + + stored_user_id, expires_before, updated_before = _find_secret_row( + db_path, "google_oauth_token" + ) + assert _parse_timestamp(expires_before) is not None + assert _parse_timestamp(updated_before) is not None + + await asyncio.sleep(0.1) + _expire_access_token(db_path, stored_user_id, "google_oauth_token") + + thread_response = await api_post(server, "/api/chat/thread/new", timeout=15) + assert thread_response.status_code == 200, thread_response.text + thread_id = thread_response.json()["id"] + + send_response = await api_post( + server, + "/api/chat/send", + json={"content": "check gmail unread", "thread_id": thread_id}, + timeout=30, + ) + assert send_response.status_code == 202, send_response.text + + history = await _wait_for_gmail_tool_call(server, thread_id) + assert any( + tool_call.get("name") == "gmail" + for turn in history.get("turns", []) + for tool_call in turn.get("tool_calls", []) + ), history + + oauth_state = await _wait_for_refresh_request(mock_base_url) + assert oauth_state["refresh_count"] == 1, oauth_state + last_refresh = oauth_state["last_refresh"] + assert last_refresh is not None, oauth_state + assert last_refresh["authorization"] == "Bearer e2e-test-token" + assert last_refresh["form"]["client_id"] == "hosted-google-client-id" + assert "client_secret" not in last_refresh["form"], last_refresh + + refreshed_user_id, expires_after, updated_after = _find_secret_row( + db_path, "google_oauth_token" + ) + assert refreshed_user_id == stored_user_id + expires_after_dt = _parse_timestamp(expires_after) + updated_after_dt = _parse_timestamp(updated_after) + updated_before_dt = _parse_timestamp(updated_before) + assert expires_after_dt is not None + assert updated_after_dt is not None + assert updated_before_dt is not None + assert expires_after_dt > datetime.now(timezone.utc) + assert updated_after_dt > updated_before_dt From 82822d7b2556a1cf29c6525d211cadd9b0a5917f Mon Sep 17 00:00:00 2001 From: Henry Park Date: Tue, 24 Mar 2026 16:11:53 -0700 Subject: [PATCH 12/20] fix: restore owner-scoped gateway startup (#1625) * fix: restore owner-scoped gateway startup * fix: split gateway owner and sender scope * fix: keep multi-user gateway sender identity * test: cover gateway sender scope regression * test: harden e2e startup teardown race * fix: align gateway owner scope across auth modes --- src/app.rs | 8 +- src/channels/web/mod.rs | 25 ++++- src/channels/web/server.rs | 32 +++--- src/channels/web/test_helpers.rs | 3 +- src/channels/web/tests/multi_tenant.rs | 39 ++++++- src/channels/web/ws.rs | 3 +- src/config/mod.rs | 12 +-- src/main.rs | 1 + tests/e2e/conftest.py | 91 +++++++++++----- tests/multi_tenant_integration.rs | 123 +++++++++++++++++++++- tests/openai_compat_integration.rs | 6 +- tests/support/gateway_workflow_harness.rs | 3 +- tests/ws_gateway_integration.rs | 3 +- 13 files changed, 278 insertions(+), 71 deletions(-) diff --git a/src/app.rs b/src/app.rs index edd547d3..074e9479 100644 --- a/src/app.rs +++ b/src/app.rs @@ -312,13 +312,7 @@ impl AppBuilder { .create_provider(&self.config.llm.nearai.base_url, self.session.clone()); // Register memory tools if database is available - let workspace_user_id = self - .config - .channels - .gateway - .as_ref() - .map(|gw| gw.user_id.as_str()) - .unwrap_or("default"); + let workspace_user_id = self.config.owner_id.as_str(); let workspace = if let Some(ref db) = self.db { let emb_cache_config = EmbeddingCacheConfig { max_entries: self.config.embeddings.cache_size, diff --git a/src/channels/web/mod.rs b/src/channels/web/mod.rs index b26a7829..a8b1ec41 100644 --- a/src/channels/web/mod.rs +++ b/src/channels/web/mod.rs @@ -98,7 +98,8 @@ impl GatewayChannel { job_manager: None, prompt_queue: None, scheduler: None, - default_user_id: config.user_id.clone(), + owner_id: config.user_id.clone(), + default_sender_id: config.user_id.clone(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())), llm_provider: None, @@ -121,6 +122,22 @@ impl GatewayChannel { } } + /// Rebind the single-user auth identity to the durable owner scope while + /// preserving the configured gateway sender/routing identity. + pub fn with_owner_scope(mut self, owner_id: impl Into) -> Self { + let owner_id = owner_id.into(); + let single_user_token = if self.config.user_tokens.is_none() { + self.auth.first_token().map(ToOwned::to_owned) + } else { + None + }; + if let Some(token) = single_user_token { + self.auth = MultiAuthState::single(token, owner_id.clone()); + } + self.rebuild_state(|s| s.owner_id = owner_id); + self + } + /// Create a gateway channel with a pre-built multi-user auth state. pub fn new_multi_auth(config: GatewayConfig, auth: MultiAuthState) -> Self { let state = Arc::new(GatewayState { @@ -137,7 +154,8 @@ impl GatewayChannel { job_manager: None, prompt_queue: None, scheduler: None, - default_user_id: config.user_id.clone(), + owner_id: config.user_id.clone(), + default_sender_id: config.user_id.clone(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())), llm_provider: None, @@ -177,7 +195,8 @@ impl GatewayChannel { job_manager: self.state.job_manager.clone(), prompt_queue: self.state.prompt_queue.clone(), scheduler: self.state.scheduler.clone(), - default_user_id: self.state.default_user_id.clone(), + owner_id: self.state.owner_id.clone(), + default_sender_id: self.state.default_sender_id.clone(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: self.state.ws_tracker.clone(), llm_provider: self.state.llm_provider.clone(), diff --git a/src/channels/web/server.rs b/src/channels/web/server.rs index fa29040e..31c2b296 100644 --- a/src/channels/web/server.rs +++ b/src/channels/web/server.rs @@ -345,8 +345,10 @@ pub struct GatewayState { pub job_manager: Option>, /// Prompt queue for Claude Code follow-up prompts. pub prompt_queue: Option, - /// Default user ID (fallback for non-request contexts like heartbeat/routines). - pub default_user_id: String, + /// Durable owner scope for persistence and unauthenticated callback flows. + pub owner_id: String, + /// Default sender/routing identity for gateway-originated messages. + pub default_sender_id: String, /// Shutdown signal sender. pub shutdown_tx: tokio::sync::RwLock>>, /// WebSocket connection tracker. @@ -775,7 +777,7 @@ async fn oauth_callback_handler( error = %error, "OAuth callback received with malformed state" ); - clear_auth_mode(&state, &state.default_user_id).await; + clear_auth_mode(&state, &state.owner_id).await; return oauth_error_page("IronClaw"); } }; @@ -1136,7 +1138,7 @@ async fn slack_relay_oauth_callback_handler( let state_key = format!("relay:{}:oauth_state", DEFAULT_RELAY_NAME); let stored_state = match ext_mgr .secrets() - .get_decrypted(&state.default_user_id, &state_key) + .get_decrypted(&state.owner_id, &state_key) .await { Ok(secret) => secret.expose().to_string(), @@ -1160,10 +1162,7 @@ async fn slack_relay_oauth_callback_handler( } // Delete the nonce (one-time use) - let _ = ext_mgr - .secrets() - .delete(&state.default_user_id, &state_key) - .await; + let _ = ext_mgr.secrets().delete(&state.owner_id, &state_key).await; let result: Result<(), String> = async { let store = state.store.as_ref().ok_or_else(|| { @@ -1174,16 +1173,12 @@ async fn slack_relay_oauth_callback_handler( // Store team_id in settings let team_id_key = format!("relay:{}:team_id", DEFAULT_RELAY_NAME); let _ = store - .set_setting( - &state.default_user_id, - &team_id_key, - &serde_json::json!(team_id), - ) + .set_setting(&state.owner_id, &team_id_key, &serde_json::json!(team_id)) .await; // Activate the relay channel ext_mgr - .activate_stored_relay(DEFAULT_RELAY_NAME, &state.default_user_id) + .activate_stored_relay(DEFAULT_RELAY_NAME, &state.owner_id) .await .map_err(|e| format!("Failed to activate relay channel: {}", e))?; @@ -1303,6 +1298,9 @@ async fn chat_send_handler( } let mut msg = IncomingMessage::new("gateway", &user.user_id, &req.content); + if state.owner_id != state.default_sender_id && user.user_id == state.owner_id { + msg = msg.with_sender_id(&state.default_sender_id); + } // Prefer timezone from JSON body, fall back to X-Timezone header let tz = req .timezone @@ -1404,6 +1402,9 @@ async fn chat_approval_handler( })?; let mut msg = IncomingMessage::new("gateway", &user.user_id, content); + if state.owner_id != state.default_sender_id && user.user_id == state.owner_id { + msg = msg.with_sender_id(&state.default_sender_id); + } if let Some(ref thread_id) = req.thread_id { msg = msg.with_thread(thread_id); @@ -2976,7 +2977,8 @@ mod tests { store: None, job_manager: None, prompt_queue: None, - default_user_id: "test".to_string(), + owner_id: "test".to_string(), + default_sender_id: "test".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: None, llm_provider: None, diff --git a/src/channels/web/test_helpers.rs b/src/channels/web/test_helpers.rs index 802512a6..0f7e5d12 100644 --- a/src/channels/web/test_helpers.rs +++ b/src/channels/web/test_helpers.rs @@ -76,7 +76,8 @@ impl TestGatewayBuilder { store: None, job_manager: None, prompt_queue: None, - default_user_id: self.user_id, + owner_id: self.user_id.clone(), + default_sender_id: self.user_id, shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: self.llm_provider, diff --git a/src/channels/web/tests/multi_tenant.rs b/src/channels/web/tests/multi_tenant.rs index 55010831..335f841c 100644 --- a/src/channels/web/tests/multi_tenant.rs +++ b/src/channels/web/tests/multi_tenant.rs @@ -16,6 +16,7 @@ use axum::routing::{delete, get, post}; use tower::ServiceExt; use uuid::Uuid; +use crate::channels::web::GatewayChannel; use crate::channels::web::auth::{ AuthenticatedUser, MultiAuthState, UserIdentity, auth_middleware, }; @@ -23,6 +24,7 @@ use crate::channels::web::server::{ ActiveConfigSnapshot, GatewayState, PerUserRateLimiter, PromptQueue, RateLimiter, WorkspacePool, }; use crate::channels::web::sse::SseManager; +use crate::config::GatewayConfig; // ── Helpers ──────────────────────────────────────────────────────────── @@ -64,7 +66,8 @@ fn build_state( store, job_manager: None, prompt_queue, - default_user_id: "test".to_string(), + owner_id: "test".to_string(), + default_sender_id: "test".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: None, llm_provider: None, @@ -82,6 +85,40 @@ fn build_state( }) } +fn gateway_config() -> GatewayConfig { + GatewayConfig { + host: "127.0.0.1".to_string(), + port: 3000, + auth_token: Some("gateway-auth".to_string()), + user_id: "gateway-sender".to_string(), + workspace_read_scopes: Vec::new(), + memory_layers: Vec::new(), + user_tokens: None, + } +} + +#[test] +fn with_owner_scope_updates_gateway_owner_scope_in_multi_user_mode() { + let mut gateway = GatewayChannel::new(gateway_config()); + gateway.auth = two_user_auth(); + gateway.config.user_tokens = Some(HashMap::new()); + let gateway = gateway.with_owner_scope("owner-scope"); + + assert_eq!(gateway.state.owner_id, "owner-scope"); + assert_eq!(gateway.state.default_sender_id, "gateway-sender"); + + let alice = gateway + .auth + .authenticate("tok-alice") + .expect("alice token should remain valid"); + let bob = gateway + .auth + .authenticate("tok-bob") + .expect("bob token should remain valid"); + assert_eq!(alice.user_id, "alice"); + assert_eq!(bob.user_id, "bob"); +} + /// Create a libSQL-backed test database in a temporary directory. /// /// Returns the database and a `TempDir` guard — the database file is diff --git a/src/channels/web/ws.rs b/src/channels/web/ws.rs index 3a601679..9d4e919c 100644 --- a/src/channels/web/ws.rs +++ b/src/channels/web/ws.rs @@ -520,7 +520,8 @@ mod tests { job_manager: None, prompt_queue: None, scheduler: None, - default_user_id: "test".to_string(), + owner_id: "test".to_string(), + default_sender_id: "test".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: None, diff --git a/src/config/mod.rs b/src/config/mod.rs index dcda0fe9..a362fd09 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -312,13 +312,11 @@ impl Config { let tunnel = TunnelConfig::resolve(settings)?; let channels = ChannelsConfig::resolve(settings, &owner_id)?; - // Resolve workspace config using the gateway user_id for default layers. - let workspace_user_id = channels - .gateway - .as_ref() - .map(|gw| gw.user_id.as_str()) - .unwrap_or("default"); - let workspace = WorkspaceConfig::resolve(workspace_user_id)?; + // Resolve the startup workspace against the durable owner scope. The + // gateway may expose a distinct sender identity, but the base runtime + // workspace stays owner-scoped and per-user gateway workspaces are + // handled separately by WorkspacePool. + let workspace = WorkspaceConfig::resolve(&owner_id)?; Ok(Self { owner_id: owner_id.clone(), diff --git a/src/main.rs b/src/main.rs index eab01264..e885cb7d 100644 --- a/src/main.rs +++ b/src/main.rs @@ -611,6 +611,7 @@ async fn async_main() -> anyhow::Result<()> { } else { GatewayChannel::new(gw_config.clone()) }; + gw = gw.with_owner_scope(config.owner_id.clone()); gw = gw.with_llm_provider(Arc::clone(&components.llm)); if let Some(ref ws) = components.workspace { gw = gw.with_workspace(Arc::clone(ws)); diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py index 1496f93f..aa8ba1cb 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -112,6 +112,30 @@ def _reserve_loopback_sockets(count: int) -> list[socket.socket]: sock.close() raise +async def _stop_process( + proc: asyncio.subprocess.Process, *, sig: int | None = None, timeout: float +) -> None: + """Signal a subprocess and wait briefly without masking exit races.""" + if proc.returncode is not None: + return + + try: + if sig is None: + proc.kill() + else: + proc.send_signal(sig) + except ProcessLookupError: + try: + await asyncio.wait_for(proc.wait(), timeout=timeout) + except asyncio.TimeoutError: + pass + return + + try: + await asyncio.wait_for(proc.wait(), timeout=timeout) + except asyncio.TimeoutError: + pass + def _forward_coverage_env(env: dict[str, str]) -> None: """Forward cargo-llvm-cov env vars into child processes when present.""" @@ -281,35 +305,39 @@ async def ironclaw_server( stderr=asyncio.subprocess.PIPE, env=env, ) + startup_kill_attempted = False base_url = f"http://127.0.0.1:{gateway_port}" try: await wait_for_ready(f"{base_url}/api/health", timeout=60) yield base_url except TimeoutError: # Dump stderr so CI logs show why the server failed to start + if proc.returncode is None: + startup_kill_attempted = True + await _stop_process(proc, timeout=2) returncode = proc.returncode stderr_bytes = b"" if proc.stderr: try: stderr_bytes = await asyncio.wait_for(proc.stderr.read(8192), timeout=2) - except (asyncio.TimeoutError, Exception): + except asyncio.TimeoutError: pass stderr_text = stderr_bytes.decode("utf-8", errors="replace") - proc.kill() pytest.fail( f"ironclaw server failed to start on port {gateway_port} " f"(returncode={returncode}).\nstderr:\n{stderr_text}" ) finally: if proc.returncode is None: - # Use SIGINT (not SIGTERM) so tokio's ctrl_c handler triggers a - # graceful shutdown. This lets the LLVM coverage runtime run its - # atexit handler and flush .profraw files for cargo-llvm-cov. - proc.send_signal(signal.SIGINT) - try: - await asyncio.wait_for(proc.wait(), timeout=10) - except asyncio.TimeoutError: - proc.kill() + if startup_kill_attempted: + await _stop_process(proc, timeout=2) + else: + # Use SIGINT (not SIGTERM) so tokio's ctrl_c handler triggers a + # graceful shutdown. This lets the LLVM coverage runtime run its + # atexit handler and flush .profraw files for cargo-llvm-cov. + await _stop_process(proc, sig=signal.SIGINT, timeout=10) + if proc.returncode is None: + await _stop_process(proc, timeout=2) @pytest.fixture(scope="session") @@ -376,6 +404,7 @@ async def hosted_oauth_refresh_server( stderr=asyncio.subprocess.PIPE, env=env, ) + startup_kill_attempted = False base_url = f"http://127.0.0.1:{gateway_port}" try: await wait_for_ready(f"{base_url}/api/health", timeout=60) @@ -386,27 +415,29 @@ async def hosted_oauth_refresh_server( "mock_llm_url": mock_llm_server, } except TimeoutError: + if proc.returncode is None: + startup_kill_attempted = True + await _stop_process(proc, timeout=2) returncode = proc.returncode stderr_bytes = b"" if proc.stderr: try: stderr_bytes = await asyncio.wait_for(proc.stderr.read(8192), timeout=2) - except (asyncio.TimeoutError, Exception): + except asyncio.TimeoutError: pass stderr_text = stderr_bytes.decode("utf-8", errors="replace") - if proc.returncode is None: - proc.kill() pytest.fail( f"hosted oauth refresh server failed to start on port {gateway_port} " f"(returncode={returncode}).\nstderr:\n{stderr_text}" ) finally: if proc.returncode is None: - proc.send_signal(signal.SIGINT) - try: - await asyncio.wait_for(proc.wait(), timeout=10) - except asyncio.TimeoutError: - proc.kill() + if startup_kill_attempted: + await _stop_process(proc, timeout=2) + else: + await _stop_process(proc, sig=signal.SIGINT, timeout=10) + if proc.returncode is None: + await _stop_process(proc, timeout=2) finally: for sock in reserved: if sock.fileno() != -1: @@ -475,6 +506,7 @@ async def http_channel_server_without_secret( stderr=asyncio.subprocess.PIPE, env=env, ) + startup_kill_attempted = False gateway_url = f"http://127.0.0.1:{gateway_port}" http_base_url = f"http://127.0.0.1:{http_port}" try: @@ -483,15 +515,17 @@ async def http_channel_server_without_secret( yield http_base_url except TimeoutError: # Dump stderr so CI logs show why the server failed to start + if proc.returncode is None: + startup_kill_attempted = True + await _stop_process(proc, timeout=2) returncode = proc.returncode stderr_bytes = b"" if proc.stderr: try: stderr_bytes = await asyncio.wait_for(proc.stderr.read(8192), timeout=2) - except (asyncio.TimeoutError, Exception): + except asyncio.TimeoutError: pass stderr_text = stderr_bytes.decode("utf-8", errors="replace") - proc.kill() pytest.fail( f"ironclaw server without webhook secret failed to start on ports " f"gateway={gateway_port}, http={http_port} " @@ -499,14 +533,15 @@ async def http_channel_server_without_secret( ) finally: if proc.returncode is None: - # Use SIGINT (not SIGTERM) so tokio's ctrl_c handler triggers a - # graceful shutdown. This lets the LLVM coverage runtime run its - # atexit handler and flush .profraw files for cargo-llvm-cov. - proc.send_signal(signal.SIGINT) - try: - await asyncio.wait_for(proc.wait(), timeout=10) - except asyncio.TimeoutError: - proc.kill() + if startup_kill_attempted: + await _stop_process(proc, timeout=2) + else: + # Use SIGINT (not SIGTERM) so tokio's ctrl_c handler triggers a + # graceful shutdown. This lets the LLVM coverage runtime run its + # atexit handler and flush .profraw files for cargo-llvm-cov. + await _stop_process(proc, sig=signal.SIGINT, timeout=10) + if proc.returncode is None: + await _stop_process(proc, timeout=2) @pytest.fixture(scope="session") diff --git a/tests/multi_tenant_integration.rs b/tests/multi_tenant_integration.rs index 02eb60e8..f2529866 100644 --- a/tests/multi_tenant_integration.rs +++ b/tests/multi_tenant_integration.rs @@ -19,10 +19,13 @@ use axum::middleware; use axum::routing::{get, post}; use tower::ServiceExt; +use ironclaw::channels::IncomingMessage; use ironclaw::channels::web::auth::{ AuthenticatedUser, MultiAuthState, UserIdentity, auth_middleware, }; -use ironclaw::channels::web::server::{GatewayState, PerUserRateLimiter, RateLimiter}; +use ironclaw::channels::web::server::{ + GatewayState, PerUserRateLimiter, RateLimiter, start_server, +}; use ironclaw::channels::web::sse::SseManager; use ironclaw::channels::web::test_helpers::TestGatewayBuilder; use ironclaw::channels::web::ws::WsConnectionTracker; @@ -37,6 +40,9 @@ const ALICE_TOKEN: &str = "tok-alice-secret"; const BOB_TOKEN: &str = "tok-bob-secret"; const ALICE_USER_ID: &str = "alice"; const BOB_USER_ID: &str = "bob"; +const OWNER_TOKEN: &str = "tok-owner-secret"; +const OWNER_SCOPE_ID: &str = "owner-scope"; +const GATEWAY_SENDER_ID: &str = "gateway-sender"; /// Build a MultiAuthState with two users. fn two_user_auth() -> MultiAuthState { @@ -537,7 +543,8 @@ fn gateway_state_has_multi_tenant_fields() { job_manager: None, prompt_queue: None, scheduler: None, - default_user_id: "fallback".to_string(), // Multi-tenant: renamed from user_id + owner_id: "fallback".to_string(), + default_sender_id: "fallback".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: None, @@ -553,7 +560,8 @@ fn gateway_state_has_multi_tenant_fields() { active_config: Default::default(), }; - assert_eq!(state.default_user_id, "fallback"); + assert_eq!(state.owner_id, "fallback"); + assert_eq!(state.default_sender_id, "fallback"); assert!(state.workspace_pool.is_none()); } @@ -572,6 +580,69 @@ async fn start_multi_user_server() -> (SocketAddr, Arc) { .expect("Failed to start multi-user test server") } +async fn start_owner_scoped_sender_server() -> ( + SocketAddr, + Arc, + tokio::sync::mpsc::Receiver, +) { + let (agent_tx, agent_rx) = tokio::sync::mpsc::channel(64); + + let mut tokens = HashMap::new(); + tokens.insert( + OWNER_TOKEN.to_string(), + UserIdentity { + user_id: OWNER_SCOPE_ID.to_string(), + workspace_read_scopes: Vec::new(), + }, + ); + tokens.insert( + BOB_TOKEN.to_string(), + UserIdentity { + user_id: BOB_USER_ID.to_string(), + workspace_read_scopes: Vec::new(), + }, + ); + + let state = Arc::new(GatewayState { + msg_tx: tokio::sync::RwLock::new(Some(agent_tx)), + sse: Arc::new(SseManager::new()), + workspace: None, + workspace_pool: None, + session_manager: None, + log_broadcaster: None, + log_level_handle: None, + extension_manager: None, + tool_registry: None, + store: None, + job_manager: None, + prompt_queue: None, + scheduler: None, + owner_id: OWNER_SCOPE_ID.to_string(), + default_sender_id: GATEWAY_SENDER_ID.to_string(), + shutdown_tx: tokio::sync::RwLock::new(None), + ws_tracker: Some(Arc::new(WsConnectionTracker::new())), + llm_provider: None, + skill_registry: None, + skill_catalog: None, + chat_rate_limiter: PerUserRateLimiter::new(30, 60), + oauth_rate_limiter: RateLimiter::new(10, 60), + webhook_rate_limiter: RateLimiter::new(10, 60), + registry_entries: Vec::new(), + cost_guard: None, + routine_engine: Arc::new(tokio::sync::RwLock::new(None)), + startup_time: std::time::Instant::now(), + active_config: Default::default(), + }); + + let auth = MultiAuthState::multi(tokens); + let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); + let bound = start_server(addr, state.clone(), auth) + .await + .expect("Failed to start owner-scoped sender test server"); + + (bound, state, agent_rx) +} + #[tokio::test] async fn full_server_alice_can_access_protected_endpoint() { let (addr, _state) = start_multi_user_server().await; @@ -677,6 +748,49 @@ async fn full_server_chat_send_accepted_for_alice() { assert_eq!(msg.channel, "gateway"); } +#[tokio::test] +async fn full_server_chat_send_rewrites_sender_only_for_owner_scope_rebind() { + let (addr, _state, mut agent_rx) = start_owner_scoped_sender_server().await; + + let client = reqwest::Client::new(); + + let owner_resp = client + .post(format!("http://{}/api/chat/send", addr)) + .header("Authorization", format!("Bearer {}", OWNER_TOKEN)) + .header("Content-Type", "application/json") + .body(r#"{"content":"hello from owner"}"#) + .send() + .await + .unwrap(); + assert_eq!(owner_resp.status(), 202); + + let owner_msg = tokio::time::timeout(Duration::from_secs(2), agent_rx.recv()) + .await + .expect("Timed out waiting for owner message") + .expect("Agent channel closed"); + assert_eq!(owner_msg.user_id, OWNER_SCOPE_ID); + assert_eq!(owner_msg.sender_id, GATEWAY_SENDER_ID); + assert_eq!(owner_msg.content, "hello from owner"); + + let other_resp = client + .post(format!("http://{}/api/chat/send", addr)) + .header("Authorization", format!("Bearer {}", BOB_TOKEN)) + .header("Content-Type", "application/json") + .body(r#"{"content":"hello from bob"}"#) + .send() + .await + .unwrap(); + assert_eq!(other_resp.status(), 202); + + let other_msg = tokio::time::timeout(Duration::from_secs(2), agent_rx.recv()) + .await + .expect("Timed out waiting for non-owner message") + .expect("Agent channel closed"); + assert_eq!(other_msg.user_id, BOB_USER_ID); + assert_eq!(other_msg.sender_id, BOB_USER_ID); + assert_eq!(other_msg.content, "hello from bob"); +} + #[tokio::test] async fn full_server_chat_send_rejected_without_auth() { let (addr, _state) = start_multi_user_server().await; @@ -888,7 +1002,8 @@ async fn start_multi_user_server_with_db() -> ( job_manager: None, prompt_queue: None, scheduler: None, - default_user_id: ALICE_USER_ID.to_string(), + owner_id: ALICE_USER_ID.to_string(), + default_sender_id: ALICE_USER_ID.to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: None, diff --git a/tests/openai_compat_integration.rs b/tests/openai_compat_integration.rs index 16568246..e1d258ed 100644 --- a/tests/openai_compat_integration.rs +++ b/tests/openai_compat_integration.rs @@ -203,7 +203,8 @@ async fn start_test_server_with_provider( job_manager: None, prompt_queue: None, scheduler: None, - default_user_id: "test-user".to_string(), + owner_id: "test-user".to_string(), + default_sender_id: "test-user".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: Some(llm_provider), @@ -701,7 +702,8 @@ async fn test_no_llm_provider_returns_503() { job_manager: None, prompt_queue: None, scheduler: None, - default_user_id: "test-user".to_string(), + owner_id: "test-user".to_string(), + default_sender_id: "test-user".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: None, // No LLM! diff --git a/tests/support/gateway_workflow_harness.rs b/tests/support/gateway_workflow_harness.rs index e4620f70..5f477de0 100644 --- a/tests/support/gateway_workflow_harness.rs +++ b/tests/support/gateway_workflow_harness.rs @@ -226,7 +226,8 @@ impl GatewayWorkflowHarness { job_manager: None, prompt_queue: None, scheduler: Some(scheduler_slot.clone()), - default_user_id: user_id.clone(), + owner_id: user_id.clone(), + default_sender_id: user_id.clone(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: Some(Arc::clone(&components.llm)), diff --git a/tests/ws_gateway_integration.rs b/tests/ws_gateway_integration.rs index 43277389..a6db5af7 100644 --- a/tests/ws_gateway_integration.rs +++ b/tests/ws_gateway_integration.rs @@ -51,7 +51,8 @@ async fn start_test_server() -> ( job_manager: None, prompt_queue: None, scheduler: None, - default_user_id: "test-user".to_string(), + owner_id: "test-user".to_string(), + default_sender_id: "test-user".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: None, From 656151783cb9aa165d9dc99e82d7855ed3943b11 Mon Sep 17 00:00:00 2001 From: Illia Polosukhin Date: Tue, 24 Mar 2026 23:01:19 -0700 Subject: [PATCH 13/20] 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 14/20] 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 15/20] 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 16/20] 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 17/20] 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 18/20] 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 19/20] 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 20/20] 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;