diff --git a/.github/workflows/e2e.yml b/.github/workflows/e2e.yml index 5b20345e..bc705df7 100644 --- a/.github/workflows/e2e.yml +++ b/.github/workflows/e2e.yml @@ -54,7 +54,7 @@ jobs: - group: features files: "tests/e2e/scenarios/test_skills.py tests/e2e/scenarios/test_tool_approval.py tests/e2e/scenarios/test_webhook.py" - group: extensions - files: "tests/e2e/scenarios/test_extensions.py tests/e2e/scenarios/test_extension_oauth.py tests/e2e/scenarios/test_telegram_token_validation.py tests/e2e/scenarios/test_telegram_hot_activation.py tests/e2e/scenarios/test_wasm_lifecycle.py tests/e2e/scenarios/test_tool_execution.py tests/e2e/scenarios/test_pairing.py tests/e2e/scenarios/test_mcp_auth_flow.py tests/e2e/scenarios/test_oauth_credential_fallback.py tests/e2e/scenarios/test_routine_oauth_credential_injection.py" + files: "tests/e2e/scenarios/test_extensions.py tests/e2e/scenarios/test_extension_oauth.py tests/e2e/scenarios/test_oauth_url_parameters.py tests/e2e/scenarios/test_telegram_token_validation.py tests/e2e/scenarios/test_telegram_hot_activation.py tests/e2e/scenarios/test_wasm_lifecycle.py tests/e2e/scenarios/test_tool_execution.py tests/e2e/scenarios/test_pairing.py tests/e2e/scenarios/test_mcp_auth_flow.py tests/e2e/scenarios/test_oauth_credential_fallback.py tests/e2e/scenarios/test_routine_oauth_credential_injection.py" - group: routines files: "tests/e2e/scenarios/test_owner_scope.py tests/e2e/scenarios/test_routine_event_batch.py" steps: diff --git a/.github/workflows/regression-test-check.yml b/.github/workflows/regression-test-check.yml index ef1a4d92..75b8eb55 100644 --- a/.github/workflows/regression-test-check.yml +++ b/.github/workflows/regression-test-check.yml @@ -121,6 +121,7 @@ jobs: fi # Whole-function context: detect edits inside existing test functions. + # Uses -W (whole function) which works when git recognises function boundaries. if git diff "${BASE_REF}...${HEAD_REF}" -W -- '*.rs' | awk ' /^@@/ { if (has_test && has_add) { found=1; exit } has_test=0; has_add=0 } /^ .*#\[test\]/ || /^ .*#\[tokio::test\]/ || /^ .*#\[cfg\(test\)\]/ || /^ .*mod tests/ { has_test=1 } @@ -132,6 +133,40 @@ jobs: exit 0 fi + # Line-level check: detect changes inside #[cfg(test)] mod blocks. + # git -W relies on function boundary detection which misses Rust mod blocks, + # so this fallback checks whether changed line numbers fall within test modules. + # We specifically match #[cfg(test)] that is followed by `mod` (same or next + # line) to avoid false positives from standalone #[cfg(test)] items like + # individual statics or functions. + CHANGED_RS=$(echo "$CHANGED_FILES" | grep '\.rs$' || true) + if [ -n "$CHANGED_RS" ]; then + while IFS= read -r rs_file; do + [ -f "$rs_file" ] || continue + + # Find the line where #[cfg(test)] precedes a `mod` declaration. + # Handles both `#[cfg(test)] mod tests` (same line) and the two-line form. + TEST_MOD_START=$(awk ' + /^[[:space:]]*#\[cfg\(test\)\].*mod / { print NR; exit } + /^[[:space:]]*#\[cfg\(test\)\][[:space:]]*$/ { pending=NR; next } + pending && /^[[:space:]]*mod / { print pending; exit } + { pending=0 } + ' "$rs_file") + [ -n "$TEST_MOD_START" ] || continue + + # Get changed line numbers in this file from the diff hunk headers. + # Each @@ line looks like: @@ -old,count +new,count @@ + while IFS= read -r hunk_line; do + line_no=$(echo "$hunk_line" | sed -E 's/^@@ -[0-9,]+ \+([0-9]+).*/\1/') + [ -n "$line_no" ] || continue + if [ "$line_no" -ge "$TEST_MOD_START" ]; then + echo "Test changes found: $rs_file has changes at line $line_no inside #[cfg(test)] mod block (starts at line $TEST_MOD_START)." + exit 0 + fi + done < <(git diff "${BASE_REF}...${HEAD_REF}" -U0 -- "$rs_file" | grep -E '^@@') + done <<< "$CHANGED_RS" + fi + if grep -qE '^tests/' <<< "$CHANGED_FILES"; then echo "Test file changes found under tests/." exit 0 diff --git a/Cargo.lock b/Cargo.lock index 83110d35..95aded9a 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -159,7 +159,7 @@ version = "1.1.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "40c48f72fd53cd289104fc64099abca73db4166ad86ea0b4341abe65af83dadc" dependencies = [ - "windows-sys 0.61.2", + "windows-sys 0.60.2", ] [[package]] @@ -170,7 +170,7 @@ checksum = "291e6a250ff86cd4a820112fb8898808a366d8f9f58ce16d1f538353ad55747d" dependencies = [ "anstyle", "once_cell_polyfill", - "windows-sys 0.61.2", + "windows-sys 0.60.2", ] [[package]] @@ -2249,7 +2249,7 @@ dependencies = [ "libc", "option-ext", "redox_users 0.5.2", - "windows-sys 0.61.2", + "windows-sys 0.59.0", ] [[package]] @@ -2436,7 +2436,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" dependencies = [ "libc", - "windows-sys 0.61.2", + "windows-sys 0.52.0", ] [[package]] @@ -4396,7 +4396,7 @@ version = "0.50.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5" dependencies = [ - "windows-sys 0.61.2", + "windows-sys 0.59.0", ] [[package]] @@ -5853,7 +5853,7 @@ dependencies = [ "errno", "libc", "linux-raw-sys 0.12.1", - "windows-sys 0.61.2", + "windows-sys 0.52.0", ] [[package]] @@ -6535,7 +6535,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3a766e1110788c36f4fa1c2b71b387a7815aa65f88ce0229841826633d93723e" dependencies = [ "libc", - "windows-sys 0.61.2", + "windows-sys 0.60.2", ] [[package]] @@ -6765,9 +6765,9 @@ checksum = "55937e1799185b12863d447f42597ed69d9928686b8d88a1df17376a097d8369" [[package]] name = "tar" -version = "0.4.44" +version = "0.4.45" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1d863878d212c87a19c1a610eb53bb01fe12951c0501cf5a0d65f724914a667a" +checksum = "22692a6476a21fa75fdfc11d452fda482af402c008cdbaf3476414e122040973" dependencies = [ "filetime", "libc", @@ -6796,7 +6796,7 @@ dependencies = [ "getrandom 0.4.2", "once_cell", "rustix 1.1.4", - "windows-sys 0.61.2", + "windows-sys 0.52.0", ] [[package]] @@ -7596,7 +7596,7 @@ checksum = "f2f6fb2847f6742cd76af783a2a2c49e9375d0a111c7bef6f71cd9e738c72d6e" dependencies = [ "memoffset", "tempfile", - "windows-sys 0.61.2", + "windows-sys 0.60.2", ] [[package]] @@ -8468,7 +8468,7 @@ version = "0.1.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" dependencies = [ - "windows-sys 0.61.2", + "windows-sys 0.48.0", ] [[package]] diff --git a/FEATURE_PARITY.md b/FEATURE_PARITY.md index a7f5fb32..ad2db551 100644 --- a/FEATURE_PARITY.md +++ b/FEATURE_PARITY.md @@ -161,7 +161,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O | `config` | ✅ | ✅ | - | Read/write config plus validate/path helpers | | `backup` | ✅ | ❌ | P3 | Create/verify local backup archives | | `channels` | ✅ | 🚧 | P2 | `list` implemented; `enable`/`disable`/`status` deferred pending config source unification | -| `models` | ✅ | 🚧 | - | Model selector in TUI | +| `models` | ✅ | 🚧 | P1 | `models list []` (`--verbose`, `--json`; fetches live model list when provider specified), `models status` (`--json`), `models set `, `models set-provider [--model model]` (alias normalization, config.toml + .env persistence). Remaining: `set` doesn't validate model against live list. | | `status` | ✅ | ✅ | - | System status (enriched session details) | | `agents` | ✅ | ❌ | P3 | Multi-agent management | | `sessions` | ✅ | ❌ | P3 | Session listing (shows subagent models) | diff --git a/README.md b/README.md index 6e14d9ea..cb759236 100644 --- a/README.md +++ b/README.md @@ -12,6 +12,9 @@ License: MIT OR Apache-2.0 Telegram: @ironclawAI Reddit: r/ironclawAI + + gitcgr +

diff --git a/channels-src/feishu/feishu.capabilities.json b/channels-src/feishu/feishu.capabilities.json index 82b1be4e..a228cc4e 100644 --- a/channels-src/feishu/feishu.capabilities.json +++ b/channels-src/feishu/feishu.capabilities.json @@ -3,11 +3,11 @@ "wit_version": "0.3.0", "type": "channel", "name": "feishu", - "description": "Feishu/Lark Bot channel for receiving and responding to Feishu messages", + "description": "Feishu/Lark Bot channel for receiving and responding to Feishu messages via Event Subscription webhooks", "auth": { "secret_name": "feishu_app_id", "display_name": "Feishu / Lark", - "instructions": "Create a bot at https://open.feishu.cn/app (Feishu) or https://open.larksuite.com/app (Lark). You need the App ID and App Secret.", + "instructions": "Create a bot at https://open.feishu.cn/app (Feishu) or https://open.larksuite.com/app (Lark). You need the App ID and App Secret. Note: IronClaw supports Event Subscription webhook delivery, but not Feishu's long-connection websocket mode.", "setup_url": "https://open.feishu.cn/app", "token_hint": "App ID looks like cli_XXXX, App Secret is a long alphanumeric string", "env_var": "FEISHU_APP_ID" @@ -16,17 +16,17 @@ "required_secrets": [ { "name": "feishu_app_id", - "prompt": "Enter your Feishu/Lark App ID (from https://open.feishu.cn/app)", + "prompt": "Enter your Feishu/Lark App ID (from https://open.feishu.cn/app). Use webhook-based Event Subscription, not long-connection websocket mode.", "optional": false }, { "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 }, { "name": "feishu_verification_token", - "prompt": "Enter your Feishu/Lark Verification Token (from Event Subscription settings)", + "prompt": "Enter your Feishu/Lark Verification Token (from Event Subscription webhook settings)", "optional": true } ], diff --git a/channels-src/feishu/src/lib.rs b/channels-src/feishu/src/lib.rs index 3094eaa0..62440d2c 100644 --- a/channels-src/feishu/src/lib.rs +++ b/channels-src/feishu/src/lib.rs @@ -5,7 +5,9 @@ //! //! This WASM component implements the channel interface for handling Feishu //! webhooks (Event Subscription v2.0) and sending messages back via the -//! Feishu/Lark Bot API. +//! Feishu/Lark Bot API. IronClaw currently does not connect to Feishu's +//! long-connection websocket subscription mode; use Event Subscription +//! webhooks for this channel. //! //! # Features //! diff --git a/src/agent/agent_loop.rs b/src/agent/agent_loop.rs index ab1c0e13..25e80cb9 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. @@ -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. @@ -235,8 +238,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/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 7fc8e0ca..a195458d 100644 --- a/src/agent/dispatcher.rs +++ b/src/agent/dispatcher.rs @@ -915,7 +915,14 @@ pub(super) async fn execute_chat_tool_standalone( params: &serde_json::Value, job_ctx: &crate::context::JobContext, ) -> Result { - crate::tools::execute::execute_tool_with_safety(tools, safety, tool_name, params, job_ctx).await + crate::tools::execute::execute_tool_with_safety( + tools, + safety, + tool_name, + params.clone(), + job_ctx, + ) + .await } /// Parsed auth result fields for emitting StatusUpdate::AuthRequired. @@ -1091,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()); @@ -1218,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( @@ -1893,7 +1909,7 @@ mod tests { Ok(ToolCompletionResponse { content: None, tool_calls: vec![ToolCall { - id: format!("call_{}", uuid::Uuid::new_v4()), + id: crate::llm::generate_tool_call_id(0, 0), name: "echo".to_string(), arguments: serde_json::json!({"message": "looping"}), }], @@ -2046,7 +2062,7 @@ mod tests { Ok(ToolCompletionResponse { content: None, tool_calls: vec![ToolCall { - id: format!("call_{}", uuid::Uuid::new_v4()), + id: crate::llm::generate_tool_call_id(0, 0), name: "nonexistent_tool".to_string(), arguments: serde_json::json!({}), }], @@ -2085,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( @@ -2205,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( @@ -2338,6 +2356,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/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/routine_engine.rs b/src/agent/routine_engine.rs index de2879b4..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" )); } @@ -1440,6 +1455,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/agent/scheduler.rs b/src/agent/scheduler.rs index 2e23b35f..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. @@ -549,11 +548,7 @@ impl Scheduler { // Delegate to shared tool execution pipeline let output_str = crate::tools::execute::execute_tool_with_safety( - &tools, - &safety, - tool_name, - &normalized_params, - &job_ctx, + &tools, &safety, tool_name, params, &job_ctx, ) .await?; diff --git a/src/agent/session.rs b/src/agent/session.rs index 745b26be..45594922 100644 --- a/src/agent/session.rs +++ b/src/agent/session.rs @@ -17,7 +17,7 @@ use serde::{Deserialize, Serialize}; use uuid::Uuid; use crate::channels::web::util::truncate_preview; -use crate::llm::{ChatMessage, ToolCall}; +use crate::llm::{ChatMessage, ToolCall, generate_tool_call_id}; /// A session containing one or more threads. #[derive(Debug, Clone, Serialize, Deserialize)] @@ -414,7 +414,12 @@ impl Thread { /// completed actions in subsequent turns. pub fn messages(&self) -> Vec { let mut messages = Vec::new(); - for turn in &self.turns { + // We use the enumeration index (`turn_idx`) rather than `turn.turn_number` + // intentionally: after `truncate_turns()`, the remaining turns are + // re-numbered starting from 0, so the enumeration index and turn_number + // are equivalent. Using the index avoids coupling to the field and keeps + // tool-call ID generation deterministic for the current message window. + for (turn_idx, turn) in self.turns.iter().enumerate() { if turn.image_content_parts.is_empty() { messages.push(ChatMessage::user(&turn.user_input)); } else { @@ -425,13 +430,23 @@ impl Thread { } if !turn.tool_calls.is_empty() { - // Build ToolCall objects with synthetic stable IDs - let tool_calls: Vec = turn + // Assign synthetic call IDs for this turn's tool calls, so that + // declarations and results can be consistently correlated. + let tool_calls_with_ids: Vec<(String, &_)> = turn .tool_calls .iter() .enumerate() - .map(|(i, tc)| ToolCall { - id: format!("turn{}_{}", turn.turn_number, i), + .map(|(tc_idx, tc)| { + // Use provider-compatible tool call IDs derived from turn/tool indices. + (generate_tool_call_id(turn_idx, tc_idx), tc) + }) + .collect(); + + // Build ToolCall objects using the synthetic call IDs. + let tool_calls: Vec = tool_calls_with_ids + .iter() + .map(|(call_id, tc)| ToolCall { + id: call_id.clone(), name: tc.name.clone(), arguments: tc.parameters.clone(), }) @@ -441,8 +456,7 @@ impl Thread { messages.push(ChatMessage::assistant_with_tool_calls(None, tool_calls)); // Individual tool result messages, truncated to limit context size. - for (i, tc) in turn.tool_calls.iter().enumerate() { - let call_id = format!("turn{}_{}", turn.turn_number, i); + for (call_id, tc) in tool_calls_with_ids { let content = if let Some(ref err) = tc.error { // .error already contains the full error text; // pass through without wrapping to avoid double-prefix. 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 b2520144..edd547d3 100644 --- a/src/app.rs +++ b/src/app.rs @@ -325,12 +325,51 @@ impl AppBuilder { }; let mut ws = Workspace::new_with_db(workspace_user_id, db.clone()) .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) + if !self.config.workspace.read_scopes.is_empty() { + ws = ws.with_additional_read_scopes(self.config.workspace.read_scopes.clone()); + tracing::info!( + user_id = workspace_user_id, + read_scopes = ?ws.read_user_ids(), + "Workspace configured with multi-scope reads" + ); } 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/bootstrap.rs b/src/bootstrap.rs index f8a283f3..a5c8ffdb 100644 --- a/src/bootstrap.rs +++ b/src/bootstrap.rs @@ -568,14 +568,12 @@ impl Drop for PidLock { #[cfg(test)] mod tests { use super::*; + use crate::config::helpers::lock_env; use std::process::Command; - use std::sync::Mutex; use std::thread; use std::time::{Duration, Instant}; use tempfile::tempdir; - static ENV_MUTEX: Mutex<()> = Mutex::new(()); - #[test] fn test_save_and_load_database_url() { let dir = tempdir().unwrap(); @@ -669,8 +667,23 @@ INJECTED="pwned"#; #[test] fn test_ironclaw_env_path() { - let path = ironclaw_env_path(); - assert!(path.ends_with(".ironclaw/.env")); + // Use compute_ironclaw_base_dir() directly to avoid LazyLock caching, + // which can be poisoned by whichever test initializes it first. + let _guard = lock_env(); + let old_val = std::env::var("IRONCLAW_BASE_DIR").ok(); + // SAFETY: Under lock_env(), no concurrent env access. + unsafe { std::env::remove_var("IRONCLAW_BASE_DIR") }; + + let path = compute_ironclaw_base_dir().join(".env"); + assert!( + path.ends_with(".ironclaw/.env"), + "expected path ending with .ironclaw/.env, got: {}", + path.display() + ); + + if let Some(val) = old_val { + unsafe { std::env::set_var("IRONCLAW_BASE_DIR", val) }; + } } #[test] @@ -836,7 +849,7 @@ INJECTED="pwned"#; #[test] fn test_libsql_autodetect_sets_backend_when_db_exists() { - let _guard = ENV_MUTEX.lock().unwrap(); + let _guard = lock_env(); let old_val = std::env::var("DATABASE_BACKEND").ok(); // SAFETY: ENV_MUTEX ensures single-threaded access to env vars in tests unsafe { std::env::remove_var("DATABASE_BACKEND") }; @@ -907,7 +920,7 @@ INJECTED="pwned"#; #[test] fn test_libsql_autodetect_does_not_override_explicit_backend() { - let _guard = ENV_MUTEX.lock().unwrap(); + let _guard = lock_env(); let old_val = std::env::var("DATABASE_BACKEND").ok(); // SAFETY: ENV_MUTEX ensures single-threaded access to env vars in tests unsafe { std::env::set_var("DATABASE_BACKEND", "postgres") }; @@ -1034,7 +1047,7 @@ INJECTED="pwned"#; fn test_ironclaw_base_dir_default() { // This test must run first (or in isolation) before the LazyLock is initialized. // It verifies that when IRONCLAW_BASE_DIR is not set, the default path is used. - let _guard = ENV_MUTEX.lock().unwrap(); + let _guard = lock_env(); let old_val = std::env::var("IRONCLAW_BASE_DIR").ok(); // SAFETY: ENV_MUTEX ensures single-threaded access to env vars in tests unsafe { std::env::remove_var("IRONCLAW_BASE_DIR") }; @@ -1054,7 +1067,7 @@ INJECTED="pwned"#; fn test_ironclaw_base_dir_env_override() { // This test verifies that when IRONCLAW_BASE_DIR is set, // the custom path is used. Must run before LazyLock is initialized. - let _guard = ENV_MUTEX.lock().unwrap(); + let _guard = lock_env(); let old_val = std::env::var("IRONCLAW_BASE_DIR").ok(); // SAFETY: ENV_MUTEX ensures single-threaded access to env vars in tests unsafe { std::env::set_var("IRONCLAW_BASE_DIR", "/custom/ironclaw/path") }; @@ -1076,7 +1089,7 @@ INJECTED="pwned"#; fn test_compute_base_dir_env_path_join() { // Verifies that ironclaw_env_path correctly joins .env to the base dir. // Uses compute_ironclaw_base_dir directly to avoid LazyLock caching. - let _guard = ENV_MUTEX.lock().unwrap(); + let _guard = lock_env(); let old_val = std::env::var("IRONCLAW_BASE_DIR").ok(); // SAFETY: ENV_MUTEX ensures single-threaded access to env vars in tests unsafe { std::env::set_var("IRONCLAW_BASE_DIR", "/my/custom/dir") }; @@ -1098,7 +1111,7 @@ INJECTED="pwned"#; #[test] fn test_ironclaw_base_dir_empty_env() { // Verifies that empty IRONCLAW_BASE_DIR falls back to default. - let _guard = ENV_MUTEX.lock().unwrap(); + let _guard = lock_env(); let old_val = std::env::var("IRONCLAW_BASE_DIR").ok(); // SAFETY: ENV_MUTEX ensures single-threaded access to env vars in tests unsafe { std::env::set_var("IRONCLAW_BASE_DIR", "") }; @@ -1120,7 +1133,7 @@ INJECTED="pwned"#; #[test] fn test_ironclaw_base_dir_special_chars() { // Verifies that paths with special characters are handled correctly. - let _guard = ENV_MUTEX.lock().unwrap(); + let _guard = lock_env(); let old_val = std::env::var("IRONCLAW_BASE_DIR").ok(); // SAFETY: ENV_MUTEX ensures single-threaded access to env vars in tests unsafe { std::env::set_var("IRONCLAW_BASE_DIR", "/tmp/test_with-special.chars") }; 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/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 7b24805c..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,210 +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 - 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, @@ -1941,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, @@ -1951,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(( @@ -1967,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() }))) } @@ -1975,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, @@ -1982,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()))?; @@ -2042,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, @@ -2062,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 @@ -2097,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) => { @@ -2105,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, @@ -2117,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); @@ -2134,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(( @@ -2141,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); @@ -2166,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()))), } @@ -2200,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 { @@ -2258,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(( @@ -2265,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()))), } @@ -2273,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(); @@ -2305,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() @@ -2336,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(( @@ -2344,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)) @@ -2366,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)> { @@ -2376,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) @@ -2395,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)) } @@ -2456,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(( @@ -2466,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 @@ -2495,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 })?; @@ -2519,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 @@ -2526,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); @@ -2543,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 { @@ -2551,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); @@ -2563,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 @@ -2570,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); @@ -2582,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 })?; @@ -2597,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 @@ -2604,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); @@ -2618,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 @@ -2864,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, @@ -2874,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![], @@ -2945,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 @@ -3023,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 @@ -3050,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, @@ -3071,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"); @@ -3233,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, @@ -3281,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; @@ -3301,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, @@ -3327,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, @@ -3404,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, @@ -3491,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, @@ -3712,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/doctor.rs b/src/cli/doctor.rs index 5d13ade6..023ac4e1 100644 --- a/src/cli/doctor.rs +++ b/src/cli/doctor.rs @@ -692,7 +692,7 @@ mod tests { } } - let _mutex = crate::config::helpers::ENV_MUTEX.lock().expect("env mutex"); + let _mutex = crate::config::helpers::lock_env(); let prev = std::env::var("LLM_BACKEND").ok(); // SAFETY: Under ENV_MUTEX, no concurrent env access. unsafe { @@ -812,7 +812,7 @@ mod tests { #[test] fn check_llm_config_shows_nearai_model_for_nearai_backend() { - let _guard = crate::config::helpers::ENV_MUTEX.lock().expect("env mutex"); + let _guard = crate::config::helpers::lock_env(); // SAFETY: Under ENV_MUTEX, no concurrent env access. unsafe { std::env::remove_var("LLM_BACKEND"); @@ -839,7 +839,7 @@ mod tests { #[test] fn check_embeddings_disabled_by_default_returns_skip() { - let _guard = crate::config::helpers::ENV_MUTEX.lock().expect("env mutex"); + let _guard = crate::config::helpers::lock_env(); // SAFETY: Under ENV_MUTEX. unsafe { std::env::remove_var("EMBEDDING_ENABLED"); @@ -861,7 +861,7 @@ mod tests { #[test] fn check_routines_enabled_by_default() { - let _guard = crate::config::helpers::ENV_MUTEX.lock().expect("env mutex"); + let _guard = crate::config::helpers::lock_env(); // SAFETY: Under ENV_MUTEX. unsafe { std::env::remove_var("ROUTINES_ENABLED"); diff --git a/src/cli/mod.rs b/src/cli/mod.rs index 9340e54f..611d7247 100644 --- a/src/cli/mod.rs +++ b/src/cli/mod.rs @@ -25,6 +25,7 @@ pub mod import; mod logs; mod mcp; pub mod memory; +mod models; pub mod oauth_defaults; mod pairing; mod registry; @@ -45,6 +46,7 @@ pub use logs::{LogsCommand, run_logs_command}; pub use mcp::{McpCommand, run_mcp_command}; pub use memory::MemoryCommand; pub use memory::run_memory_command_with_db; +pub use models::{ModelsCommand, run_models_command}; pub use pairing::{PairingCommand, run_pairing_command, run_pairing_command_with_store}; pub use registry::{RegistryCommand, run_registry_command}; pub use routines::{RoutinesCommand, run_routines_command}; @@ -217,6 +219,14 @@ pub enum Command { )] Hooks(HooksCommand), + /// Manage LLM providers and models + #[command( + subcommand, + about = "Manage LLM providers and models", + long_about = "List providers, view current configuration, and set active provider/model.\nExamples:\n ironclaw models list\n ironclaw models list openai --verbose\n ironclaw models status\n ironclaw models set gpt-4o\n ironclaw models set-provider anthropic --model claude-sonnet-4-6-20250514" + )] + Models(ModelsCommand), + /// Probe external dependencies and validate configuration #[command( about = "Run diagnostics", diff --git a/src/cli/models.rs b/src/cli/models.rs new file mode 100644 index 00000000..e24c324a --- /dev/null +++ b/src/cli/models.rs @@ -0,0 +1,864 @@ +//! Models management CLI commands. +//! +//! Provides subcommands for listing providers, viewing current model +//! configuration, and setting the active provider/model. Settings are +//! persisted to both `config.toml` and `~/.ironclaw/.env` so changes +//! take effect immediately (no DB connection required). + +use clap::Subcommand; +use std::path::Path; + +use crate::llm::registry::ProviderRegistry; +use crate::settings::Settings; + +#[derive(Subcommand, Debug, Clone)] +pub enum ModelsCommand { + /// List providers (or available models for a specific provider) + List { + /// Show only a specific provider (by ID or alias) + provider: Option, + + /// Show detailed information (env vars, base URL, protocol) + #[arg(short, long)] + verbose: bool, + + /// Output as JSON + #[arg(long)] + json: bool, + }, + + /// Show current model configuration + Status { + /// Output as JSON + #[arg(long)] + json: bool, + }, + + /// Set the default model + Set { + /// Model name (e.g., "gpt-5-mini", "claude-sonnet-4-6-20250514") + model: String, + }, + + /// Set the LLM provider + SetProvider { + /// Provider ID or alias (e.g., "openai", "anthropic", "ollama") + provider: String, + + /// Also set the model (defaults to provider's default model) + #[arg(long)] + model: Option, + }, +} + +/// Run the models CLI subcommand. +pub async fn run_models_command( + cmd: ModelsCommand, + config_path: Option<&Path>, +) -> anyhow::Result<()> { + match cmd { + ModelsCommand::List { + provider, + verbose, + json, + } => { + if let Some(ref id) = provider { + cmd_show_provider(id, verbose, json, config_path).await + } else { + cmd_list_providers(verbose, json, config_path).await + } + } + ModelsCommand::Status { json } => cmd_status(json, config_path), + ModelsCommand::Set { model } => cmd_set_model(&model, config_path), + ModelsCommand::SetProvider { provider, model } => { + cmd_set_provider(&provider, model.as_deref(), config_path) + } + } +} + +// ─── Shared helpers ─────────────────────────────────────────────── + +/// Resolve the currently active backend and model from env + settings. +fn resolve_active(config_path: Option<&Path>) -> (String, String) { + let settings = load_settings(config_path); + resolve_active_from_settings(&settings) +} + +/// Resolve active backend + model from a pre-loaded Settings. +fn resolve_active_from_settings(settings: &Settings) -> (String, String) { + let backend = std::env::var("LLM_BACKEND") + .ok() + .or_else(|| settings.llm_backend.clone()) + .unwrap_or_else(|| "nearai".to_string()); + + let registry = ProviderRegistry::load(); + + let canonical_backend = registry + .find(&backend) + .map(|d| d.id.clone()) + .unwrap_or_else(|| backend.clone()); + + let model = if canonical_backend == "nearai" { + std::env::var("NEARAI_MODEL") + .ok() + .or_else(|| settings.selected_model.clone()) + .unwrap_or_else(|| "qwen2.5-72b-instruct:free".to_string()) + } else if let Some(def) = registry.find(&canonical_backend) { + std::env::var(&def.model_env) + .ok() + .or_else(|| settings.selected_model.clone()) + .unwrap_or_else(|| def.default_model.clone()) + } else { + settings + .selected_model + .clone() + .unwrap_or_else(|| "unknown".to_string()) + }; + + (canonical_backend, model) +} + +fn load_settings(config_path: Option<&Path>) -> Settings { + if let Some(path) = config_path { + Settings::load_toml(path).ok().flatten().unwrap_or_default() + } else { + let toml_path = config_toml_path(); + if toml_path.exists() { + Settings::load_toml(&toml_path) + .ok() + .flatten() + .unwrap_or_default() + } else { + Settings::load() + } + } +} + +fn save_settings(settings: &Settings, config_path: Option<&Path>) -> anyhow::Result<()> { + let path = config_path + .map(|p| p.to_path_buf()) + .unwrap_or_else(config_toml_path); + + settings + .save_toml(&path) + .map_err(|e| anyhow::anyhow!("{}", e))?; + + Ok(()) +} + +fn config_toml_path() -> std::path::PathBuf { + crate::bootstrap::ironclaw_base_dir().join("config.toml") +} + +/// Try to fetch the live model list from a provider. +/// +/// Best-effort: returns `None` if config loading, provider creation, or the +/// `list_models()` call fails (missing API key, network error, etc.). +async fn try_fetch_models(provider_id: &str, config_path: Option<&Path>) -> Option> { + let config = crate::config::Config::from_env_with_toml(config_path) + .await + .ok()?; + + // Override backend to the requested provider so create_llm_provider + // constructs the right one. + let mut llm_config = config.llm.clone(); + llm_config.backend = provider_id.to_string(); + + // For registry providers, resolve the RegistryProviderConfig if not + // already set for this backend. + if provider_id != "nearai" && provider_id != "bedrock" { + let registry = ProviderRegistry::load(); + if let Some(def) = registry.find(provider_id) + && llm_config + .provider + .as_ref() + .is_none_or(|p| p.provider_id != def.id) + { + // Build a minimal RegistryProviderConfig from env + registry + let api_key = def + .api_key_env + .as_ref() + .and_then(|env| std::env::var(env).ok()); + if def.api_key_required && api_key.is_none() { + return None; + } + let base_url = def.default_base_url.clone().unwrap_or_default(); + llm_config.provider = Some(crate::llm::RegistryProviderConfig { + protocol: def.protocol, + provider_id: def.id.clone(), + model: def.default_model.clone(), + api_key: api_key.map(secrecy::SecretString::from), + base_url, + extra_headers: Vec::new(), + oauth_token: None, + is_codex_chatgpt: false, + refresh_token: None, + auth_path: None, + cache_retention: Default::default(), + unsupported_params: def.unsupported_params.clone(), + }); + } + } + + let session = crate::llm::create_session_manager(config.llm.session.clone()).await; + let provider = crate::llm::create_llm_provider(&llm_config, session) + .await + .ok()?; + provider.list_models().await.ok().filter(|m| !m.is_empty()) +} + +/// Print available models section (text output). +fn print_model_list(models: &Option>, active_model: Option<&String>) { + match models { + Some(models) => { + println!("\n Available models ({}):", models.len()); + for m in models { + let marker = active_model + .filter(|a| a.as_str() == m) + .map(|_| " (active)") + .unwrap_or(""); + println!(" {}{}", m, marker); + } + } + None => { + println!( + "\n Could not fetch model list (missing credentials or provider unavailable)." + ); + } + } +} + +/// Also update `~/.ironclaw/.env` so changes take effect immediately. +/// +/// Skipped when `config_path` is `Some` (custom `--config`), because the user +/// is explicitly targeting a different config file and we must not pollute the +/// default profile's `.env`. +fn sync_to_dotenv(config_path: Option<&Path>, vars: &[(&str, &str)]) { + if config_path.is_some() { + return; + } + if let Err(e) = crate::bootstrap::upsert_bootstrap_vars(vars) { + eprintln!("Warning: failed to update .env: {}", e); + } +} + +// ─── status ─────────────────────────────────────────────────────── + +fn cmd_status(json: bool, config_path: Option<&Path>) -> anyhow::Result<()> { + let settings = load_settings(config_path); + let (backend, model) = resolve_active_from_settings(&settings); + let registry = ProviderRegistry::load(); + + let fallback = std::env::var("NEARAI_FALLBACK_MODEL").ok(); + let cheap = std::env::var("NEARAI_CHEAP_MODEL").ok(); + + let description = if backend == "nearai" { + "NEAR AI inference (default)".to_string() + } else { + registry + .find(&backend) + .map(|d| d.description.clone()) + .unwrap_or_default() + }; + + if json { + let v = serde_json::json!({ + "provider": backend, + "model": model, + "description": description, + "fallback_model": fallback, + "cheap_model": cheap, + }); + println!( + "{}", + serde_json::to_string_pretty(&v).unwrap_or_else(|_| "{}".to_string()) + ); + return Ok(()); + } + + println!("Provider: {} ({})", backend, description); + println!("Model: {}", model); + if let Some(ref fb) = fallback { + println!("Fallback: {}", fb); + } + if let Some(ref ch) = cheap { + println!("Cheap: {}", ch); + } + + Ok(()) +} + +// ─── set ────────────────────────────────────────────────────────── + +fn cmd_set_model(model: &str, config_path: Option<&Path>) -> anyhow::Result<()> { + let trimmed = model.trim(); + if trimmed.is_empty() { + anyhow::bail!("Model name cannot be empty"); + } + + let mut settings = load_settings(config_path); + let registry = ProviderRegistry::load(); + + // Warn if model name doesn't match any known provider's default model + let known_model = registry.all().iter().any(|d| d.default_model == trimmed) + || trimmed.contains("qwen") // nearai models + || trimmed.contains("llama") + || trimmed.contains("gpt") + || trimmed.contains("claude") + || trimmed.contains("gemini") + || trimmed.contains("mistral"); + if !known_model { + eprintln!( + "Warning: '{}' is not a recognized model name. Proceeding anyway.", + trimmed + ); + } + + settings.selected_model = Some(trimmed.to_string()); + save_settings(&settings, config_path)?; + + let backend = std::env::var("LLM_BACKEND") + .ok() + .or_else(|| settings.llm_backend.clone()) + .unwrap_or_else(|| "nearai".to_string()); + + // Also write to .env so the change takes effect immediately + let model_env = if backend == "nearai" { + "NEARAI_MODEL".to_string() + } else { + registry + .find(&backend) + .map(|d| d.model_env.clone()) + .unwrap_or_default() + }; + if !model_env.is_empty() { + sync_to_dotenv(config_path, &[(&model_env, trimmed)]); + } + + println!("Model set to '{}' (provider: {})", trimmed, backend); + println!( + "Saved to {}", + config_path + .map(|p| p.display().to_string()) + .unwrap_or_else(|| config_toml_path().display().to_string()) + ); + + Ok(()) +} + +// ─── set-provider ───────────────────────────────────────────────── + +fn cmd_set_provider( + provider: &str, + model: Option<&str>, + config_path: Option<&Path>, +) -> anyhow::Result<()> { + let registry = ProviderRegistry::load(); + + // Validate and normalize provider + let canonical_id = if provider == "nearai" || provider == "near_ai" || provider == "near" { + "nearai".to_string() + } else { + let def = registry.find(provider).ok_or_else(|| { + let known: Vec<&str> = std::iter::once("nearai") + .chain(registry.all().iter().map(|d| d.id.as_str())) + .collect(); + anyhow::anyhow!( + "Unknown provider '{}'. Known providers: {}", + provider, + known.join(", ") + ) + })?; + def.id.clone() + }; + + // Resolve model: explicit > provider default + let resolved_model = if let Some(m) = model { + m.to_string() + } else if canonical_id == "nearai" { + "qwen2.5-72b-instruct:free".to_string() + } else if let Some(def) = registry.find(&canonical_id) { + def.default_model.clone() + } else { + "default".to_string() + }; + + let mut settings = load_settings(config_path); + settings.llm_backend = Some(canonical_id.clone()); + settings.selected_model = Some(resolved_model.clone()); + save_settings(&settings, config_path)?; + + // Also write to .env so the change takes effect immediately + let model_env = if canonical_id == "nearai" { + "NEARAI_MODEL".to_string() + } else { + registry + .find(&canonical_id) + .map(|d| d.model_env.clone()) + .unwrap_or_default() + }; + let mut vars: Vec<(&str, &str)> = vec![("LLM_BACKEND", &canonical_id)]; + if !model_env.is_empty() { + vars.push((&model_env, &resolved_model)); + } + sync_to_dotenv(config_path, &vars); + + println!( + "Provider set to '{}', model set to '{}'", + canonical_id, resolved_model + ); + println!( + "Saved to {}", + config_path + .map(|p| p.display().to_string()) + .unwrap_or_else(|| config_toml_path().display().to_string()) + ); + + Ok(()) +} + +// ─── list ───────────────────────────────────────────────────────── + +/// List all providers with their default models. +async fn cmd_list_providers( + verbose: bool, + json: bool, + config_path: Option<&Path>, +) -> anyhow::Result<()> { + let registry = ProviderRegistry::load(); + let (active_backend, active_model) = resolve_active(config_path); + + if json { + let mut entries: Vec = Vec::new(); + + // NEAR AI (not in registry) + let nearai_active = active_backend == "nearai"; + entries.push(serde_json::json!({ + "id": "nearai", + "description": "NEAR AI inference (default)", + "default_model": "qwen2.5-72b-instruct:free", + "active": nearai_active, + "active_model": if nearai_active { Some(&active_model) } else { None }, + })); + + for def in registry.all() { + let is_active = active_backend == def.id; + let mut v = serde_json::json!({ + "id": def.id, + "description": def.description, + "default_model": def.default_model, + "protocol": format!("{:?}", def.protocol), + "active": is_active, + }); + if is_active { + v["active_model"] = serde_json::json!(active_model); + } + if verbose { + v["aliases"] = serde_json::json!(def.aliases); + v["model_env"] = serde_json::json!(def.model_env); + v["api_key_env"] = serde_json::json!(def.api_key_env); + v["api_key_required"] = serde_json::json!(def.api_key_required); + if let Some(ref url) = def.default_base_url { + v["base_url"] = serde_json::json!(url); + } + if let Some(ref setup) = def.setup { + v["can_list_models"] = serde_json::json!(setup.can_list_models()); + } + } + entries.push(v); + } + + println!( + "{}", + serde_json::to_string_pretty(&entries).unwrap_or_else(|_| "[]".to_string()) + ); + return Ok(()); + } + + let providers = registry.all(); + + println!("Active: {} (model: {})\n", active_backend, active_model); + println!( + "{} provider(s) available:\n", + providers.len() + 1 // +1 for NEAR AI + ); + + // NEAR AI (not in registry) + let nearai_marker = if active_backend == "nearai" { " *" } else { "" }; + if verbose { + println!(" nearai{}", nearai_marker); + println!(" Description: NEAR AI inference (default)"); + println!(" Default model: qwen2.5-72b-instruct:free"); + println!(" Model env: NEARAI_MODEL"); + if active_backend == "nearai" { + println!(" Active model: {}", active_model); + } + println!(); + } else { + println!( + " {:<22} {:<40} NEAR AI inference (default)", + format!("nearai{nearai_marker}"), + "qwen2.5-72b-instruct:free" + ); + } + + for def in providers { + let is_active = active_backend == def.id; + let marker = if is_active { " *" } else { "" }; + + if verbose { + println!(" {}{}", def.id, marker); + println!(" Description: {}", def.description); + println!(" Default model: {}", def.default_model); + println!(" Protocol: {:?}", def.protocol); + println!(" Model env: {}", def.model_env); + if let Some(ref env) = def.api_key_env { + println!( + " API key env: {} ({})", + env, + if def.api_key_required { + "required" + } else { + "optional" + } + ); + } + if let Some(ref url) = def.default_base_url { + println!(" Base URL: {}", url); + } + if !def.aliases.is_empty() { + println!(" Aliases: {}", def.aliases.join(", ")); + } + if is_active { + println!(" Active model: {}", active_model); + } + println!(); + } else { + let model_display = if is_active { + active_model.clone() + } else { + def.default_model.clone() + }; + println!( + " {:<22} {:<40} {}", + format!("{}{marker}", def.id), + model_display, + def.description, + ); + } + } + + if !verbose { + println!(); + println!("* = active provider. Use --verbose for details."); + } + + Ok(()) +} + +/// Show details for a specific provider. +async fn cmd_show_provider( + id: &str, + verbose: bool, + json: bool, + config_path: Option<&Path>, +) -> anyhow::Result<()> { + let registry = ProviderRegistry::load(); + let (active_backend, active_model) = resolve_active(config_path); + + // Resolve canonical ID for model fetching + let canonical_id = if id == "nearai" || id == "near_ai" || id == "near" { + "nearai".to_string() + } else { + registry + .find(id) + .map(|d| d.id.clone()) + .unwrap_or_else(|| id.to_string()) + }; + + // Try to fetch live model list from the provider + let live_models = try_fetch_models(&canonical_id, config_path).await; + + // Check NEAR AI first (not in registry) + if id == "nearai" || id == "near_ai" || id == "near" { + let is_active = active_backend == "nearai"; + if json { + let mut v = serde_json::json!({ + "id": "nearai", + "description": "NEAR AI inference (default)", + "default_model": "qwen2.5-72b-instruct:free", + "model_env": "NEARAI_MODEL", + "active": is_active, + }); + if is_active { + v["active_model"] = serde_json::json!(active_model); + } + if let Some(ref models) = live_models { + v["available_models"] = serde_json::json!(models); + } + println!( + "{}", + serde_json::to_string_pretty(&v).unwrap_or_else(|_| "{}".to_string()) + ); + } else { + println!("Provider: nearai"); + println!(" Description: NEAR AI inference (default)"); + println!(" Default model: qwen2.5-72b-instruct:free"); + println!(" Model env: NEARAI_MODEL"); + println!(" Active: {}", if is_active { "yes" } else { "no" }); + if is_active { + println!(" Active model: {}", active_model); + } + print_model_list(&live_models, is_active.then_some(&active_model)); + } + return Ok(()); + } + + let def = registry.find(id).ok_or_else(|| { + let known: Vec<&str> = std::iter::once("nearai") + .chain(registry.all().iter().map(|d| d.id.as_str())) + .collect(); + anyhow::anyhow!( + "Unknown provider '{}'. Known providers: {}", + id, + known.join(", ") + ) + })?; + + let is_active = active_backend == def.id; + + if json { + let mut v = serde_json::json!({ + "id": def.id, + "description": def.description, + "protocol": format!("{:?}", def.protocol), + "default_model": def.default_model, + "model_env": def.model_env, + "api_key_env": def.api_key_env, + "api_key_required": def.api_key_required, + "aliases": def.aliases, + "active": is_active, + }); + if let Some(ref url) = def.default_base_url { + v["base_url"] = serde_json::json!(url); + } + if let Some(ref setup) = def.setup { + v["can_list_models"] = serde_json::json!(setup.can_list_models()); + v["display_name"] = serde_json::json!(setup.display_name()); + } + if is_active { + v["active_model"] = serde_json::json!(active_model); + } + if verbose && !def.unsupported_params.is_empty() { + v["unsupported_params"] = serde_json::json!(def.unsupported_params); + } + if let Some(ref models) = live_models { + v["available_models"] = serde_json::json!(models); + } + println!( + "{}", + serde_json::to_string_pretty(&v).unwrap_or_else(|_| "{}".to_string()) + ); + return Ok(()); + } + + println!("Provider: {}", def.id); + println!(" Description: {}", def.description); + println!(" Protocol: {:?}", def.protocol); + println!(" Default model: {}", def.default_model); + println!(" Model env: {}", def.model_env); + if let Some(ref env) = def.api_key_env { + println!( + " API key env: {} ({})", + env, + if def.api_key_required { + "required" + } else { + "optional" + } + ); + } + if let Some(ref url) = def.default_base_url { + println!(" Base URL: {}", url); + } + if !def.aliases.is_empty() { + println!(" Aliases: {}", def.aliases.join(", ")); + } + if let Some(ref setup) = def.setup { + println!( + " List models: {}", + if setup.can_list_models() { + "supported" + } else { + "not supported" + } + ); + println!(" Display name: {}", setup.display_name()); + } + if !def.unsupported_params.is_empty() { + println!(" Unsupported: {}", def.unsupported_params.join(", ")); + } + println!(" Active: {}", if is_active { "yes" } else { "no" }); + if is_active { + println!(" Active model: {}", active_model); + } + print_model_list(&live_models, is_active.then_some(&active_model)); + + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn resolve_active_defaults_to_nearai() { + let settings = Settings::default(); + assert!(settings.llm_backend.is_none()); + assert!(settings.selected_model.is_none()); + } + + #[test] + fn registry_loads_all_providers() { + let registry = ProviderRegistry::load(); + let all = registry.all(); + assert!( + all.len() >= 10, + "should have at least 10 built-in providers, got {}", + all.len() + ); + } + + #[test] + fn registry_find_by_alias() { + let registry = ProviderRegistry::load(); + let def = registry + .find("claude") + .expect("claude alias should resolve"); + assert_eq!(def.id, "anthropic"); + } + + #[test] + fn all_providers_have_description() { + let registry = ProviderRegistry::load(); + for def in registry.all() { + assert!( + !def.description.is_empty(), + "provider {} should have a description", + def.id + ); + } + } + + #[test] + fn set_model_persists_to_toml() { + let dir = tempfile::tempdir().expect("create temp dir"); + let toml_path = dir.path().join("config.toml"); + + cmd_set_model("gpt-5-mini", Some(&toml_path)).expect("set model"); + + let settings = Settings::load_toml(&toml_path) + .expect("read toml") + .expect("should have settings"); + assert_eq!(settings.selected_model.as_deref(), Some("gpt-5-mini")); + } + + #[test] + fn set_provider_validates_unknown() { + let dir = tempfile::tempdir().expect("create temp dir"); + let toml_path = dir.path().join("config.toml"); + + let result = cmd_set_provider("nonexistent_provider", None, Some(&toml_path)); + assert!(result.is_err()); + let err = result.unwrap_err().to_string(); + assert!( + err.contains("Unknown provider"), + "should mention unknown provider: {}", + err + ); + } + + #[test] + fn set_provider_persists_to_toml() { + let dir = tempfile::tempdir().expect("create temp dir"); + let toml_path = dir.path().join("config.toml"); + + cmd_set_provider("groq", None, Some(&toml_path)).expect("set provider"); + + let settings = Settings::load_toml(&toml_path) + .expect("read toml") + .expect("should have settings"); + assert_eq!(settings.llm_backend.as_deref(), Some("groq")); + assert_eq!( + settings.selected_model.as_deref(), + Some("llama-3.3-70b-versatile") + ); + } + + #[test] + fn set_provider_with_custom_model() { + let dir = tempfile::tempdir().expect("create temp dir"); + let toml_path = dir.path().join("config.toml"); + + cmd_set_provider("anthropic", Some("claude-opus-4-6"), Some(&toml_path)) + .expect("set provider with model"); + + let settings = Settings::load_toml(&toml_path) + .expect("read toml") + .expect("should have settings"); + assert_eq!(settings.llm_backend.as_deref(), Some("anthropic")); + assert_eq!(settings.selected_model.as_deref(), Some("claude-opus-4-6")); + } + + #[test] + fn custom_config_does_not_pollute_default_dotenv() { + let dir = tempfile::tempdir().expect("create temp dir"); + let toml_path = dir.path().join("config.toml"); + + // With a custom config path, sync_to_dotenv should be a no-op + // (it returns early when config_path is Some). + // We verify by checking that cmd_set_provider succeeds without + // trying to write to the default ~/.ironclaw/.env. + cmd_set_provider("groq", None, Some(&toml_path)).expect("set provider with custom config"); + + let settings = Settings::load_toml(&toml_path) + .expect("read toml") + .expect("should have settings"); + assert_eq!(settings.llm_backend.as_deref(), Some("groq")); + // The key assertion is that no error was thrown trying to write + // to the default .env — sync_to_dotenv skipped it. + } + + #[test] + fn set_model_rejects_empty_name() { + let dir = tempfile::tempdir().expect("create temp dir"); + let toml_path = dir.path().join("config.toml"); + + let result = cmd_set_model("", Some(&toml_path)); + assert!(result.is_err()); + assert!( + result.unwrap_err().to_string().contains("cannot be empty"), + "should reject empty model name" + ); + + let result2 = cmd_set_model(" ", Some(&toml_path)); + assert!(result2.is_err()); + } + + #[test] + fn set_provider_normalizes_alias() { + let dir = tempfile::tempdir().expect("create temp dir"); + let toml_path = dir.path().join("config.toml"); + + cmd_set_provider("claude", None, Some(&toml_path)).expect("set via alias"); + + let settings = Settings::load_toml(&toml_path) + .expect("read toml") + .expect("should have settings"); + assert_eq!( + settings.llm_backend.as_deref(), + Some("anthropic"), + "alias should be normalized to canonical ID" + ); + } +} diff --git a/src/cli/oauth_defaults.rs b/src/cli/oauth_defaults.rs index b4e93704..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. @@ -758,7 +758,7 @@ mod tests { use crate::cli::oauth_defaults::{ builtin_credentials, callback_host, callback_url, is_loopback_host, landing_html, }; - use crate::config::helpers::ENV_MUTEX; + use crate::config::helpers::lock_env; #[test] fn test_is_loopback_host() { @@ -775,7 +775,7 @@ mod tests { #[test] fn test_callback_host_default() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); let original = std::env::var("OAUTH_CALLBACK_HOST").ok(); // SAFETY: Under ENV_MUTEX, no concurrent env access. unsafe { @@ -792,7 +792,7 @@ mod tests { #[test] fn test_callback_host_env_override() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); let original_host = std::env::var("OAUTH_CALLBACK_HOST").ok(); let original_url = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok(); // SAFETY: Under ENV_MUTEX, no concurrent env access. @@ -819,7 +819,7 @@ mod tests { #[test] fn test_callback_url_default() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); // Clear both env vars to test default behavior let original_url = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok(); let original_host = std::env::var("OAUTH_CALLBACK_HOST").ok(); @@ -843,7 +843,7 @@ mod tests { #[test] fn test_callback_url_env_override() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok(); // SAFETY: Under ENV_MUTEX, no concurrent env access. unsafe { @@ -1008,7 +1008,7 @@ mod tests { #[test] fn test_use_gateway_callback_false_by_default() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok(); // SAFETY: Under ENV_MUTEX, no concurrent env access. unsafe { @@ -1024,7 +1024,7 @@ mod tests { #[test] fn test_use_gateway_callback_true_for_hosted() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok(); // SAFETY: Under ENV_MUTEX, no concurrent env access. unsafe { @@ -1045,7 +1045,7 @@ mod tests { #[test] fn test_use_gateway_callback_false_for_localhost() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok(); // SAFETY: Under ENV_MUTEX, no concurrent env access. unsafe { @@ -1063,7 +1063,7 @@ mod tests { #[test] fn test_use_gateway_callback_false_for_empty() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok(); // SAFETY: Under ENV_MUTEX, no concurrent env access. unsafe { @@ -1083,7 +1083,7 @@ mod tests { fn test_build_platform_state_with_instance() { use crate::cli::oauth_defaults::{build_platform_state, decode_hosted_oauth_state}; - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); let original = std::env::var("IRONCLAW_INSTANCE_NAME").ok(); // SAFETY: Under ENV_MUTEX, no concurrent env access. unsafe { @@ -1107,7 +1107,7 @@ mod tests { fn test_build_platform_state_without_instance() { use crate::cli::oauth_defaults::{build_platform_state, decode_hosted_oauth_state}; - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); let original = std::env::var("IRONCLAW_INSTANCE_NAME").ok(); let original_oc = std::env::var("OPENCLAW_INSTANCE_NAME").ok(); // SAFETY: Under ENV_MUTEX, no concurrent env access. @@ -1134,7 +1134,7 @@ mod tests { fn test_build_platform_state_with_openclaw_instance() { use crate::cli::oauth_defaults::{build_platform_state, decode_hosted_oauth_state}; - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); let original_ic = std::env::var("IRONCLAW_INSTANCE_NAME").ok(); let original_oc = std::env::var("OPENCLAW_INSTANCE_NAME").ok(); // SAFETY: Under ENV_MUTEX, no concurrent env access. 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/cli/snapshots/ironclaw__cli__tests__help_output.snap b/src/cli/snapshots/ironclaw__cli__tests__help_output.snap index 13a45bb5..e946381f 100644 --- a/src/cli/snapshots/ironclaw__cli__tests__help_output.snap +++ b/src/cli/snapshots/ironclaw__cli__tests__help_output.snap @@ -20,6 +20,7 @@ Commands: service Manage OS service skills Manage skills hooks Manage lifecycle hooks + models Manage LLM providers and models doctor Run diagnostics logs View and manage gateway logs status Show system status diff --git a/src/cli/snapshots/ironclaw__cli__tests__help_output_without_import.snap b/src/cli/snapshots/ironclaw__cli__tests__help_output_without_import.snap index 52177b76..8fcec25e 100644 --- a/src/cli/snapshots/ironclaw__cli__tests__help_output_without_import.snap +++ b/src/cli/snapshots/ironclaw__cli__tests__help_output_without_import.snap @@ -20,6 +20,7 @@ Commands: service Manage OS service skills Manage skills hooks Manage lifecycle hooks + models Manage LLM providers and models doctor Run diagnostics logs View and manage gateway logs status Show system status diff --git a/src/cli/snapshots/ironclaw__cli__tests__long_help_output.snap b/src/cli/snapshots/ironclaw__cli__tests__long_help_output.snap index 9f0dbfb7..63dcbb04 100644 --- a/src/cli/snapshots/ironclaw__cli__tests__long_help_output.snap +++ b/src/cli/snapshots/ironclaw__cli__tests__long_help_output.snap @@ -23,6 +23,7 @@ Commands: service Manage OS service skills Manage skills hooks Manage lifecycle hooks + models Manage LLM providers and models doctor Run diagnostics logs View and manage gateway logs status Show system status diff --git a/src/cli/snapshots/ironclaw__cli__tests__long_help_output_without_import.snap b/src/cli/snapshots/ironclaw__cli__tests__long_help_output_without_import.snap index efef7eac..cb799ce7 100644 --- a/src/cli/snapshots/ironclaw__cli__tests__long_help_output_without_import.snap +++ b/src/cli/snapshots/ironclaw__cli__tests__long_help_output_without_import.snap @@ -23,6 +23,7 @@ Commands: service Manage OS service skills Manage skills hooks Manage lifecycle hooks + models Manage LLM providers and models doctor Run diagnostics logs View and manage gateway logs status Show system status diff --git a/src/config/builder.rs b/src/config/builder.rs index 088db90c..f7bad12c 100644 --- a/src/config/builder.rs +++ b/src/config/builder.rs @@ -63,12 +63,12 @@ impl BuilderModeConfig { #[cfg(test)] mod tests { use super::*; - use crate::config::helpers::ENV_MUTEX; + use crate::config::helpers::lock_env; use crate::settings::Settings; #[test] fn resolve_falls_back_to_settings() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); let mut settings = Settings::default(); settings.builder.max_iterations = 99; settings.builder.auto_register = false; @@ -80,7 +80,7 @@ mod tests { #[test] fn env_overrides_settings() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); let mut settings = Settings::default(); settings.builder.timeout_secs = 123; diff --git a/src/config/channels.rs b/src/config/channels.rs index bc704445..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). @@ -113,8 +134,120 @@ impl ChannelsConfig { let gateway = if gateway_enabled { let user_id = optional_env("GATEWAY_USER_ID")? .or_else(|| cs.gateway_user_id.clone()) - .unwrap_or_else(|| "default".to_string()); + .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 @@ -236,7 +372,7 @@ fn default_channels_dir() -> PathBuf { #[cfg(test)] mod tests { use crate::config::channels::*; - use crate::config::helpers::ENV_MUTEX; + use crate::config::helpers::lock_env; use crate::settings::Settings; #[test] @@ -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()); } @@ -395,7 +537,7 @@ mod tests { #[test] fn resolve_uses_settings_channel_values_with_owner_scope_user_ids() { - let _guard = ENV_MUTEX.lock().unwrap_or_else(|e| e.into_inner()); + let _guard = lock_env(); let mut settings = Settings::default(); settings.channels.http_enabled = true; settings.channels.http_host = Some("127.0.0.2".to_string()); diff --git a/src/config/embeddings.rs b/src/config/embeddings.rs index 68b0ff2c..98183976 100644 --- a/src/config/embeddings.rs +++ b/src/config/embeddings.rs @@ -196,7 +196,7 @@ impl EmbeddingsConfig { #[cfg(test)] mod tests { use super::*; - use crate::config::helpers::ENV_MUTEX; + use crate::config::helpers::lock_env; use crate::settings::{EmbeddingsSettings, Settings}; use crate::testing::credentials::*; @@ -215,7 +215,7 @@ mod tests { #[test] fn embeddings_disabled_not_overridden_by_openai_key() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_embedding_env(); // SAFETY: Under ENV_MUTEX, no concurrent env access. unsafe { @@ -245,7 +245,7 @@ mod tests { #[test] fn embeddings_enabled_from_settings() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_embedding_env(); let settings = Settings { @@ -265,7 +265,7 @@ mod tests { #[test] fn embeddings_env_override_takes_precedence() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_embedding_env(); // SAFETY: Under ENV_MUTEX. unsafe { @@ -294,7 +294,7 @@ mod tests { #[test] fn embedding_base_url_parsed_from_env() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_embedding_env(); // SAFETY: Under ENV_MUTEX, no concurrent env access. @@ -313,7 +313,7 @@ mod tests { #[test] fn embedding_base_url_defaults_to_none() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_embedding_env(); let settings = Settings::default(); @@ -326,7 +326,7 @@ mod tests { #[test] fn cache_size_zero_rejected() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_embedding_env(); // SAFETY: Under ENV_MUTEX. unsafe { diff --git a/src/config/helpers.rs b/src/config/helpers.rs index dc40fc9f..ff5ee706 100644 --- a/src/config/helpers.rs +++ b/src/config/helpers.rs @@ -14,6 +14,16 @@ use crate::config::INJECTED_VARS; #[cfg(test)] pub(crate) static ENV_MUTEX: std::sync::Mutex<()> = std::sync::Mutex::new(()); +/// Acquire the env-var mutex, recovering from poison. +/// +/// A poisoned mutex means a previous test panicked while holding the lock. +/// The env state might be slightly stale, but cascading every subsequent +/// test into a `PoisonError` panic is far worse. Recover and carry on. +#[cfg(test)] +pub(crate) fn lock_env() -> std::sync::MutexGuard<'static, ()> { + ENV_MUTEX.lock().unwrap_or_else(|e| e.into_inner()) +} + /// Thread-safe mutable overlay for env vars set at runtime. /// /// Unlike `INJECTED_VARS` (which is set once at startup from the secrets @@ -353,7 +363,7 @@ mod tests { #[test] fn real_env_var_takes_priority_over_runtime_override() { - let _guard = ENV_MUTEX.lock().unwrap(); + let _guard = lock_env(); let key = "IRONCLAW_TEST_ENV_PRIORITY_42"; // Set runtime override @@ -372,6 +382,26 @@ mod tests { assert_eq!(env_or_override(key), Some("override_value".to_string())); } + // --- lock_env poison recovery (regression for env mutex cascade) --- + + #[test] + fn lock_env_recovers_from_poisoned_mutex() { + // Simulate a poisoned mutex: spawn a thread that panics while holding the lock. + let _ = std::thread::spawn(|| { + let _guard = ENV_MUTEX.lock().unwrap(); + panic!("intentional poison"); + }) + .join(); + + // The mutex is now poisoned. lock_env() should recover, not cascade. + assert!(ENV_MUTEX.lock().is_err(), "mutex should be poisoned"); + let _guard = lock_env(); // must not panic + drop(_guard); + + // Clean up so this test doesn't leave ENV_MUTEX permanently poisoned. + ENV_MUTEX.clear_poison(); + } + // --- validate_base_url tests (regression for #1103) --- #[test] diff --git a/src/config/llm.rs b/src/config/llm.rs index 0976051f..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), }); } @@ -532,10 +535,15 @@ pub fn default_session_path() -> PathBuf { #[cfg(test)] mod tests { use super::*; - use crate::config::helpers::ENV_MUTEX; + use crate::config::helpers::lock_env; 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. @@ -548,7 +556,7 @@ mod tests { #[test] fn openai_compatible_uses_selected_model_when_llm_model_unset() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_openai_compatible_env(); let settings = Settings { @@ -566,7 +574,7 @@ mod tests { #[test] fn openai_compatible_llm_model_env_overrides_selected_model() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_openai_compatible_env(); // SAFETY: Under ENV_MUTEX. unsafe { @@ -690,7 +698,7 @@ mod tests { #[test] fn ollama_uses_selected_model_when_ollama_model_unset() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_ollama_env(); let settings = Settings { @@ -707,7 +715,7 @@ mod tests { #[test] fn ollama_model_env_overrides_selected_model() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_ollama_env(); // SAFETY: Under ENV_MUTEX. unsafe { @@ -733,7 +741,7 @@ mod tests { #[test] fn openai_compatible_preserves_dotted_model_name() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_openai_compatible_env(); let settings = Settings { @@ -754,7 +762,7 @@ mod tests { #[test] fn registry_provider_resolves_groq() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); // SAFETY: Under ENV_MUTEX. unsafe { std::env::remove_var("LLM_BACKEND"); @@ -779,7 +787,7 @@ mod tests { #[test] fn registry_provider_resolves_tinfoil() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); // SAFETY: Under ENV_MUTEX. unsafe { std::env::remove_var("LLM_BACKEND"); @@ -807,7 +815,7 @@ mod tests { #[test] fn registry_provider_alias_resolves_zai() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); // SAFETY: Under ENV_MUTEX. unsafe { std::env::remove_var("LLM_BACKEND"); @@ -832,7 +840,7 @@ mod tests { #[test] fn registry_provider_resolves_github_copilot_alias() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); // SAFETY: Under ENV_MUTEX. unsafe { std::env::set_var("LLM_BACKEND", "github-copilot"); @@ -880,7 +888,7 @@ mod tests { #[test] fn nearai_backend_has_no_registry_provider() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); // SAFETY: Under ENV_MUTEX. unsafe { std::env::remove_var("LLM_BACKEND"); @@ -894,7 +902,7 @@ mod tests { #[test] fn backend_alias_normalized_to_canonical_id() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_openai_compatible_env(); // SAFETY: Under ENV_MUTEX. unsafe { @@ -920,7 +928,7 @@ mod tests { #[test] fn unknown_backend_falls_back_to_openai_compatible() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_openai_compatible_env(); // SAFETY: Under ENV_MUTEX. unsafe { @@ -944,7 +952,7 @@ mod tests { #[test] fn nearai_aliases_all_resolve_to_nearai() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); for alias in &["nearai", "near_ai", "near"] { // SAFETY: Under ENV_MUTEX. @@ -971,7 +979,7 @@ mod tests { #[test] fn base_url_resolution_priority() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_openai_compatible_env(); // SAFETY: Under ENV_MUTEX. @@ -1029,7 +1037,7 @@ mod tests { fn anthropic_oauth_token_sets_placeholder_api_key() { use secrecy::ExposeSecret; - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_anthropic_env(); // SAFETY: Under ENV_MUTEX. unsafe { @@ -1067,7 +1075,7 @@ mod tests { fn anthropic_api_key_takes_priority_over_oauth() { use secrecy::ExposeSecret; - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_anthropic_env(); // SAFETY: Under ENV_MUTEX. unsafe { @@ -1100,7 +1108,7 @@ mod tests { #[test] fn non_anthropic_provider_has_no_oauth_token() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_anthropic_env(); // SAFETY: Under ENV_MUTEX. unsafe { @@ -1208,7 +1216,7 @@ mod tests { #[test] fn test_request_timeout_defaults_to_120() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); // SAFETY: Under ENV_MUTEX. unsafe { std::env::remove_var("LLM_REQUEST_TIMEOUT_SECS"); @@ -1219,7 +1227,7 @@ mod tests { #[test] fn test_request_timeout_configurable() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); // SAFETY: Under ENV_MUTEX. unsafe { std::env::set_var("LLM_REQUEST_TIMEOUT_SECS", "300"); @@ -1246,7 +1254,7 @@ mod tests { #[test] fn openai_codex_resolves_config() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_openai_codex_env(); let settings = Settings { @@ -1266,7 +1274,7 @@ mod tests { #[test] fn openai_codex_model_env_resolution() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_openai_codex_env(); // SAFETY: Under ENV_MUTEX. unsafe { @@ -1290,7 +1298,7 @@ mod tests { #[test] fn openai_codex_falls_back_to_openai_model() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_openai_codex_env(); // SAFETY: Under ENV_MUTEX. unsafe { @@ -1314,7 +1322,7 @@ mod tests { #[test] fn openai_codex_falls_back_to_selected_model() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_openai_codex_env(); let settings = Settings { @@ -1331,7 +1339,7 @@ mod tests { /// Regression: SSRF validation on OPENAI_CODEX_API_URL (#1103). #[test] fn openai_codex_rejects_ssrf_api_url() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_openai_codex_env(); // SAFETY: Under ENV_MUTEX. unsafe { @@ -1362,7 +1370,7 @@ mod tests { /// Regression: SSRF validation on OPENAI_CODEX_AUTH_URL (#1103). #[test] fn openai_codex_rejects_ssrf_auth_url() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_openai_codex_env(); // SAFETY: Under ENV_MUTEX. unsafe { diff --git a/src/config/mod.rs b/src/config/mod.rs index 68b23ab2..dcda0fe9 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -24,7 +24,7 @@ mod skills; mod transcription; mod tunnel; mod wasm; -mod workspace; +pub(crate) mod workspace; use std::collections::HashMap; use std::sync::{LazyLock, Mutex, Once}; @@ -178,9 +178,7 @@ impl Config { }, transcription: TranscriptionConfig::default(), search: WorkspaceSearchConfig::default(), - workspace: WorkspaceConfig { - memory_layers: vec![], - }, + workspace: WorkspaceConfig::default(), observability: crate::observability::ObservabilityConfig::default(), relay: None, } @@ -313,11 +311,14 @@ 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.clone()) - .unwrap_or_else(|| "default".to_string()); + .map(|gw| gw.user_id.as_str()) + .unwrap_or("default"); + let workspace = WorkspaceConfig::resolve(workspace_user_id)?; Ok(Self { owner_id: owner_id.clone(), @@ -339,7 +340,7 @@ impl Config { skills: SkillsConfig::resolve()?, transcription: TranscriptionConfig::resolve(settings)?, search: WorkspaceSearchConfig::resolve()?, - workspace: WorkspaceConfig::resolve(&workspace_user_id)?, + workspace, observability: crate::observability::ObservabilityConfig { backend: std::env::var("OBSERVABILITY_BACKEND").unwrap_or_else(|_| "none".into()), }, diff --git a/src/config/safety.rs b/src/config/safety.rs index ff9e900a..edeceee0 100644 --- a/src/config/safety.rs +++ b/src/config/safety.rs @@ -19,12 +19,12 @@ pub(crate) fn resolve_safety_config( #[cfg(test)] mod tests { use super::*; - use crate::config::helpers::ENV_MUTEX; + use crate::config::helpers::lock_env; use crate::settings::Settings; #[test] fn resolve_falls_back_to_settings() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); let mut settings = Settings::default(); settings.safety.max_output_length = 42; settings.safety.injection_check_enabled = false; @@ -36,7 +36,7 @@ mod tests { #[test] fn env_overrides_settings() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); let mut settings = Settings::default(); settings.safety.max_output_length = 42; diff --git a/src/config/sandbox.rs b/src/config/sandbox.rs index 8c0eb689..01a8c327 100644 --- a/src/config/sandbox.rs +++ b/src/config/sandbox.rs @@ -594,9 +594,7 @@ mod tests { #[test] fn sandbox_resolve_falls_back_to_settings() { - let _guard = crate::config::helpers::ENV_MUTEX - .lock() - .expect("env mutex poisoned"); + let _guard = crate::config::helpers::lock_env(); let mut settings = crate::settings::Settings::default(); settings.sandbox.cpu_shares = 99; settings.sandbox.auto_pull_image = false; @@ -610,9 +608,7 @@ mod tests { #[test] fn sandbox_env_overrides_settings() { - let _guard = crate::config::helpers::ENV_MUTEX - .lock() - .expect("env mutex poisoned"); + let _guard = crate::config::helpers::lock_env(); let mut settings = crate::settings::Settings::default(); settings.sandbox.timeout_secs = 999; @@ -628,9 +624,7 @@ mod tests { #[test] fn claude_code_resolve_uses_settings_enabled() { - let _guard = crate::config::helpers::ENV_MUTEX - .lock() - .expect("env mutex poisoned"); + let _guard = crate::config::helpers::lock_env(); let mut settings = crate::settings::Settings::default(); settings.sandbox.claude_code_enabled = true; @@ -640,9 +634,7 @@ mod tests { #[test] fn claude_code_resolve_defaults_disabled() { - let _guard = crate::config::helpers::ENV_MUTEX - .lock() - .expect("env mutex poisoned"); + let _guard = crate::config::helpers::lock_env(); let settings = crate::settings::Settings::default(); let cfg = ClaudeCodeConfig::resolve(&settings).expect("resolve"); assert!(!cfg.enabled); @@ -650,9 +642,7 @@ mod tests { #[test] fn claude_code_env_overrides_settings() { - let _guard = crate::config::helpers::ENV_MUTEX - .lock() - .expect("env mutex poisoned"); + let _guard = crate::config::helpers::lock_env(); let mut settings = crate::settings::Settings::default(); settings.sandbox.claude_code_enabled = true; diff --git a/src/config/search.rs b/src/config/search.rs index 9555fecc..e6b663cf 100644 --- a/src/config/search.rs +++ b/src/config/search.rs @@ -92,7 +92,7 @@ impl WorkspaceSearchConfig { #[cfg(test)] mod tests { use super::*; - use crate::config::helpers::ENV_MUTEX; + use crate::config::helpers::lock_env; fn clear_search_env() { // SAFETY: Only called under ENV_MUTEX in tests. @@ -106,7 +106,7 @@ mod tests { #[test] fn defaults_when_no_env() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_search_env(); let config = WorkspaceSearchConfig::resolve().expect("should resolve"); @@ -118,7 +118,7 @@ mod tests { #[test] fn env_overrides() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_search_env(); // SAFETY: Under ENV_MUTEX. @@ -140,7 +140,7 @@ mod tests { #[test] fn invalid_strategy_rejected() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_search_env(); // SAFETY: Under ENV_MUTEX. @@ -156,7 +156,7 @@ mod tests { #[test] fn weighted_strategy_defaults() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_search_env(); // SAFETY: Under ENV_MUTEX. @@ -175,7 +175,7 @@ mod tests { #[test] fn weighted_both_zero_rejected() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_search_env(); // SAFETY: Under ENV_MUTEX. @@ -193,7 +193,7 @@ mod tests { #[test] fn rrf_both_zero_allowed() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); clear_search_env(); // SAFETY: Under ENV_MUTEX. diff --git a/src/config/wasm.rs b/src/config/wasm.rs index a9bfbd35..4c494a38 100644 --- a/src/config/wasm.rs +++ b/src/config/wasm.rs @@ -95,12 +95,12 @@ impl WasmConfig { #[cfg(test)] mod tests { use super::*; - use crate::config::helpers::ENV_MUTEX; + use crate::config::helpers::lock_env; use crate::settings::Settings; #[test] fn resolve_falls_back_to_settings() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); let mut settings = Settings::default(); settings.wasm.default_memory_limit = 42; settings.wasm.cache_compiled = false; @@ -112,7 +112,7 @@ mod tests { #[test] fn env_overrides_settings() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); let mut settings = Settings::default(); settings.wasm.default_fuel_limit = 42; diff --git a/src/config/workspace.rs b/src/config/workspace.rs index 5f89c655..27bc06f0 100644 --- a/src/config/workspace.rs +++ b/src/config/workspace.rs @@ -2,18 +2,29 @@ use crate::config::helpers::optional_env; use crate::error::ConfigError; use crate::workspace::layer::MemoryLayer; -/// Workspace memory configuration. +/// Workspace-level configuration (memory layers, read scopes). /// -/// Controls memory layer definitions for privacy-aware writes. -/// Layers are parsed from the `MEMORY_LAYERS` env var (JSON array) -/// or default to a single private layer scoped to the gateway user. -#[derive(Debug, Clone)] +/// Parsed from environment variables. Lives outside of `GatewayConfig` +/// so that non-gateway channels can eventually use the same settings. +#[derive(Debug, Clone, Default)] pub struct WorkspaceConfig { + /// Memory layer definitions (JSON in `MEMORY_LAYERS` env var, or defaults). pub memory_layers: Vec, + /// Additional user scopes for workspace reads. + /// + /// When set, the workspace can read (search, read, list) from these + /// additional user scopes while writes remain isolated to the primary + /// `user_id`. Parsed from `WORKSPACE_READ_SCOPES` (comma-separated). + pub read_scopes: Vec, } impl WorkspaceConfig { - pub(crate) fn resolve(user_id: &str) -> Result { + /// Resolve workspace config from environment variables. + /// + /// `user_id` is used to derive default memory layers when `MEMORY_LAYERS` + /// is not set. + pub fn resolve(user_id: &str) -> Result { + // --- Memory layers --- let memory_layers: Vec = match optional_env("MEMORY_LAYERS")? { Some(json_str) => { serde_json::from_str(&json_str).map_err(|e| ConfigError::InvalidValue { @@ -57,6 +68,20 @@ impl WorkspaceConfig { message: format!("layer '{}' has an empty scope", layer.name), }); } + if !layer + .scope + .chars() + .all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '-') + { + return Err(ConfigError::InvalidValue { + key: "MEMORY_LAYERS".to_string(), + message: format!( + "layer '{}' scope '{}' contains invalid characters \ + (allowed: a-z, A-Z, 0-9, _, -)", + layer.name, layer.scope + ), + }); + } } // Check for duplicate layer names @@ -72,20 +97,53 @@ impl WorkspaceConfig { } } - Ok(Self { memory_layers }) + // --- Read scopes --- + let 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 &read_scopes { + if scope.len() > 128 { + let prefix: String = scope.chars().take(32).collect(); + return Err(ConfigError::InvalidValue { + key: "WORKSPACE_READ_SCOPES".to_string(), + message: format!("scope '{prefix}...' exceeds 128 characters"), + }); + } + if !scope + .chars() + .all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '-') + { + return Err(ConfigError::InvalidValue { + key: "WORKSPACE_READ_SCOPES".to_string(), + message: format!( + "scope '{}' contains invalid characters \ + (allowed: a-z, A-Z, 0-9, _, -)", + scope + ), + }); + } + } + + Ok(Self { + memory_layers, + read_scopes, + }) } } #[cfg(test)] mod tests { use super::*; - use std::sync::Mutex; - - // Serialize env-var-dependent tests to avoid races. - static ENV_LOCK: Mutex<()> = Mutex::new(()); + use crate::config::helpers::lock_env; fn with_env(key: &str, val: Option<&str>, f: impl FnOnce()) { - let _guard = ENV_LOCK.lock().unwrap(); + let _guard = lock_env(); let prev = std::env::var(key).ok(); match val { Some(v) => unsafe { std::env::set_var(key, v) }, 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/libsql/workspace.rs b/src/db/libsql/workspace.rs index d43f1277..5680e435 100644 --- a/src/db/libsql/workspace.rs +++ b/src/db/libsql/workspace.rs @@ -1017,7 +1017,7 @@ mod tests { mod resolve_dimension { use super::*; - use crate::config::helpers::ENV_MUTEX; + use crate::config::helpers::lock_env; fn clear_embedding_env() { // SAFETY: called under ENV_MUTEX @@ -1030,14 +1030,14 @@ mod tests { #[test] fn returns_none_when_disabled() { - let _guard = ENV_MUTEX.lock().expect("env mutex"); + let _guard = lock_env(); clear_embedding_env(); assert!(resolve_embedding_dimension().is_none()); } #[test] fn returns_explicit_dimension() { - let _guard = ENV_MUTEX.lock().expect("env mutex"); + let _guard = lock_env(); clear_embedding_env(); // SAFETY: under ENV_MUTEX unsafe { @@ -1053,7 +1053,7 @@ mod tests { #[test] fn infers_from_model() { - let _guard = ENV_MUTEX.lock().expect("env mutex"); + let _guard = lock_env(); clear_embedding_env(); // SAFETY: under ENV_MUTEX unsafe { @@ -1069,7 +1069,7 @@ mod tests { #[test] fn defaults_to_1536_for_unknown_model() { - let _guard = ENV_MUTEX.lock().expect("env mutex"); + let _guard = lock_env(); clear_embedding_env(); // SAFETY: under ENV_MUTEX unsafe { diff --git a/src/db/mod.rs b/src/db/mod.rs index 900d1810..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>; @@ -644,6 +652,103 @@ pub trait WorkspaceStore: Send + Sync { embedding: Option<&[f32]>, config: &SearchConfig, ) -> Result, WorkspaceError>; + + // ==================== Multi-scope read methods ==================== + // + // Default implementations loop over user_ids calling single-scope methods, + // then merge results. Backends can override with efficient SQL (e.g., + // `WHERE user_id = ANY($1::text[])`). + + /// Hybrid search across multiple user scopes, merging results by score. + /// + /// **Note:** The default implementation calls `hybrid_search` per scope and + /// merges by raw score. Because RRF scores are normalized independently + /// within each scope, scores are not directly comparable across scopes. + /// The Postgres backend overrides this with a single combined query that + /// applies RRF once to the unified result set. + async fn hybrid_search_multi( + &self, + user_ids: &[String], + agent_id: Option, + query: &str, + embedding: Option<&[f32]>, + config: &SearchConfig, + ) -> Result, WorkspaceError> { + if user_ids.len() > 1 { + tracing::debug!( + scope_count = user_ids.len(), + "hybrid_search_multi: using default per-scope RRF merge; \ + cross-scope score comparison may be unreliable" + ); + } + let mut all_results = Vec::new(); + for uid in user_ids { + let results = self + .hybrid_search(uid, agent_id, query, embedding, config) + .await?; + all_results.extend(results); + } + // Re-sort by score descending and truncate to limit + all_results.sort_by(|a, b| { + b.score + .partial_cmp(&a.score) + .unwrap_or(std::cmp::Ordering::Equal) + }); + all_results.truncate(config.limit); + Ok(all_results) + } + + /// List all file paths across multiple user scopes. + async fn list_all_paths_multi( + &self, + user_ids: &[String], + agent_id: Option, + ) -> Result, WorkspaceError> { + let mut all_paths = Vec::new(); + for uid in user_ids { + let paths = self.list_all_paths(uid, agent_id).await?; + all_paths.extend(paths); + } + all_paths.sort(); + all_paths.dedup(); + Ok(all_paths) + } + + /// Get a document by path, searching across multiple user scopes. + /// + /// Returns the first match found (tries each user_id in order). + async fn get_document_by_path_multi( + &self, + user_ids: &[String], + agent_id: Option, + path: &str, + ) -> Result { + for uid in user_ids { + match self.get_document_by_path(uid, agent_id, path).await { + Ok(doc) => return Ok(doc), + Err(WorkspaceError::DocumentNotFound { .. }) => continue, + Err(e) => return Err(e), + } + } + Err(WorkspaceError::DocumentNotFound { + doc_type: path.to_string(), + user_id: format!("[{}]", user_ids.join(", ")), + }) + } + + /// List directory contents across multiple user scopes. + async fn list_directory_multi( + &self, + user_ids: &[String], + agent_id: Option, + directory: &str, + ) -> Result, WorkspaceError> { + let mut all_entries = Vec::new(); + for uid in user_ids { + all_entries.extend(self.list_directory(uid, agent_id, directory).await?); + } + Ok(crate::workspace::merge_workspace_entries(all_entries)) + } } /// Backend-agnostic database supertrait. diff --git a/src/db/postgres.rs b/src/db/postgres.rs index e77452db..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, @@ -717,4 +731,49 @@ impl WorkspaceStore for PgBackend { .hybrid_search(user_id, agent_id, query, embedding, config) .await } + + // Optimized multi-scope overrides using `ANY($1::text[])` SQL. + + async fn hybrid_search_multi( + &self, + user_ids: &[String], + agent_id: Option, + query: &str, + embedding: Option<&[f32]>, + config: &SearchConfig, + ) -> Result, WorkspaceError> { + self.repo + .hybrid_search_multi(user_ids, agent_id, query, embedding, config) + .await + } + + async fn list_all_paths_multi( + &self, + user_ids: &[String], + agent_id: Option, + ) -> Result, WorkspaceError> { + self.repo.list_all_paths_multi(user_ids, agent_id).await + } + + async fn get_document_by_path_multi( + &self, + user_ids: &[String], + agent_id: Option, + path: &str, + ) -> Result { + self.repo + .get_document_by_path_multi(user_ids, agent_id, path) + .await + } + + async fn list_directory_multi( + &self, + user_ids: &[String], + agent_id: Option, + directory: &str, + ) -> Result, WorkspaceError> { + self.repo + .list_directory_multi(user_ids, agent_id, directory) + .await + } } diff --git a/src/error.rs b/src/error.rs index 30ec58f4..e4f1b957 100644 --- a/src/error.rs +++ b/src/error.rs @@ -304,9 +304,6 @@ pub enum WorkspaceError { #[error("I/O error: {reason}")] IoError { reason: String }, - #[error("Not found: {path}")] - NotFound { path: String }, - #[error("Layer not found: {name}")] LayerNotFound { name: String }, diff --git a/src/extensions/manager.rs b/src/extensions/manager.rs index 3ecf3657..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!( @@ -7305,9 +7393,7 @@ mod tests { #[test] fn should_use_gateway_mode_true_for_tunnel_url() { - let _guard = crate::config::helpers::ENV_MUTEX - .lock() - .expect("env mutex poisoned"); + let _guard = crate::config::helpers::lock_env(); let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok(); // SAFETY: Under ENV_MUTEX, no concurrent env access. unsafe { @@ -7329,9 +7415,7 @@ mod tests { #[test] fn should_use_gateway_mode_false_without_tunnel() { - let _guard = crate::config::helpers::ENV_MUTEX - .lock() - .expect("env mutex poisoned"); + let _guard = crate::config::helpers::lock_env(); let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok(); unsafe { std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL"); @@ -7352,9 +7436,7 @@ mod tests { #[test] fn should_use_gateway_mode_false_for_loopback_tunnel() { - let _guard = crate::config::helpers::ENV_MUTEX - .lock() - .expect("env mutex poisoned"); + let _guard = crate::config::helpers::lock_env(); let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok(); unsafe { std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL"); @@ -7382,9 +7464,7 @@ mod tests { impl EnvGuard { fn new() -> Self { - let guard = crate::config::helpers::ENV_MUTEX - .lock() - .expect("env mutex poisoned"); + let guard = crate::config::helpers::lock_env(); let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok(); // SAFETY: Under ENV_MUTEX, no concurrent env access. unsafe { @@ -7442,9 +7522,7 @@ mod tests { #[test] fn gateway_callback_redirect_uri_does_not_duplicate_callback_path_from_env() { - let _guard = crate::config::helpers::ENV_MUTEX - .lock() - .expect("env mutex poisoned"); + let _guard = crate::config::helpers::lock_env(); let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok(); unsafe { std::env::set_var( @@ -7470,9 +7548,7 @@ mod tests { #[test] fn gateway_callback_redirect_uri_trims_trailing_slash_from_env_callback() { - let _guard = crate::config::helpers::ENV_MUTEX - .lock() - .expect("env mutex poisoned"); + let _guard = crate::config::helpers::lock_env(); let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok(); unsafe { std::env::set_var( @@ -7580,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. @@ -7620,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 @@ -7663,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 @@ -7796,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/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/llm/mod.rs b/src/llm/mod.rs index 64ecd519..308b3983 100644 --- a/src/llm/mod.rs +++ b/src/llm/mod.rs @@ -59,7 +59,7 @@ pub use openai_codex_session::{OpenAiCodexSession, OpenAiCodexSessionManager}; pub use provider::{ ChatMessage, CompletionRequest, CompletionResponse, ContentPart, FinishReason, ImageUrl, LlmProvider, ModelMetadata, Role, ToolCall, ToolCompletionRequest, ToolCompletionResponse, - ToolDefinition, ToolResult, + ToolDefinition, ToolResult, generate_tool_call_id, }; pub use reasoning::{ ActionPlan, Reasoning, ReasoningContext, RespondOutput, RespondResult, SILENT_REPLY_TOKEN, diff --git a/src/llm/oauth_helpers.rs b/src/llm/oauth_helpers.rs index 2881e60e..daaf1b42 100644 --- a/src/llm/oauth_helpers.rs +++ b/src/llm/oauth_helpers.rs @@ -361,7 +361,7 @@ pub fn landing_html(provider_name: &str, success: bool) -> String { #[cfg(test)] mod tests { use super::*; - use crate::config::helpers::ENV_MUTEX; + use crate::config::helpers::lock_env; #[test] fn loopback_detection() { @@ -390,7 +390,7 @@ mod tests { #[allow(clippy::await_holding_lock)] #[tokio::test] async fn bind_rejects_wildcard_ipv4() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); let original = std::env::var("OAUTH_CALLBACK_HOST").ok(); // SAFETY: Under ENV_MUTEX, no concurrent env access. unsafe { std::env::set_var("OAUTH_CALLBACK_HOST", "0.0.0.0") }; @@ -414,7 +414,7 @@ mod tests { #[allow(clippy::await_holding_lock)] #[tokio::test] async fn bind_rejects_wildcard_ipv6() { - let _guard = ENV_MUTEX.lock().expect("env mutex poisoned"); + let _guard = lock_env(); let original = std::env::var("OAUTH_CALLBACK_HOST").ok(); // SAFETY: Under ENV_MUTEX, no concurrent env access. unsafe { std::env::set_var("OAUTH_CALLBACK_HOST", "::") }; diff --git a/src/llm/provider.rs b/src/llm/provider.rs index 8a213031..bb45ec68 100644 --- a/src/llm/provider.rs +++ b/src/llm/provider.rs @@ -233,6 +233,32 @@ pub struct ToolCall { pub arguments: serde_json::Value, } +/// Generate a tool-call ID that satisfies all providers. +/// +/// Mistral requires exactly 9 alphanumeric characters (`[a-zA-Z0-9]{9}`). +/// Other providers accept any non-empty string. By default we produce a +/// 9-char base-62 string derived from two seed values so the ID is both +/// deterministic (for replayed history) and provider-compatible. +pub fn generate_tool_call_id(seed_a: usize, seed_b: usize) -> String { + // Mix the two seeds into a single u64 using a simple hash-like combine. + let combined = (seed_a as u64) + .wrapping_mul(6364136223846793005) + .wrapping_add(seed_b as u64); + // Format as 9-char zero-padded base-62 (0-9, a-z, A-Z). + let mut buf = [b'0'; 9]; + let mut val = combined; + for b in buf.iter_mut().rev() { + let digit = (val % 62) as u8; + *b = match digit { + 0..=9 => b'0' + digit, + 10..=35 => b'a' + (digit - 10), + _ => b'A' + (digit - 36), + }; + val /= 62; + } + buf.iter().map(|&b| b as char).collect::() +} + /// Result of a tool execution to send back to the LLM. #[derive(Debug, Clone)] pub struct ToolResult { @@ -533,6 +559,77 @@ pub fn strip_unsupported_tool_params( #[cfg(test)] mod tests { use super::*; + use std::collections::HashSet; + + #[test] + fn generate_tool_call_id_has_valid_format() { + let samples = [ + (0usize, 0usize), + (1usize, 2usize), + (42usize, 999usize), + (usize::MAX, usize::MAX), + ]; + + for (a, b) in samples { + let id = generate_tool_call_id(a, b); + assert_eq!( + id.len(), + 9, + "tool-call ID must be exactly 9 characters for seeds ({a}, {b})" + ); + assert!( + id.chars().all(|c| c.is_ascii_alphanumeric()), + "tool-call ID must be ASCII alphanumeric for seeds ({a}, {b}), got: {id}" + ); + } + } + + #[test] + fn generate_tool_call_id_is_deterministic_for_same_seeds() { + let pairs = [ + (0usize, 0usize), + (1usize, 2usize), + (123usize, 456usize), + (usize::MAX, 0usize), + ]; + + for (a, b) in pairs { + let id1 = generate_tool_call_id(a, b); + let id2 = generate_tool_call_id(a, b); + let id3 = generate_tool_call_id(a, b); + assert_eq!( + id1, id2, + "tool-call ID must be deterministic for seeds ({a}, {b})" + ); + assert_eq!( + id2, id3, + "tool-call ID must be deterministic across multiple calls for seeds ({a}, {b})" + ); + } + } + + #[test] + fn generate_tool_call_id_differs_for_different_seeds_in_small_sample() { + let seed_pairs = [ + (0usize, 1usize), + (1usize, 0usize), + (1usize, 2usize), + (2usize, 3usize), + (10usize, 20usize), + (100usize, 200usize), + ]; + + let mut ids = HashSet::new(); + for (a, b) in seed_pairs { + let id = generate_tool_call_id(a, b); + let inserted = ids.insert(id.clone()); + assert!( + inserted, + "expected distinct tool-call IDs for different seeds, \ + but duplicate ID '{id}' found for seeds ({a}, {b})" + ); + } + } #[test] fn test_sanitize_preserves_valid_pairs() { diff --git a/src/llm/reasoning.rs b/src/llm/reasoning.rs index b00948ae..cbec297b 100644 --- a/src/llm/reasoning.rs +++ b/src/llm/reasoning.rs @@ -23,6 +23,13 @@ You said you would perform an action, but you did not include any tool calls.\n\ Do NOT describe what you intend to do — actually call the tool now.\n\ Use the tool_calls mechanism to invoke the appropriate tool."; +/// Seed value used as the second argument to `generate_tool_call_id` when +/// recovering tool calls from malformed LLM text responses. This must differ +/// from the `0` seed used in `rig_adapter::normalized_tool_call_id` to avoid +/// ID collisions between provider-generated and text-recovered tool calls at +/// the same positional index. +const RECOVERED_TOOL_CALL_SEED: usize = 99; + /// Detect when an LLM response expresses intent to call a tool without /// actually issuing tool calls. Returns `true` if the text contains phrases /// like "Let me search …" or "I'll fetch …" outside of fenced/indented code blocks. @@ -1337,7 +1344,10 @@ fn recover_tool_calls_from_content( .cloned() .unwrap_or(serde_json::Value::Object(Default::default())); calls.push(ToolCall { - id: format!("recovered_{}", calls.len()), + id: super::provider::generate_tool_call_id( + calls.len(), + RECOVERED_TOOL_CALL_SEED, + ), name: name.to_string(), arguments, }); @@ -1348,7 +1358,10 @@ fn recover_tool_calls_from_content( let name = inner.trim(); if tool_names.contains(name) { calls.push(ToolCall { - id: format!("recovered_{}", calls.len()), + id: super::provider::generate_tool_call_id( + calls.len(), + RECOVERED_TOOL_CALL_SEED, + ), name: name.to_string(), arguments: serde_json::Value::Object(Default::default()), }); @@ -1382,7 +1395,10 @@ fn recover_tool_calls_from_content( let arguments = serde_json::from_str::(args_str) .unwrap_or(serde_json::Value::Object(Default::default())); calls.push(ToolCall { - id: format!("recovered_{}", calls.len()), + id: super::provider::generate_tool_call_id( + calls.len(), + RECOVERED_TOOL_CALL_SEED, + ), name: name.to_string(), arguments, }); @@ -1393,7 +1409,7 @@ fn recover_tool_calls_from_content( // No arguments or malformed — call with empty args calls.push(ToolCall { - id: format!("recovered_{}", calls.len()), + id: super::provider::generate_tool_call_id(calls.len(), RECOVERED_TOOL_CALL_SEED), name: name.to_string(), arguments: serde_json::Value::Object(Default::default()), }); diff --git a/src/llm/rig_adapter.rs b/src/llm/rig_adapter.rs index 1741e860..a9030929 100644 --- a/src/llm/rig_adapter.rs +++ b/src/llm/rig_adapter.rs @@ -20,6 +20,7 @@ use rust_decimal_macros::dec; use serde::Serialize; use serde::de::DeserializeOwned; use serde_json::Value as JsonValue; +use sha2::{Digest, Sha256}; use std::collections::HashSet; @@ -400,11 +401,48 @@ fn convert_messages(messages: &[ChatMessage]) -> (Option, Vec, seed: usize) -> String { - match raw.map(str::trim).filter(|id| !id.is_empty()) { - Some(id) => id.to_string(), - None => format!("generated_tool_call_{seed}"), + // Trim and treat empty as None. + let trimmed = raw.and_then(|s| { + let t = s.trim(); + if t.is_empty() { None } else { Some(t) } + }); + + if let Some(id) = trimmed { + // If the ID already satisfies `[a-zA-Z0-9]{9}`, pass it through unchanged. + if id.len() == 9 && id.chars().all(|c| c.is_ascii_alphanumeric()) { + return id.to_string(); + } + + // Otherwise, deterministically hash the raw ID and feed the hash-derived + // seed into the provider-level generator so that the encoding and any + // provider-specific constraints remain centralized in one place. + let digest = Sha256::digest(id.as_bytes()); + // Derive a 64-bit value from the first 8 bytes of the digest, then + // split it into two usize seeds so we preserve all 64 bits of entropy + // even on 32-bit targets. + let hash64 = { + // SHA-256 always produces 32 bytes, so indexing the first 8 is safe. + let bytes: [u8; 8] = [ + digest[0], digest[1], digest[2], digest[3], digest[4], digest[5], digest[6], + digest[7], + ]; + u64::from_be_bytes(bytes) + }; + let hi_seed: usize = (hash64 >> 32) as usize; + let lo_seed: usize = (hash64 & 0xFFFF_FFFF) as usize; + return super::provider::generate_tool_call_id(hi_seed, lo_seed); } + + // Fallback for missing/empty raw IDs: use the provider-level generator, + // which already produces compliant IDs. + super::provider::generate_tool_call_id(seed, 0) } /// Convert IronClaw tool definitions to rig-core format. @@ -813,8 +851,9 @@ mod tests { #[test] fn test_convert_messages_tool_result() { + // Use a conforming 9-char alphanumeric ID so it passes through unchanged. let messages = vec![ChatMessage::tool_result( - "call_123", + "abcDE1234", "search", "result text", )]; @@ -825,8 +864,8 @@ mod tests { match &history[0] { RigMessage::User { content } => match content.first() { UserContent::ToolResult(r) => { - assert_eq!(r.id, "call_123"); - assert_eq!(r.call_id.as_deref(), Some("call_123")); + assert_eq!(r.id, "abcDE1234"); + assert_eq!(r.call_id.as_deref(), Some("abcDE1234")); } other => panic!("Expected tool result content, got: {:?}", other), }, @@ -836,8 +875,9 @@ mod tests { #[test] fn test_convert_messages_assistant_with_tool_calls() { + // Use a conforming 9-char alphanumeric ID so it passes through unchanged. let tc = IronToolCall { - id: "call_1".to_string(), + id: "Xt7mK9pQ2".to_string(), name: "search".to_string(), arguments: serde_json::json!({"query": "test"}), }; @@ -851,7 +891,7 @@ mod tests { assert!(content.iter().count() >= 2); for item in content.iter() { if let AssistantContent::ToolCall(tc) = item { - assert_eq!(tc.call_id.as_deref(), Some("call_1")); + assert_eq!(tc.call_id.as_deref(), Some("Xt7mK9pQ2")); } } } @@ -873,7 +913,14 @@ mod tests { match &history[0] { RigMessage::User { content } => match content.first() { UserContent::ToolResult(r) => { - assert!(r.id.starts_with("generated_tool_call_")); + // Missing ID → normalized_tool_call_id generates a 9-char alphanumeric ID. + assert_eq!( + r.id.len(), + 9, + "fallback ID should be 9 chars, got: {}", + r.id + ); + assert!(r.id.chars().all(|c| c.is_ascii_alphanumeric())); assert_eq!(r.call_id.as_deref(), Some(r.id.as_str())); } other => panic!("Expected tool result content, got: {:?}", other), @@ -961,12 +1008,14 @@ mod tests { _ => None, }); let tc = tool_call.expect("should have a tool call"); - assert!(!tc.id.is_empty(), "tool call id must not be empty"); - assert!( - tc.id.starts_with("generated_tool_call_"), - "empty id should be replaced with generated id, got: {}", + // Empty ID → normalized_tool_call_id generates a 9-char alphanumeric ID. + assert_eq!( + tc.id.len(), + 9, + "generated id should be 9 chars, got: {}", tc.id ); + assert!(tc.id.chars().all(|c| c.is_ascii_alphanumeric())); assert_eq!(tc.call_id.as_deref(), Some(tc.id.as_str())); } other => panic!("Expected Assistant message, got: {:?}", other), @@ -990,11 +1039,14 @@ mod tests { _ => None, }); let tc = tool_call.expect("should have a tool call"); - assert!( - tc.id.starts_with("generated_tool_call_"), - "whitespace-only id should be replaced, got: {:?}", + // Whitespace-only ID → normalized_tool_call_id generates a 9-char alphanumeric ID. + assert_eq!( + tc.id.len(), + 9, + "generated id should be 9 chars, got: {}", tc.id ); + assert!(tc.id.chars().all(|c| c.is_ascii_alphanumeric())); } other => panic!("Expected Assistant message, got: {:?}", other), } @@ -1381,4 +1433,67 @@ mod tests { // Should be 2 separate User messages (text user + tool result user) assert_eq!(history.len(), 2); } + + // -- normalized_tool_call_id tests -- + + #[test] + fn test_normalized_tool_call_id_conforming_passthrough() { + // A 9-char alphanumeric ID should pass through unchanged. + let id = normalized_tool_call_id(Some("abcDE1234"), 42); + assert_eq!(id, "abcDE1234"); + } + + #[test] + fn test_normalized_tool_call_id_non_conforming_hashed() { + // An ID that doesn't match [a-zA-Z0-9]{9} should be hashed into one. + let id = normalized_tool_call_id(Some("call_abc_long_id"), 0); + assert_eq!(id.len(), 9); + assert!(id.chars().all(|c| c.is_ascii_alphanumeric())); + // Should NOT be the raw input. + assert_ne!(id, "call_abc_l"); + } + + #[test] + fn test_normalized_tool_call_id_empty_input() { + let id = normalized_tool_call_id(Some(""), 5); + assert_eq!(id.len(), 9); + assert!(id.chars().all(|c| c.is_ascii_alphanumeric())); + } + + #[test] + fn test_normalized_tool_call_id_whitespace_input() { + let id = normalized_tool_call_id(Some(" "), 5); + assert_eq!(id.len(), 9); + assert!(id.chars().all(|c| c.is_ascii_alphanumeric())); + // Empty and whitespace-only with the same seed should produce identical results. + let id_empty = normalized_tool_call_id(Some(""), 5); + assert_eq!(id, id_empty); + } + + #[test] + fn test_normalized_tool_call_id_none_input() { + let id = normalized_tool_call_id(None, 7); + assert_eq!(id.len(), 9); + assert!(id.chars().all(|c| c.is_ascii_alphanumeric())); + // None and empty string with same seed should produce identical results. + let id_empty = normalized_tool_call_id(Some(""), 7); + assert_eq!(id, id_empty); + } + + #[test] + fn test_normalized_tool_call_id_deterministic() { + let id1 = normalized_tool_call_id(Some("call_xyz_123"), 0); + let id2 = normalized_tool_call_id(Some("call_xyz_123"), 0); + assert_eq!(id1, id2, "same input must produce same output"); + } + + #[test] + fn test_normalized_tool_call_id_different_inputs_differ() { + let id_a = normalized_tool_call_id(Some("call_aaa"), 0); + let id_b = normalized_tool_call_id(Some("call_bbb"), 0); + assert_ne!( + id_a, id_b, + "different raw IDs should produce different hashed IDs" + ); + } } diff --git a/src/main.rs b/src/main.rs index 23224d0f..eab01264 100644 --- a/src/main.rs +++ b/src/main.rs @@ -142,6 +142,11 @@ async fn async_main() -> anyhow::Result<()> { init_cli_tracing(); return ironclaw::cli::run_logs_command(logs_cmd.clone(), cli.config.as_deref()).await; } + Some(Command::Models(models_cmd)) => { + init_cli_tracing(); + return ironclaw::cli::run_models_command(models_cmd.clone(), cli.config.as_deref()) + .await; + } Some(Command::Doctor) => { init_cli_tracing(); return ironclaw::cli::run_doctor_command().await; @@ -584,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)); @@ -643,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); + } } }); } @@ -686,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; } @@ -749,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() @@ -769,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, @@ -799,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 @@ -844,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( @@ -862,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/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 b72f90ee..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 { @@ -164,19 +165,15 @@ pub async fn setup_orchestrator( #[cfg(test)] mod tests { - use std::sync::Mutex; - use super::*; - - /// Serialize access to `ORCHESTRATOR_PORT` env var across test threads. - static ENV_LOCK: Mutex<()> = Mutex::new(()); + use crate::config::helpers::lock_env; #[test] fn resolve_orchestrator_port_from_env() { - let _guard = ENV_LOCK.lock().unwrap(); + let _guard = lock_env(); // Safety: env-var mutation requires unsafe in edition 2024; - // ENV_LOCK serializes concurrent access from other test threads. + // lock_env() serializes concurrent access from other test threads. // Absent env var → default 50051 unsafe { std::env::remove_var("ORCHESTRATOR_PORT") }; 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/setup/wizard.rs b/src/setup/wizard.rs index b7669070..7ad86610 100644 --- a/src/setup/wizard.rs +++ b/src/setup/wizard.rs @@ -3736,7 +3736,7 @@ mod tests { use tempfile::tempdir; use super::*; - use crate::config::helpers::ENV_MUTEX; + use crate::config::helpers::lock_env; #[test] fn test_wizard_creation() { @@ -3760,7 +3760,7 @@ mod tests { #[test] fn test_wizard_owner_id_uses_resolved_env_scope() { - let _guard = ENV_MUTEX.lock().unwrap_or_else(|e| e.into_inner()); + let _guard = lock_env(); let _owner = EnvGuard::set("IRONCLAW_OWNER_ID", " wizard-owner "); let wizard = SetupWizard::new(); @@ -3769,7 +3769,7 @@ mod tests { #[test] fn test_wizard_owner_id_uses_toml_scope() { - let _guard = ENV_MUTEX.lock().unwrap_or_else(|e| e.into_inner()); + let _guard = lock_env(); let _owner = EnvGuard::clear("IRONCLAW_OWNER_ID"); let dir = tempdir().unwrap(); // safety: test-only tempdir setup let path = dir.path().join("config.toml"); @@ -3785,7 +3785,7 @@ mod tests { fn test_try_with_config_and_toml_propagates_invalid_owner_env() { use std::os::unix::ffi::OsStringExt; - let _guard = ENV_MUTEX.lock().unwrap_or_else(|e| e.into_inner()); + let _guard = lock_env(); let original = std::env::var_os("IRONCLAW_OWNER_ID"); unsafe { std::env::set_var("IRONCLAW_OWNER_ID", OsString::from_vec(vec![0x66, 0x80])); @@ -4245,7 +4245,7 @@ mod tests { fn test_build_nearai_model_fetch_config_picks_up_api_key_env() { use secrecy::ExposeSecret; - let _lock = ENV_MUTEX.lock().unwrap(); + let _lock = lock_env(); let _guard = EnvGuard::set("NEARAI_API_KEY", "test-cloud-api-key-12345"); let _guard2 = EnvGuard::clear("NEARAI_BASE_URL"); @@ -4269,7 +4269,7 @@ mod tests { /// the config should have `api_key: None` (session token path). #[test] fn test_build_nearai_model_fetch_config_none_when_no_api_key() { - let _lock = ENV_MUTEX.lock().unwrap(); + let _lock = lock_env(); let _guard = EnvGuard::clear("NEARAI_API_KEY"); let _guard2 = EnvGuard::clear("NEARAI_BASE_URL"); @@ -4288,7 +4288,7 @@ mod tests { /// Regression test for #799: empty NEARAI_API_KEY should be treated as absent. #[test] fn test_build_nearai_model_fetch_config_none_when_empty_api_key() { - let _lock = ENV_MUTEX.lock().unwrap(); + let _lock = lock_env(); let _guard = EnvGuard::set("NEARAI_API_KEY", ""); let config = build_nearai_model_fetch_config(); @@ -4306,7 +4306,7 @@ mod tests { fn test_model_discovery_picks_up_injected_var() { use secrecy::ExposeSecret; - let _lock = ENV_MUTEX.lock().unwrap(); + let _lock = lock_env(); let _guard = EnvGuard::clear("NEARAI_API_KEY"); let _guard2 = EnvGuard::clear("NEARAI_BASE_URL"); @@ -4337,7 +4337,7 @@ mod tests { /// the NEAR AI authentication menu. #[test] fn test_build_nearai_model_fetch_config_picks_up_runtime_env() { - let _lock = ENV_MUTEX.lock().unwrap(); + let _lock = lock_env(); // Ensure the real env var is unset so the only source is the overlay. let _guard = EnvGuard::clear("NEARAI_API_KEY"); diff --git a/src/testing/mod.rs b/src/testing/mod.rs index 953cbfcd..e580b169 100644 --- a/src/testing/mod.rs +++ b/src/testing/mod.rs @@ -28,7 +28,7 @@ use std::sync::atomic::{AtomicBool, AtomicU32, Ordering}; use async_trait::async_trait; use rust_decimal::Decimal; -use tokio::sync::mpsc; +use tokio::sync::{Mutex as AsyncMutex, mpsc}; use crate::agent::AgentDeps; use crate::channels::{ @@ -361,6 +361,75 @@ impl Channel for StubChannel { } } +/// Captured broadcast deliveries keyed by the target user or chat identifier. +pub type BroadcastCapture = Arc>>; + +/// A lightweight channel double that only records `broadcast()` traffic. +/// +/// This is useful for unit tests that need to assert message routing without +/// spinning up a full interactive channel harness. +pub struct RecordingBroadcastChannel { + name: &'static str, + captures: BroadcastCapture, +} + +impl RecordingBroadcastChannel { + pub fn new(name: &'static str) -> (Self, BroadcastCapture) { + let captures = Arc::new(AsyncMutex::new(Vec::new())); + ( + Self { + name, + captures: Arc::clone(&captures), + }, + captures, + ) + } +} + +#[async_trait] +impl Channel for RecordingBroadcastChannel { + fn name(&self) -> &str { + self.name + } + + async fn start(&self) -> Result { + let (_tx, rx) = mpsc::channel::(1); + Ok(Box::pin(tokio_stream::wrappers::ReceiverStream::new(rx))) + } + + async fn respond( + &self, + _msg: &IncomingMessage, + _response: OutgoingResponse, + ) -> Result<(), ChannelError> { + Ok(()) + } + + async fn send_status( + &self, + _status: StatusUpdate, + _metadata: &serde_json::Value, + ) -> Result<(), ChannelError> { + Ok(()) + } + + async fn broadcast( + &self, + user_id: &str, + response: OutgoingResponse, + ) -> Result<(), ChannelError> { + self.captures + .lock() + .await + .push((user_id.to_string(), response)); + Ok(()) + } + + async fn health_check(&self) -> Result<(), ChannelError> { + Ok(()) + } +} + /// Assembled test components. pub struct TestHarness { /// The agent dependencies, ready for use. @@ -494,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/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 1c27b539..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", @@ -271,12 +316,13 @@ impl Tool for MemoryWriteTool { .and_then(|v| v.as_bool()) .unwrap_or(false); + // Parse timezone once for targets that need it (daily_log). + let tz = crate::timezone::parse_timezone(&ctx.user_timezone).unwrap_or(chrono_tz::Tz::UTC); + // Resolve the target to a workspace path let resolved_path = match target { "memory" => paths::MEMORY.to_string(), "daily_log" => { - let tz = crate::timezone::parse_timezone(&ctx.user_timezone) - .unwrap_or(chrono_tz::Tz::UTC); let now = chrono::Utc::now().with_timezone(&tz); format!("daily/{}.md", now.format("%Y-%m-%d")) } @@ -288,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)? @@ -306,12 +352,12 @@ 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)?; @@ -320,19 +366,19 @@ impl Tool for MemoryWriteTool { "daily_log" => { let tz = crate::timezone::parse_timezone(&ctx.user_timezone) .unwrap_or(chrono_tz::Tz::UTC); - self.workspace + 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)?; @@ -362,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 @@ -417,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)), + } } } @@ -457,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(); @@ -471,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)))?; @@ -496,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, @@ -518,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)))?; @@ -534,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 { @@ -585,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(); @@ -597,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( @@ -651,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()); @@ -669,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"); @@ -682,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"); @@ -699,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"); @@ -712,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!({ @@ -734,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/message.rs b/src/tools/builtin/message.rs index 83041b80..08029d6f 100644 --- a/src/tools/builtin/message.rs +++ b/src/tools/builtin/message.rs @@ -80,6 +80,12 @@ fn metadata_notify_user(metadata: &serde_json::Value) -> Option { metadata_string(metadata, "notify_user").filter(|value| value != "default") } +// Autonomous runs include `owner_id` when the job is executing on behalf of a +// durable owner scope instead of an interactive channel actor. +fn metadata_owner_id(metadata: &serde_json::Value) -> Option { + metadata_string(metadata, "owner_id") +} + fn channel_matches_source(resolved_channel: Option<&str>, source_channel: Option<&str>) -> bool { match (resolved_channel, source_channel) { (None, _) => true, @@ -91,11 +97,13 @@ fn channel_matches_source(resolved_channel: Option<&str>, source_channel: Option async fn resolve_channel_fallback_target( extension_manager: Option<&Arc>, channel: Option<&str>, + owner_scope_target: Option<&str>, ctx_user_id: &str, ) -> Option { - let channel_name = channel?; - - if let Some(extension_manager) = extension_manager + // Prefer an explicit channel binding when the extension manager knows the + // durable delivery target (for example, a bound Telegram chat ID). + if let Some(channel_name) = channel + && let Some(extension_manager) = extension_manager && let Some(target) = extension_manager .notification_target_for_channel(channel_name) .await @@ -103,13 +111,19 @@ async fn resolve_channel_fallback_target( return Some(target); } - Some(ctx_user_id.to_string()) + // `owner_id` is only present for autonomous owner-scoped executions. + // Interactive chat turns intentionally fall back to `ctx.user_id`, which is + // already the active conversation target for the current channel. + owner_scope_target + .map(ToOwned::to_owned) + .or_else(|| Some(ctx_user_id.to_string())) } struct MessageTargetResolution<'a> { extension_manager: Option<&'a Arc>, explicit_target: Option, metadata_target: Option, + owner_scope_target: Option, default_target: Option, channel: Option<&'a str>, metadata_channel: Option<&'a str>, @@ -133,6 +147,7 @@ async fn resolve_message_target(inputs: MessageTargetResolution<'_>) -> Option) -> Option>>; - - struct RecordingChannel { - name: &'static str, - captures: BroadcastCapture, - } - - impl RecordingChannel { - fn new(name: &'static str) -> (Self, BroadcastCapture) { - let captures = Arc::new(Mutex::new(Vec::new())); - ( - Self { - name, - captures: Arc::clone(&captures), - }, - captures, - ) - } - } - - #[async_trait] - impl Channel for RecordingChannel { - fn name(&self) -> &str { - self.name - } - - async fn start(&self) -> Result { - let (_tx, rx) = mpsc::channel::(1); - Ok(Box::pin(tokio_stream::wrappers::ReceiverStream::new(rx))) - } - - async fn respond( - &self, - _msg: &IncomingMessage, - _response: OutgoingResponse, - ) -> Result<(), ChannelError> { - Ok(()) - } - - async fn send_status( - &self, - _status: StatusUpdate, - _metadata: &serde_json::Value, - ) -> Result<(), ChannelError> { - Ok(()) - } - - async fn broadcast( - &self, - user_id: &str, - response: OutgoingResponse, - ) -> Result<(), ChannelError> { - self.captures - .lock() - .await - .push((user_id.to_string(), response)); - Ok(()) - } - - async fn health_check(&self) -> Result<(), ChannelError> { - Ok(()) - } - } + use crate::testing::{BroadcastCapture, RecordingBroadcastChannel}; async fn message_tool_with_recording_channels() -> (MessageTool, BroadcastCapture, BroadcastCapture) { let channel_manager = ChannelManager::new(); - let (gateway, gateway_captures) = RecordingChannel::new("gateway"); - let (telegram, telegram_captures) = RecordingChannel::new("telegram"); + let (gateway, gateway_captures) = RecordingBroadcastChannel::new("gateway"); + let (telegram, telegram_captures) = RecordingBroadcastChannel::new("telegram"); channel_manager.add(Box::new(gateway)).await; channel_manager.add(Box::new(telegram)).await; @@ -870,28 +820,63 @@ mod tests { } #[tokio::test] - async fn message_tool_falls_back_to_ctx_user_when_channel_known() { - // Regression for owner-scoped notifications: a channel can be known - // even when the concrete delivery target is omitted, so the message - // tool should pass ctx.user_id through to the channel layer. - let tool = MessageTool::new(Arc::new(ChannelManager::new())); + async fn message_tool_falls_back_to_owner_scope_when_channel_known() { + let (tool, gateway_captures, telegram_captures) = + message_tool_with_recording_channels().await; let mut ctx = - crate::context::JobContext::with_user("owner-scope", "routine-job", "price alert"); + crate::context::JobContext::with_user("telegram", "routine-job", "price alert"); + ctx.metadata = serde_json::json!({ + "notify_channel": "telegram", + "owner_id": "owner-scope", + }); + + let result = tool + .execute(serde_json::json!({"content": "NEAR price is $5"}), &ctx) + .await + .expect("message tool should use owner scope before ctx.user_id"); + + assert_eq!( + result.result.as_str(), + Some("Sent message to telegram:owner-scope") + ); + assert!(gateway_captures.lock().await.is_empty()); + let telegram = telegram_captures.lock().await.clone(); + assert_eq!(telegram.len(), 1); + assert_eq!(telegram[0].0, "owner-scope"); + assert_eq!(telegram[0].1.content, "NEAR price is $5"); + } + + #[tokio::test] + async fn message_tool_falls_back_to_ctx_user_when_owner_scope_absent() { + let (tool, gateway_captures, telegram_captures) = + message_tool_with_recording_channels().await; + + let mut ctx = crate::context::JobContext::with_user( + "interactive-chat-user", + "routine-job", + "price alert", + ); ctx.metadata = serde_json::json!({ "notify_channel": "telegram", }); let result = tool .execute(serde_json::json!({"content": "NEAR price is $5"}), &ctx) - .await; + .await + .expect( + "message tool should fall back to ctx.user_id when owner scope metadata is absent", + ); - assert!(result.is_err()); // safety: test-only assertion - let err = result.unwrap_err().to_string(); - let mentions_missing_target = err.contains("No target specified"); - assert!(!mentions_missing_target); // safety: test-only assertion - let mentions_missing_channel = err.contains("No channel specified"); - assert!(!mentions_missing_channel); // safety: test-only assertion + assert_eq!( + result.result.as_str(), + Some("Sent message to telegram:interactive-chat-user") + ); + assert!(gateway_captures.lock().await.is_empty()); + let telegram = telegram_captures.lock().await.clone(); + assert_eq!(telegram.len(), 1); + assert_eq!(telegram[0].0, "interactive-chat-user"); + assert_eq!(telegram[0].1.content, "NEAR price is $5"); } #[tokio::test] 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/builtin/routine.rs b/src/tools/builtin/routine.rs index c197fe25..f4313483 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(), @@ -605,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 { @@ -852,11 +856,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 +896,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); @@ -1007,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 { @@ -1863,6 +1872,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 +2260,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/src/tools/execute.rs b/src/tools/execute.rs index 86da157b..69c72e46 100644 --- a/src/tools/execute.rs +++ b/src/tools/execute.rs @@ -19,7 +19,7 @@ pub async fn execute_tool_with_safety( tools: &ToolRegistry, safety: &SafetyLayer, tool_name: &str, - params: &serde_json::Value, + params: serde_json::Value, job_ctx: &JobContext, ) -> Result { if tool_name.is_empty() { @@ -35,7 +35,7 @@ pub async fn execute_tool_with_safety( name: tool_name.to_string(), })?; - let normalized_params = prepare_tool_params(tool.as_ref(), params); + let normalized_params = prepare_tool_params(tool.as_ref(), ¶ms); // Validate tool parameters let validation = safety.validator().validate_tool_params(&normalized_params); @@ -63,10 +63,7 @@ pub async fn execute_tool_with_safety( // Execute with per-tool timeout let timeout = tool.execution_timeout(); let start = std::time::Instant::now(); - let result = tokio::time::timeout(timeout, async { - tool.execute(normalized_params.clone(), job_ctx).await - }) - .await; + let result = tokio::time::timeout(timeout, tool.execute(normalized_params, job_ctx)).await; let elapsed = start.elapsed(); match &result { @@ -149,7 +146,7 @@ pub async fn execute_tool_simple( tools: &ToolRegistry, safety: &SafetyLayer, tool_name: &str, - params: &serde_json::Value, + params: serde_json::Value, job_ctx: &JobContext, ) -> Result { execute_tool_with_safety(tools, safety, tool_name, params, job_ctx) @@ -308,7 +305,7 @@ mod tests { ®istry, &safety, "", - &serde_json::json!({}), + serde_json::json!({}), &test_job_ctx(), ) .await; @@ -331,7 +328,7 @@ mod tests { let params = serde_json::json!({"message": "hello"}); let result = - execute_tool_with_safety(®istry, &safety, "echo", ¶ms, &test_job_ctx()).await; + execute_tool_with_safety(®istry, &safety, "echo", params, &test_job_ctx()).await; assert!(result.is_ok(), "Echo tool should succeed"); let output = result.unwrap(); @@ -350,7 +347,7 @@ mod tests { ®istry, &safety, "nonexistent", - &serde_json::json!({}), + serde_json::json!({}), &test_job_ctx(), ) .await; @@ -373,7 +370,7 @@ mod tests { ®istry, &safety, "fail_tool", - &serde_json::json!({}), + serde_json::json!({}), &test_job_ctx(), ) .await; @@ -397,7 +394,7 @@ mod tests { ®istry, &safety, "slow_tool", - &serde_json::json!({}), + serde_json::json!({}), &test_job_ctx(), ) .await; @@ -425,7 +422,7 @@ mod tests { ®istry, &safety, "array_echo", - &serde_json::json!({"values": "[\"1\", \"2\", 3]"}), + serde_json::json!({"values": "[\"1\", \"2\", 3]"}), &test_job_ctx(), ) .await diff --git a/src/tools/mcp/http_transport.rs b/src/tools/mcp/http_transport.rs index ec7139c9..59873ce4 100644 --- a/src/tools/mcp/http_transport.rs +++ b/src/tools/mcp/http_transport.rs @@ -130,6 +130,16 @@ impl McpTransport for HttpMcpTransport { ))); } + // MCP notifications commonly acknowledge with 202 Accepted and no body. + if response.status() == reqwest::StatusCode::ACCEPTED { + return Ok(McpResponse { + jsonrpc: "2.0".to_string(), + id: request.id, + result: None, + error: None, + }); + } + // Determine response format from Content-Type. let content_type = response .headers() @@ -506,4 +516,55 @@ mod tests { let echoed = response.result.unwrap(); assert_eq!(echoed["authorization"], "Bearer custom-token"); } + + async fn spawn_accepted_server() -> (String, tokio::task::JoinHandle<()>) { + use axum::{Router, routing::post}; + use tokio::net::TcpListener; + + async fn accepted() -> axum::http::StatusCode { + axum::http::StatusCode::ACCEPTED + } + + let app = Router::new().route("/", post(accepted)); + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("Failed to bind to an ephemeral port"); + let addr = listener + .local_addr() + .expect("Failed to get listener's local address"); + let url = format!("http://127.0.0.1:{}", addr.port()); + + let handle = tokio::spawn(async move { + axum::serve(listener, app) + .await + .expect("Test server failed to run"); + }); + + (url, handle) + } + + fn notification_request(method: &str) -> McpRequest { + McpRequest { + jsonrpc: "2.0".to_string(), + id: None, + method: method.to_string(), + params: None, + } + } + + #[tokio::test] + async fn test_accepted_notification_returns_empty_response() { + let (url, _handle) = spawn_accepted_server().await; + let transport = HttpMcpTransport::new(&url, "accepted-test"); + let request = notification_request("notifications/initialized"); + + let response = transport + .send(&request, &HashMap::new()) + .await + .expect("202 notification response"); + assert_eq!(response.jsonrpc, "2.0"); + assert_eq!(response.id, request.id); + assert!(response.result.is_none()); + assert!(response.error.is_none()); + } } 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/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/src/worker/container.rs b/src/worker/container.rs index 920cc2ce..e0933975 100644 --- a/src/worker/container.rs +++ b/src/worker/container.rs @@ -462,9 +462,14 @@ impl LoopDelegate for ContainerDelegate { ..Default::default() }; - let result = - execute_tool_simple(&self.tools, &self.safety, &tc.name, &tc.arguments, &job_ctx) - .await; + let result = execute_tool_simple( + &self.tools, + &self.safety, + &tc.name, + tc.arguments.clone(), + &job_ctx, + ) + .await; self.post_event( "tool_result", diff --git a/src/worker/job.rs b/src/worker/job.rs index 436a23ce..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); } } } @@ -1438,6 +1438,9 @@ impl From for Result { #[cfg(test)] mod tests { + use std::sync::Arc; + + use crate::channels::ChannelManager; use crate::llm::ToolSelection; use super::*; @@ -1448,6 +1451,8 @@ mod tests { ToolCompletionResponse, }; use crate::safety::SafetyLayer; + use crate::testing::{BroadcastCapture, RecordingBroadcastChannel}; + use crate::tools::builtin::MessageTool; use crate::tools::{Tool, ToolError as ToolExecError, ToolOutput}; /// A test tool that sleeps for a configurable duration before returning. @@ -1539,6 +1544,20 @@ mod tests { Worker::new(job_id, deps) } + async fn make_worker_with_message_tool() + -> (Worker, Arc, BroadcastCapture, BroadcastCapture) { + let channel_manager = ChannelManager::new(); + let (gateway, gateway_captures) = RecordingBroadcastChannel::new("gateway"); + let (telegram, telegram_captures) = RecordingBroadcastChannel::new("telegram"); + channel_manager.add(Box::new(gateway)).await; + channel_manager.add(Box::new(telegram)).await; + + let message_tool = Arc::new(MessageTool::new(Arc::new(channel_manager))); + let worker = make_worker(vec![message_tool.clone()]).await; + + (worker, message_tool, gateway_captures, telegram_captures) + } + #[test] fn test_tool_selection_preserves_call_id() { let selection = ToolSelection { @@ -2147,4 +2166,50 @@ mod tests { assert_eq!(ctx.metadata, original); // safety: test } + + #[tokio::test] + async fn autonomous_message_tool_ignores_stale_gateway_context_when_routine_metadata_targets_telegram() + { + let (worker, message_tool, gateway_captures, telegram_captures) = + make_worker_with_message_tool().await; + + message_tool + .set_context( + Some("gateway".to_string()), + Some("stale-gateway-target".to_string()), + ) + .await; + + worker + .context_manager() + .update_context(worker.job_id, |ctx| { + ctx.user_id = "telegram".to_string(); + ctx.metadata = serde_json::json!({ + "notify_channel": "telegram", + "owner_id": "owner-scope", + }); + Ok::<(), String>(()) + }) + .await + .unwrap() // safety: test + .unwrap(); // safety: test + + let result = worker + .execute_tool( + "message", + &serde_json::json!({"content": "hello from routine"}), + ) + .await + .unwrap(); // safety: test + assert!( + result.contains("telegram:owner-scope"), + "expected telegram owner-scope routing, got: {result}" + ); + + assert!(gateway_captures.lock().await.is_empty()); + let telegram = telegram_captures.lock().await.clone(); + assert_eq!(telegram.len(), 1); + assert_eq!(telegram[0].0, "owner-scope"); + assert_eq!(telegram[0].1.content, "hello from routine"); + } } diff --git a/src/workspace/README.md b/src/workspace/README.md index 67b9907f..061a5564 100644 --- a/src/workspace/README.md +++ b/src/workspace/README.md @@ -91,6 +91,27 @@ Default k=60. Results from both methods are combined, with documents appearing i - **PostgreSQL:** `ts_rank_cd` for FTS, pgvector cosine distance for vectors, full RRF - **libSQL:** FTS5 for keyword search + vector search via `libsql_vector_idx` (dimension set dynamically by `ensure_vector_index()` during startup) +## Multi-Scope Reads & Identity Isolation + +When a workspace has additional read scopes (via `with_additional_read_scopes`), read operations can span multiple user scopes — a user with scopes `["alice", "shared"]` can read documents from both. + +**Identity files are exempt from multi-scope reads.** The system prompt reads identity and configuration files from the **primary scope only** (`read_primary()`), never from secondary scopes: + +| File | Read method | Rationale | +|------|------------|-----------| +| AGENTS.md | `read_primary()` | Agent instructions are per-user | +| SOUL.md | `read_primary()` | Core values are per-user | +| USER.md | `read_primary()` | User context is per-user | +| IDENTITY.md | `read_primary()` | Identity is per-user | +| TOOLS.md | `read_primary()` | Tool config is per-user | +| BOOTSTRAP.md | `read_primary()` | Onboarding is per-user | +| MEMORY.md | `read()` | Shared memory is a feature | +| daily/*.md | `read()` | Shared daily logs are a feature | + +**Why:** Without this, a user with read access to another scope could silently inherit that scope's identity if their own copy is missing. The agent would present itself as the wrong user — a correctness and security issue. + +**Design rule:** If you want shared identity across users, seed the same content into each user's scope at setup time. Don't rely on multi-scope fallback for identity files. + ## Heartbeat System Proactive periodic execution (default: 30 minutes): diff --git a/src/workspace/document.rs b/src/workspace/document.rs index 3396b677..b1fa176a 100644 --- a/src/workspace/document.rs +++ b/src/workspace/document.rs @@ -37,6 +37,25 @@ pub mod paths { pub const ASSISTANT_DIRECTIVES: &str = "context/assistant-directives.md"; } +/// Paths treated as identity documents for multi-scope isolation. +/// +/// These files are always read from the primary scope only — never from +/// secondary read scopes. This prevents silent identity inheritance +/// (e.g., user A accidentally presenting as user B). +pub const IDENTITY_PATHS: &[&str] = &[ + paths::IDENTITY, + paths::SOUL, + paths::AGENTS, + paths::USER, + paths::TOOLS, + paths::BOOTSTRAP, +]; + +/// Check if a path is an identity document that must be isolated to primary scope. +pub fn is_identity_path(path: &str) -> bool { + IDENTITY_PATHS.contains(&path) +} + /// A memory document stored in the database. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct MemoryDocument { @@ -101,10 +120,7 @@ impl MemoryDocument { /// Check if this is a well-known identity document. pub fn is_identity_document(&self) -> bool { - matches!( - self.path.as_str(), - paths::IDENTITY | paths::SOUL | paths::AGENTS | paths::USER - ) + is_identity_path(&self.path) } } @@ -128,6 +144,42 @@ impl WorkspaceEntry { } } +/// Merge workspace entries from multiple scopes into a deduplicated, sorted list. +/// +/// When the same path appears in multiple scopes: +/// - Keeps the most recent `updated_at` +/// - If any scope marks it as a directory, the merged entry is a directory +pub fn merge_workspace_entries( + entries: impl IntoIterator, +) -> Vec { + let mut seen = std::collections::HashMap::new(); + for entry in entries { + seen.entry(entry.path.clone()) + .and_modify(|existing: &mut WorkspaceEntry| { + // Keep the most recent updated_at (and its content_preview) + if let (Some(existing_ts), Some(new_ts)) = (&existing.updated_at, &entry.updated_at) + { + if new_ts > existing_ts { + existing.updated_at = Some(*new_ts); + existing.content_preview = entry.content_preview.clone(); + } + } else if existing.updated_at.is_none() { + existing.updated_at = entry.updated_at; + existing.content_preview = entry.content_preview.clone(); + } + // If either is a directory, mark as directory + if entry.is_directory { + existing.is_directory = true; + existing.content_preview = None; + } + }) + .or_insert(entry); + } + let mut result: Vec = seen.into_values().collect(); + result.sort_by(|a, b| a.path.cmp(&b.path)); + result +} + /// A chunk of a memory document for search indexing. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct MemoryChunk { @@ -226,4 +278,115 @@ mod tests { }; assert_eq!(entry.name(), "alpha"); } + + #[test] + fn test_merge_workspace_entries_empty() { + let result = merge_workspace_entries(vec![]); + assert!(result.is_empty()); + } + + #[test] + fn test_merge_workspace_entries_keeps_newer_timestamp_and_preview() { + use chrono::TimeZone; + let old_ts = chrono::Utc.with_ymd_and_hms(2025, 1, 1, 0, 0, 0).unwrap(); + let new_ts = chrono::Utc.with_ymd_and_hms(2025, 6, 1, 0, 0, 0).unwrap(); + + let entries = vec![ + WorkspaceEntry { + path: "notes.md".to_string(), + is_directory: false, + updated_at: Some(old_ts), + content_preview: Some("old".to_string()), + }, + WorkspaceEntry { + path: "notes.md".to_string(), + is_directory: false, + updated_at: Some(new_ts), + content_preview: Some("new".to_string()), + }, + ]; + + let result = merge_workspace_entries(entries); + assert_eq!(result.len(), 1); + assert_eq!(result[0].updated_at, Some(new_ts)); + assert_eq!(result[0].content_preview, Some("new".to_string())); + } + + #[test] + fn test_merge_workspace_entries_directory_wins() { + let entries = vec![ + WorkspaceEntry { + path: "projects".to_string(), + is_directory: false, + updated_at: None, + content_preview: Some("file content".to_string()), + }, + WorkspaceEntry { + path: "projects".to_string(), + is_directory: true, + updated_at: None, + content_preview: None, + }, + ]; + + let result = merge_workspace_entries(entries); + assert_eq!(result.len(), 1); + assert!(result[0].is_directory); + assert!(result[0].content_preview.is_none()); + } + + #[test] + fn test_merge_workspace_entries_fills_missing_timestamp() { + use chrono::TimeZone; + let ts = chrono::Utc.with_ymd_and_hms(2025, 3, 1, 0, 0, 0).unwrap(); + + let entries = vec![ + WorkspaceEntry { + path: "a.md".to_string(), + is_directory: false, + updated_at: None, + content_preview: None, + }, + WorkspaceEntry { + path: "a.md".to_string(), + is_directory: false, + updated_at: Some(ts), + content_preview: None, + }, + ]; + + let result = merge_workspace_entries(entries); + assert_eq!(result.len(), 1); + assert_eq!(result[0].updated_at, Some(ts)); + } + + #[test] + fn test_merge_workspace_entries_sorted_by_path() { + let entries = vec![ + WorkspaceEntry { + path: "z.md".to_string(), + is_directory: false, + updated_at: None, + content_preview: None, + }, + WorkspaceEntry { + path: "a.md".to_string(), + is_directory: false, + updated_at: None, + content_preview: None, + }, + WorkspaceEntry { + path: "m.md".to_string(), + is_directory: false, + updated_at: None, + content_preview: None, + }, + ]; + + let result = merge_workspace_entries(entries); + assert_eq!(result.len(), 3); + assert_eq!(result[0].path, "a.md"); + assert_eq!(result[1].path, "m.md"); + assert_eq!(result[2].path, "z.md"); + } } 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); } diff --git a/src/workspace/mod.rs b/src/workspace/mod.rs index 5aac2500..0242047f 100644 --- a/src/workspace/mod.rs +++ b/src/workspace/mod.rs @@ -52,7 +52,10 @@ mod repository; mod search; pub use chunker::{ChunkConfig, chunk_document}; -pub use document::{MemoryChunk, MemoryDocument, WorkspaceEntry, paths}; +pub use document::{ + IDENTITY_PATHS, MemoryChunk, MemoryDocument, WorkspaceEntry, is_identity_path, + merge_workspace_entries, paths, +}; pub use embedding_cache::{CachedEmbeddingProvider, EmbeddingCacheConfig}; pub use embeddings::{ EmbeddingProvider, MockEmbeddings, NearAiEmbeddings, OllamaEmbeddings, OpenAiEmbeddings, @@ -320,6 +323,48 @@ impl WorkspaceStorage { } } } + + // ==================== Multi-scope read methods ==================== + + async fn hybrid_search_multi( + &self, + user_ids: &[String], + agent_id: Option, + query: &str, + embedding: Option<&[f32]>, + config: &SearchConfig, + ) -> Result, WorkspaceError> { + match self { + #[cfg(feature = "postgres")] + Self::Repo(repo) => { + repo.hybrid_search_multi(user_ids, agent_id, query, embedding, config) + .await + } + Self::Db(db) => { + db.hybrid_search_multi(user_ids, agent_id, query, embedding, config) + .await + } + } + } + + async fn get_document_by_path_multi( + &self, + user_ids: &[String], + agent_id: Option, + path: &str, + ) -> Result { + match self { + #[cfg(feature = "postgres")] + Self::Repo(repo) => { + repo.get_document_by_path_multi(user_ids, agent_id, path) + .await + } + Self::Db(db) => { + db.get_document_by_path_multi(user_ids, agent_id, path) + .await + } + } + } } /// Default template seeded into HEARTBEAT.md on first access. @@ -340,9 +385,20 @@ const BOOTSTRAP_SEED: &str = include_str!("seeds/BOOTSTRAP.md"); /// Each workspace is scoped to a user (and optionally an agent). /// Documents are persisted to the database and indexed for search. /// Supports both PostgreSQL (via Repository) and libSQL (via Database trait). +/// +/// ## Multi-scope reads +/// +/// By default, a workspace reads from and writes to a single `user_id`. +/// With `with_additional_read_scopes`, read operations (search, read, list) +/// can span multiple user scopes while writes remain isolated to the primary +/// `user_id`. This enables cross-tenant read access (e.g., a user reading +/// from both their own workspace and a "shared" workspace). pub struct Workspace { - /// User identifier (from channel). + /// User identifier (from channel). All writes go to this scope. user_id: String, + /// User identifiers for read operations. Includes `user_id` as the first + /// element, plus any additional scopes added via `with_additional_read_scopes`. + read_user_ids: Vec, /// Optional agent ID for multi-agent isolation. agent_id: Option, /// Database storage backend. @@ -371,6 +427,7 @@ impl Workspace { let user_id_str = user_id.into(); let memory_layers = crate::workspace::layer::MemoryLayer::default_for_user(&user_id_str); Self { + read_user_ids: vec![user_id_str.clone()], user_id: user_id_str, agent_id: None, storage: WorkspaceStorage::Repo(Repository::new(pool)), @@ -390,6 +447,7 @@ impl Workspace { let user_id_str = user_id.into(); let memory_layers = crate::workspace::layer::MemoryLayer::default_for_user(&user_id_str); Self { + read_user_ids: vec![user_id_str.clone()], user_id: user_id_str, agent_id: None, storage: WorkspaceStorage::Db(db), @@ -474,6 +532,12 @@ impl Workspace { /// /// Also updates read_user_ids to include all layer scopes. pub fn with_memory_layers(mut self, layers: Vec) -> Self { + // Add layer scopes to read_user_ids (same dedup logic as with_additional_read_scopes) + for layer in &layers { + if !self.read_user_ids.contains(&layer.scope) { + self.read_user_ids.push(layer.scope.clone()); + } + } self.memory_layers = layers; self } @@ -496,11 +560,37 @@ impl Workspace { &self.memory_layers } - /// Get the user ID. + /// Add additional user scopes for read operations. + /// + /// The primary `user_id` is always included. Additional scopes allow + /// read operations (search, read, list) to span multiple tenants while + /// writes remain isolated to the primary scope. + /// + /// Duplicate scopes are ignored. + pub fn with_additional_read_scopes(mut self, scopes: Vec) -> Self { + for scope in scopes { + if !self.read_user_ids.contains(&scope) { + self.read_user_ids.push(scope); + } + } + self + } + + /// Get the user ID (primary scope for writes). pub fn user_id(&self) -> &str { &self.user_id } + /// Get the user IDs used for read operations. + pub fn read_user_ids(&self) -> &[String] { + &self.read_user_ids + } + + /// Whether this workspace has multiple read scopes. + fn is_multi_scope(&self) -> bool { + self.read_user_ids.len() > 1 + } + /// Get the agent ID. pub fn agent_id(&self) -> Option { self.agent_id @@ -518,6 +608,33 @@ impl Workspace { /// println!("{}", doc.content); /// ``` pub async fn read(&self, path: &str) -> Result { + let path = normalize_path(path); + if self.is_multi_scope() && is_identity_path(&path) { + // Identity files must only come from the primary scope. + self.storage + .get_document_by_path(&self.user_id, self.agent_id, &path) + .await + } else if self.is_multi_scope() { + self.storage + .get_document_by_path_multi(&self.read_user_ids, self.agent_id, &path) + .await + } else { + self.storage + .get_document_by_path(&self.user_id, self.agent_id, &path) + .await + } + } + + /// Read a file from the **primary scope only**, ignoring additional read scopes. + /// + /// Use this for identity and configuration files (AGENTS.md, SOUL.md, USER.md, + /// IDENTITY.md, TOOLS.md, BOOTSTRAP.md) where inheriting content from another + /// scope would be a correctness/security issue — the agent must never silently + /// present itself as the wrong user. + /// + /// For memory files that should span scopes (MEMORY.md, daily logs), use + /// [`read`] instead. + pub async fn read_primary(&self, path: &str) -> Result { let path = normalize_path(path); self.storage .get_document_by_path(&self.user_id, self.agent_id, &path) @@ -556,6 +673,9 @@ impl Workspace { /// Uses a single `\n` separator (suitable for log-style entries). /// For semantic separation (e.g., memory entries), use `append_memory()` /// which uses `\n\n`. + /// + /// Uses a read-modify-write pattern that is not concurrency-safe: + /// concurrent appends to the same path may lose writes. pub async fn append(&self, path: &str, content: &str) -> Result<(), WorkspaceError> { let path = normalize_path(path); // Scan system-prompt-injected files for prompt injection. @@ -676,6 +796,20 @@ impl Workspace { } /// Write to a layer, with append semantics. + /// + /// Note: privacy classification only examines the new `content`, not the + /// full document after concatenation. See [`PatternPrivacyClassifier`] + /// limitations for details. + /// + /// When a privacy redirect occurs, the append targets a **separate + /// document** in the private scope at the same path — the shared-scope + /// document is left unmodified. Subsequent multi-scope reads will return + /// the private copy (primary scope wins), effectively shadowing the + /// shared document at that path. The `WriteResult::redirected` flag + /// indicates when this has happened. + /// + /// Uses a read-modify-write pattern that is not concurrency-safe: + /// concurrent appends to the same path may lose writes. pub async fn append_to_layer( &self, layer_name: &str, @@ -706,13 +840,25 @@ impl Workspace { } /// Check if a file exists. + /// + /// When multi-scope reads are configured, checks across all read scopes. pub async fn exists(&self, path: &str) -> Result { let path = normalize_path(path); - match self - .storage - .get_document_by_path(&self.user_id, self.agent_id, &path) - .await - { + let result = if self.is_multi_scope() && is_identity_path(&path) { + // Identity files only checked in primary scope. + self.storage + .get_document_by_path(&self.user_id, self.agent_id, &path) + .await + } else if self.is_multi_scope() { + self.storage + .get_document_by_path_multi(&self.read_user_ids, self.agent_id, &path) + .await + } else { + self.storage + .get_document_by_path(&self.user_id, self.agent_id, &path) + .await + }; + match result { Ok(_) => Ok(true), Err(WorkspaceError::DocumentNotFound { .. }) => Ok(false), Err(e) => Err(e), @@ -747,16 +893,55 @@ impl Workspace { /// ``` pub async fn list(&self, directory: &str) -> Result, WorkspaceError> { let directory = normalize_directory(directory); - self.storage - .list_directory(&self.user_id, self.agent_id, &directory) - .await + if self.is_multi_scope() { + // Iterate per-scope rather than using list_directory_multi because + // we need to filter identity paths from secondary scopes only — the + // merged _multi result loses scope attribution. + let primary = self + .storage + .list_directory(&self.user_id, self.agent_id, &directory) + .await?; + let mut all_entries = primary; + for scope in &self.read_user_ids[1..] { + let entries = self + .storage + .list_directory(scope, self.agent_id, &directory) + .await?; + all_entries.extend(entries.into_iter().filter(|e| !is_identity_path(&e.path))); + } + Ok(merge_workspace_entries(all_entries)) + } else { + self.storage + .list_directory(&self.user_id, self.agent_id, &directory) + .await + } } /// List all files recursively (flat list of all paths). + /// + /// When multi-scope reads are configured, lists across all read scopes. pub async fn list_all(&self) -> Result, WorkspaceError> { - self.storage - .list_all_paths(&self.user_id, self.agent_id) - .await + if self.is_multi_scope() { + // Iterate per-scope rather than using list_all_paths_multi because + // we need to filter identity paths from secondary scopes only. + // Primary scope: all paths. Secondary scopes: filter identity paths. + let mut all_paths = self + .storage + .list_all_paths(&self.user_id, self.agent_id) + .await?; + for scope in &self.read_user_ids[1..] { + let paths = self.storage.list_all_paths(scope, self.agent_id).await?; + all_paths.extend(paths.into_iter().filter(|p| !is_identity_path(p))); + } + // Deduplicate and sort + all_paths.sort(); + all_paths.dedup(); + Ok(all_paths) + } else { + self.storage + .list_all_paths(&self.user_id, self.agent_id) + .await + } } // ==================== Convenience Methods ==================== @@ -791,7 +976,7 @@ impl Workspace { /// comments, which the heartbeat runner treats as "effectively empty" /// and skips the LLM call. pub async fn heartbeat_checklist(&self) -> Result, WorkspaceError> { - match self.read(paths::HEARTBEAT).await { + match self.read_primary(paths::HEARTBEAT).await { Ok(doc) => Ok(Some(doc.content)), Err(WorkspaceError::DocumentNotFound { .. }) => Ok(Some(HEARTBEAT_SEED.to_string())), Err(e) => Err(e), @@ -799,7 +984,29 @@ impl Workspace { } /// Helper to read or create a file. + /// + /// When multi-scope reads are configured, checks all read scopes before + /// creating. If the file exists in any scope, returns it. If not found in + /// any scope, creates it in the primary (write) scope. + /// + /// **Important:** In multi-scope mode, the returned document may belong to + /// a secondary scope. Callers that intend to **write** to the document + /// (via `update_document(doc.id, ...)`) must NOT use this method — use + /// `storage.get_or_create_document_by_path(&self.user_id, ...)` instead + /// to guarantee writes target the primary scope. See `append_memory` for + /// the correct pattern. async fn read_or_create(&self, path: &str) -> Result { + if self.is_multi_scope() { + match self + .storage + .get_document_by_path_multi(&self.read_user_ids, self.agent_id, path) + .await + { + Ok(doc) => return Ok(doc), + Err(WorkspaceError::DocumentNotFound { .. }) => {} + Err(e) => return Err(e), + } + } self.storage .get_or_create_document_by_path(&self.user_id, self.agent_id, path) .await @@ -811,9 +1018,18 @@ impl Workspace { /// /// This is for important facts, decisions, and preferences worth /// remembering long-term. + /// + /// Uses `get_or_create_document_by_path` with the primary `user_id` + /// instead of `self.memory()` to guarantee writes always target the + /// primary (write) scope. `self.memory()` delegates to `read_or_create`, + /// which in multi-scope mode may return a document owned by a secondary + /// scope; writing to that document by UUID would violate write isolation. pub async fn append_memory(&self, entry: &str) -> Result<(), WorkspaceError> { - // Use double newline for memory entries (semantic separation) - let doc = self.memory().await?; + // Always get/create in the primary scope to preserve write isolation. + let doc = self + .storage + .get_or_create_document_by_path(&self.user_id, self.agent_id, paths::MEMORY) + .await?; let new_content = if doc.content.is_empty() { entry.to_string() } else { @@ -905,9 +1121,16 @@ impl Workspace { // Safety net: if `profile_onboarding_completed` was already set (the // LLM completed onboarding but forgot to delete BOOTSTRAP.md), skip // injection to avoid repeating the first-run ritual. + // + // Identity and config files use read_primary() to prevent cross-scope + // bleed in multi-scope workspaces. Without this, a user with read access + // to other scopes could silently inherit another user's identity if their + // own copy is missing — the agent would present as the wrong person. + // Memory files (MEMORY.md, daily logs) intentionally use multi-scope + // read() since sharing memory across scopes is a feature. let bootstrap_injected = if self.is_bootstrap_completed() { if self - .read(paths::BOOTSTRAP) + .read_primary(paths::BOOTSTRAP) .await .is_ok_and(|d| !d.content.is_empty()) { @@ -917,7 +1140,7 @@ impl Workspace { ); } false - } else if let Ok(doc) = self.read(paths::BOOTSTRAP).await + } else if let Ok(doc) = self.read_primary(paths::BOOTSTRAP).await && !doc.content.is_empty() { parts.push(format!("## First-Run Bootstrap\n\n{}", doc.content)); @@ -926,7 +1149,8 @@ impl Workspace { false }; - // Load identity files in order of importance + // Load identity files in order of importance. + // These MUST use read_primary() — see comment above. let identity_files = [ (paths::AGENTS, "## Agent Instructions"), (paths::SOUL, "## Core Values"), @@ -935,7 +1159,7 @@ impl Workspace { ]; for (path, header) in identity_files { - if let Ok(doc) = self.read(path).await + if let Ok(doc) = self.read_primary(path).await && !doc.content.is_empty() { parts.push(format!("{}\n\n{}", header, doc.content)); @@ -944,7 +1168,8 @@ impl Workspace { // Tool notes: environment-specific guidance the agent or user has written. // TOOLS.md does not control tool availability; it is guidance only. - if let Ok(doc) = self.read(paths::TOOLS).await + // Uses read_primary() — tool config is per-user, not inherited. + if let Ok(doc) = self.read_primary(paths::TOOLS).await && !doc.content.is_empty() { parts.push(format!("## Tool Notes\n\n{}", doc.content)); @@ -1235,6 +1460,8 @@ impl Workspace { } /// Search with custom configuration. + /// + /// When multi-scope reads are configured, searches across all read scopes. pub async fn search_with_config( &self, query: &str, @@ -1254,15 +1481,46 @@ impl Workspace { None }; - self.storage - .hybrid_search( - &self.user_id, - self.agent_id, - query, - embedding.as_deref(), - &config, - ) - .await + if self.is_multi_scope() { + let results = self + .storage + .hybrid_search_multi( + &self.read_user_ids, + self.agent_id, + query, + embedding.as_deref(), + &config, + ) + .await?; + // Post-filter: exclude identity documents from secondary scopes. + // Collect document IDs that are identity paths in secondary scopes. + let mut excluded_doc_ids = std::collections::HashSet::new(); + for result in &results { + if is_identity_path(&result.document_path) { + // Check if this document belongs to a secondary scope + match self.storage.get_document_by_id(result.document_id).await { + Ok(doc) if doc.user_id != self.user_id => { + excluded_doc_ids.insert(result.document_id); + } + _ => {} + } + } + } + Ok(results + .into_iter() + .filter(|r| !excluded_doc_ids.contains(&r.document_id)) + .collect()) + } else { + self.storage + .hybrid_search( + &self.user_id, + self.agent_id, + query, + embedding.as_deref(), + &config, + ) + .await + } } // ==================== Indexing ==================== @@ -1323,13 +1581,13 @@ impl Workspace { // Check freshness BEFORE seeding identity files, otherwise the // seeded files make the workspace look non-fresh and BOOTSTRAP.md // never gets created. - let is_fresh_workspace = if self.read(paths::BOOTSTRAP).await.is_ok() { + let is_fresh_workspace = if self.read_primary(paths::BOOTSTRAP).await.is_ok() { false // BOOTSTRAP already exists } else { let (agents_res, soul_res, user_res) = tokio::join!( - self.read(paths::AGENTS), - self.read(paths::SOUL), - self.read(paths::USER), + self.read_primary(paths::AGENTS), + self.read_primary(paths::SOUL), + self.read_primary(paths::USER), ); matches!(agents_res, Err(WorkspaceError::DocumentNotFound { .. })) && matches!(soul_res, Err(WorkspaceError::DocumentNotFound { .. })) @@ -1338,8 +1596,10 @@ impl Workspace { let mut count = 0; for (path, content) in seed_files { - // Skip files that already exist (never overwrite user edits) - match self.read(path).await { + // Skip files that already exist in the primary scope (never overwrite user edits). + // Uses read_primary to avoid false positives from secondary scopes — + // a file in another scope should not suppress seeding in this scope. + match self.read_primary(path).await { Ok(_) => continue, Err(WorkspaceError::DocumentNotFound { .. }) => {} Err(e) => { @@ -1360,7 +1620,8 @@ impl Workspace { // may already have a profile from a previous install and doesn't need // onboarding). This prevents existing users from getting a spurious // first-run ritual after upgrading. - let has_profile = self.read(paths::PROFILE).await.is_ok_and(|d| { + // Uses read_primary() to avoid false positives from secondary scopes. + let has_profile = self.read_primary(paths::PROFILE).await.is_ok_and(|d| { !d.content.trim().is_empty() && serde_json::from_str::(&d.content).is_ok() }); @@ -1791,4 +2052,67 @@ mod seed_tests { "BOOTSTRAP.md should NOT have been seeded with existing profile" ); } + + #[test] + fn test_default_single_scope() { + // Verify backward compatibility: default workspace has single read scope + // matching user_id. + let user_id = "alice"; + let read_user_ids = [user_id.to_string()]; + assert_eq!(read_user_ids.len(), 1); + assert_eq!(read_user_ids[0], user_id); + } + + #[test] + fn test_additional_read_scopes() { + // Verify that additional read scopes are added correctly. + let user_id = "alice".to_string(); + let mut read_user_ids = Vec::from([user_id.clone()]); + + // Simulate with_additional_read_scopes logic + let scopes = ["shared", "team"]; + for scope in scopes { + let s = scope.to_string(); + if !read_user_ids.contains(&s) { + read_user_ids.push(s); + } + } + + assert_eq!(read_user_ids.len(), 3); + assert_eq!(read_user_ids[0], "alice"); + assert_eq!(read_user_ids[1], "shared"); + assert_eq!(read_user_ids[2], "team"); + } + + #[test] + fn test_additional_read_scopes_dedup() { + // Verify that duplicate scopes are ignored. + let user_id = "alice".to_string(); + let mut read_user_ids = Vec::from([user_id.clone()]); + + let scopes = ["shared", "alice", "shared"]; + for scope in scopes { + let s = scope.to_string(); + if !read_user_ids.contains(&s) { + read_user_ids.push(s); + } + } + + assert_eq!(read_user_ids.len(), 2); + assert_eq!(read_user_ids[0], "alice"); + assert_eq!(read_user_ids[1], "shared"); + } + + #[test] + fn test_is_multi_scope_logic() { + // Test the multi-scope detection logic: > 1 means multi-scope + let single_count = 1_usize; + let multi_count = 2_usize; + + // Single scope: not multi + assert!(single_count <= 1); + + // Multi scope: is multi + assert!(multi_count > 1); + } } diff --git a/src/workspace/repository.rs b/src/workspace/repository.rs index 82e4f949..78ddfec5 100644 --- a/src/workspace/repository.rs +++ b/src/workspace/repository.rs @@ -502,4 +502,203 @@ impl Repository { }) .collect()) } + + // ==================== Multi-scope search (optimized SQL) ==================== + + /// Hybrid search across multiple user scopes with efficient SQL. + /// + /// Uses `user_id = ANY($1::text[])` instead of N separate queries. + pub async fn hybrid_search_multi( + &self, + user_ids: &[String], + agent_id: Option, + query: &str, + embedding: Option<&[f32]>, + config: &SearchConfig, + ) -> Result, WorkspaceError> { + let fts_results = if config.use_fts { + self.fts_search_multi(user_ids, agent_id, query, config.pre_fusion_limit) + .await? + } else { + Vec::new() + }; + + let vector_results = if config.use_vector { + if let Some(embedding) = embedding { + self.vector_search_multi(user_ids, agent_id, embedding, config.pre_fusion_limit) + .await? + } else { + Vec::new() + } + } else { + Vec::new() + }; + + Ok(fuse_results(fts_results, vector_results, config)) + } + + /// FTS search across multiple user scopes. + async fn fts_search_multi( + &self, + user_ids: &[String], + agent_id: Option, + query: &str, + limit: usize, + ) -> Result, WorkspaceError> { + let conn = self.conn().await?; + + let rows = conn + .query( + r#" + SELECT c.id as chunk_id, c.document_id, d.path as document_path, + c.content, + ts_rank_cd(c.content_tsv, plainto_tsquery('english', $3)) as rank + FROM memory_chunks c + JOIN memory_documents d ON d.id = c.document_id + WHERE d.user_id = ANY($1::text[]) AND d.agent_id IS NOT DISTINCT FROM $2 + AND c.content_tsv @@ plainto_tsquery('english', $3) + ORDER BY rank DESC + LIMIT $4 + "#, + &[&user_ids, &agent_id, &query, &(limit as i64)], + ) + .await + .map_err(|e| WorkspaceError::SearchFailed { + reason: format!("FTS multi-scope query failed: {}", e), + })?; + + Ok(rows + .iter() + .enumerate() + .map(|(i, row)| RankedResult { + chunk_id: row.get("chunk_id"), + document_id: row.get("document_id"), + document_path: row.get("document_path"), + content: row.get("content"), + rank: (i + 1) as u32, + }) + .collect()) + } + + /// Vector search across multiple user scopes. + async fn vector_search_multi( + &self, + user_ids: &[String], + agent_id: Option, + embedding: &[f32], + limit: usize, + ) -> Result, WorkspaceError> { + let conn = self.conn().await?; + let embedding_vec = Vector::from(embedding.to_vec()); + + let rows = conn + .query( + r#" + SELECT c.id as chunk_id, c.document_id, d.path as document_path, + c.content, 1 - (c.embedding <=> $3) as similarity + FROM memory_chunks c + JOIN memory_documents d ON d.id = c.document_id + WHERE d.user_id = ANY($1::text[]) AND d.agent_id IS NOT DISTINCT FROM $2 + AND c.embedding IS NOT NULL + ORDER BY c.embedding <=> $3 + LIMIT $4 + "#, + &[&user_ids, &agent_id, &embedding_vec, &(limit as i64)], + ) + .await + .map_err(|e| WorkspaceError::SearchFailed { + reason: format!("Vector multi-scope query failed: {}", e), + })?; + + Ok(rows + .iter() + .enumerate() + .map(|(i, row)| RankedResult { + chunk_id: row.get("chunk_id"), + document_id: row.get("document_id"), + document_path: row.get("document_path"), + content: row.get("content"), + rank: (i + 1) as u32, + }) + .collect()) + } + + /// List all file paths across multiple user scopes with a single query. + pub async fn list_all_paths_multi( + &self, + user_ids: &[String], + agent_id: Option, + ) -> Result, WorkspaceError> { + let conn = self.conn().await?; + + let rows = conn + .query( + r#" + SELECT DISTINCT path FROM memory_documents + WHERE user_id = ANY($1::text[]) AND agent_id IS NOT DISTINCT FROM $2 + ORDER BY path + "#, + &[&user_ids, &agent_id], + ) + .await + .map_err(|e| WorkspaceError::SearchFailed { + reason: format!("List paths multi-scope failed: {}", e), + })?; + + Ok(rows.iter().map(|row| row.get("path")).collect()) + } + + /// Get a document by path across multiple user scopes. + /// + /// Returns the first match (ordered by the input user_ids priority). + pub async fn get_document_by_path_multi( + &self, + user_ids: &[String], + agent_id: Option, + path: &str, + ) -> Result { + let conn = self.conn().await?; + + let row = conn + .query_opt( + r#" + SELECT id, user_id, agent_id, path, content, + created_at, updated_at, metadata + FROM memory_documents + WHERE user_id = ANY($1::text[]) AND agent_id IS NOT DISTINCT FROM $2 AND path = $3 + ORDER BY array_position($1::text[], user_id) + LIMIT 1 + "#, + &[&user_ids, &agent_id, &path], + ) + .await + .map_err(|e| WorkspaceError::SearchFailed { + reason: format!("get_document_by_path_multi failed: {}", e), + })?; + + match row { + Some(row) => Ok(self.row_to_document(&row)), + None => Err(WorkspaceError::DocumentNotFound { + doc_type: path.to_string(), + user_id: format!("[{}]", user_ids.join(", ")), + }), + } + } + + /// List directory contents across multiple user scopes. + /// + /// Iterates per scope and merges results. A future migration could add an + /// optimised SQL function, at which point this method can call it directly. + pub async fn list_directory_multi( + &self, + user_ids: &[String], + agent_id: Option, + directory: &str, + ) -> Result, WorkspaceError> { + let mut all_entries = Vec::new(); + for uid in user_ids { + all_entries.extend(self.list_directory(uid, agent_id, directory).await?); + } + Ok(crate::workspace::merge_workspace_entries(all_entries)) + } } diff --git a/tests/e2e/scenarios/test_oauth_url_parameters.py b/tests/e2e/scenarios/test_oauth_url_parameters.py new file mode 100644 index 00000000..0dae3e53 --- /dev/null +++ b/tests/e2e/scenarios/test_oauth_url_parameters.py @@ -0,0 +1,249 @@ +"""OAuth URL parameter validation e2e tests. + +Tests for bug #992: Google OAuth URL broken when initiated from Telegram. +Specifically verifies that OAuth query parameters are correctly formatted: +- "client_id" (with underscore) NOT "clientid" (without underscore) +- All standard OAuth parameters are present and correctly encoded +- URLs are consistent across channels (web, Telegram, etc.) + +The test verifies: +1. OAuth URL is generated with correct parameters +2. URL works with the OAuth provider (Google) +3. Extra parameters (access_type, prompt) are preserved +""" + +from urllib.parse import parse_qs, urlparse +import pytest + +from helpers import api_post, api_get + + +async def _extract_oauth_params(auth_url: str) -> dict: + """Extract and validate OAuth query parameters from auth_url. + + Returns dict with parsed parameters: + { + 'client_id': '...', + 'redirect_uri': '...', + 'response_type': 'code', + 'scope': '...', + 'state': '...', + 'access_type': '...', + 'prompt': '...', + ... + } + """ + parsed = urlparse(auth_url) + qs = parse_qs(parsed.query) + + # Convert lists to single values for easier testing + params = {k: v[0] if len(v) > 0 else v for k, v in qs.items()} + return params + + +async def _get_extension(ironclaw_server, name): + """Get a specific extension from the extensions list, or None.""" + r = await api_get(ironclaw_server, "/api/extensions") + for ext in r.json().get("extensions", []): + if ext["name"] == name: + return ext + return None + + +@pytest.fixture +async def installed_gmail(ironclaw_server): + """Installs the 'gmail' extension before a test and removes it after. + + This fixture handles the setup and teardown of the Gmail extension, + ensuring a clean state for each test. + """ + # Ensure Gmail is not installed before test + ext = await _get_extension(ironclaw_server, "gmail") + if ext: + r = await api_post(ironclaw_server, "/api/extensions/gmail/remove", timeout=30) + assert r.status_code == 200 + + # Install Gmail + r = await api_post( + ironclaw_server, + "/api/extensions/install", + json={"name": "gmail"}, + timeout=180, + ) + assert r.status_code == 200, f"Gmail install failed: {r.text}" + assert r.json().get("success") is True, f"Install failed: {r.json().get('message', '')}" + + yield + + # Teardown: remove gmail + r = await api_post(ironclaw_server, "/api/extensions/gmail/remove", timeout=30) + assert r.status_code == 200, f"Gmail removal failed: {r.text}" + + +@pytest.fixture +async def auth_url(ironclaw_server, installed_gmail): + """Generate and return an OAuth auth URL. + + Requires Gmail to be installed (depends on installed_gmail fixture). + """ + r = await api_post( + ironclaw_server, + "/api/extensions/gmail/setup", + json={"secrets": {}}, + timeout=30, + ) + assert r.status_code == 200 + data = r.json() + assert data.get("success") is True, f"Setup failed: {data.get('message', '')}" + + url = data.get("auth_url") + assert url is not None, f"Expected auth_url in response: {data}" + assert "accounts.google.com" in url, f"auth_url should point to Google: {url}" + + return url + + +@pytest.fixture +async def oauth_params(auth_url): + """Extract and return OAuth parameters from auth_url. + + Depends on auth_url fixture. + """ + return await _extract_oauth_params(auth_url) + + +# ─ OAuth URL parameter validation tests ──────────────────────────────── + +async def test_oauth_url_has_client_id_not_clientid(oauth_params, auth_url): + """Verify OAuth URL has 'client_id' (with underscore), NOT 'clientid'. + + Bug #992: Ensure the parameter name is correct across all channels. + """ + params = oauth_params + + # The bug: "clientid" appears instead of "client_id" + # Verify the CORRECT parameter name exists + assert "client_id" in params, ( + f"OAuth URL missing 'client_id' parameter. " + f"URL: {auth_url}\nParams: {params}" + ) + assert params["client_id"], "client_id should have a value" + + # Verify the INCORRECT parameter name does NOT exist + assert "clientid" not in params, ( + f"OAuth URL should NOT have 'clientid' (without underscore). " + f"Bug #992: URL: {auth_url}\nParams: {params}" + ) + + +async def test_oauth_url_has_required_parameters(oauth_params): + """Verify all required OAuth 2.0 parameters are present.""" + params = oauth_params + + # Required OAuth 2.0 parameters + required = ["client_id", "response_type", "redirect_uri", "scope", "state"] + for param in required: + assert param in params, ( + f"Missing required OAuth parameter: {param}. " + f"Params: {params}" + ) + assert params[param], f"Parameter '{param}' should have a non-empty value" + + # Validate specific values + assert params["response_type"] == "code", "Should use authorization_code flow" + assert "oauth" in params["redirect_uri"], "Redirect URI should be an OAuth callback" + + +async def test_oauth_url_has_extra_params(oauth_params): + """Verify extra_params from capabilities.json are included.""" + params = oauth_params + + # Google-specific extra_params from gmail-tool.capabilities.json + assert "access_type" in params, ( + "Should include 'access_type' from extra_params" + ) + assert params["access_type"] == "offline", ( + "access_type should be 'offline' for Gmail" + ) + + assert "prompt" in params, ( + "Should include 'prompt' from extra_params" + ) + assert params["prompt"] == "consent", ( + "prompt should be 'consent' for Gmail" + ) + + +async def test_oauth_url_is_valid_google_oauth(auth_url): + """Verify the URL is a valid Google OAuth 2.0 authorization URL.""" + # Verify scheme and host + parsed = urlparse(auth_url) + assert parsed.scheme == "https", "OAuth URL must use HTTPS" + assert "accounts.google.com" in parsed.netloc, "Must be Google's OAuth endpoint" + assert parsed.path == "/o/oauth2/v2/auth", "Must use Google OAuth 2.0 endpoint" + + +async def test_oauth_url_state_is_unique(ironclaw_server, installed_gmail, oauth_params, auth_url): + """Verify CSRF state is present and unique per request.""" + # Get a new OAuth URL + r = await api_post( + ironclaw_server, + "/api/extensions/gmail/setup", + json={"secrets": {}}, + timeout=30, + ) + assert r.status_code == 200 + new_auth_url = r.json().get("auth_url") + assert new_auth_url is not None + + # Extract state from both URLs + original_params = oauth_params + new_params = await _extract_oauth_params(new_auth_url) + + original_state = original_params.get("state") + new_state = new_params.get("state") + + assert original_state is not None, "Should have state parameter" + assert new_state is not None, "New request should have state parameter" + assert original_state != new_state, ( + "CSRF state should be unique per request (for security)" + ) + + +async def test_oauth_url_escaping(auth_url): + """Verify URL query parameters are properly escaped.""" + # Verify special characters in values are URL-encoded + # For example, scopes contain spaces which should be %20 + assert "%20" in auth_url or "+" in auth_url or "%2B" in auth_url or " " not in auth_url, ( + "OAuth URL should properly encode special characters in parameters" + ) + + +# ─ Telegram-specific tests (when Telegram channel is available) ────────── + +class TestOAuthURLViaTelegram: + """Test OAuth URL generation specifically via Telegram channel. + + These tests would verify that the same OAuth URL works correctly when + transmitted through the Telegram WASM channel (as opposed to web gateway). + + Currently marked as xfail pending Telegram channel setup in E2E tests. + """ + + @pytest.mark.skip(reason="Telegram channel E2E setup not yet implemented") + async def test_telegram_oauth_url_has_correct_parameters(self): + """Verify OAuth URL sent via Telegram has correct parameter names.""" + # This test would: + # 1. Send a message via Telegram that triggers OAuth + # 2. Capture the status update sent to Telegram + # 3. Extract the auth_url from the message + # 4. Verify it has "client_id" not "clientid" + pass + + @pytest.mark.skip(reason="Telegram channel E2E setup not yet implemented") + async def test_telegram_oauth_url_can_be_regenerated(self): + """Verify OAuth URL can be regenerated when requested via Telegram.""" + # This test would verify that the bug #992 symptom + # "URL cannot be regenerated when asked" is fixed. + # If the URL is cached incorrectly, regeneration would fail. + pass 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/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/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/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 + } + } + ] +} diff --git a/tests/identity_scope_isolation.rs b/tests/identity_scope_isolation.rs new file mode 100644 index 00000000..314e87f3 --- /dev/null +++ b/tests/identity_scope_isolation.rs @@ -0,0 +1,195 @@ +//! Tests for identity file scope isolation in multi-scope workspaces. +//! +//! When a workspace has multiple read scopes (e.g., Andrew can read from +//! "andrew", "grace", "household"), identity files (SOUL.md, USER.md, +//! IDENTITY.md, AGENTS.md) must ONLY come from the primary scope. +//! +//! Multi-scope reads are designed for memory sharing (MEMORY.md, daily logs), +//! not identity inheritance. Silently inheriting identity from another scope +//! is a correctness and security issue — the agent would present itself as +//! the wrong user. +//! +//! These tests verify that: +//! 1. Identity files are read from primary scope only +//! 2. If the primary scope's identity file is missing, it's absent from the +//! system prompt — never falls back to another scope +//! 3. Memory files (MEMORY.md) still benefit from multi-scope reads +#![cfg(feature = "libsql")] + +use std::sync::Arc; + +use ironclaw::db::Database; +use ironclaw::db::libsql::LibSqlBackend; +use ironclaw::workspace::{Workspace, paths}; + +async fn setup() -> (Arc, tempfile::TempDir) { + let dir = tempfile::tempdir().expect("create temp dir"); + let db_path = dir.path().join("test.db"); + let backend = LibSqlBackend::new_local(&db_path).await.expect("create db"); + backend.run_migrations().await.expect("run migrations"); + let db: Arc = Arc::new(backend); + (db, dir) +} + +/// Seed a document into a specific user's workspace scope. +async fn seed(db: &Arc, user_id: &str, path: &str, content: &str) { + let ws = Workspace::new_with_db(user_id, db.clone()); + ws.write(path, content) + .await + .unwrap_or_else(|e| panic!("Failed to seed {path} for {user_id}: {e}")); +} + +// ─── Test 1: Primary scope identity appears in system prompt ─────────── + +#[tokio::test] +async fn system_prompt_uses_primary_scope_identity() { + let (db, _dir) = setup().await; + + // Seed Alice's identity files in her own scope + seed(&db, "alice", paths::SOUL, "Alice is kind and curious.").await; + seed( + &db, + "alice", + paths::USER, + "You are talking to Alice, a software engineer.", + ) + .await; + + // Seed Bob's identity files in his scope + seed(&db, "bob", paths::SOUL, "Bob is analytical and precise.").await; + seed( + &db, + "bob", + paths::USER, + "You are talking to Bob, a marine biologist.", + ) + .await; + + // Create Alice's workspace WITH multi-scope reads including Bob + let ws = Workspace::new_with_db("alice", db.clone()) + .with_additional_read_scopes(vec!["bob".to_string()]); + + let prompt = ws + .system_prompt_for_context(false) + .await + .expect("system_prompt_for_context failed"); + + // Alice's identity must appear + assert!( + prompt.contains("Alice is kind and curious"), + "Primary scope SOUL.md should appear in system prompt.\nPrompt:\n{prompt}" + ); + assert!( + prompt.contains("Alice, a software engineer"), + "Primary scope USER.md should appear in system prompt.\nPrompt:\n{prompt}" + ); + + // Bob's identity must NOT appear + assert!( + !prompt.contains("Bob is analytical"), + "Secondary scope SOUL.md must NOT appear in system prompt.\nPrompt:\n{prompt}" + ); + assert!( + !prompt.contains("Bob, a marine biologist"), + "Secondary scope USER.md must NOT appear in system prompt.\nPrompt:\n{prompt}" + ); +} + +// ─── Test 2: Missing primary identity does NOT fall back to other scope ─ + +#[tokio::test] +async fn missing_primary_identity_does_not_fallback_to_other_scope() { + let (db, _dir) = setup().await; + + // Only seed Bob's identity — Alice has no identity files + seed(&db, "bob", paths::SOUL, "Bob is analytical and precise.").await; + seed( + &db, + "bob", + paths::USER, + "You are talking to Bob, a marine biologist.", + ) + .await; + + // Create Alice's workspace with multi-scope reads including Bob + let ws = Workspace::new_with_db("alice", db.clone()) + .with_additional_read_scopes(vec!["bob".to_string()]); + + let prompt = ws + .system_prompt_for_context(false) + .await + .expect("system_prompt_for_context failed"); + + // Bob's identity must NOT appear — Alice's missing identity should stay missing, + // not silently inherit from Bob's scope + assert!( + !prompt.contains("Bob"), + "When primary scope identity is missing, must NOT fall back to secondary scope.\n\ + This would cause the agent to present itself as the wrong user.\nPrompt:\n{prompt}" + ); +} + +// ─── Test 3: MEMORY.md still benefits from multi-scope reads ──────────── + +#[tokio::test] +async fn memory_files_still_use_multi_scope_reads() { + let (db, _dir) = setup().await; + + // Seed shared memory in the "shared" scope (not Alice's primary) + seed( + &db, + "shared", + paths::MEMORY, + "Shared grocery list: milk, eggs, bread.", + ) + .await; + + // Create Alice's workspace with read access to shared scope + let ws = Workspace::new_with_db("alice", db.clone()) + .with_additional_read_scopes(vec!["shared".to_string()]); + + let prompt = ws + .system_prompt_for_context(false) + .await + .expect("system_prompt_for_context failed"); + + // Shared memory SHOULD appear — multi-scope reads are correct for memory + assert!( + prompt.contains("grocery list"), + "MEMORY.md should still use multi-scope reads.\nPrompt:\n{prompt}" + ); +} + +// ─── Test 4: All identity files are scope-isolated ────────────────────── + +#[tokio::test] +async fn all_identity_files_are_scope_isolated() { + let (db, _dir) = setup().await; + + // Seed identity files ONLY in the "other" scope, not in Alice's + seed(&db, "other", paths::AGENTS, "You are Other's agent.").await; + seed(&db, "other", paths::SOUL, "Other's soul values.").await; + seed(&db, "other", paths::USER, "You are talking to Other.").await; + seed(&db, "other", paths::IDENTITY, "Other's identity.").await; + + // Also seed BOOTSTRAP.md and TOOLS.md in other scope + seed(&db, "other", "BOOTSTRAP.md", "Other's bootstrap.").await; + seed(&db, "other", "TOOLS.md", "Other's tool notes.").await; + + // Create Alice's workspace with read access to "other" + let ws = Workspace::new_with_db("alice", db.clone()) + .with_additional_read_scopes(vec!["other".to_string()]); + + let prompt = ws + .system_prompt_for_context(false) + .await + .expect("system_prompt_for_context failed"); + + // None of Other's identity/config files should appear + assert!( + !prompt.contains("Other"), + "No identity or config files from secondary scope should appear.\n\ + Every identity file (AGENTS.md, SOUL.md, USER.md, IDENTITY.md, \ + BOOTSTRAP.md, TOOLS.md) must read from primary scope only.\nPrompt:\n{prompt}" + ); +} 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_scope_functional.rs b/tests/multi_scope_functional.rs new file mode 100644 index 00000000..77829b9d --- /dev/null +++ b/tests/multi_scope_functional.rs @@ -0,0 +1,451 @@ +#![cfg(feature = "libsql")] +//! Integration tests for multi-scope workspace reads using file-backed libSQL. +//! +//! Guards the PR2 contract: workspaces can read from multiple user scopes +//! while writes remain isolated to the primary scope. + +use std::sync::Arc; + +use ironclaw::db::Database; +use ironclaw::db::libsql::LibSqlBackend; +use ironclaw::workspace::Workspace; + +async fn setup() -> (Arc, tempfile::TempDir) { + let dir = tempfile::tempdir().expect("create temp dir"); + let db_path = dir.path().join("test.db"); + let backend = LibSqlBackend::new_local(&db_path).await.expect("create db"); + backend.run_migrations().await.expect("run migrations"); + let db: Arc = Arc::new(backend); + (db, dir) +} + +#[tokio::test] +async fn read_across_scopes() { + let (db, _dir) = setup().await; + + // Write docs as the "shared" user + let ws_shared = Workspace::new_with_db("shared", Arc::clone(&db)); + ws_shared + .write("docs/team-standup.md", "Team standup notes from Monday") + .await + .expect("shared write failed"); + + // Alice's workspace with "shared" as an additional read scope + let ws_alice = Workspace::new_with_db("alice", Arc::clone(&db)) + .with_additional_read_scopes(vec!["shared".to_string()]); + + // Alice can read shared docs + let doc = ws_alice + .read("docs/team-standup.md") + .await + .expect("cross-scope read failed"); + assert_eq!(doc.content, "Team standup notes from Monday"); +} + +#[tokio::test] +async fn write_stays_in_primary_scope() { + let (db, _dir) = setup().await; + + // Alice has "shared" as a read scope + let ws_alice = Workspace::new_with_db("alice", Arc::clone(&db)) + .with_additional_read_scopes(vec!["shared".to_string()]); + + // Alice writes a personal note + ws_alice + .write("notes/personal.md", "Alice's private note") + .await + .expect("alice write failed"); + + // The "shared" workspace should NOT see Alice's note + let ws_shared = Workspace::new_with_db("shared", Arc::clone(&db)); + let result = ws_shared.read("notes/personal.md").await; + assert!(result.is_err(), "Shared scope should not see Alice's note"); +} + +#[tokio::test] +async fn list_paths_merges_across_scopes() { + let (db, _dir) = setup().await; + + // Write as alice + let ws_alice_plain = Workspace::new_with_db("alice", Arc::clone(&db)); + ws_alice_plain + .write("notes/personal.md", "My notes") + .await + .expect("alice write failed"); + + // Write as shared + let ws_shared = Workspace::new_with_db("shared", Arc::clone(&db)); + ws_shared + .write("docs/shared-doc.md", "Shared document") + .await + .expect("shared write failed"); + + // Alice with multi-scope should see both + let ws_alice = Workspace::new_with_db("alice", Arc::clone(&db)) + .with_additional_read_scopes(vec!["shared".to_string()]); + + let all_paths = ws_alice.list_all().await.expect("list_all failed"); + assert!( + all_paths.contains(&"notes/personal.md".to_string()), + "Should contain alice's note: {:?}", + all_paths + ); + assert!( + all_paths.contains(&"docs/shared-doc.md".to_string()), + "Should contain shared doc: {:?}", + all_paths + ); +} + +#[tokio::test] +async fn list_directory_merges_across_scopes() { + let (db, _dir) = setup().await; + + // Alice writes to docs/ + let ws_alice_plain = Workspace::new_with_db("alice", Arc::clone(&db)); + ws_alice_plain + .write("docs/alice-doc.md", "Alice's doc") + .await + .expect("alice write failed"); + + // Shared writes to docs/ + let ws_shared = Workspace::new_with_db("shared", Arc::clone(&db)); + ws_shared + .write("docs/shared-doc.md", "Shared doc") + .await + .expect("shared write failed"); + + // Alice with multi-scope lists docs/ + let ws_alice = Workspace::new_with_db("alice", Arc::clone(&db)) + .with_additional_read_scopes(vec!["shared".to_string()]); + + let entries = ws_alice.list("docs").await.expect("list failed"); + let paths: Vec<&str> = entries.iter().map(|e| e.path.as_str()).collect(); + assert!( + paths.contains(&"docs/alice-doc.md"), + "Should contain alice's doc: {:?}", + paths + ); + assert!( + paths.contains(&"docs/shared-doc.md"), + "Should contain shared doc: {:?}", + paths + ); +} + +#[tokio::test] +async fn search_spans_scopes() { + let (db, _dir) = setup().await; + + // Write searchable content in shared scope + let ws_shared = Workspace::new_with_db("shared", Arc::clone(&db)); + ws_shared + .write( + "docs/architecture.md", + "The microservice architecture uses gRPC for inter-service communication", + ) + .await + .expect("shared write failed"); + + // Write searchable content in alice scope + let ws_alice_plain = Workspace::new_with_db("alice", Arc::clone(&db)); + ws_alice_plain + .write("notes/ideas.md", "Consider switching to GraphQL federation") + .await + .expect("alice write failed"); + + // Alice with multi-scope searches + let ws_alice = Workspace::new_with_db("alice", Arc::clone(&db)) + .with_additional_read_scopes(vec!["shared".to_string()]); + + // Search for content in the shared scope + let results = ws_alice + .search("microservice architecture gRPC", 10) + .await + .expect("search failed"); + assert!(!results.is_empty(), "Should find results from shared scope"); +} + +#[tokio::test] +async fn read_priority_primary_first() { + let (db, _dir) = setup().await; + + // Write same path in both scopes + let ws_shared = Workspace::new_with_db("shared", Arc::clone(&db)); + ws_shared + .write("config/settings.md", "Shared settings v1") + .await + .expect("shared write failed"); + + let ws_alice_plain = Workspace::new_with_db("alice", Arc::clone(&db)); + ws_alice_plain + .write("config/settings.md", "Alice's settings override") + .await + .expect("alice write failed"); + + // Alice with multi-scope should get her own version (primary scope wins) + let ws_alice = Workspace::new_with_db("alice", Arc::clone(&db)) + .with_additional_read_scopes(vec!["shared".to_string()]); + + let doc = ws_alice + .read("config/settings.md") + .await + .expect("read failed"); + assert_eq!( + doc.content, "Alice's settings override", + "Primary scope should take priority" + ); +} + +#[tokio::test] +async fn exists_spans_scopes() { + let (db, _dir) = setup().await; + + // Write a doc as "shared" + let ws_shared = Workspace::new_with_db("shared", Arc::clone(&db)); + ws_shared + .write("docs/shared-only.md", "Shared content") + .await + .expect("shared write failed"); + + // Alice without multi-scope should NOT see it + let ws_alice_plain = Workspace::new_with_db("alice", Arc::clone(&db)); + assert!( + !ws_alice_plain + .exists("docs/shared-only.md") + .await + .expect("exists failed"), + "Alice without multi-scope should not see shared doc" + ); + + // Alice with multi-scope should see it + let ws_alice = Workspace::new_with_db("alice", Arc::clone(&db)) + .with_additional_read_scopes(vec!["shared".to_string()]); + assert!( + ws_alice + .exists("docs/shared-only.md") + .await + .expect("exists failed"), + "Alice with multi-scope should see shared doc" + ); +} + +#[tokio::test] +async fn append_stays_in_primary_scope() { + let (db, _dir) = setup().await; + + // Write a document as "shared" + let ws_shared = Workspace::new_with_db("shared", Arc::clone(&db)); + ws_shared + .write("notes/log.md", "shared original content") + .await + .expect("shared write failed"); + + // Alice has "shared" as a read scope and appends to the same path + let ws_alice = Workspace::new_with_db("alice", Arc::clone(&db)) + .with_additional_read_scopes(vec!["shared".to_string()]); + ws_alice + .append("notes/log.md", "alice appended line") + .await + .expect("alice append failed"); + + // Shared document must be unchanged (write isolation) + let shared_doc = ws_shared + .read("notes/log.md") + .await + .expect("shared read failed"); + assert_eq!( + shared_doc.content, "shared original content", + "Append must not modify the secondary scope's document" + ); + + // Alice should have her own copy with the appended content + let ws_alice_plain = Workspace::new_with_db("alice", Arc::clone(&db)); + let alice_doc = ws_alice_plain + .read("notes/log.md") + .await + .expect("alice read failed"); + assert_eq!( + alice_doc.content, "alice appended line", + "Append should create a new document in alice's scope" + ); +} + +#[tokio::test] +async fn append_memory_stays_in_primary_scope() { + let (db, _dir) = setup().await; + + // Write MEMORY.md as "shared" + let ws_shared = Workspace::new_with_db("shared", Arc::clone(&db)); + ws_shared + .write("MEMORY.md", "shared memory baseline") + .await + .expect("shared write failed"); + + // Alice has "shared" as a read scope and appends a memory entry + let ws_alice = Workspace::new_with_db("alice", Arc::clone(&db)) + .with_additional_read_scopes(vec!["shared".to_string()]); + ws_alice + .append_memory("alice remembers this") + .await + .expect("alice append_memory failed"); + + // Shared MEMORY.md must be unchanged + let shared_doc = ws_shared + .read("MEMORY.md") + .await + .expect("shared read failed"); + assert_eq!( + shared_doc.content, "shared memory baseline", + "append_memory must not modify the secondary scope's document" + ); + + // Alice should have her own MEMORY.md + let ws_alice_plain = Workspace::new_with_db("alice", Arc::clone(&db)); + let alice_doc = ws_alice_plain + .read("MEMORY.md") + .await + .expect("alice read failed"); + assert_eq!( + alice_doc.content, "alice remembers this", + "append_memory should create in alice's scope" + ); +} + +// ==================== Identity isolation tests ==================== + +#[tokio::test] +async fn identity_files_not_readable_from_secondary_scope() { + let (db, _dir) = setup().await; + + let ws_other = Workspace::new_with_db("other-user", Arc::clone(&db)); + ws_other + .write("IDENTITY.md", "I am the other user") + .await + .expect("write failed"); + ws_other + .write("SOUL.md", "Other user soul overlay") + .await + .expect("write failed"); + ws_other + .write("USER.md", "Other user profile") + .await + .expect("write failed"); + ws_other + .write("AGENTS.md", "Other user agent config") + .await + .expect("write failed"); + + let ws_primary = Workspace::new_with_db("primary", Arc::clone(&db)) + .with_additional_read_scopes(vec!["other-user".to_string()]); + + for path in &["IDENTITY.md", "SOUL.md", "USER.md", "AGENTS.md"] { + let result = ws_primary.read(path).await; + assert!( + result.is_err(), + "Primary should NOT read other user's {} via secondary scope", + path + ); + } +} + +#[tokio::test] +async fn identity_files_not_in_search_from_secondary_scope() { + let (db, _dir) = setup().await; + + let ws_other = Workspace::new_with_db("other-user", Arc::clone(&db)); + ws_other + .write("SOUL.md", "Other user loves xylophone music passionately") + .await + .expect("write failed"); + ws_other + .write( + "notes/music.md", + "Other user played xylophone at the concert", + ) + .await + .expect("write failed"); + + let ws_primary = Workspace::new_with_db("primary", Arc::clone(&db)) + .with_additional_read_scopes(vec!["other-user".to_string()]); + + let results = ws_primary + .search("xylophone", 10) + .await + .expect("search failed"); + let has_concert = results.iter().any(|r| r.content.contains("concert")); + assert!( + has_concert, + "Should find non-identity content from secondary scope" + ); + let has_soul = results.iter().any(|r| r.content.contains("passionately")); + assert!( + !has_soul, + "SOUL.md content from secondary scope should not appear in search results" + ); +} + +#[tokio::test] +async fn identity_files_not_in_list_from_secondary_scope() { + let (db, _dir) = setup().await; + + let ws_other = Workspace::new_with_db("other-user", Arc::clone(&db)); + ws_other + .write("IDENTITY.md", "I am the other user") + .await + .expect("write failed"); + ws_other + .write("notes/shared-note.md", "A shared note") + .await + .expect("write failed"); + + let ws_primary = Workspace::new_with_db("primary", Arc::clone(&db)) + .with_additional_read_scopes(vec!["other-user".to_string()]); + + let paths = ws_primary.list_all().await.expect("list failed"); + assert!( + !paths.contains(&"IDENTITY.md".to_string()), + "IDENTITY.md from secondary scope should not appear" + ); + assert!( + paths.contains(&"notes/shared-note.md".to_string()), + "Non-identity files should be listed" + ); +} + +#[tokio::test] +async fn empty_read_scopes_reads_primary_only() { + let (db, _dir) = setup().await; + + let ws_shared = Workspace::new_with_db("shared", Arc::clone(&db)); + ws_shared + .write("docs/note.md", "Shared note") + .await + .expect("write failed"); + + let ws_primary = + Workspace::new_with_db("primary", Arc::clone(&db)).with_additional_read_scopes(vec![]); + + let result = ws_primary.read("docs/note.md").await; + assert!( + result.is_err(), + "Empty read scopes should not grant cross-scope access" + ); +} + +#[tokio::test] +async fn duplicate_read_scopes_handled() { + let (db, _dir) = setup().await; + + let ws_shared = Workspace::new_with_db("shared", Arc::clone(&db)); + ws_shared + .write("docs/note.md", "One note") + .await + .expect("write failed"); + + let ws_primary = Workspace::new_with_db("primary", Arc::clone(&db)) + .with_additional_read_scopes(vec!["shared".to_string(), "shared".to_string()]); + + let doc = ws_primary.read("docs/note.md").await.expect("read failed"); + assert_eq!(doc.content, "One note"); +} 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..e4620f70 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,12 +258,13 @@ 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, sandbox_readiness: ironclaw::agent::SandboxReadiness::DisabledByConfig, builder: None, + llm_backend: "nearai".to_string(), }, channels, None, @@ -288,10 +293,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/support/test_rig.rs b/tests/support/test_rig.rs index be2b3bb2..624bb054 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)); } @@ -768,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. diff --git a/tests/workspace_integration.rs b/tests/workspace_integration.rs index 2182fc38..2184d8f2 100644 --- a/tests/workspace_integration.rs +++ b/tests/workspace_integration.rs @@ -407,3 +407,333 @@ async fn test_workspace_system_prompt() { cleanup_user(&pool, user_id).await; } + +// ── Multi-scope workspace read tests ────────────────────────────────── +// +// These exercise the PostgreSQL-optimized `_multi` query paths +// (repository.rs) that the libSQL backend covers via default trait impls. + +#[tokio::test] +async fn test_multi_scope_read_across_scopes() { + let pool = get_pool(); + if try_connect(&pool).await.is_none() { + return; + } + let shared_id = "ms_shared_read"; + let alice_id = "ms_alice_read"; + cleanup_user(&pool, shared_id).await; + cleanup_user(&pool, alice_id).await; + + // Write a doc as "shared" + let ws_shared = Workspace::new(shared_id, pool.clone()); + ws_shared + .write("docs/team-standup.md", "Team standup notes from Monday") + .await + .expect("shared write failed"); + + // Alice with "shared" as an additional read scope + let ws_alice = Workspace::new(alice_id, pool.clone()) + .with_additional_read_scopes(vec![shared_id.to_string()]); + + let doc = ws_alice + .read("docs/team-standup.md") + .await + .expect("cross-scope read failed"); + assert_eq!(doc.content, "Team standup notes from Monday"); + + cleanup_user(&pool, shared_id).await; + cleanup_user(&pool, alice_id).await; +} + +#[tokio::test] +async fn test_multi_scope_write_stays_in_primary() { + let pool = get_pool(); + if try_connect(&pool).await.is_none() { + return; + } + let shared_id = "ms_shared_write"; + let alice_id = "ms_alice_write"; + cleanup_user(&pool, shared_id).await; + cleanup_user(&pool, alice_id).await; + + let ws_alice = Workspace::new(alice_id, pool.clone()) + .with_additional_read_scopes(vec![shared_id.to_string()]); + + ws_alice + .write("notes/personal.md", "Alice's private note") + .await + .expect("alice write failed"); + + // Shared workspace should NOT see Alice's note + let ws_shared = Workspace::new(shared_id, pool.clone()); + let result = ws_shared.read("notes/personal.md").await; + assert!(result.is_err(), "Shared scope should not see Alice's note"); + + cleanup_user(&pool, shared_id).await; + cleanup_user(&pool, alice_id).await; +} + +#[tokio::test] +async fn test_multi_scope_list_all_merges() { + let pool = get_pool(); + if try_connect(&pool).await.is_none() { + return; + } + let shared_id = "ms_shared_list"; + let alice_id = "ms_alice_list"; + cleanup_user(&pool, shared_id).await; + cleanup_user(&pool, alice_id).await; + + // Write as alice (plain, no multi-scope) + let ws_alice_plain = Workspace::new(alice_id, pool.clone()); + ws_alice_plain + .write("notes/personal.md", "My notes") + .await + .expect("alice write failed"); + + // Write as shared + let ws_shared = Workspace::new(shared_id, pool.clone()); + ws_shared + .write("docs/shared-doc.md", "Shared document") + .await + .expect("shared write failed"); + + // Alice with multi-scope should see both + let ws_alice = Workspace::new(alice_id, pool.clone()) + .with_additional_read_scopes(vec![shared_id.to_string()]); + + let all_paths = ws_alice.list_all().await.expect("list_all failed"); + assert!( + all_paths.contains(&"notes/personal.md".to_string()), + "Should contain alice's note: {:?}", + all_paths + ); + assert!( + all_paths.contains(&"docs/shared-doc.md".to_string()), + "Should contain shared doc: {:?}", + all_paths + ); + + cleanup_user(&pool, shared_id).await; + cleanup_user(&pool, alice_id).await; +} + +#[tokio::test] +async fn test_multi_scope_list_directory_merges() { + let pool = get_pool(); + if try_connect(&pool).await.is_none() { + return; + } + let shared_id = "ms_shared_dir"; + let alice_id = "ms_alice_dir"; + cleanup_user(&pool, shared_id).await; + cleanup_user(&pool, alice_id).await; + + let ws_alice_plain = Workspace::new(alice_id, pool.clone()); + ws_alice_plain + .write("docs/alice-doc.md", "Alice's doc") + .await + .expect("alice write failed"); + + let ws_shared = Workspace::new(shared_id, pool.clone()); + ws_shared + .write("docs/shared-doc.md", "Shared doc") + .await + .expect("shared write failed"); + + let ws_alice = Workspace::new(alice_id, pool.clone()) + .with_additional_read_scopes(vec![shared_id.to_string()]); + + let entries = ws_alice.list("docs").await.expect("list failed"); + let paths: Vec<&str> = entries.iter().map(|e| e.path.as_str()).collect(); + assert!( + paths.contains(&"docs/alice-doc.md"), + "Should contain alice's doc: {:?}", + paths + ); + assert!( + paths.contains(&"docs/shared-doc.md"), + "Should contain shared doc: {:?}", + paths + ); + + cleanup_user(&pool, shared_id).await; + cleanup_user(&pool, alice_id).await; +} + +#[tokio::test] +async fn test_multi_scope_read_priority_primary_first() { + let pool = get_pool(); + if try_connect(&pool).await.is_none() { + return; + } + let shared_id = "ms_shared_prio"; + let alice_id = "ms_alice_prio"; + cleanup_user(&pool, shared_id).await; + cleanup_user(&pool, alice_id).await; + + // Write same path in both scopes + let ws_shared = Workspace::new(shared_id, pool.clone()); + ws_shared + .write("config/settings.md", "Shared settings v1") + .await + .expect("shared write failed"); + + let ws_alice_plain = Workspace::new(alice_id, pool.clone()); + ws_alice_plain + .write("config/settings.md", "Alice's settings override") + .await + .expect("alice write failed"); + + // Alice with multi-scope should get her own version (primary scope wins) + let ws_alice = Workspace::new(alice_id, pool.clone()) + .with_additional_read_scopes(vec![shared_id.to_string()]); + + let doc = ws_alice + .read("config/settings.md") + .await + .expect("read failed"); + assert_eq!( + doc.content, "Alice's settings override", + "Primary scope should take priority" + ); + + cleanup_user(&pool, shared_id).await; + cleanup_user(&pool, alice_id).await; +} + +#[tokio::test] +async fn test_multi_scope_exists_spans_scopes() { + let pool = get_pool(); + if try_connect(&pool).await.is_none() { + return; + } + let shared_id = "ms_shared_exists"; + let alice_id = "ms_alice_exists"; + cleanup_user(&pool, shared_id).await; + cleanup_user(&pool, alice_id).await; + + let ws_shared = Workspace::new(shared_id, pool.clone()); + ws_shared + .write("docs/shared-only.md", "Shared content") + .await + .expect("shared write failed"); + + // Alice without multi-scope should NOT see it + let ws_alice_plain = Workspace::new(alice_id, pool.clone()); + assert!( + !ws_alice_plain + .exists("docs/shared-only.md") + .await + .expect("exists failed"), + "Alice without multi-scope should not see shared doc" + ); + + // Alice with multi-scope should see it + let ws_alice = Workspace::new(alice_id, pool.clone()) + .with_additional_read_scopes(vec![shared_id.to_string()]); + assert!( + ws_alice + .exists("docs/shared-only.md") + .await + .expect("exists failed"), + "Alice with multi-scope should see shared doc" + ); + + cleanup_user(&pool, shared_id).await; + cleanup_user(&pool, alice_id).await; +} + +#[tokio::test] +async fn test_multi_scope_search_spans_scopes() { + let pool = get_pool(); + if try_connect(&pool).await.is_none() { + return; + } + let shared_id = "ms_shared_search"; + let alice_id = "ms_alice_search"; + cleanup_user(&pool, shared_id).await; + cleanup_user(&pool, alice_id).await; + + let ws_shared = Workspace::new(shared_id, pool.clone()); + ws_shared + .write( + "docs/architecture.md", + "The microservice architecture uses gRPC for inter-service communication", + ) + .await + .expect("shared write failed"); + + let ws_alice_plain = Workspace::new(alice_id, pool.clone()); + ws_alice_plain + .write("notes/ideas.md", "Consider switching to GraphQL federation") + .await + .expect("alice write failed"); + + let ws_alice = Workspace::new(alice_id, pool.clone()) + .with_additional_read_scopes(vec![shared_id.to_string()]); + + // Search for content in the shared scope + let results = ws_alice + .search_with_config( + "microservice gRPC architecture", + SearchConfig::default().fts_only(), + ) + .await + .expect("search failed"); + assert!(!results.is_empty(), "Should find results from shared scope"); + + cleanup_user(&pool, shared_id).await; + cleanup_user(&pool, alice_id).await; +} + +#[tokio::test] +async fn test_multi_scope_append_stays_in_primary() { + let pool = get_pool(); + if try_connect(&pool).await.is_none() { + return; + } + let shared_id = "ms_shared_append"; + let alice_id = "ms_alice_append"; + cleanup_user(&pool, shared_id).await; + cleanup_user(&pool, alice_id).await; + + // Write a document as "shared" + let ws_shared = Workspace::new(shared_id, pool.clone()); + ws_shared + .write("notes/log.md", "shared original content") + .await + .expect("shared write failed"); + + // Alice has "shared" as a read scope and appends to the same path + let ws_alice = Workspace::new(alice_id, pool.clone()) + .with_additional_read_scopes(vec![shared_id.to_string()]); + ws_alice + .append("notes/log.md", "alice appended line") + .await + .expect("alice append failed"); + + // Shared document must be unchanged (write isolation) + let shared_doc = ws_shared + .read("notes/log.md") + .await + .expect("shared read failed"); + assert_eq!( + shared_doc.content, "shared original content", + "Append must not modify the secondary scope's document" + ); + + // Alice should have her own copy with the appended content + let ws_alice_plain = Workspace::new(alice_id, pool.clone()); + let alice_doc = ws_alice_plain + .read("notes/log.md") + .await + .expect("alice read failed"); + assert_eq!( + alice_doc.content, "alice appended line", + "Append should create a new document in alice's scope" + ); + + cleanup_user(&pool, shared_id).await; + cleanup_user(&pool, alice_id).await; +} 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"); 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": [