diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 00488c70..5d4eabc0 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -12,6 +12,7 @@ jobs: tests: name: Tests (${{ matrix.name }}) runs-on: ubuntu-latest + timeout-minutes: 45 strategy: fail-fast: false matrix: @@ -40,11 +41,14 @@ jobs: - name: Build WASM channels (for integration tests) run: ./scripts/build-wasm-extensions.sh --channels - name: Run Tests - run: cargo test ${{ matrix.flags }} -- --nocapture + run: | + timeout --signal=INT --kill-after=30s 40m \ + cargo test ${{ matrix.flags }} -- --nocapture heavy-integration-tests: name: Heavy Integration Tests runs-on: ubuntu-latest + timeout-minutes: 20 steps: - name: Checkout repository uses: actions/checkout@v6 @@ -58,9 +62,13 @@ jobs: - name: Build Telegram WASM channel run: cargo build --manifest-path channels-src/telegram/Cargo.toml --target wasm32-wasip2 --release - name: Run thread scheduling integration tests - run: cargo test --no-default-features --features libsql,integration --test e2e_thread_scheduling -- --nocapture + run: | + timeout --signal=INT --kill-after=30s 15m \ + cargo test --no-default-features --features libsql,integration --test e2e_thread_scheduling -- --nocapture - name: Run Telegram thread-scope regression test - run: cargo test --features integration --test telegram_auth_integration test_private_messages_use_chat_id_as_thread_scope -- --exact + run: | + timeout --signal=INT --kill-after=30s 10m \ + cargo test --features integration --test telegram_auth_integration test_private_messages_use_chat_id_as_thread_scope -- --exact telegram-tests: name: Telegram Channel Tests @@ -68,6 +76,7 @@ jobs: github.event_name != 'pull_request' || github.base_ref != 'staging' runs-on: ubuntu-latest + timeout-minutes: 15 steps: - name: Checkout repository uses: actions/checkout@v6 @@ -75,7 +84,9 @@ jobs: uses: dtolnay/rust-toolchain@stable - uses: Swatinem/rust-cache@v2 - name: Run Telegram Channel Tests - run: cargo test --manifest-path channels-src/telegram/Cargo.toml -- --nocapture + run: | + timeout --signal=INT --kill-after=30s 10m \ + cargo test --manifest-path channels-src/telegram/Cargo.toml -- --nocapture windows-build: name: Windows Build (${{ matrix.name }}) @@ -110,6 +121,7 @@ jobs: github.event_name != 'pull_request' || github.base_ref != 'staging' runs-on: ubuntu-latest + timeout-minutes: 30 steps: - name: Checkout repository uses: actions/checkout@v6 @@ -125,7 +137,9 @@ jobs: - name: Build all WASM extensions against current WIT run: ./scripts/build-wasm-extensions.sh - name: Instantiation test (host linker compatibility) - run: cargo test --all-features wit_compat -- --nocapture + run: | + timeout --signal=INT --kill-after=30s 20m \ + cargo test --all-features wit_compat -- --nocapture bench-compile: name: Benchmark Compilation diff --git a/CHANGELOG.md b/CHANGELOG.md index 6aad4993..9acc56ad 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,138 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +## [0.22.0](https://github.com/nearai/ironclaw/compare/ironclaw-v0.21.0...ironclaw-v0.22.0) - 2026-03-25 + +### Added + +- *(agent)* thread per-tool reasoning through provider, session, and all surfaces ([#1513](https://github.com/nearai/ironclaw/pull/1513)) +- *(cli)* show credential auth status in tool info ([#1572](https://github.com/nearai/ironclaw/pull/1572)) +- multi-tenant auth with per-user workspace isolation ([#1118](https://github.com/nearai/ironclaw/pull/1118)) +- *(cli)* add ironclaw models subcommands (list/status/set/set-provider) ([#1043](https://github.com/nearai/ironclaw/pull/1043)) +- *(workspace)* multi-scope workspace reads ([#1117](https://github.com/nearai/ironclaw/pull/1117)) +- *(ux)* complete UX overhaul — design system, onboarding, web polish ([#1277](https://github.com/nearai/ironclaw/pull/1277)) +- *(gemini_oauth)* full Gemini CLI OAuth integration with Cloud Code API ([#1356](https://github.com/nearai/ironclaw/pull/1356)) +- *(shell)* add Low/Medium/High risk levels for graduated command approval (closes #172) ([#368](https://github.com/nearai/ironclaw/pull/368)) +- *(agent)* queue and merge messages during active turns ([#1412](https://github.com/nearai/ironclaw/pull/1412)) +- *(cli)* add `ironclaw hooks list` subcommand ([#1023](https://github.com/nearai/ironclaw/pull/1023)) +- *(extensions)* support text setup fields in web configure modal ([#496](https://github.com/nearai/ironclaw/pull/496)) +- *(llm)* add GitHub Copilot as LLM provider ([#1512](https://github.com/nearai/ironclaw/pull/1512)) +- *(workspace)* layered memory with sensitivity-based privacy redirect ([#1112](https://github.com/nearai/ironclaw/pull/1112)) +- *(webhooks)* add public webhook trigger endpoint for routines ([#736](https://github.com/nearai/ironclaw/pull/736)) +- *(llm)* Add OpenAI Codex (ChatGPT subscription) as LLM provider ([#1461](https://github.com/nearai/ironclaw/pull/1461)) +- *(web)* add light theme with dark/light/system toggle ([#1457](https://github.com/nearai/ironclaw/pull/1457)) +- *(agent)* activate stuck_threshold for time-based stuck job detection ([#1234](https://github.com/nearai/ironclaw/pull/1234)) +- chat onboarding and routine advisor ([#927](https://github.com/nearai/ironclaw/pull/927)) + +### Fixed + +- ensure LLM calls always end with user message (closes #763) ([#1259](https://github.com/nearai/ironclaw/pull/1259)) +- restore owner-scoped gateway startup ([#1625](https://github.com/nearai/ironclaw/pull/1625)) +- remove stale stream_token gate from channel-relay activation ([#1623](https://github.com/nearai/ironclaw/pull/1623)) +- *(agent)* case-insensitive channel match and user_id filter for event triggers ([#1211](https://github.com/nearai/ironclaw/pull/1211)) +- *(routines)* normalize status display across web and CLI ([#1469](https://github.com/nearai/ironclaw/pull/1469)) +- *(tunnel)* managed tunnels target wrong port and die from SIGPIPE ([#1093](https://github.com/nearai/ironclaw/pull/1093)) +- *(agent)* persist /model selection to .env, TOML, and DB ([#1581](https://github.com/nearai/ironclaw/pull/1581)) +- post-merge review sweep — 8 fixes across security, perf, and correctness ([#1550](https://github.com/nearai/ironclaw/pull/1550)) +- generate Mistral-compatible 9-char alphanumeric tool call IDs ([#1242](https://github.com/nearai/ironclaw/pull/1242)) +- *(mcp)* handle empty 202 notification acknowledgements ([#1539](https://github.com/nearai/ironclaw/pull/1539)) +- *(tests)* eliminate env mutex poison cascade ([#1558](https://github.com/nearai/ironclaw/pull/1558)) +- *(safety)* escape tool output XML content and remove misleading sanitized attr ([#1067](https://github.com/nearai/ironclaw/pull/1067)) +- *(oauth)* reject malformed ic2.* states in decode_hosted_oauth_state ([#1441](https://github.com/nearai/ironclaw/pull/1441)) ([#1454](https://github.com/nearai/ironclaw/pull/1454)) +- parameter coercion and validation for oneOf/anyOf/allOf schemas ([#1397](https://github.com/nearai/ironclaw/pull/1397)) +- persist startup-loaded MCP clients in ExtensionManager ([#1509](https://github.com/nearai/ironclaw/pull/1509)) +- *(deps)* patch rustls-webpki vulnerability (RUSTSEC-2026-0049) +- *(routines)* add missing extension_manager field in trigger_manual EngineContext +- *(ci)* serialize env-mutating OAuth wildcard tests with ENV_MUTEX ([#1280](https://github.com/nearai/ironclaw/pull/1280)) ([#1468](https://github.com/nearai/ironclaw/pull/1468)) +- *(setup)* remove redundant LLM config and API keys from bootstrap .env ([#1448](https://github.com/nearai/ironclaw/pull/1448)) +- resolve wasm broadcast merge conflicts with staging ([#395](https://github.com/nearai/ironclaw/pull/395)) ([#1460](https://github.com/nearai/ironclaw/pull/1460)) +- skip credential validation for Bedrock backend ([#1011](https://github.com/nearai/ironclaw/pull/1011)) +- register sandbox jobs in ContextManager for query tool visibility ([#1426](https://github.com/nearai/ironclaw/pull/1426)) +- prefer execution-local message routing metadata ([#1449](https://github.com/nearai/ironclaw/pull/1449)) +- *(security)* validate embedding base URLs to prevent SSRF ([#1221](https://github.com/nearai/ironclaw/pull/1221)) +- f32→f64 precision artifact in temperature causes provider 400 errors ([#1450](https://github.com/nearai/ironclaw/pull/1450)) +- *(routines)* surface errors when sandbox unavailable for full_job routines ([#769](https://github.com/nearai/ironclaw/pull/769)) +- restore libSQL vector search with dynamic dimensions ([#1393](https://github.com/nearai/ironclaw/pull/1393)) +- staging CI triage — consolidate retry parsing, fix flaky tests, add docs ([#1427](https://github.com/nearai/ironclaw/pull/1427)) + +### Other + +- Merge branch 'main' into staging-promote/455f543b-23329172268 +- Merge pull request #1655 from nearai/codex/fix-staging-promotion-1451-version-bumps +- Merge pull request #1499 from nearai/staging-promote/9603fefd-23364438978 +- Fix libsql prompt scope regressions ([#1651](https://github.com/nearai/ironclaw/pull/1651)) +- Normalize cron schedules on routine create ([#1648](https://github.com/nearai/ironclaw/pull/1648)) +- Fix MCP lifecycle trace user scope ([#1646](https://github.com/nearai/ironclaw/pull/1646)) +- Fix REPL single-message hang and cap CI test duration ([#1643](https://github.com/nearai/ironclaw/pull/1643)) +- extract AppEvent to crates/ironclaw_common ([#1615](https://github.com/nearai/ironclaw/pull/1615)) +- Fix hosted OAuth refresh via proxy ([#1602](https://github.com/nearai/ironclaw/pull/1602)) +- *(agent)* optimize approval thread resolution (UUID parsing + lock contention) ([#1592](https://github.com/nearai/ironclaw/pull/1592)) +- *(tools)* auto-compact WASM tool schemas, add descriptions, improve credential prompts ([#1525](https://github.com/nearai/ironclaw/pull/1525)) +- Default new lightweight routines to tools-enabled ([#1573](https://github.com/nearai/ironclaw/pull/1573)) +- Google OAuth URL broken when initiated from Telegram channel ([#1165](https://github.com/nearai/ironclaw/pull/1165)) +- add gitcgr code graph badge ([#1563](https://github.com/nearai/ironclaw/pull/1563)) +- Fix owner-scoped message routing fallbacks ([#1574](https://github.com/nearai/ironclaw/pull/1574)) +- *(tools)* remove unconditional params clone in shared execution (fix #893) ([#926](https://github.com/nearai/ironclaw/pull/926)) +- *(llm)* move transcription module into src/llm/ ([#1559](https://github.com/nearai/ironclaw/pull/1559)) +- *(agent)* avoid preview allocations for non-truncated strings (fix #894) ([#924](https://github.com/nearai/ironclaw/pull/924)) +- Expand AGENTS.md with coding agents guidance ([#1392](https://github.com/nearai/ironclaw/pull/1392)) +- Fix CI approval flows and stale fixtures ([#1478](https://github.com/nearai/ironclaw/pull/1478)) +- Use live owner tool scope for autonomous routines and jobs ([#1453](https://github.com/nearai/ironclaw/pull/1453)) +- use Arc in embedding cache to avoid clones on miss path ([#1438](https://github.com/nearai/ironclaw/pull/1438)) +- Add owner-scoped permissions for full-job routines ([#1440](https://github.com/nearai/ironclaw/pull/1440)) + +## [0.21.0](https://github.com/nearai/ironclaw/compare/v0.20.0...v0.21.0) - 2026-03-20 + +### Added + +- structured fallback deliverables for failed/stuck jobs ([#236](https://github.com/nearai/ironclaw/pull/236)) +- LRU embedding cache for workspace search ([#1423](https://github.com/nearai/ironclaw/pull/1423)) +- receive relay events via webhook callbacks ([#1254](https://github.com/nearai/ironclaw/pull/1254)) + +### Fixed + +- bump Feishu channel version for promotion +- *(approval)* make "always" auto-approve work for credentialed HTTP requests ([#1257](https://github.com/nearai/ironclaw/pull/1257)) +- skip NEAR AI session check when backend is not nearai ([#1413](https://github.com/nearai/ironclaw/pull/1413)) + +### Other + +- Make hosted OAuth and MCP auth generic ([#1375](https://github.com/nearai/ironclaw/pull/1375)) + +## [0.20.0](https://github.com/nearai/ironclaw/compare/v0.19.0...v0.20.0) - 2026-03-19 + +### Added + +- *(self-repair)* wire stuck_threshold, store, and builder ([#712](https://github.com/nearai/ironclaw/pull/712)) +- *(testing)* add FaultInjector framework for StubLlm ([#1233](https://github.com/nearai/ironclaw/pull/1233)) +- *(gateway)* unified settings page with subtabs ([#1191](https://github.com/nearai/ironclaw/pull/1191)) +- upgrade MiniMax default model to M2.7 ([#1357](https://github.com/nearai/ironclaw/pull/1357)) + +### Fixed + +- navigate telegram E2E tests to channels subtab ([#1408](https://github.com/nearai/ironclaw/pull/1408)) +- add missing `builder` field and update E2E extensions tab navigation ([#1400](https://github.com/nearai/ironclaw/pull/1400)) +- remove debug_assert guards that panic on valid error paths ([#1385](https://github.com/nearai/ironclaw/pull/1385)) +- address valid review comments from PR #1359 ([#1380](https://github.com/nearai/ironclaw/pull/1380)) +- full_job routine runs stay running until linked job completion ([#1374](https://github.com/nearai/ironclaw/pull/1374)) +- full_job routine concurrency tracks linked job lifetime ([#1372](https://github.com/nearai/ironclaw/pull/1372)) +- remove -x from coverage pytest to prevent suite-blocking failures ([#1360](https://github.com/nearai/ironclaw/pull/1360)) +- add debug_assert invariant guards to critical code paths ([#1312](https://github.com/nearai/ironclaw/pull/1312)) +- *(mcp)* retry after missing session id errors ([#1355](https://github.com/nearai/ironclaw/pull/1355)) +- *(telegram)* preserve polling after secret-blocked updates ([#1353](https://github.com/nearai/ironclaw/pull/1353)) +- *(llm)* cap retry-after delays ([#1351](https://github.com/nearai/ironclaw/pull/1351)) +- *(setup)* remove nonexistent webhook secret command hint ([#1349](https://github.com/nearai/ironclaw/pull/1349)) +- Rate limiter returns retry after None instead of a duration ([#1269](https://github.com/nearai/ironclaw/pull/1269)) + +### Other + +- bump telegram channel version to 0.2.5 ([#1410](https://github.com/nearai/ironclaw/pull/1410)) +- *(ci)* enforce test requirement for state machine and resilience changes ([#1230](https://github.com/nearai/ironclaw/pull/1230)) ([#1304](https://github.com/nearai/ironclaw/pull/1304)) +- Fix duplicate LLM responses for matched event routines ([#1275](https://github.com/nearai/ironclaw/pull/1275)) +- add Japanese README ([#1306](https://github.com/nearai/ironclaw/pull/1306)) +- *(ci)* add coverage gates via codecov.yml ([#1228](https://github.com/nearai/ironclaw/pull/1228)) ([#1291](https://github.com/nearai/ironclaw/pull/1291)) +- Redesign routine create requests for LLMs ([#1147](https://github.com/nearai/ironclaw/pull/1147)) + ## [0.19.0](https://github.com/nearai/ironclaw/compare/v0.18.0...v0.19.0) - 2026-03-17 ### Added diff --git a/Cargo.lock b/Cargo.lock index a813ef2b..c3747590 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2323,7 +2323,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" dependencies = [ "libc", - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] @@ -3390,7 +3390,7 @@ dependencies = [ [[package]] name = "ironclaw" -version = "0.19.0" +version = "0.22.0" dependencies = [ "aes-gcm", "aho-corasick", @@ -3428,6 +3428,7 @@ dependencies = [ "hyper-util", "iana-time-zone", "insta", + "ironclaw_common", "ironclaw_safety", "json5", "libsql", @@ -3486,8 +3487,16 @@ dependencies = [ ] [[package]] -name = "ironclaw_safety" +name = "ironclaw_common" version = "0.1.0" +dependencies = [ + "serde", + "serde_json", +] + +[[package]] +name = "ironclaw_safety" +version = "0.2.0" dependencies = [ "aho-corasick", "regex", @@ -5472,7 +5481,7 @@ dependencies = [ "errno", "libc", "linux-raw-sys 0.12.1", - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] @@ -6379,7 +6388,7 @@ dependencies = [ "getrandom 0.4.2", "once_cell", "rustix 1.1.4", - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index 99992a40..41895b16 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,5 +1,5 @@ [workspace] -members = [".", "crates/ironclaw_safety"] +members = [".", "crates/ironclaw_common", "crates/ironclaw_safety"] exclude = [ "channels-src/discord", "channels-src/telegram", @@ -20,7 +20,7 @@ exclude = [ [package] name = "ironclaw" -version = "0.19.0" +version = "0.22.0" edition = "2024" rust-version = "1.92" description = "Secure personal AI assistant that protects your data and expands its capabilities on the fly" @@ -100,8 +100,11 @@ tower-http = { version = "0.6", features = ["trace", "cors", "set-header"] } # Cron scheduling for routines cron = "0.13" +# Shared types +ironclaw_common = { path = "crates/ironclaw_common", version = "0.1.0" } + # Safety/sanitization -ironclaw_safety = { path = "crates/ironclaw_safety", version = "0.1.0" } +ironclaw_safety = { path = "crates/ironclaw_safety", version = "0.2.0" } regex = "1" aho-corasick = "1" diff --git a/crates/ironclaw_common/Cargo.toml b/crates/ironclaw_common/Cargo.toml new file mode 100644 index 00000000..6e7db5a4 --- /dev/null +++ b/crates/ironclaw_common/Cargo.toml @@ -0,0 +1,17 @@ +[package] +name = "ironclaw_common" +version = "0.1.0" +edition = "2024" +rust-version = "1.92" +description = "Shared types and utilities for the IronClaw workspace" +authors = ["NEAR AI "] +license = "MIT OR Apache-2.0" +homepage = "https://github.com/nearai/ironclaw" +repository = "https://github.com/nearai/ironclaw" + +[package.metadata.dist] +dist = false + +[dependencies] +serde = { version = "1", features = ["derive"] } +serde_json = "1" diff --git a/crates/ironclaw_common/src/event.rs b/crates/ironclaw_common/src/event.rs new file mode 100644 index 00000000..256aba3d --- /dev/null +++ b/crates/ironclaw_common/src/event.rs @@ -0,0 +1,393 @@ +//! Application-wide event types. +//! +//! `AppEvent` is the real-time event protocol used across the entire +//! application. The web gateway serialises these to SSE / WebSocket +//! frames, but other subsystems (agent loop, orchestrator, extensions) +//! produce and consume them too. + +use serde::{Deserialize, Serialize}; + +/// A single tool decision in a reasoning update (SSE DTO). +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ToolDecisionDto { + pub tool_name: String, + pub rationale: String, +} + +impl ToolDecisionDto { + /// Parse a list of tool decisions from a JSON array value. + pub fn from_json_array(value: &serde_json::Value) -> Vec { + value + .as_array() + .map(|arr| { + arr.iter() + .filter_map(|d| { + Some(Self { + tool_name: d.get("tool_name")?.as_str()?.to_string(), + rationale: d.get("rationale")?.as_str()?.to_string(), + }) + }) + .collect() + }) + .unwrap_or_default() + } +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "type")] +pub enum AppEvent { + #[serde(rename = "response")] + Response { content: String, thread_id: String }, + #[serde(rename = "thinking")] + Thinking { + message: String, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + #[serde(rename = "tool_started")] + ToolStarted { + name: String, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + #[serde(rename = "tool_completed")] + ToolCompleted { + name: String, + success: bool, + #[serde(skip_serializing_if = "Option::is_none")] + error: Option, + #[serde(skip_serializing_if = "Option::is_none")] + parameters: Option, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + #[serde(rename = "tool_result")] + ToolResult { + name: String, + preview: String, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + #[serde(rename = "stream_chunk")] + StreamChunk { + content: String, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + #[serde(rename = "status")] + Status { + message: String, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + #[serde(rename = "job_started")] + JobStarted { + job_id: String, + title: String, + browse_url: String, + }, + #[serde(rename = "approval_needed")] + ApprovalNeeded { + request_id: String, + tool_name: String, + description: String, + parameters: String, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + /// Whether the "always" auto-approve option should be shown. + allow_always: bool, + }, + #[serde(rename = "auth_required")] + AuthRequired { + extension_name: String, + #[serde(skip_serializing_if = "Option::is_none")] + instructions: Option, + #[serde(skip_serializing_if = "Option::is_none")] + auth_url: Option, + #[serde(skip_serializing_if = "Option::is_none")] + setup_url: Option, + }, + #[serde(rename = "auth_completed")] + AuthCompleted { + extension_name: String, + success: bool, + message: String, + }, + #[serde(rename = "error")] + Error { + message: String, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + #[serde(rename = "heartbeat")] + Heartbeat, + + // Sandbox job streaming events (worker + Claude Code bridge) + #[serde(rename = "job_message")] + JobMessage { + job_id: String, + role: String, + content: String, + }, + #[serde(rename = "job_tool_use")] + JobToolUse { + job_id: String, + tool_name: String, + input: serde_json::Value, + }, + #[serde(rename = "job_tool_result")] + JobToolResult { + job_id: String, + tool_name: String, + output: String, + }, + #[serde(rename = "job_status")] + JobStatus { job_id: String, message: String }, + #[serde(rename = "job_result")] + JobResult { + job_id: String, + status: String, + #[serde(skip_serializing_if = "Option::is_none")] + session_id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + fallback_deliverable: Option, + }, + + /// An image was generated by a tool. + #[serde(rename = "image_generated")] + ImageGenerated { + data_url: String, + #[serde(skip_serializing_if = "Option::is_none")] + path: Option, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + + /// Suggested follow-up messages for the user. + #[serde(rename = "suggestions")] + Suggestions { + suggestions: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + + /// Per-turn token usage and cost summary. + #[serde(rename = "turn_cost")] + TurnCost { + input_tokens: u64, + output_tokens: u64, + cost_usd: String, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + + /// Extension activation status change (WASM channels). + #[serde(rename = "extension_status")] + ExtensionStatus { + extension_name: String, + status: String, + #[serde(skip_serializing_if = "Option::is_none")] + message: Option, + }, + + /// Agent reasoning update (why it chose specific tools). + #[serde(rename = "reasoning_update")] + ReasoningUpdate { + narrative: String, + decisions: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + thread_id: Option, + }, + + /// Reasoning update for a sandbox job. + #[serde(rename = "job_reasoning")] + JobReasoning { + job_id: String, + narrative: String, + decisions: Vec, + }, +} + +impl AppEvent { + /// The wire-format event type string (matches the `#[serde(rename)]` value). + pub fn event_type(&self) -> &'static str { + match self { + Self::Response { .. } => "response", + Self::Thinking { .. } => "thinking", + Self::ToolStarted { .. } => "tool_started", + Self::ToolCompleted { .. } => "tool_completed", + Self::ToolResult { .. } => "tool_result", + Self::StreamChunk { .. } => "stream_chunk", + Self::Status { .. } => "status", + Self::JobStarted { .. } => "job_started", + Self::ApprovalNeeded { .. } => "approval_needed", + Self::AuthRequired { .. } => "auth_required", + Self::AuthCompleted { .. } => "auth_completed", + Self::Error { .. } => "error", + Self::Heartbeat => "heartbeat", + Self::JobMessage { .. } => "job_message", + Self::JobToolUse { .. } => "job_tool_use", + Self::JobToolResult { .. } => "job_tool_result", + Self::JobStatus { .. } => "job_status", + Self::JobResult { .. } => "job_result", + Self::ImageGenerated { .. } => "image_generated", + Self::Suggestions { .. } => "suggestions", + Self::TurnCost { .. } => "turn_cost", + Self::ExtensionStatus { .. } => "extension_status", + Self::ReasoningUpdate { .. } => "reasoning_update", + Self::JobReasoning { .. } => "job_reasoning", + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + /// Verify that `event_type()` returns the same string as the serde + /// `"type"` field for every variant. This catches drift between the + /// `#[serde(rename)]` attributes and the manual match arms. + #[test] + fn event_type_matches_serde_type_field() { + let variants: Vec = vec![ + AppEvent::Response { + content: String::new(), + thread_id: String::new(), + }, + AppEvent::Thinking { + message: String::new(), + thread_id: None, + }, + AppEvent::ToolStarted { + name: String::new(), + thread_id: None, + }, + AppEvent::ToolCompleted { + name: String::new(), + success: true, + error: None, + parameters: None, + thread_id: None, + }, + AppEvent::ToolResult { + name: String::new(), + preview: String::new(), + thread_id: None, + }, + AppEvent::StreamChunk { + content: String::new(), + thread_id: None, + }, + AppEvent::Status { + message: String::new(), + thread_id: None, + }, + AppEvent::JobStarted { + job_id: String::new(), + title: String::new(), + browse_url: String::new(), + }, + AppEvent::ApprovalNeeded { + request_id: String::new(), + tool_name: String::new(), + description: String::new(), + parameters: String::new(), + thread_id: None, + allow_always: false, + }, + AppEvent::AuthRequired { + extension_name: String::new(), + instructions: None, + auth_url: None, + setup_url: None, + }, + AppEvent::AuthCompleted { + extension_name: String::new(), + success: true, + message: String::new(), + }, + AppEvent::Error { + message: String::new(), + thread_id: None, + }, + AppEvent::Heartbeat, + AppEvent::JobMessage { + job_id: String::new(), + role: String::new(), + content: String::new(), + }, + AppEvent::JobToolUse { + job_id: String::new(), + tool_name: String::new(), + input: serde_json::Value::Null, + }, + AppEvent::JobToolResult { + job_id: String::new(), + tool_name: String::new(), + output: String::new(), + }, + AppEvent::JobStatus { + job_id: String::new(), + message: String::new(), + }, + AppEvent::JobResult { + job_id: String::new(), + status: String::new(), + session_id: None, + fallback_deliverable: None, + }, + AppEvent::ImageGenerated { + data_url: String::new(), + path: None, + thread_id: None, + }, + AppEvent::Suggestions { + suggestions: vec![], + thread_id: None, + }, + AppEvent::TurnCost { + input_tokens: 0, + output_tokens: 0, + cost_usd: String::new(), + thread_id: None, + }, + AppEvent::ExtensionStatus { + extension_name: String::new(), + status: String::new(), + message: None, + }, + AppEvent::ReasoningUpdate { + narrative: String::new(), + decisions: vec![], + thread_id: None, + }, + AppEvent::JobReasoning { + job_id: String::new(), + narrative: String::new(), + decisions: vec![], + }, + ]; + + for variant in &variants { + let json: serde_json::Value = serde_json::to_value(variant).unwrap(); + let serde_type = json["type"].as_str().unwrap(); + assert_eq!( + variant.event_type(), + serde_type, + "event_type() mismatch for variant: {:?}", + variant + ); + } + } + + #[test] + fn round_trip_deserialize() { + let original = AppEvent::Response { + content: "hello".to_string(), + thread_id: "t1".to_string(), + }; + let json = serde_json::to_string(&original).unwrap(); + let deserialized: AppEvent = serde_json::from_str(&json).unwrap(); + assert_eq!(deserialized.event_type(), "response"); + } +} diff --git a/crates/ironclaw_common/src/lib.rs b/crates/ironclaw_common/src/lib.rs new file mode 100644 index 00000000..f52dc0aa --- /dev/null +++ b/crates/ironclaw_common/src/lib.rs @@ -0,0 +1,7 @@ +//! Shared types and utilities for the IronClaw workspace. + +mod event; +mod util; + +pub use event::{AppEvent, ToolDecisionDto}; +pub use util::truncate_preview; diff --git a/crates/ironclaw_common/src/util.rs b/crates/ironclaw_common/src/util.rs new file mode 100644 index 00000000..4f054671 --- /dev/null +++ b/crates/ironclaw_common/src/util.rs @@ -0,0 +1,100 @@ +//! Shared utility functions. + +/// Truncate a string to at most `max_bytes` bytes at a char boundary, appending "...". +/// +/// If the input is wrapped in `...` and truncation +/// removes the closing tag, the tag is re-appended so downstream XML parsers +/// never see an unclosed element. +pub fn truncate_preview(s: &str, max_bytes: usize) -> String { + if s.len() <= max_bytes { + return s.to_string(); + } + // Walk backwards from max_bytes to find a valid char boundary + let mut end = max_bytes; + while end > 0 && !s.is_char_boundary(end) { + end -= 1; + } + let mut result = format!("{}...", &s[..end]); + + // Re-close if truncation cut through the closing tag. + if s.starts_with("") { + result.push_str("\n"); + } + + result +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_truncate_preview_short_string() { + assert_eq!(truncate_preview("hello", 10), "hello"); + } + + #[test] + fn test_truncate_preview_exact_boundary() { + assert_eq!(truncate_preview("hello", 5), "hello"); + } + + #[test] + fn test_truncate_preview_truncates_ascii() { + assert_eq!(truncate_preview("hello world", 5), "hello..."); + } + + #[test] + fn test_truncate_preview_empty_string() { + assert_eq!(truncate_preview("", 10), ""); + } + + #[test] + fn test_truncate_preview_multibyte_char_boundary() { + let s = "a\u{20AC}b"; + let result = truncate_preview(s, 3); + assert_eq!(result, "a..."); + } + + #[test] + fn test_truncate_preview_emoji() { + let s = "hi\u{1F980}"; + let result = truncate_preview(s, 4); + assert_eq!(result, "hi..."); + } + + #[test] + fn test_truncate_preview_cjk() { + let s = "\u{4F60}\u{597D}\u{4E16}\u{754C}"; + let result = truncate_preview(s, 7); + assert_eq!(result, "\u{4F60}\u{597D}..."); + } + + #[test] + fn test_truncate_preview_zero_max_bytes() { + assert_eq!(truncate_preview("hello", 0), "..."); + } + + #[test] + fn test_truncate_preview_closes_tool_output_tag() { + let s = "\nSome very long content here\n"; + let result = truncate_preview(s, 60); + assert!(result.ends_with("")); + assert!(result.contains("...")); + } + + #[test] + fn test_truncate_preview_no_extra_close_when_intact() { + let s = "\nshort\n"; + let result = truncate_preview(s, 500); + assert_eq!(result, s); + assert_eq!(result.matches("").count(), 1); + } + + #[test] + fn test_truncate_preview_non_xml_unaffected() { + let s = "Just a plain long string that gets truncated"; + let result = truncate_preview(s, 10); + assert_eq!(result, "Just a pla..."); + assert!(!result.contains("")); + } +} diff --git a/crates/ironclaw_safety/Cargo.toml b/crates/ironclaw_safety/Cargo.toml index d12aa909..38b8718a 100644 --- a/crates/ironclaw_safety/Cargo.toml +++ b/crates/ironclaw_safety/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "ironclaw_safety" -version = "0.1.0" +version = "0.2.0" edition = "2024" rust-version = "1.92" description = "Prompt injection defense, input validation, secret leak detection, and safety policy enforcement" @@ -8,7 +8,6 @@ authors = ["NEAR AI "] license = "MIT OR Apache-2.0" homepage = "https://github.com/nearai/ironclaw" repository = "https://github.com/nearai/ironclaw" -publish = false [package.metadata.dist] dist = false diff --git a/registry/channels/feishu.json b/registry/channels/feishu.json index 66cecf1d..a7530943 100644 --- a/registry/channels/feishu.json +++ b/registry/channels/feishu.json @@ -2,7 +2,7 @@ "name": "feishu", "display_name": "Feishu / Lark Channel", "kind": "channel", - "version": "0.1.1", + "version": "0.1.3", "wit_version": "0.3.0", "description": "Talk to your agent through a Feishu or Lark bot", "keywords": [ @@ -19,8 +19,8 @@ }, "artifacts": { "wasm32-wasip2": { - "sha256": "5fca74022264d1c8e78a0853766276f7ffa3cf0d8065b2f51ca10985acad4714", - "url": "https://github.com/nearai/ironclaw/releases/download/v0.19.0/channel-feishu-0.1.1-wasm32-wasip2.tar.gz" + "sha256": "a66ff0dafb67d2216d8161bb7e96e724a94acb0ab993b85d2782d30412f8fe94", + "url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/channel-feishu-0.1.3-wasm32-wasip2.tar.gz" } }, "auth_summary": { diff --git a/registry/channels/telegram.json b/registry/channels/telegram.json index 85d793ed..52f66ce3 100644 --- a/registry/channels/telegram.json +++ b/registry/channels/telegram.json @@ -18,8 +18,8 @@ }, "artifacts": { "wasm32-wasip2": { - "url": "https://github.com/nearai/ironclaw/releases/download/v0.19.0/channel-telegram-0.2.4-wasm32-wasip2.tar.gz", - "sha256": "a7cb300ec1c946831cfceaa95c1dc8f30d0f42a3924f3cb5de8098821573f4b8" + "url": "https://github.com/nearai/ironclaw/releases/download/v0.20.0/channel-telegram-0.2.5-wasm32-wasip2.tar.gz", + "sha256": "1ef20a538f55b379e049356e4d6758006251846bc3365ceaa1c87eba8379a329" } }, "auth_summary": { diff --git a/registry/tools/github.json b/registry/tools/github.json index e760c4df..bb351259 100644 --- a/registry/tools/github.json +++ b/registry/tools/github.json @@ -2,7 +2,7 @@ "name": "github", "display_name": "GitHub", "kind": "tool", - "version": "0.2.1", + "version": "0.2.2", "wit_version": "0.3.0", "description": "GitHub integration for issues, PRs, repos, and code search", "keywords": [ @@ -19,8 +19,8 @@ }, "artifacts": { "wasm32-wasip2": { - "url": "https://github.com/nearai/ironclaw/releases/download/v0.19.0/tool-github-0.2.1-wasm32-wasip2.tar.gz", - "sha256": "92c530b3ad172e2372d819744b5233f1d8f65768e26eb5a6c213eba3ce1de758" + "url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-github-0.2.2-wasm32-wasip2.tar.gz", + "sha256": "70b55af593193d8fa495c0f702ea23284d83a624124f8a5f7564916ec5032c3f" } }, "auth_summary": { diff --git a/registry/tools/gmail.json b/registry/tools/gmail.json index 08913ce6..c4772129 100644 --- a/registry/tools/gmail.json +++ b/registry/tools/gmail.json @@ -2,7 +2,7 @@ "name": "gmail", "display_name": "Gmail", "kind": "tool", - "version": "0.2.0", + "version": "0.2.1", "wit_version": "0.3.0", "description": "Read, send, and manage Gmail messages and threads", "keywords": [ @@ -18,8 +18,8 @@ }, "artifacts": { "wasm32-wasip2": { - "url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/gmail-0.2.0-wasm32-wasip2.tar.gz", - "sha256": "ee9574e02e92bc1d481f1310eb88afd99ee52bf6971074ab33bd76bf99b34b1d" + "url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-gmail-0.2.1-wasm32-wasip2.tar.gz", + "sha256": "79025b40ee70ce1120acc4320bae50da095d7afb0ef67bd56d99b064b72ea779" } }, "auth_summary": { diff --git a/registry/tools/google-calendar.json b/registry/tools/google-calendar.json index c43112d3..73065a67 100644 --- a/registry/tools/google-calendar.json +++ b/registry/tools/google-calendar.json @@ -2,7 +2,7 @@ "name": "google-calendar", "display_name": "Google Calendar", "kind": "tool", - "version": "0.2.0", + "version": "0.2.1", "wit_version": "0.3.0", "description": "Create, read, update, and delete Google Calendar events", "keywords": [ @@ -18,8 +18,8 @@ }, "artifacts": { "wasm32-wasip2": { - "url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-calendar-0.2.0-wasm32-wasip2.tar.gz", - "sha256": "2fa47150ea222e787c122182ad6f4dfa2ffaf5fe490d05e8de887a76445f8d2d" + "url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-google-calendar-0.2.1-wasm32-wasip2.tar.gz", + "sha256": "86bcc075010b08f5ab2f98f504cec1c6c9e0ca144857d185cbecf72a11f504bf" } }, "auth_summary": { diff --git a/registry/tools/google-docs.json b/registry/tools/google-docs.json index 9f1ab133..02cc94fe 100644 --- a/registry/tools/google-docs.json +++ b/registry/tools/google-docs.json @@ -2,7 +2,7 @@ "name": "google-docs", "display_name": "Google Docs", "kind": "tool", - "version": "0.2.0", + "version": "0.2.1", "wit_version": "0.3.0", "description": "Create and edit Google Docs documents", "keywords": [ @@ -18,8 +18,8 @@ }, "artifacts": { "wasm32-wasip2": { - "url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-docs-0.2.0-wasm32-wasip2.tar.gz", - "sha256": "40e134a1c1564f832ca861c3396895d4e33ec67b99313fc1f97baf8d971423a9" + "url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-google-docs-0.2.1-wasm32-wasip2.tar.gz", + "sha256": "39d476029764949498a53a6a223f9952b5f4df151be7b8b19bf3fe4d401a57cd" } }, "auth_summary": { diff --git a/registry/tools/google-drive.json b/registry/tools/google-drive.json index 9766e555..719690f7 100644 --- a/registry/tools/google-drive.json +++ b/registry/tools/google-drive.json @@ -2,7 +2,7 @@ "name": "google-drive", "display_name": "Google Drive", "kind": "tool", - "version": "0.2.0", + "version": "0.2.1", "wit_version": "0.3.0", "description": "Upload, download, search, and manage Google Drive files and folders", "keywords": [ @@ -18,8 +18,8 @@ }, "artifacts": { "wasm32-wasip2": { - "url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-drive-0.2.0-wasm32-wasip2.tar.gz", - "sha256": "002a341a1d58125563a7c69561b26fbc2629b04ea723cade744102bdc0fbb71f" + "url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-google-drive-0.2.1-wasm32-wasip2.tar.gz", + "sha256": "6e9a700fab93865c852af718666af64c5b534ad6a419fb4b736e07740188f494" } }, "auth_summary": { diff --git a/registry/tools/google-sheets.json b/registry/tools/google-sheets.json index b63265e1..09aae574 100644 --- a/registry/tools/google-sheets.json +++ b/registry/tools/google-sheets.json @@ -2,7 +2,7 @@ "name": "google-sheets", "display_name": "Google Sheets", "kind": "tool", - "version": "0.2.0", + "version": "0.2.1", "wit_version": "0.3.0", "description": "Read and write Google Sheets spreadsheet data", "keywords": [ @@ -18,8 +18,8 @@ }, "artifacts": { "wasm32-wasip2": { - "url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-sheets-0.2.0-wasm32-wasip2.tar.gz", - "sha256": "8aa2c9d52f033edea3a6c2311b0ec694ccb6d0a54ef07e94d72bf8be1ce8009a" + "url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-google-sheets-0.2.1-wasm32-wasip2.tar.gz", + "sha256": "1f8c381799a916be83263cac9d497d52946e21b1b588592a3a42ca94a73b7051" } }, "auth_summary": { diff --git a/registry/tools/google-slides.json b/registry/tools/google-slides.json index 54187531..64bc0e45 100644 --- a/registry/tools/google-slides.json +++ b/registry/tools/google-slides.json @@ -2,7 +2,7 @@ "name": "google-slides", "display_name": "Google Slides", "kind": "tool", - "version": "0.2.0", + "version": "0.2.1", "wit_version": "0.3.0", "description": "Create and edit Google Slides presentations", "keywords": [ @@ -17,8 +17,8 @@ }, "artifacts": { "wasm32-wasip2": { - "url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-slides-0.2.0-wasm32-wasip2.tar.gz", - "sha256": "e931a97d4fd0b0b938e464dc7c7f2be6ea6b4d1508f5ea3cd931d44db23f05f5" + "url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-google-slides-0.2.1-wasm32-wasip2.tar.gz", + "sha256": "e2528be5da02f1b8cfc8ee9b0cdd849516c53d412e2f75c6175b3bded7f512cb" } }, "auth_summary": { diff --git a/registry/tools/llm-context.json b/registry/tools/llm-context.json index e4e9808c..422f2e18 100644 --- a/registry/tools/llm-context.json +++ b/registry/tools/llm-context.json @@ -2,7 +2,7 @@ "name": "llm-context", "display_name": "LLM Context", "kind": "tool", - "version": "0.1.0", + "version": "0.1.1", "wit_version": "0.3.0", "description": "Fetch pre-extracted web content from Brave Search for grounding LLM answers (RAG, fact-checking)", "keywords": [ @@ -21,8 +21,8 @@ }, "artifacts": { "wasm32-wasip2": { - "url": "https://github.com/nearai/ironclaw/releases/download/v0.19.0/tool-llm-context-0.1.0-wasm32-wasip2.tar.gz", - "sha256": "d9ced2b1226b879135891e0ee40e072c7c95412e1b2462925a23853e1f92497e" + "url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-llm-context-0.1.1-wasm32-wasip2.tar.gz", + "sha256": "9b19e2fd05dbbbe3c8bd55309a91db09124e8415eb0f767828b6e10b55771e63" } }, "auth_summary": { diff --git a/registry/tools/slack.json b/registry/tools/slack.json index 8e1df989..236062a4 100644 --- a/registry/tools/slack.json +++ b/registry/tools/slack.json @@ -2,7 +2,7 @@ "name": "slack-tool", "display_name": "Slack Tool", "kind": "tool", - "version": "0.2.0", + "version": "0.2.1", "wit_version": "0.3.0", "description": "Your agent uses Slack to post and read messages in your workspace", "keywords": [ @@ -17,8 +17,8 @@ }, "artifacts": { "wasm32-wasip2": { - "url": "https://github.com/nearai/ironclaw/releases/download/v0.19.0/tool-slack-0.2.0-wasm32-wasip2.tar.gz", - "sha256": "ccfb0415d7a04f9497726c712d15216de36e86f498b849101283c017f5ab4efb" + "url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-slack-0.2.1-wasm32-wasip2.tar.gz", + "sha256": "927519e5b7734beeb022d3b8bbd152e0e6b9f67c9452a8ad47809d3c4221a137" } }, "auth_summary": { diff --git a/registry/tools/telegram.json b/registry/tools/telegram.json index 12e58c68..e684ca94 100644 --- a/registry/tools/telegram.json +++ b/registry/tools/telegram.json @@ -2,7 +2,7 @@ "name": "telegram-mtproto", "display_name": "Telegram Tool", "kind": "tool", - "version": "0.2.0", + "version": "0.2.1", "wit_version": "0.3.0", "description": "Your agent uses your Telegram account to read and send messages", "keywords": [ @@ -18,8 +18,8 @@ }, "artifacts": { "wasm32-wasip2": { - "url": "https://github.com/nearai/ironclaw/releases/download/v0.19.0/tool-telegram-0.2.0-wasm32-wasip2.tar.gz", - "sha256": "c17065ca41fae5f2a7c43b36144686718cd310a2f22442313bb1aa82bbad0ae4" + "url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-telegram-0.2.1-wasm32-wasip2.tar.gz", + "sha256": "1e57d0755fc9c7b3ec013d079f30168898b484a6919f9edd105f0cd80131c1cd" } }, "auth_summary": { diff --git a/registry/tools/web-search.json b/registry/tools/web-search.json index 5c1dedef..014466af 100644 --- a/registry/tools/web-search.json +++ b/registry/tools/web-search.json @@ -2,7 +2,7 @@ "name": "web-search", "display_name": "Web Search", "kind": "tool", - "version": "0.2.1", + "version": "0.2.2", "wit_version": "0.3.0", "description": "Search the web using Brave Search API", "keywords": [ @@ -18,8 +18,8 @@ }, "artifacts": { "wasm32-wasip2": { - "url": "https://github.com/nearai/ironclaw/releases/download/v0.19.0/tool-web-search-0.2.1-wasm32-wasip2.tar.gz", - "sha256": "bad275ca4ec314adea5241d6b92c44ccf9cebcbca8e30ba2493cc0bcb4b57218" + "url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-web-search-0.2.2-wasm32-wasip2.tar.gz", + "sha256": "47382b50c1ea7525b20d59dc02fab04e336d018665826c2f24710bdf460779ae" } }, "auth_summary": { diff --git a/release-plz.toml b/release-plz.toml index b003952d..e8e0670f 100644 --- a/release-plz.toml +++ b/release-plz.toml @@ -1,7 +1,2 @@ [workspace] git_release_enable = false - -[[package]] -name = "ironclaw_safety" -publish = false -release = false diff --git a/src/agent/agent_loop.rs b/src/agent/agent_loop.rs index 7961250d..4ee846f7 100644 --- a/src/agent/agent_loop.rs +++ b/src/agent/agent_loop.rs @@ -13,9 +13,10 @@ use futures::StreamExt; use uuid::Uuid; use crate::agent::context_monitor::ContextMonitor; -use crate::agent::heartbeat::spawn_heartbeat; +use crate::agent::heartbeat::{spawn_heartbeat, spawn_multi_user_heartbeat}; use crate::agent::routine_engine::{RoutineEngine, spawn_cron_ticker}; use crate::agent::self_repair::{DefaultSelfRepair, RepairResult, SelfRepair}; +use crate::agent::session::ThreadState; use crate::agent::session_manager::SessionManager; use crate::agent::submission::{Submission, SubmissionParser, SubmissionResult}; use crate::agent::{HeartbeatConfig as AgentHeartbeatConfig, Router, Scheduler, SchedulerDeps}; @@ -84,6 +85,15 @@ fn resolve_owner_scope_notification_user( trimmed_option(explicit_user).or_else(|| trimmed_option(owner_fallback)) } +fn is_single_message_repl(message: &IncomingMessage) -> bool { + message.channel == "repl" + && message + .metadata + .get("single_message_mode") + .and_then(|value| value.as_bool()) + .unwrap_or(false) +} + async fn resolve_channel_notification_user( extension_manager: Option<&Arc>, channel: Option<&str>, @@ -172,6 +182,8 @@ pub struct AgentDeps { /// Resolved LLM backend identifier (e.g., "nearai", "openai", "groq"). /// Used by `/model` persistence to determine which env var to update. pub llm_backend: String, + /// Per-tenant rate limiting registry (lazily creates rate state per user). + pub tenant_rates: Arc, } /// The main agent that coordinates all components. @@ -234,7 +246,10 @@ impl Agent { SchedulerDeps { tools: deps.tools.clone(), extension_manager: deps.extension_manager.clone(), - store: deps.store.clone(), + store: deps + .store + .as_ref() + .map(|db| crate::tenant::AdminScope::new(Arc::clone(db))), hooks: deps.hooks.clone(), }, ); @@ -315,6 +330,50 @@ impl Agent { &self.deps.cost_guard } + /// Build a tenant-scoped execution context for the given user. + /// + /// This is the standard entry point for per-user operations. The returned + /// [`TenantCtx`] provides a [`TenantScope`] that auto-binds `user_id` on + /// every database operation and a per-user rate limiter. + pub(super) async fn tenant_ctx(&self, user_id: &str) -> crate::tenant::TenantCtx { + let rate = self.deps.tenant_rates.get_or_create(user_id).await; + + let store = self + .deps + .store + .as_ref() + .map(|db| crate::tenant::TenantScope::new(user_id, Arc::clone(db))); + + // Reuse the owner workspace if user matches, otherwise create per-user. + let workspace = match &self.deps.workspace { + Some(ws) if ws.user_id() == user_id => Some(Arc::clone(ws)), + _ => self + .deps + .store + .as_ref() + .map(|db| Arc::new(Workspace::new_with_db(user_id, Arc::clone(db)))), + }; + + crate::tenant::TenantCtx::new( + user_id, + store, + workspace, + Arc::clone(&self.deps.cost_guard), + rate, + ) + } + + /// Get an admin-scoped database accessor for cross-tenant operations. + /// + /// Only for system-level components (heartbeat, routine engine, self-repair, + /// scheduler). Handler code should use [`tenant_ctx()`](Self::tenant_ctx) instead. + pub(super) fn admin_store(&self) -> Option { + self.deps + .store + .as_ref() + .map(|db| crate::tenant::AdminScope::new(Arc::clone(db))) + } + pub(super) fn skill_registry(&self) -> Option<&Arc>> { self.deps.skill_registry.as_ref() } @@ -400,8 +459,8 @@ impl Agent { self.config.stuck_threshold, self.config.max_repair_attempts, ); - if let Some(ref store) = self.deps.store { - self_repair = self_repair.with_store(Arc::clone(store)); + if let Some(admin) = self.admin_store() { + self_repair = self_repair.with_store(admin); } if let Some(ref builder) = self.deps.builder { self_repair = self_repair.with_builder(Arc::clone(builder), Arc::clone(self.tools())); @@ -508,6 +567,7 @@ impl Agent { .with_interval(std::time::Duration::from_secs(hb_config.interval_secs)); config.quiet_hours_start = hb_config.quiet_hours_start; config.quiet_hours_end = hb_config.quiet_hours_end; + config.multi_tenant = hb_config.multi_tenant; config.timezone = hb_config .timezone .clone() @@ -537,30 +597,52 @@ impl Agent { .await; let notify_user = heartbeat_notify_user; let channels = self.channels.clone(); + let is_multi_tenant = hb_config.multi_tenant; tokio::spawn(async move { while let Some(response) = notify_rx.recv().await { + // In multi-tenant mode, extract the owning user_id from + // the response metadata so notifications reach the + // correct user rather than the agent's owner. + // This intentionally overrides the configured notify_target + // because each user's heartbeat should notify that user. + let effective_user = if is_multi_tenant { + response + .metadata + .get("owner_id") + .and_then(|v| v.as_str()) + .map(String::from) + } else { + None + }; + // Try the configured channel first, fall back to // broadcasting on all channels. - let targeted_ok = if let Some(ref channel) = notify_channel - && let Some(ref user) = notify_target - { - channels - .broadcast(channel, user, response.clone()) - .await - .is_ok() + let targeted_ok = if let Some(ref channel) = notify_channel { + let target = effective_user.as_deref().or(notify_target.as_deref()); + if let Some(user) = target { + channels + .broadcast(channel, user, response.clone()) + .await + .is_ok() + } else { + false + } } else { false }; - if !targeted_ok && let Some(ref user) = notify_user { - let results = channels.broadcast_all(user, response).await; - for (ch, result) in results { - if let Err(e) = result { - tracing::warn!( - "Failed to broadcast heartbeat to {}: {}", - ch, - e - ); + if !targeted_ok { + let fallback = effective_user.as_deref().or(notify_user.as_deref()); + if let Some(user) = fallback { + let results = channels.broadcast_all(user, response).await; + for (ch, result) in results { + if let Err(e) = result { + tracing::warn!( + "Failed to broadcast heartbeat to {}: {}", + ch, + e + ); + } } } } @@ -573,14 +655,29 @@ impl Agent { .map(|h| h.to_workspace_config()) .unwrap_or_default(); - Some(spawn_heartbeat( - config, - hygiene, - workspace.clone(), - self.cheap_llm().clone(), - Some(notify_tx), - self.store().map(Arc::clone), - )) + if config.multi_tenant { + if let Some(admin) = self.admin_store() { + Some(spawn_multi_user_heartbeat( + config, + hygiene, + self.cheap_llm().clone(), + Some(notify_tx), + admin, + )) + } else { + tracing::warn!("Multi-tenant heartbeat requires a database store"); + None + } + } else { + Some(spawn_heartbeat( + config, + hygiene, + workspace.clone(), + self.cheap_llm().clone(), + Some(notify_tx), + self.admin_store(), + )) + } } else { tracing::warn!("Heartbeat enabled but no workspace available"); None @@ -602,7 +699,7 @@ impl Agent { let engine = Arc::new(RoutineEngine::new( rt_config.clone(), - Arc::clone(store), + crate::tenant::AdminScope::new(Arc::clone(store)), self.llm().clone(), Arc::clone(workspace), notify_tx, @@ -1055,10 +1152,11 @@ impl Agent { } else { drop(sess); self.session_manager - .resolve_thread( + .resolve_thread_with_parsed_uuid( &message.user_id, &message.channel, message.conversation_scope(), + approval_thread_uuid, ) .await } @@ -1139,9 +1237,14 @@ impl Agent { && let Submission::UserInput { ref content } = submission && let Some(engine) = self.routine_engine().await { + let single_message_repl = is_single_message_repl(message); // Use post-hook content so that BeforeInbound hooks that rewrite // input are respected by event trigger matching. - let fired = engine.check_event_triggers(message, content).await; + let fired = if single_message_repl { + engine.check_event_triggers_and_wait(message, content).await + } else { + engine.check_event_triggers(message, content).await + }; if fired > 0 { tracing::debug!( channel = %message.channel, @@ -1149,15 +1252,30 @@ impl Agent { fired, "Consumed inbound user message with matching event-triggered routine(s)" ); - return Ok(Some(String::new())); + return if single_message_repl { + Ok(None) + } else { + Ok(Some(String::new())) + }; } } + // Build per-tenant execution context once; threaded through all handlers. + let tenant = self.tenant_ctx(&message.user_id).await; + + let session_for_empty_exit = Arc::clone(&session); + // Process based on submission type let result = match submission { Submission::UserInput { content } => { let mut result = self - .process_user_input(message, session.clone(), thread_id, &content) + .process_user_input( + message, + tenant.clone(), + session.clone(), + thread_id, + &content, + ) .await; // Drain any messages queued during processing. @@ -1224,7 +1342,13 @@ impl Agent { let mut queued_msg = message.clone(); queued_msg.attachments.clear(); result = self - .process_user_input(&queued_msg, session.clone(), thread_id, &next_content) + .process_user_input( + &queued_msg, + tenant.clone(), + session.clone(), + thread_id, + &next_content, + ) .await; // If processing failed, re-queue the drained content so it @@ -1249,8 +1373,30 @@ impl Agent { command, message.channel ); + // /reasoning is special-cased here (not in handle_system_command) + // because it needs the session + thread_id to read turn reasoning + // data, which handle_system_command's signature doesn't provide. + if command == "reasoning" { + let result = self + .handle_reasoning_command(&args, &session, thread_id) + .await; + return match result { + SubmissionResult::Response { content } => Ok(Some(content)), + SubmissionResult::Ok { message } => Ok(message), + SubmissionResult::Error { message } => { + Ok(Some(format!("Error: {}", message))) + } + _ => { + if is_single_message_repl(message) { + Ok(None) + } else { + Ok(Some(String::new())) + } + } + }; + } // Authorization checks (including restart channel check) are enforced in handle_system_command - self.handle_system_command(&command, &args, &message.channel) + self.handle_system_command(&command, &args, &message.channel, &tenant) .await } Submission::Undo => self.process_undo(session, thread_id).await, @@ -1263,12 +1409,9 @@ impl Agent { Submission::Summarize => self.process_summarize(session, thread_id).await, Submission::Suggest => self.process_suggest(session, thread_id).await, Submission::JobStatus { job_id } => { - self.process_job_status(&message.user_id, job_id.as_deref()) - .await - } - Submission::JobCancel { job_id } => { - self.process_job_cancel(&message.user_id, &job_id).await + self.process_job_status(&tenant, job_id.as_deref()).await } + Submission::JobCancel { job_id } => self.process_job_cancel(&tenant, &job_id).await, Submission::Quit => return Ok(None), Submission::SwitchThread { thread_id: target } => { self.process_switch_thread(message, target).await @@ -1308,7 +1451,26 @@ impl Agent { Ok(Some(content)) } } - SubmissionResult::Ok { message } => Ok(message), + SubmissionResult::Ok { + message: output_message, + } => { + let should_exit = + if output_message.as_deref() == Some("") && is_single_message_repl(message) { + let sess = session_for_empty_exit.lock().await; + sess.threads + .get(&thread_id) + .map(|thread| thread.state != ThreadState::AwaitingApproval) + .unwrap_or(true) + } else { + false + }; + + if should_exit { + Ok(None) + } else { + Ok(output_message) + } + } SubmissionResult::Error { message } => Ok(Some(format!("Error: {}", message))), SubmissionResult::Interrupted => Ok(Some("Interrupted.".into())), SubmissionResult::NeedApproval { .. } => { @@ -1324,7 +1486,7 @@ impl Agent { #[cfg(test)] mod tests { use super::{ - chat_tool_execution_metadata, resolve_routine_notification_user, + chat_tool_execution_metadata, is_single_message_repl, resolve_routine_notification_user, should_fallback_routine_notification, truncate_for_preview, }; use crate::channels::IncomingMessage; @@ -1486,4 +1648,17 @@ mod tests { assert!(should_fallback_routine_notification(&error)); // safety: test-only assertion } + + #[test] + fn single_message_repl_detection_requires_repl_channel_and_metadata_flag() { + let repl = IncomingMessage::new("repl", "owner-scope", "hello") + .with_metadata(serde_json::json!({ "single_message_mode": true })); + let gateway = IncomingMessage::new("gateway", "owner-scope", "hello") + .with_metadata(serde_json::json!({ "single_message_mode": true })); + let plain_repl = IncomingMessage::new("repl", "owner-scope", "hello"); + + assert!(is_single_message_repl(&repl)); // safety: test-only assertion + assert!(!is_single_message_repl(&gateway)); // safety: test-only assertion + assert!(!is_single_message_repl(&plain_repl)); // safety: test-only assertion + } } diff --git a/src/agent/agentic_loop.rs b/src/agent/agentic_loop.rs index cc6fd486..27c2ab72 100644 --- a/src/agent/agentic_loop.rs +++ b/src/agent/agentic_loop.rs @@ -10,7 +10,7 @@ use std::borrow::Cow; use crate::agent::session::PendingApproval; use crate::error::Error; -use crate::llm::{ChatMessage, Reasoning, ReasoningContext, RespondResult}; +use crate::llm::{ChatMessage, FinishReason, Reasoning, ReasoningContext, RespondResult}; /// Signal from the delegate indicating how the loop should proceed. pub enum LoopSignal { @@ -134,6 +134,9 @@ pub async fn run_agentic_loop( config: &AgenticLoopConfig, ) -> Result { let mut consecutive_tool_intent_nudges: u32 = 0; + // Accumulates across all iterations (not reset by text responses) so + // non-consecutive truncations still escalate to force_text. + let mut truncation_count: u32 = 0; for iteration in 1..=config.max_iterations { // Check for external signals (stop, cancellation, user messages) @@ -215,7 +218,35 @@ pub async fn run_agentic_loop( tool_calls, content, } => { + // If the response was truncated, tool call parameters are likely + // incomplete. Discard them and tell the LLM to try a different + // approach rather than executing malformed tool calls. + if output.finish_reason == FinishReason::Length { + truncation_count += 1; + let names: Vec<&str> = tool_calls.iter().map(|tc| tc.name.as_str()).collect(); + tracing::warn!( + iteration, + tools = ?names, + truncation_count, + "Discarding truncated tool calls (finish_reason=Length)" + ); + if let Some(ref text) = content { + reason_ctx.messages.push(ChatMessage::assistant(text)); + } + reason_ctx + .messages + .push(ChatMessage::user(crate::llm::TRUNCATED_TOOL_CALL_NOTICE)); + // After repeated truncations, force text-only mode so the LLM + // stops attempting tool calls it can't fit in the output budget. + if truncation_count >= 3 { + reason_ctx.force_text = true; + } + delegate.after_iteration(iteration).await; + continue; + } + consecutive_tool_intent_nudges = 0; + truncation_count = 0; if let Some(outcome) = delegate .execute_tool_calls(tool_calls, content, reason_ctx) @@ -271,6 +302,7 @@ mod tests { RespondOutput { result: RespondResult::Text(text.to_string()), usage: zero_usage(), + finish_reason: FinishReason::Stop, } } @@ -281,6 +313,7 @@ mod tests { content: None, }, usage: zero_usage(), + finish_reason: FinishReason::ToolUse, } } @@ -414,6 +447,7 @@ mod tests { id: "call_1".to_string(), name: "echo".to_string(), arguments: serde_json::json!({}), + reasoning: None, }; let delegate = MockDelegate::new(vec![ tool_calls_output(vec![tool_call]), @@ -621,4 +655,95 @@ mod tests { let result = truncate_for_preview("café", 4); assert_eq!(result, "caf..."); } + + #[tokio::test] + async fn test_truncated_tool_calls_discarded_on_length() { + let truncated_tool_call = ToolCall { + id: "call_1".to_string(), + name: "memory_write".to_string(), + arguments: serde_json::json!({}), // empty — truncated + reasoning: None, + }; + let truncated_output = RespondOutput { + result: RespondResult::ToolCalls { + tool_calls: vec![truncated_tool_call], + content: Some("I'll write the report.".to_string()), + }, + usage: zero_usage(), + finish_reason: FinishReason::Length, // response was truncated + }; + let delegate = MockDelegate::new(vec![truncated_output, text_output("Summarized it.")]); + let reasoning = stub_reasoning(); + let mut ctx = ReasoningContext::new(); + let config = AgenticLoopConfig { + max_iterations: 5, + ..Default::default() + }; + + let outcome = run_agentic_loop(&delegate, &reasoning, &mut ctx, &config) + .await + .unwrap(); + + // Tool calls should NOT have been executed + assert_eq!(delegate.tool_exec_count.load(Ordering::SeqCst), 0); + // The loop should have continued and returned the text response + assert!(matches!(outcome, LoopOutcome::Response(ref t) if t == "Summarized it.")); + // A truncation notice should have been injected into context + assert!( + ctx.messages + .iter() + .any(|m| m.role == crate::llm::Role::User && m.content.contains("truncated")), + "Should inject truncation notice into context" + ); + // The partial assistant content should have been preserved + assert!( + ctx.messages + .iter() + .any(|m| m.role == crate::llm::Role::Assistant + && m.content.contains("write the report")), + "Should preserve partial assistant content" + ); + } + + #[tokio::test] + async fn test_repeated_truncations_force_text_mode() { + let make_truncated = || RespondOutput { + result: RespondResult::ToolCalls { + tool_calls: vec![ToolCall { + id: "call_1".to_string(), + name: "memory_write".to_string(), + arguments: serde_json::json!({}), + reasoning: None, + }], + content: None, + }, + usage: zero_usage(), + finish_reason: FinishReason::Length, + }; + // Three truncated responses, then a text response + let delegate = MockDelegate::new(vec![ + make_truncated(), + make_truncated(), + make_truncated(), + text_output("Gave up on tool calls."), + ]); + let reasoning = stub_reasoning(); + let mut ctx = ReasoningContext::new(); + let config = AgenticLoopConfig { + max_iterations: 5, + ..Default::default() + }; + + let outcome = run_agentic_loop(&delegate, &reasoning, &mut ctx, &config) + .await + .unwrap(); + + assert!(matches!(outcome, LoopOutcome::Response(_))); + assert_eq!(delegate.tool_exec_count.load(Ordering::SeqCst), 0); + // After 3 truncations, force_text should be set + assert!( + ctx.force_text, + "Should escalate to force_text after repeated truncations" + ); + } } diff --git a/src/agent/commands.rs b/src/agent/commands.rs index b6aff3c0..643d8c7c 100644 --- a/src/agent/commands.rs +++ b/src/agent/commands.rs @@ -33,6 +33,7 @@ impl Agent { &self, intent: MessageIntent, message: &IncomingMessage, + tenant: &crate::tenant::TenantCtx, ) -> Result { // Send thinking status for non-trivial operations if let MessageIntent::CreateJob { .. } = &intent { @@ -52,24 +53,18 @@ impl Agent { description, category, } => { - self.handle_create_job(&message.user_id, title, description, category) + self.handle_create_job(tenant, title, description, category) .await? } MessageIntent::CheckJobStatus { job_id } => { - self.handle_check_status(&message.user_id, job_id).await? - } - MessageIntent::CancelJob { job_id } => { - self.handle_cancel_job(&message.user_id, &job_id).await? - } - MessageIntent::ListJobs { filter } => { - self.handle_list_jobs(&message.user_id, filter).await? - } - MessageIntent::HelpJob { job_id } => { - self.handle_help_job(&message.user_id, &job_id).await? + self.handle_check_status(tenant, job_id).await? } + MessageIntent::CancelJob { job_id } => self.handle_cancel_job(tenant, &job_id).await?, + MessageIntent::ListJobs { filter } => self.handle_list_jobs(tenant, filter).await?, + MessageIntent::HelpJob { job_id } => self.handle_help_job(tenant, &job_id).await?, MessageIntent::Command { command, args } => { match self - .handle_command(&command, &args, &message.channel) + .handle_command(&command, &args, &message.channel, tenant) .await? { Some(s) => s, @@ -83,14 +78,14 @@ impl Agent { async fn handle_create_job( &self, - user_id: &str, + tenant: &crate::tenant::TenantCtx, title: String, description: String, category: Option, ) -> Result { let job_id = self .scheduler - .dispatch_job(user_id, &title, &description, None) + .dispatch_job(tenant.user_id(), &title, &description, None) .await?; // Set the dedicated category field (not stored in metadata) @@ -113,7 +108,7 @@ impl Agent { async fn handle_check_status( &self, - user_id: &str, + tenant: &crate::tenant::TenantCtx, job_id: Option, ) -> Result { match job_id { @@ -122,7 +117,8 @@ impl Agent { .map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?; // Try DB first for persistent state, fall back to ContextManager. - if let Some(store) = self.store() + // TenantScope.get_job() auto-filters by ownership — no manual check needed. + if let Some(store) = tenant.store() && let Ok(Some(ctx)) = store.get_job(uuid).await { return Ok(format!( @@ -138,7 +134,7 @@ impl Agent { } let ctx = self.context_manager.get_context(uuid).await?; - if ctx.user_id != user_id { + if ctx.user_id != tenant.user_id() { return Err(crate::error::JobError::NotFound { id: uuid }.into()); } @@ -155,7 +151,8 @@ impl Agent { } None => { // Show summary from DB for consistency with Jobs tab. - if let Some(store) = self.store() { + // TenantScope methods auto-scope to user — no user_id parameter needed. + if let Some(store) = tenant.store() { let mut total = 0; let mut in_progress = 0; let mut completed = 0; @@ -183,7 +180,7 @@ impl Agent { } // Fallback to ContextManager if no DB. - let summary = self.context_manager.summary_for(user_id).await; + let summary = self.context_manager.summary_for(tenant.user_id()).await; Ok(format!( "Jobs summary: Total: {} In Progress: {} Completed: {} Failed: {} Stuck: {}", summary.total, @@ -196,19 +193,24 @@ impl Agent { } } - async fn handle_cancel_job(&self, user_id: &str, job_id: &str) -> Result { + async fn handle_cancel_job( + &self, + tenant: &crate::tenant::TenantCtx, + job_id: &str, + ) -> Result { let uuid = Uuid::parse_str(job_id) .map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?; let ctx = self.context_manager.get_context(uuid).await?; - if ctx.user_id != user_id { + if ctx.user_id != tenant.user_id() { return Err(crate::error::JobError::NotFound { id: uuid }.into()); } self.scheduler.stop(uuid).await?; // Also update DB so the Jobs tab reflects cancellation immediately. - if let Some(store) = self.store() + // Use TenantScope — ownership already verified above. + if let Some(store) = tenant.store() && let Err(e) = store .update_job_status(uuid, JobState::Cancelled, Some("Cancelled by user")) .await @@ -221,11 +223,12 @@ impl Agent { async fn handle_list_jobs( &self, - user_id: &str, + tenant: &crate::tenant::TenantCtx, _filter: Option, ) -> Result { // List from DB for consistency with Jobs tab. - if let Some(store) = self.store() { + // TenantScope methods auto-scope to user. + if let Some(store) = tenant.store() { let agent_jobs = match store.list_agent_jobs().await { Ok(jobs) => jobs, Err(e) => { @@ -256,7 +259,7 @@ impl Agent { } // Fallback to ContextManager if no DB. - let jobs = self.context_manager.all_jobs_for(user_id).await; + let jobs = self.context_manager.all_jobs_for(tenant.user_id()).await; if jobs.is_empty() { return Ok("No jobs found.".to_string()); } @@ -270,12 +273,16 @@ impl Agent { Ok(output) } - async fn handle_help_job(&self, user_id: &str, job_id: &str) -> Result { + async fn handle_help_job( + &self, + tenant: &crate::tenant::TenantCtx, + job_id: &str, + ) -> Result { let uuid = Uuid::parse_str(job_id) .map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?; let ctx = self.context_manager.get_context(uuid).await?; - if ctx.user_id != user_id { + if ctx.user_id != tenant.user_id() { return Err(crate::error::JobError::NotFound { id: uuid }.into()); } @@ -308,11 +315,11 @@ impl Agent { /// Show job status inline — either all jobs (no id) or a specific job. pub(super) async fn process_job_status( &self, - user_id: &str, + tenant: &crate::tenant::TenantCtx, job_id: Option<&str>, ) -> Result { match self - .handle_check_status(user_id, job_id.map(|s| s.to_string())) + .handle_check_status(tenant, job_id.map(|s| s.to_string())) .await { Ok(text) => Ok(SubmissionResult::response(text)), @@ -323,10 +330,10 @@ impl Agent { /// Cancel a job by ID. pub(super) async fn process_job_cancel( &self, - user_id: &str, + tenant: &crate::tenant::TenantCtx, job_id: &str, ) -> Result { - match self.handle_cancel_job(user_id, job_id).await { + match self.handle_cancel_job(tenant, job_id).await { Ok(text) => Ok(SubmissionResult::response(text)), Err(e) => Ok(SubmissionResult::error(format!("Cancel error: {}", e))), } @@ -465,12 +472,101 @@ impl Agent { } } + /// Handle `/reasoning [N|all]` — show reasoning history for the active thread. + pub(super) async fn handle_reasoning_command( + &self, + args: &[String], + session: &Arc>, + thread_id: Uuid, + ) -> SubmissionResult { + // Clone the turn data we need, then drop the session lock. + let turns_snapshot: Vec<( + usize, + Option, + Vec, + )>; + { + let sess = session.lock().await; + let thread = match sess.threads.get(&thread_id) { + Some(t) => t, + None => return SubmissionResult::error("No active thread."), + }; + + if thread.turns.is_empty() { + return SubmissionResult::ok_with_message("No turns yet."); + } + + // Parse argument: default=last turn, "all"=all turns, N=specific turn (1-based). + let selected: Vec<&crate::agent::session::Turn> = match args.first().map(|s| s.as_str()) + { + Some("all") => thread.turns.iter().collect(), + Some(n) => match n.parse::() { + Ok(0) => return SubmissionResult::error("Turn numbers start at 1."), + Ok(num) if num > thread.turns.len() => { + return SubmissionResult::error(format!( + "Turn {} does not exist (max: {}).", + num, + thread.turns.len() + )); + } + Ok(num) => vec![&thread.turns[num - 1]], + Err(_) => return SubmissionResult::error("Usage: /reasoning [N|all]"), + }, + None => { + // Default: last turn that has tool calls + match thread.turns.iter().rev().find(|t| !t.tool_calls.is_empty()) { + Some(t) => vec![t], + None => { + return SubmissionResult::ok_with_message("No turns with tool calls."); + } + } + } + }; + + turns_snapshot = selected + .into_iter() + .map(|t| (t.turn_number, t.narrative.clone(), t.tool_calls.clone())) + .collect(); + } + // Session lock is now dropped — format output without holding it. + + let mut output = String::new(); + for (turn_number, narrative, tool_calls) in &turns_snapshot { + output.push_str(&format!("--- Turn {} ---\n", turn_number + 1)); + if let Some(narrative) = narrative { + output.push_str(&format!("Reasoning: {}\n", narrative)); + } + if tool_calls.is_empty() { + output.push_str(" (no tool calls)\n"); + } else { + for tc in tool_calls { + let status = if tc.error.is_some() { + "error" + } else if tc.result.is_some() { + "ok" + } else { + "pending" + }; + output.push_str(&format!(" {} [{}]", tc.name, status)); + if let Some(ref rationale) = tc.rationale { + output.push_str(&format!(" — {}", rationale)); + } + output.push('\n'); + } + } + output.push('\n'); + } + + SubmissionResult::response(output.trim_end()) + } + /// Handle system commands that bypass thread-state checks entirely. pub(super) async fn handle_system_command( &self, command: &str, args: &[String], channel: &str, + tenant: &crate::tenant::TenantCtx, ) -> Result { match command { "help" => Ok(SubmissionResult::response(concat!( @@ -480,6 +576,7 @@ impl Agent { " /version Show version info\n", " /tools List available tools\n", " /debug Toggle debug mode\n", + " /reasoning [N|all] Show agent reasoning for turns\n", " /ping Connectivity check\n", "\n", "Jobs:\n", @@ -663,19 +760,32 @@ impl Agent { } } - match self.llm().set_model(requested) { - Ok(()) => { - // Persist the model choice so it survives restarts. - self.persist_selected_model(requested).await; - Ok(SubmissionResult::response(format!( - "Switched model to: {}", - requested - ))) + if self.config.multi_tenant { + // Multi-tenant: only persist to per-user DB settings. + // Do NOT call set_model() on the shared provider — that + // would change the default for all users. The per-request + // model_override in the dispatcher reads from the same + // "selected_model" setting and applies it per-user. + self.persist_selected_model(tenant, requested).await; + Ok(SubmissionResult::response(format!( + "Model preference set to: {} (per-user)", + requested + ))) + } else { + match self.llm().set_model(requested) { + Ok(()) => { + // Persist the model choice so it survives restarts. + self.persist_selected_model(tenant, requested).await; + Ok(SubmissionResult::response(format!( + "Switched model to: {}", + requested + ))) + } + Err(e) => Ok(SubmissionResult::error(format!( + "Failed to switch model: {}", + e + ))), } - Err(e) => Ok(SubmissionResult::error(format!( - "Failed to switch model: {}", - e - ))), } } } @@ -817,10 +927,14 @@ impl Agent { command: &str, args: &[String], channel: &str, + tenant: &crate::tenant::TenantCtx, ) -> Result, Error> { // System commands are now handled directly via Submission::SystemCommand, // but the router may still send us unknown /commands. - match self.handle_system_command(command, args, channel).await? { + match self + .handle_system_command(command, args, channel, tenant) + .await? + { SubmissionResult::Response { content } => Ok(Some(content)), SubmissionResult::Ok { message } => Ok(message), SubmissionResult::Error { message } => Ok(Some(format!("Error: {}", message))), @@ -832,23 +946,33 @@ impl Agent { /// /// Best-effort: logs warnings on failure but does not propagate errors, /// since the in-memory model switch already succeeded. - async fn persist_selected_model(&self, model: &str) { - // 1. Persist to DB if available. - if let Some(store) = self.store() { + /// + /// In multi-tenant mode, only the per-user DB setting is written — global + /// .env and TOML files are shared across users and must not be mutated. + async fn persist_selected_model(&self, tenant: &crate::tenant::TenantCtx, model: &str) { + // 1. Persist to DB if available (per-user scoped via TenantScope). + if let Some(store) = tenant.store() { let value = serde_json::Value::String(model.to_string()); - if let Err(e) = store - .set_setting(self.owner_id(), "selected_model", &value) - .await - { + if let Err(e) = store.set_setting("selected_model", &value).await { tracing::warn!("Failed to persist model to DB: {}", e); } else { - tracing::debug!("Persisted selected_model to DB: {}", model); + tracing::debug!( + user_id = tenant.user_id(), + "Persisted selected_model to DB: {}", + model + ); } } else { tracing::warn!("No database store available — model choice will not persist to DB"); } - // 2. Update .env and TOML config file (sync I/O in spawn_blocking). + // 2. In multi-tenant mode, skip .env/TOML writes — these are global + // files shared by all users. The per-user DB setting is sufficient. + if self.config.multi_tenant { + return; + } + + // 3. Update .env and TOML config file (sync I/O in spawn_blocking). let model_owned = model.to_string(); let backend = self.deps.llm_backend.clone(); if let Err(e) = tokio::task::spawn_blocking(move || { diff --git a/src/agent/cost_guard.rs b/src/agent/cost_guard.rs index 4563bbbe..4885364b 100644 --- a/src/agent/cost_guard.rs +++ b/src/agent/cost_guard.rs @@ -21,6 +21,9 @@ pub struct CostGuardConfig { pub max_cost_per_day_cents: Option, /// Maximum LLM calls per hour. None = unlimited. pub max_actions_per_hour: Option, + /// Maximum spend per user per day in cents. None = unlimited. + /// Applied independently per user alongside the global budget. + pub max_cost_per_user_per_day_cents: Option, } /// Error returned when a cost limit is exceeded. @@ -30,6 +33,12 @@ pub enum CostLimitExceeded { DailyBudget { spent_cents: u64, limit_cents: u64 }, /// Hourly action rate limit reached. HourlyRate { actions: u64, limit: u64 }, + /// Per-user daily spending cap reached. + UserDailyBudget { + user_id: String, + spent_cents: u64, + limit_cents: u64, + }, } impl std::fmt::Display for CostLimitExceeded { @@ -49,6 +58,17 @@ impl std::fmt::Display for CostLimitExceeded { "Hourly action limit exceeded: {} actions of {} allowed per hour", actions, limit ), + Self::UserDailyBudget { + user_id, + spent_cents, + limit_cents, + } => write!( + f, + "User '{}' daily cost limit exceeded: spent ${:.2} of ${:.2} allowed", + user_id, + *spent_cents as f64 / 100.0, + *limit_cents as f64 / 100.0 + ), } } } @@ -78,6 +98,9 @@ pub struct CostGuard { /// Per-model token usage since startup. model_tokens: Mutex>, + + /// Per-user daily cost tracking. Each entry resets independently at midnight UTC. + per_user_daily_cost: Mutex>, } struct DailyCost { @@ -97,6 +120,7 @@ impl CostGuard { action_window: Mutex::new(VecDeque::new()), budget_exceeded: AtomicBool::new(false), model_tokens: Mutex::new(HashMap::new()), + per_user_daily_cost: Mutex::new(HashMap::new()), } } @@ -203,6 +227,11 @@ impl CostGuard { daily.reset_date = today; self.budget_exceeded.store(false, Ordering::Relaxed); tracing::info!("Cost guard: daily counter reset for {}", today); + + // Prune per-user entries from previous days to prevent + // unbounded HashMap growth in long-lived deployments. + let mut per_user = self.per_user_daily_cost.lock().await; + per_user.retain(|_, entry| entry.reset_date == today); } daily.total += cost; @@ -248,6 +277,85 @@ impl CostGuard { cost } + /// Record an LLM call with per-user attribution. + /// + /// Delegates to `record_llm_call` for global tracking, then additionally + /// records the cost against the user's daily budget. + #[allow(clippy::too_many_arguments)] + pub async fn record_llm_call_for_user( + &self, + user_id: &str, + model: &str, + input_tokens: u32, + output_tokens: u32, + cache_read_input_tokens: u32, + cache_creation_input_tokens: u32, + cache_read_discount: Decimal, + cache_write_multiplier: Decimal, + cost_per_token: Option<(Decimal, Decimal)>, + ) -> Decimal { + let cost = self + .record_llm_call( + model, + input_tokens, + output_tokens, + cache_read_input_tokens, + cache_creation_input_tokens, + cache_read_discount, + cache_write_multiplier, + cost_per_token, + ) + .await; + + // Track per-user daily cost + { + let today = chrono::Utc::now().date_naive(); + let mut per_user = self.per_user_daily_cost.lock().await; + let entry = per_user + .entry(user_id.to_string()) + .or_insert_with(|| DailyCost { + total: Decimal::ZERO, + reset_date: today, + }); + if today != entry.reset_date { + entry.total = Decimal::ZERO; + entry.reset_date = today; + } + entry.total += cost; + } + + cost + } + + /// Check whether the next action is allowed for a specific user. + /// + /// Checks the global limits first (via `check_allowed`), then additionally + /// checks the per-user daily budget if configured. + pub async fn check_allowed_for_user(&self, user_id: &str) -> Result<(), CostLimitExceeded> { + // Check global limits first + self.check_allowed().await?; + + // Check per-user daily budget + if let Some(limit_cents) = self.config.max_cost_per_user_per_day_cents { + let today = chrono::Utc::now().date_naive(); + let per_user = self.per_user_daily_cost.lock().await; + if let Some(entry) = per_user.get(user_id) + && entry.reset_date == today + { + let spent_cents = to_cents(entry.total); + if spent_cents >= limit_cents { + return Err(CostLimitExceeded::UserDailyBudget { + user_id: user_id.to_string(), + spent_cents, + limit_cents, + }); + } + } + } + + Ok(()) + } + /// Current daily spend in USD (as Decimal). pub async fn daily_spend(&self) -> Decimal { let daily = self.daily_cost.lock().await; @@ -259,6 +367,16 @@ impl CostGuard { } } + /// Current daily spend for a specific user in USD (as Decimal). + pub async fn daily_spend_for_user(&self, user_id: &str) -> Decimal { + let today = chrono::Utc::now().date_naive(); + let per_user = self.per_user_daily_cost.lock().await; + match per_user.get(user_id) { + Some(entry) if entry.reset_date == today => entry.total, + _ => Decimal::ZERO, + } + } + /// Number of actions in the current hourly window. pub async fn actions_this_hour(&self) -> u64 { let mut window = self.action_window.lock().await; @@ -314,7 +432,7 @@ mod tests { async fn test_daily_budget_enforcement() { let guard = CostGuard::new(CostGuardConfig { max_cost_per_day_cents: Some(1), // $0.01 limit - max_actions_per_hour: None, + ..CostGuardConfig::default() }); // First call allowed @@ -350,8 +468,8 @@ mod tests { #[tokio::test] async fn test_hourly_rate_enforcement() { let guard = CostGuard::new(CostGuardConfig { - max_cost_per_day_cents: None, max_actions_per_hour: Some(3), + ..CostGuardConfig::default() }); // First 3 actions allowed @@ -633,8 +751,8 @@ mod tests { // A fresh CostGuard with rate limits should not panic even if // checked_sub returns None (simulating short uptime). let guard = CostGuard::new(CostGuardConfig { - max_cost_per_day_cents: None, max_actions_per_hour: Some(100), + ..CostGuardConfig::default() }); // These must not panic regardless of system uptime @@ -656,4 +774,119 @@ mod tests { let result = Instant::now().checked_sub(std::time::Duration::MAX); assert!(result.is_none()); } + + #[tokio::test] + async fn test_per_user_daily_budget_enforcement() { + let guard = CostGuard::new(CostGuardConfig { + max_cost_per_day_cents: None, + max_actions_per_hour: None, + max_cost_per_user_per_day_cents: Some(1), // $0.01 per user + }); + + // Both users initially allowed + assert!(guard.check_allowed_for_user("alice").await.is_ok()); + assert!(guard.check_allowed_for_user("bob").await.is_ok()); + + // Alice makes an expensive call + guard + .record_llm_call_for_user( + "alice", + "gpt-4o", + 10_000, + 10_000, + 0, + 0, + Decimal::ONE, + Decimal::ONE, + None, + ) + .await; + + // Alice should be blocked, Bob should still be allowed + let result = guard.check_allowed_for_user("alice").await; + assert!(result.is_err()); + match result.unwrap_err() { + CostLimitExceeded::UserDailyBudget { + user_id, + limit_cents, + .. + } => { + assert_eq!(user_id, "alice"); + assert_eq!(limit_cents, 1); + } + other => panic!("Expected UserDailyBudget, got {:?}", other), + } + assert!(guard.check_allowed_for_user("bob").await.is_ok()); + } + + #[tokio::test] + async fn test_per_user_daily_spend_tracking() { + let guard = CostGuard::new(CostGuardConfig::default()); + + assert_eq!(guard.daily_spend_for_user("alice").await, Decimal::ZERO); + assert_eq!(guard.daily_spend_for_user("bob").await, Decimal::ZERO); + + let cost = guard + .record_llm_call_for_user( + "alice", + "gpt-4o", + 1000, + 500, + 0, + 0, + Decimal::ONE, + Decimal::ONE, + None, + ) + .await; + + assert_eq!(guard.daily_spend_for_user("alice").await, cost); + assert_eq!(guard.daily_spend_for_user("bob").await, Decimal::ZERO); + // Global spend should also be tracked + assert_eq!(guard.daily_spend().await, cost); + } + + #[tokio::test] + async fn test_per_user_budget_independent_of_global() { + let guard = CostGuard::new(CostGuardConfig { + max_cost_per_day_cents: Some(100_000), // $1000 global limit + max_actions_per_hour: None, + max_cost_per_user_per_day_cents: Some(1), // $0.01 per user + }); + + // User hits their personal limit + guard + .record_llm_call_for_user( + "alice", + "gpt-4o", + 10_000, + 10_000, + 0, + 0, + Decimal::ONE, + Decimal::ONE, + None, + ) + .await; + + // Alice blocked by per-user limit, not global + assert!(guard.check_allowed_for_user("alice").await.is_err()); + // Global limit is far from reached + assert!(guard.check_allowed().await.is_ok()); + // Bob is unaffected + assert!(guard.check_allowed_for_user("bob").await.is_ok()); + } + + #[test] + fn test_user_cost_limit_display() { + let limit = CostLimitExceeded::UserDailyBudget { + user_id: "alice".to_string(), + spent_cents: 150, + limit_cents: 100, + }; + let msg = limit.to_string(); + assert!(msg.contains("alice")); + assert!(msg.contains("$1.50")); + assert!(msg.contains("$1.00")); + } } diff --git a/src/agent/dispatcher.rs b/src/agent/dispatcher.rs index c3584ae1..c3ea321e 100644 --- a/src/agent/dispatcher.rs +++ b/src/agent/dispatcher.rs @@ -42,6 +42,7 @@ impl Agent { pub(super) async fn run_agentic_loop( &self, message: &IncomingMessage, + tenant: crate::tenant::TenantCtx, session: Arc>, thread_id: Uuid, initial_messages: Vec, @@ -72,7 +73,12 @@ impl Agent { ); let system_prompt = if let Some(ws) = self.workspace() { - match ws + let scoped_workspace = if ws.user_id() == message.user_id { + Arc::clone(ws) + } else { + Arc::new(ws.scoped_to_user(&message.user_id)) + }; + match scoped_workspace .system_prompt_for_context_tz(is_group_chat, user_tz) .await { @@ -172,6 +178,7 @@ impl Agent { let delegate = ChatDelegate { agent: self, + tenant, session: session.clone(), thread_id, message, @@ -244,6 +251,7 @@ impl Agent { /// auth intercept, and cost tracking. struct ChatDelegate<'a> { agent: &'a Agent, + tenant: crate::tenant::TenantCtx, session: Arc>, thread_id: Uuid, message: &'a IncomingMessage, @@ -307,6 +315,8 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { // Update context for this iteration reason_ctx.available_tools = tool_defs; + // Preserve force_text if already set (e.g. by truncation escalation). + let force_text = force_text || reason_ctx.force_text; reason_ctx.system_prompt = Some(if force_text { self.cached_prompt_no_tools.clone() } else { @@ -340,8 +350,8 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { reason_ctx: &mut ReasoningContext, iteration: usize, ) -> Result { - // Enforce cost guardrails before the LLM call - if let Err(limit) = self.agent.cost_guard().check_allowed().await { + // Enforce cost guardrails before the LLM call (global + per-user) + if let Err(limit) = self.tenant.check_cost_allowed().await { return Err(crate::error::LlmError::InvalidResponse { provider: "agent".to_string(), reason: limit.to_string(), @@ -349,6 +359,21 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { .into()); } + // Apply per-user model override from settings (first iteration only + // to avoid repeated DB lookups within the same agentic loop). + // Uses "selected_model" — the same key the /model command persists to + // via SettingsStore (per-user scoped via TenantScope). + if iteration == 0 + && let Some(store) = self.tenant.store() + && let Ok(Some(value)) = store.get_setting("selected_model").await + && let Some(model) = value.as_str() + { + let model = model.trim(); + if !model.is_empty() { + reason_ctx.model_override = Some(model.to_string()); + } + } + let output = match reasoning.respond_with_tools(reason_ctx).await { Ok(output) => output, Err(crate::error::LlmError::ContextLengthExceeded { used, limit }) => { @@ -383,13 +408,22 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { Err(e) => return Err(e.into()), }; - // Record cost and track token usage - let model_name = self.agent.llm().active_model_name(); + // Record cost and track token usage (global + per-user). + // When a model override is active, use the override name for attribution + // and let CostGuard look up pricing via costs::model_cost() instead of + // using the default provider's cost_per_token (which reflects the wrong model). + let (model_name, cost_per_token) = if let Some(ref ovr) = reason_ctx.model_override { + (ovr.clone(), None) + } else { + ( + self.agent.llm().active_model_name(), + Some(self.agent.llm().cost_per_token()), + ) + }; let read_discount = self.agent.llm().cache_read_discount(); let write_multiplier = self.agent.llm().cache_write_multiplier(); let call_cost = self - .agent - .cost_guard() + .tenant .record_llm_call( &model_name, output.usage.input_tokens, @@ -398,7 +432,7 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { output.usage.cache_creation_input_tokens, read_discount, write_multiplier, - Some(self.agent.llm().cost_per_token()), + cost_per_token, ) .await; tracing::debug!( @@ -429,6 +463,19 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { content: Option, reason_ctx: &mut ReasoningContext, ) -> Result, Error> { + // Extract and sanitize the narrative before consuming `content`. + let narrative = content + .as_deref() + .filter(|c| !c.trim().is_empty()) + .map(|c| { + let sanitized = self + .agent + .safety() + .sanitize_tool_output("agent_narrative", c); + sanitized.content + }) + .filter(|c| !c.trim().is_empty()); + // Add the assistant message with tool_calls to context. // OpenAI protocol requires this before tool-result messages. reason_ctx @@ -449,6 +496,41 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { ) .await; + // Build per-tool decisions for the reasoning update. + // Sanitize each rationale through SafetyLayer (parity with JobDelegate). + let decisions: Vec = tool_calls + .iter() + .filter_map(|tc| { + tc.reasoning.as_ref().map(|r| { + let sanitized = self + .agent + .safety() + .sanitize_tool_output("tool_rationale", r) + .content; + crate::channels::ToolDecision { + tool_name: tc.name.clone(), + rationale: sanitized, + } + }) + }) + .collect(); + + // Emit reasoning update to channels. + if narrative.is_some() || !decisions.is_empty() { + let _ = self + .agent + .channels + .send_status( + &self.message.channel, + StatusUpdate::ReasoningUpdate { + narrative: narrative.clone().unwrap_or_default(), + decisions: decisions.clone(), + }, + &self.message.metadata, + ) + .await; + } + // Record tool calls in the thread with sensitive params redacted. { let mut redacted_args: Vec = Vec::with_capacity(tool_calls.len()); @@ -464,8 +546,23 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { if let Some(thread) = sess.threads.get_mut(&self.thread_id) && let Some(turn) = thread.last_turn_mut() { + // Set turn-level narrative. + if turn.narrative.is_none() { + turn.narrative = narrative; + } for (tc, safe_args) in tool_calls.iter().zip(redacted_args) { - turn.record_tool_call(&tc.name, safe_args); + let sanitized_rationale = tc.reasoning.as_ref().map(|r| { + self.agent + .safety() + .sanitize_tool_output("tool_rationale", r) + .content + }); + turn.record_tool_call_with_reasoning( + &tc.name, + safe_args, + sanitized_rationale, + Some(tc.id.clone()), + ); } } } @@ -735,7 +832,7 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { if let Some(thread) = sess.threads.get_mut(&self.thread_id) && let Some(turn) = thread.last_turn_mut() { - turn.record_tool_error(error_msg.clone()); + turn.record_tool_error_for(&tc.id, error_msg.clone()); } } reason_ctx @@ -861,16 +958,19 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { Err(e) => format!("Tool '{}' failed: {}", tc.name, e), }; - // Record sanitized result in thread + // Record sanitized result in thread (identity-based matching). { let mut sess = self.session.lock().await; if let Some(thread) = sess.threads.get_mut(&self.thread_id) && let Some(turn) = thread.last_turn_mut() { if is_tool_error { - turn.record_tool_error(result_content.clone()); + turn.record_tool_error_for(&tc.id, result_content.clone()); } else { - turn.record_tool_result(serde_json::json!(result_content)); + turn.record_tool_result_for( + &tc.id, + serde_json::json!(result_content), + ); } } } @@ -1243,6 +1343,7 @@ mod tests { sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, builder: None, llm_backend: "nearai".to_string(), + tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)), }; Agent::new( @@ -1258,10 +1359,14 @@ mod tests { allow_local_tools: false, max_cost_per_day_cents: None, max_actions_per_hour: None, + max_cost_per_user_per_day_cents: None, max_tool_iterations: 50, auto_approve_tools: false, default_timezone: "UTC".to_string(), max_tokens_per_job: 0, + multi_tenant: false, + max_llm_concurrent_per_user: None, + max_jobs_concurrent_per_user: None, }, deps, Arc::new(ChannelManager::new()), @@ -1471,11 +1576,13 @@ mod tests { id: "call_2".to_string(), name: "http".to_string(), arguments: serde_json::json!({"url": "https://example.com"}), + reasoning: None, }, ToolCall { id: "call_3".to_string(), name: "echo".to_string(), arguments: serde_json::json!({"message": "done"}), + reasoning: None, }, ], user_timezone: None, @@ -1661,6 +1768,7 @@ mod tests { id: "call_1".to_string(), name: "echo".to_string(), arguments: serde_json::json!({"message": "hi"}), + reasoning: None, }], ), ChatMessage::tool_result("call_1", "echo", "hi"), @@ -1753,11 +1861,13 @@ mod tests { id: "c1".to_string(), name: "http".to_string(), arguments: serde_json::json!({}), + reasoning: None, }, ToolCall { id: "c2".to_string(), name: "echo".to_string(), arguments: serde_json::json!({}), + reasoning: None, }, ], ), @@ -1791,6 +1901,7 @@ mod tests { id: "c1".to_string(), name: "echo".to_string(), arguments: serde_json::json!({}), + reasoning: None, }], ), ChatMessage::tool_result("c1", "echo", "done"), @@ -1921,6 +2032,7 @@ mod tests { id: crate::llm::generate_tool_call_id(0, 0), name: "echo".to_string(), arguments: serde_json::json!({"message": "looping"}), + reasoning: None, }], input_tokens: 0, output_tokens: 5, @@ -2074,6 +2186,7 @@ mod tests { id: crate::llm::generate_tool_call_id(0, 0), name: "nonexistent_tool".to_string(), arguments: serde_json::json!({}), + reasoning: None, }], input_tokens: 0, output_tokens: 5, @@ -2111,6 +2224,7 @@ mod tests { sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, builder: None, llm_backend: "nearai".to_string(), + tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)), }; Agent::new( @@ -2126,10 +2240,14 @@ mod tests { allow_local_tools: false, max_cost_per_day_cents: None, max_actions_per_hour: None, + max_cost_per_user_per_day_cents: None, max_tool_iterations, auto_approve_tools: true, default_timezone: "UTC".to_string(), max_tokens_per_job: 0, + multi_tenant: false, + max_llm_concurrent_per_user: None, + max_jobs_concurrent_per_user: None, }, deps, Arc::new(ChannelManager::new()), @@ -2164,13 +2282,14 @@ mod tests { let message = IncomingMessage::new("test", "test-user", "do something"); let initial_messages = vec![ChatMessage::user("do something")]; + let tenant = agent.tenant_ctx("test-user").await; // The dispatcher must terminate within 5 seconds. If there is an // infinite loop bug (e.g., index not advancing on tool failure), the // timeout will fire and the test will fail. let result = tokio::time::timeout( Duration::from_secs(5), - agent.run_agentic_loop(&message, session, thread_id, initial_messages), + agent.run_agentic_loop(&message, tenant, session, thread_id, initial_messages), ) .await; @@ -2232,6 +2351,7 @@ mod tests { sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, builder: None, llm_backend: "nearai".to_string(), + tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)), }; Agent::new( @@ -2247,10 +2367,14 @@ mod tests { allow_local_tools: false, max_cost_per_day_cents: None, max_actions_per_hour: None, + max_cost_per_user_per_day_cents: None, max_tool_iterations: max_iter, auto_approve_tools: true, default_timezone: "UTC".to_string(), max_tokens_per_job: 0, + multi_tenant: false, + max_llm_concurrent_per_user: None, + max_jobs_concurrent_per_user: None, }, deps, Arc::new(ChannelManager::new()), @@ -2270,13 +2394,14 @@ mod tests { let message = IncomingMessage::new("test", "test-user", "keep calling tools"); let initial_messages = vec![ChatMessage::user("keep calling tools")]; + let tenant = agent.tenant_ctx("test-user").await; // Even with an LLM that always wants to call tools, the dispatcher // must terminate within the timeout thanks to force_text at // max_tool_iterations. let result = tokio::time::timeout( Duration::from_secs(5), - agent.run_agentic_loop(&message, session, thread_id, initial_messages), + agent.run_agentic_loop(&message, tenant, session, thread_id, initial_messages), ) .await; diff --git a/src/agent/heartbeat.rs b/src/agent/heartbeat.rs index ec4cd5e9..f7a8f869 100644 --- a/src/agent/heartbeat.rs +++ b/src/agent/heartbeat.rs @@ -31,8 +31,8 @@ use chrono_tz::Tz; use tokio::sync::mpsc; use crate::channels::OutgoingResponse; -use crate::db::Database; use crate::llm::{ChatMessage, CompletionRequest, LlmProvider, Reasoning}; +use crate::tenant::AdminScope; use crate::workspace::Workspace; use crate::workspace::hygiene::HygieneConfig; @@ -57,6 +57,9 @@ pub struct HeartbeatConfig { pub quiet_hours_end: Option, /// Timezone for fire_at and quiet hours evaluation (IANA name). pub timezone: Option, + /// When true, cycle through all users with routines instead of + /// running heartbeat for a single user. Requires a database store. + pub multi_tenant: bool, } impl Default for HeartbeatConfig { @@ -71,6 +74,7 @@ impl Default for HeartbeatConfig { quiet_hours_start: None, quiet_hours_end: None, timezone: None, + multi_tenant: false, } } } @@ -178,7 +182,7 @@ pub struct HeartbeatRunner { workspace: Arc, llm: Arc, response_tx: Option>, - store: Option>, + store: Option, consecutive_failures: u32, } @@ -207,8 +211,8 @@ impl HeartbeatRunner { self } - /// Set the database store for persistent heartbeat conversations. - pub fn with_store(mut self, store: Arc) -> Self { + /// Set the admin-scoped database store for persistent heartbeat conversations. + pub fn with_store(mut self, store: AdminScope) -> Self { self.store = Some(store); self } @@ -396,7 +400,7 @@ impl HeartbeatRunner { } /// Send a notification about heartbeat findings. - async fn send_notification(&self, message: &str) { + pub(crate) async fn send_notification(&self, message: &str) { let Some(ref tx) = self.response_tx else { tracing::debug!("No response channel configured for heartbeat notifications"); return; @@ -493,7 +497,7 @@ pub fn spawn_heartbeat( workspace: Arc, llm: Arc, response_tx: Option>, - store: Option>, + store: Option, ) -> tokio::task::JoinHandle<()> { let mut runner = HeartbeatRunner::new(config, hygiene_config, workspace, llm); if let Some(tx) = response_tx { @@ -508,6 +512,179 @@ pub fn spawn_heartbeat( }) } +/// Spawn a multi-user heartbeat runner that cycles through all users that +/// own routines (enabled or not). Each tick, it queries the DB for distinct +/// user_ids, creates a per-user workspace, and runs a heartbeat check for +/// each user concurrently. Per-user failure counts are tracked independently. +pub fn spawn_multi_user_heartbeat( + config: HeartbeatConfig, + hygiene_config: HygieneConfig, + llm: Arc, + response_tx: Option>, + store: AdminScope, +) -> tokio::task::JoinHandle<()> { + tokio::spawn(async move { + if !config.enabled { + tracing::info!("Multi-user heartbeat is disabled"); + return; + } + + let mut tick_interval = if config.fire_at.is_none() { + let mut iv = tokio::time::interval(config.interval); + iv.tick().await; // skip immediate tick + Some(iv) + } else { + None + }; + + // Track consecutive failures per user so we can disable heartbeat + // for persistently-failing users (same semantics as single-user mode). + let mut user_failures: std::collections::HashMap = + std::collections::HashMap::new(); + + tracing::info!("Starting multi-user heartbeat loop"); + + loop { + if let Some(fire_at) = config.fire_at { + let sleep_dur = duration_until_next_fire(fire_at, config.resolved_tz()); + tokio::time::sleep(sleep_dur).await; + } else if let Some(ref mut iv) = tick_interval { + iv.tick().await; + } + + if config.is_quiet_hours() { + continue; + } + + // Get distinct user_ids from routines + let user_ids = match store.list_all_routines().await { + Ok(routines) => { + let mut ids: Vec = routines + .iter() + .map(|r| r.user_id.clone()) + .collect::>() + .into_iter() + .collect(); + ids.sort(); + ids + } + Err(e) => { + tracing::error!("Multi-user heartbeat: failed to list routines: {}", e); + continue; + } + }; + + // Run user heartbeats concurrently so one slow LLM call doesn't + // block others. Cap concurrency to avoid flooding the LLM provider. + const MAX_CONCURRENT_HEARTBEATS: usize = 8; + let mut join_set = tokio::task::JoinSet::new(); + + for user_id in &user_ids { + // Skip users that have exceeded max_failures + let failures = user_failures.get(user_id).copied().unwrap_or(0); + if failures >= config.max_failures { + continue; + } + + let workspace = Arc::new(Workspace::new_with_db(user_id, Arc::clone(store.db()))); + + // Run memory hygiene per user (same as single-user heartbeat). + let hygiene_ws = Arc::clone(&workspace); + let hygiene_cfg = hygiene_config.clone(); + let hygiene_user = user_id.clone(); + tokio::spawn(async move { + let report = + crate::workspace::hygiene::run_if_due(&hygiene_ws, &hygiene_cfg).await; + if report.had_work() { + tracing::info!( + user_id = hygiene_user, + daily_logs_deleted = report.daily_logs_deleted, + conversation_docs_deleted = report.conversation_docs_deleted, + "multi-user heartbeat: memory hygiene deleted stale documents" + ); + } + }); + + // Drain completed tasks to stay within the concurrency cap. + while join_set.len() >= MAX_CONCURRENT_HEARTBEATS { + if let Some(join_result) = join_set.join_next().await { + collect_heartbeat_result(join_result, &mut user_failures, &config); + } + } + + let uid = user_id.clone(); + let cfg = config.clone(); + let hyg = hygiene_config.clone(); + let llm_clone = llm.clone(); + let tx = response_tx.clone(); + let admin = store.clone(); + + join_set.spawn(async move { + let mut runner = HeartbeatRunner::new(cfg, hyg, workspace, llm_clone); + if let Some(tx) = tx { + runner = runner.with_response_channel(tx); + } + runner = runner.with_store(admin); + + let result = runner.check_heartbeat().await; + if let HeartbeatResult::NeedsAttention(msg) = &result { + runner.send_notification(msg).await; + } + (uid, result) + }); + } + + // Collect remaining results and update failure counts + while let Some(join_result) = join_set.join_next().await { + collect_heartbeat_result(join_result, &mut user_failures, &config); + } + } + }) +} + +/// Process a single JoinSet result from the multi-user heartbeat loop. +fn collect_heartbeat_result( + join_result: Result<(String, HeartbeatResult), tokio::task::JoinError>, + user_failures: &mut std::collections::HashMap, + config: &HeartbeatConfig, +) { + let (uid, result) = match join_result { + Ok(pair) => pair, + Err(e) => { + tracing::error!("Multi-user heartbeat task panicked: {}", e); + return; + } + }; + match result { + HeartbeatResult::Ok => { + tracing::trace!(user_id = uid, "Multi-user heartbeat OK"); + user_failures.remove(&uid); + } + HeartbeatResult::NeedsAttention(_) => { + tracing::info!(user_id = uid, "Multi-user heartbeat needs attention"); + user_failures.remove(&uid); + } + HeartbeatResult::Skipped => {} + HeartbeatResult::Failed(err) => { + let count = user_failures.entry(uid.clone()).or_insert(0); + *count += 1; + tracing::error!( + user_id = uid, + consecutive_failures = *count, + "Multi-user heartbeat failed: {}", + err + ); + if *count >= config.max_failures { + tracing::error!( + user_id = uid, + "Multi-user heartbeat disabled for user after {} consecutive failures", + count + ); + } + } + } +} + #[cfg(test)] mod tests { use super::*; @@ -726,7 +903,7 @@ mod tests { Arc, Arc, Option>, - Option>, + Option, ) -> tokio::task::JoinHandle<()> = spawn_heartbeat; let _ = _fn_ptr; } diff --git a/src/agent/job_monitor.rs b/src/agent/job_monitor.rs index 02f5e3e2..e102dfbf 100644 --- a/src/agent/job_monitor.rs +++ b/src/agent/job_monitor.rs @@ -21,8 +21,8 @@ use tokio::task::JoinHandle; use uuid::Uuid; use crate::channels::IncomingMessage; -use crate::channels::web::types::SseEvent; use crate::context::{ContextManager, JobState}; +use ironclaw_common::AppEvent; /// Route context for forwarding job monitor events back to the user's channel. #[derive(Debug, Clone)] @@ -36,15 +36,15 @@ pub struct JobMonitorRoute { /// injects assistant messages into the agent loop. /// /// The monitor forwards: -/// - `SseEvent::JobMessage` (assistant role): injected as incoming messages so +/// - `AppEvent::JobMessage` (assistant role): injected as incoming messages so /// the main agent can read and relay to the user. -/// - `SseEvent::JobResult`: injected as a completion notice, then the task exits. +/// - `AppEvent::JobResult`: injected as a completion notice, then the task exits. /// /// Tool use/result and status events are intentionally skipped (too noisy for /// the main agent's context window). pub fn spawn_job_monitor( job_id: Uuid, - event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>, + event_rx: broadcast::Receiver<(Uuid, String, AppEvent)>, inject_tx: mpsc::Sender, route: JobMonitorRoute, ) -> JoinHandle<()> { @@ -56,7 +56,7 @@ pub fn spawn_job_monitor( /// jobs don't stay `InProgress` forever in the `ContextManager`. pub fn spawn_job_monitor_with_context( job_id: Uuid, - mut event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>, + mut event_rx: broadcast::Receiver<(Uuid, String, AppEvent)>, inject_tx: mpsc::Sender, route: JobMonitorRoute, context_manager: Option>, @@ -74,7 +74,7 @@ pub fn spawn_job_monitor_with_context( } match event { - SseEvent::JobMessage { role, content, .. } if role == "assistant" => { + AppEvent::JobMessage { role, content, .. } if role == "assistant" => { let mut msg = IncomingMessage::new( route.channel.clone(), route.user_id.clone(), @@ -92,7 +92,7 @@ pub fn spawn_job_monitor_with_context( break; } } - SseEvent::JobResult { status, .. } => { + AppEvent::JobResult { status, .. } => { // Transition in-memory state so the job frees its // max_jobs slot and query tools show the final state. if let Some(ref cm) = context_manager { @@ -162,7 +162,7 @@ pub fn spawn_job_monitor_with_context( /// inject messages into) but we still need to free the `max_jobs` slot. pub fn spawn_completion_watcher( job_id: Uuid, - mut event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>, + mut event_rx: broadcast::Receiver<(Uuid, String, AppEvent)>, context_manager: Arc, ) -> JoinHandle<()> { let short_id = job_id.to_string()[..8].to_string(); @@ -170,7 +170,7 @@ pub fn spawn_completion_watcher( tokio::spawn(async move { loop { match event_rx.recv().await { - Ok((ev_job_id, _user_id, SseEvent::JobResult { status, .. })) + Ok((ev_job_id, _user_id, AppEvent::JobResult { status, .. })) if ev_job_id == job_id => { let target = if status == "completed" { @@ -229,7 +229,7 @@ mod tests { #[tokio::test] async fn test_monitor_forwards_assistant_messages() { - let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let job_id = Uuid::new_v4(); @@ -240,7 +240,7 @@ mod tests { .send(( job_id, "test-user".to_string(), - SseEvent::JobMessage { + AppEvent::JobMessage { job_id: job_id.to_string(), role: "assistant".to_string(), content: "I found a bug".to_string(), @@ -262,7 +262,7 @@ mod tests { #[tokio::test] async fn test_monitor_ignores_other_jobs() { - let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let job_id = Uuid::new_v4(); @@ -274,7 +274,7 @@ mod tests { .send(( other_job_id, "test-user".to_string(), - SseEvent::JobMessage { + AppEvent::JobMessage { job_id: other_job_id.to_string(), role: "assistant".to_string(), content: "wrong job".to_string(), @@ -293,7 +293,7 @@ mod tests { #[tokio::test] async fn test_monitor_exits_on_job_result() { - let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let job_id = Uuid::new_v4(); @@ -304,7 +304,7 @@ mod tests { .send(( job_id, "test-user".to_string(), - SseEvent::JobResult { + AppEvent::JobResult { job_id: job_id.to_string(), status: "completed".to_string(), session_id: None, @@ -329,7 +329,7 @@ mod tests { #[tokio::test] async fn test_monitor_skips_tool_events() { - let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let job_id = Uuid::new_v4(); @@ -340,7 +340,7 @@ mod tests { .send(( job_id, "test-user".to_string(), - SseEvent::JobToolUse { + AppEvent::JobToolUse { job_id: job_id.to_string(), tool_name: "shell".to_string(), input: serde_json::json!({"command": "ls"}), @@ -353,7 +353,7 @@ mod tests { .send(( job_id, "test-user".to_string(), - SseEvent::JobMessage { + AppEvent::JobMessage { job_id: job_id.to_string(), role: "user".to_string(), content: "user prompt".to_string(), @@ -409,7 +409,7 @@ mod tests { .await .unwrap(); - let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let handle = spawn_job_monitor_with_context( @@ -425,7 +425,7 @@ mod tests { .send(( job_id, "test-user".to_string(), - SseEvent::JobResult { + AppEvent::JobResult { job_id: job_id.to_string(), status: "completed".to_string(), session_id: None, @@ -458,7 +458,7 @@ mod tests { .await .unwrap(); - let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let handle = spawn_job_monitor_with_context( @@ -474,7 +474,7 @@ mod tests { .send(( job_id, "test-user".to_string(), - SseEvent::JobResult { + AppEvent::JobResult { job_id: job_id.to_string(), status: "failed".to_string(), session_id: None, @@ -507,14 +507,14 @@ mod tests { .await .unwrap(); - let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let handle = spawn_completion_watcher(job_id, event_tx.subscribe(), Arc::clone(&cm)); event_tx .send(( job_id, "test-user".to_string(), - SseEvent::JobResult { + AppEvent::JobResult { job_id: job_id.to_string(), status: "completed".to_string(), session_id: None, diff --git a/src/agent/mod.rs b/src/agent/mod.rs index 84155666..e7242845 100644 --- a/src/agent/mod.rs +++ b/src/agent/mod.rs @@ -36,7 +36,9 @@ pub(crate) use agent_loop::truncate_for_preview; pub use agent_loop::{Agent, AgentDeps}; pub use compaction::{CompactionResult, ContextCompactor}; pub use context_monitor::{CompactionStrategy, ContextBreakdown, ContextMonitor}; -pub use heartbeat::{HeartbeatConfig, HeartbeatResult, HeartbeatRunner, spawn_heartbeat}; +pub use heartbeat::{ + HeartbeatConfig, HeartbeatResult, HeartbeatRunner, spawn_heartbeat, spawn_multi_user_heartbeat, +}; pub use router::{MessageIntent, Router}; pub use routine::{Routine, RoutineAction, RoutineRun, Trigger}; pub use routine_engine::{RoutineEngine, SandboxReadiness}; diff --git a/src/agent/routine_engine.rs b/src/agent/routine_engine.rs index 39acb83d..64c3b94c 100644 --- a/src/agent/routine_engine.rs +++ b/src/agent/routine_engine.rs @@ -18,6 +18,7 @@ use std::time::Duration; use chrono::Utc; use regex::Regex; use tokio::sync::{RwLock, mpsc}; +use tokio::task::JoinHandle; use uuid::Uuid; use crate::agent::Scheduler; @@ -27,12 +28,12 @@ use crate::agent::routine::{ use crate::channels::{IncomingMessage, OutgoingResponse}; use crate::config::RoutineConfig; use crate::context::{JobContext, JobState}; -use crate::db::Database; use crate::error::RoutineError; use crate::extensions::ExtensionManager; use crate::llm::{ ChatMessage, CompletionRequest, FinishReason, LlmProvider, ToolCall, ToolCompletionRequest, }; +use crate::tenant::AdminScope; use crate::tools::{ ToolError, ToolRegistry, autonomous_allowed_tool_names, autonomous_unavailable_message, prepare_tool_params, @@ -45,6 +46,11 @@ enum EventMatcher { System { routine: Routine }, } +struct TriggeredRoutine { + routine: Routine, + detail: String, +} + /// Distinguishes why sandbox is unavailable so error messages are accurate. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum SandboxReadiness { @@ -93,7 +99,7 @@ pub(crate) fn routine_matches_message(routine: &Routine, message: &IncomingMessa /// The routine execution engine. pub struct RoutineEngine { config: RoutineConfig, - store: Arc, + store: AdminScope, llm: Arc, workspace: Arc, /// Sender for notifications (routed to channel manager). @@ -122,7 +128,7 @@ impl RoutineEngine { #[allow(clippy::too_many_arguments)] pub fn new( config: RoutineConfig, - store: Arc, + store: AdminScope, llm: Arc, workspace: Arc, notify_tx: mpsc::Sender, @@ -202,6 +208,44 @@ impl RoutineEngine { /// Check incoming message against event triggers. Returns number of routines fired. pub async fn check_event_triggers(&self, message: &IncomingMessage, content: &str) -> usize { + let triggered = self.matching_event_triggers(message, content).await; + let fired = triggered.len(); + for triggered in triggered { + std::mem::drop(self.spawn_fire(triggered.routine, "event", Some(triggered.detail))); + } + fired + } + + /// Fire matching event-triggered routines and wait for them to complete. + /// + /// Used by single-message REPL mode so the process does not exit before + /// background event-triggered routines finish. + pub async fn check_event_triggers_and_wait( + &self, + message: &IncomingMessage, + content: &str, + ) -> usize { + let triggered = self.matching_event_triggers(message, content).await; + let fired = triggered.len(); + let handles: Vec> = triggered + .into_iter() + .map(|triggered| self.spawn_fire(triggered.routine, "event", Some(triggered.detail))) + .collect(); + + for handle in handles { + if let Err(e) = handle.await { + tracing::warn!(error = %e, "Event-triggered routine task failed"); + } + } + + fired + } + + async fn matching_event_triggers( + &self, + message: &IncomingMessage, + content: &str, + ) -> Vec { let cache = self.event_cache.read().await; // Early return if there are no message matchers at all. @@ -209,10 +253,9 @@ impl RoutineEngine { .iter() .any(|m| matches!(m, EventMatcher::Message { .. })) { - return 0; + return Vec::new(); } - - let mut fired = 0; + let mut triggered = Vec::new(); // Collect routine IDs for batch query let routine_ids: Vec = cache @@ -224,13 +267,13 @@ impl RoutineEngine { .collect(); if routine_ids.is_empty() { - return 0; + return Vec::new(); } // Single batch query instead of N queries let concurrent_counts = match self.batch_concurrent_counts(&routine_ids).await { Some(counts) => counts, - None => return 0, + None => return Vec::new(), }; for matcher in cache.iter() { @@ -285,11 +328,13 @@ impl RoutineEngine { } let detail = truncate(content, 200); - self.spawn_fire(routine.clone(), "event", Some(detail)); - fired += 1; + triggered.push(TriggeredRoutine { + routine: routine.clone(), + detail, + }); } - fired + triggered } /// Emit a structured event to system-event routines. @@ -737,12 +782,22 @@ impl RoutineEngine { }); } + // Per-user workspace (same pattern as spawn_fire). + let routine_workspace = if routine.user_id == self.workspace.user_id() { + self.workspace.clone() + } else { + Arc::new(Workspace::new_with_db( + &routine.user_id, + Arc::clone(self.store.db()), + )) + }; + // Execute inline for manual triggers (caller wants to wait) let engine = EngineContext { config: self.config.clone(), store: self.store.clone(), llm: self.llm.clone(), - workspace: self.workspace.clone(), + workspace: routine_workspace, notify_tx: self.notify_tx.clone(), running_count: self.running_count.clone(), scheduler: self.scheduler.clone(), @@ -845,7 +900,12 @@ impl RoutineEngine { } /// Spawn a fire in a background task. - fn spawn_fire(&self, routine: Routine, trigger_type: &str, trigger_detail: Option) { + fn spawn_fire( + &self, + routine: Routine, + trigger_type: &str, + trigger_detail: Option, + ) -> JoinHandle<()> { let run = RoutineRun { id: Uuid::new_v4(), routine_id: routine.id, @@ -860,11 +920,23 @@ impl RoutineEngine { created_at: Utc::now(), }; + // Use per-user workspace so each routine executes in the correct + // user's context. Fall back to the engine-wide workspace when the + // routine belongs to the same user (avoids unnecessary allocation). + let routine_workspace = if routine.user_id == self.workspace.user_id() { + self.workspace.clone() + } else { + Arc::new(Workspace::new_with_db( + &routine.user_id, + Arc::clone(self.store.db()), + )) + }; + let engine = EngineContext { config: self.config.clone(), store: self.store.clone(), llm: self.llm.clone(), - workspace: self.workspace.clone(), + workspace: routine_workspace, notify_tx: self.notify_tx.clone(), running_count: self.running_count.clone(), scheduler: self.scheduler.clone(), @@ -882,7 +954,7 @@ impl RoutineEngine { return; } execute_routine(engine, routine, run).await; - }); + }) } fn check_cooldown(&self, routine: &Routine) -> bool { @@ -917,7 +989,7 @@ impl RoutineEngine { /// an active state (Pending/InProgress/Stuck). Maps the final `JobState` to /// a `RunStatus` for the routine run. struct FullJobWatcher { - store: Arc, + store: AdminScope, job_id: Uuid, routine_name: String, } @@ -928,7 +1000,7 @@ impl FullJobWatcher { /// Safety ceiling: 24 hours, derived from POLL_INTERVAL. const MAX_POLLS: u32 = (24 * 60 * 60) / Self::POLL_INTERVAL.as_secs() as u32; - fn new(store: Arc, job_id: Uuid, routine_name: String) -> Self { + fn new(store: AdminScope, job_id: Uuid, routine_name: String) -> Self { Self { store, job_id, @@ -1000,7 +1072,7 @@ impl FullJobWatcher { /// Shared context passed to the execution function. struct EngineContext { config: RoutineConfig, - store: Arc, + store: AdminScope, llm: Arc, workspace: Arc, notify_tx: mpsc::Sender, @@ -1541,7 +1613,10 @@ async fn execute_lightweight_with_tools( let force_text = iteration >= max_iterations; if force_text { - // Final iteration: no tools, just get text response + // Final iteration: no tools, just get text response. + // Claude 4.6 rejects assistant prefill; NEAR AI rejects any non-user-ending + // conversation. Ensure the last message is user-role. + crate::util::ensure_ends_with_user_message(&mut messages); let request = CompletionRequest::new(messages) .with_max_tokens(effective_max_tokens) .with_temperature(0.3); diff --git a/src/agent/scheduler.rs b/src/agent/scheduler.rs index 02953a4b..88eb2a64 100644 --- a/src/agent/scheduler.rs +++ b/src/agent/scheduler.rs @@ -11,12 +11,12 @@ use uuid::Uuid; use crate::agent::task::{Task, TaskContext, TaskOutput}; use crate::config::AgentConfig; use crate::context::{ContextManager, JobContext, JobState}; -use crate::db::Database; use crate::error::{Error, JobError}; use crate::extensions::ExtensionManager; use crate::hooks::HookRegistry; use crate::llm::LlmProvider; use crate::safety::SafetyLayer; +use crate::tenant::AdminScope; use crate::tools::{ ApprovalContext, ToolRegistry, autonomous_allowed_tool_names, autonomous_unavailable_error, prepare_tool_params, @@ -52,7 +52,7 @@ struct ScheduledSubtask { pub struct SchedulerDeps { pub tools: Arc, pub extension_manager: Option>, - pub store: Option>, + pub store: Option, pub hooks: Arc, } @@ -64,7 +64,7 @@ pub struct Scheduler { safety: Arc, tools: Arc, extension_manager: Option>, - store: Option>, + store: Option, hooks: Arc, /// SSE manager for live job event streaming. sse_tx: Option>, @@ -780,10 +780,14 @@ mod tests { allow_local_tools: true, max_cost_per_day_cents: None, max_actions_per_hour: None, + max_cost_per_user_per_day_cents: None, max_tool_iterations: 10, auto_approve_tools: true, default_timezone: "UTC".to_string(), max_tokens_per_job, + multi_tenant: false, + max_llm_concurrent_per_user: None, + max_jobs_concurrent_per_user: None, }; let cm = Arc::new(ContextManager::new(5)); let llm: Arc = Arc::new(StubLlm); diff --git a/src/agent/self_repair.rs b/src/agent/self_repair.rs index 4e58cb15..050c2e90 100644 --- a/src/agent/self_repair.rs +++ b/src/agent/self_repair.rs @@ -8,8 +8,8 @@ use chrono::{DateTime, Utc}; use uuid::Uuid; use crate::context::{ContextManager, JobState}; -use crate::db::Database; use crate::error::RepairError; +use crate::tenant::AdminScope; use crate::tools::{BuildRequirement, Language, SoftwareBuilder, SoftwareType, ToolRegistry}; /// A job that has been detected as stuck. @@ -69,7 +69,7 @@ pub struct DefaultSelfRepair { /// Jobs in `InProgress` longer than this are treated as stuck. stuck_threshold: Duration, max_repair_attempts: u32, - store: Option>, + store: Option, builder: Option>, tools: Option>, } @@ -91,8 +91,8 @@ impl DefaultSelfRepair { } } - /// Add a Store for tool failure tracking. - pub fn with_store(mut self, store: Arc) -> Self { + /// Add an admin-scoped store for tool failure tracking. + pub fn with_store(mut self, store: AdminScope) -> Self { self.store = Some(store); self } @@ -806,7 +806,7 @@ mod tests { // Create self-repair with zero threshold (detect immediately), // wired with store, builder, and tools. let repair = DefaultSelfRepair::new(Arc::clone(&cm), Duration::from_secs(0), 3) - .with_store(Arc::clone(&db)) + .with_store(crate::tenant::AdminScope::new(Arc::clone(&db))) .with_builder( Arc::clone(&builder) as Arc, tools, diff --git a/src/agent/session.rs b/src/agent/session.rs index 45594922..6c873e46 100644 --- a/src/agent/session.rs +++ b/src/agent/session.rs @@ -16,8 +16,8 @@ use chrono::{DateTime, TimeDelta, Utc}; use serde::{Deserialize, Serialize}; use uuid::Uuid; -use crate::channels::web::util::truncate_preview; use crate::llm::{ChatMessage, ToolCall, generate_tool_call_id}; +use ironclaw_common::truncate_preview; /// A session containing one or more threads. #[derive(Debug, Clone, Serialize, Deserialize)] @@ -449,6 +449,7 @@ impl Thread { id: call_id.clone(), name: tc.name.clone(), arguments: tc.parameters.clone(), + reasoning: None, }) .collect(); @@ -522,7 +523,12 @@ impl Thread { && let Some(ref tcs) = assistant_msg.tool_calls { for tc in tcs { - turn.record_tool_call(&tc.name, tc.arguments.clone()); + turn.record_tool_call_with_reasoning( + &tc.name, + tc.arguments.clone(), + tc.reasoning.clone(), + Some(tc.id.clone()), + ); } } @@ -602,6 +608,10 @@ pub struct Turn { pub completed_at: Option>, /// Error message (if failed). pub error: Option, + /// Agent's reasoning narrative for this turn. + /// Cleaned via `clean_response` and sanitized through `SafetyLayer` before storage. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub narrative: Option, /// Transient image content parts for multimodal LLM input. /// Not serialized — images are only needed for the current LLM call. /// The text description in `user_input` persists for compaction/context. @@ -621,6 +631,7 @@ impl Turn { started_at: Utc::now(), completed_at: None, error: None, + narrative: None, image_content_parts: Vec::new(), } } @@ -656,6 +667,26 @@ impl Turn { parameters: params, result: None, error: None, + rationale: None, + tool_call_id: None, + }); + } + + /// Record a tool call with reasoning context. + pub fn record_tool_call_with_reasoning( + &mut self, + name: impl Into, + params: serde_json::Value, + rationale: Option, + tool_call_id: Option, + ) { + self.tool_calls.push(TurnToolCall { + name: name.into(), + parameters: params, + result: None, + error: None, + rationale, + tool_call_id, }); } @@ -672,6 +703,60 @@ impl Turn { call.error = Some(error.into()); } } + + /// Record a tool result by tool_call_id, with fallback to first pending call. + pub fn record_tool_result_for(&mut self, tool_call_id: &str, result: serde_json::Value) { + if let Some(call) = self + .tool_calls + .iter_mut() + .find(|c| c.tool_call_id.as_deref() == Some(tool_call_id)) + { + call.result = Some(result); + } else if let Some(call) = self + .tool_calls + .iter_mut() + .find(|c| c.result.is_none() && c.error.is_none()) + { + tracing::debug!( + tool_call_id = %tool_call_id, + fallback_tool = %call.name, + "tool_call_id not found, falling back to first pending call" + ); + call.result = Some(result); + } else { + tracing::warn!( + tool_call_id = %tool_call_id, + "Tool result dropped: no matching or pending tool call" + ); + } + } + + /// Record a tool error by tool_call_id, with fallback to first pending call. + pub fn record_tool_error_for(&mut self, tool_call_id: &str, error: impl Into) { + if let Some(call) = self + .tool_calls + .iter_mut() + .find(|c| c.tool_call_id.as_deref() == Some(tool_call_id)) + { + call.error = Some(error.into()); + } else if let Some(call) = self + .tool_calls + .iter_mut() + .find(|c| c.result.is_none() && c.error.is_none()) + { + tracing::debug!( + tool_call_id = %tool_call_id, + fallback_tool = %call.name, + "tool_call_id not found, falling back to first pending call" + ); + call.error = Some(error.into()); + } else { + tracing::warn!( + tool_call_id = %tool_call_id, + "Tool error dropped: no matching or pending tool call" + ); + } + } } /// Record of a tool call made during a turn. @@ -685,6 +770,12 @@ pub struct TurnToolCall { pub result: Option, /// Error from the tool (if failed). pub error: Option, + /// Agent's reasoning for choosing this tool. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub rationale: Option, + /// The tool_call_id from the LLM, for identity-based result matching. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tool_call_id: Option, } #[cfg(test)] @@ -1309,6 +1400,7 @@ mod tests { id: "call_0".to_string(), name: "search".to_string(), arguments: serde_json::json!({"q": "test"}), + reasoning: None, }; let messages = vec![ ChatMessage::user("Find test"), @@ -1339,6 +1431,7 @@ mod tests { id: "call_0".to_string(), name: "http".to_string(), arguments: serde_json::json!({}), + reasoning: None, }; let messages = vec![ ChatMessage::user("Fetch URL"), @@ -1404,11 +1497,13 @@ mod tests { id: "call_a".to_string(), name: "search".to_string(), arguments: serde_json::json!({"q": "data"}), + reasoning: None, }; let tc2 = ToolCall { id: "call_b".to_string(), name: "write".to_string(), arguments: serde_json::json!({"path": "out.txt"}), + reasoning: None, }; let messages = vec![ ChatMessage::user("Find and save"), @@ -1620,4 +1715,100 @@ mod tests { let merged = thread.drain_pending_messages().unwrap(); assert_eq!(merged, "failed batch\nnew msg"); } + + #[test] + fn test_record_tool_result_for_by_id() { + let mut turn = Turn::new(0, "test"); + turn.record_tool_call_with_reasoning( + "tool_a", + serde_json::json!({}), + None, + Some("id_a".into()), + ); + turn.record_tool_call_with_reasoning( + "tool_b", + serde_json::json!({}), + None, + Some("id_b".into()), + ); + + // Record result for second tool by ID + turn.record_tool_result_for("id_b", serde_json::json!("result_b")); + assert!(turn.tool_calls[0].result.is_none()); + assert_eq!( + turn.tool_calls[1].result.as_ref().unwrap(), + &serde_json::json!("result_b") + ); + } + + #[test] + fn test_record_tool_error_for_by_id() { + let mut turn = Turn::new(0, "test"); + turn.record_tool_call_with_reasoning( + "tool_a", + serde_json::json!({}), + None, + Some("id_a".into()), + ); + turn.record_tool_call_with_reasoning( + "tool_b", + serde_json::json!({}), + None, + Some("id_b".into()), + ); + + turn.record_tool_error_for("id_a", "failed"); + assert_eq!(turn.tool_calls[0].error.as_deref(), Some("failed")); + assert!(turn.tool_calls[1].error.is_none()); + } + + #[test] + fn test_record_tool_result_for_fallback_to_pending() { + let mut turn = Turn::new(0, "test"); + turn.record_tool_call_with_reasoning( + "tool_a", + serde_json::json!({}), + None, + Some("id_a".into()), + ); + turn.record_tool_call_with_reasoning( + "tool_b", + serde_json::json!({}), + None, + Some("id_b".into()), + ); + + // First tool already has a result + turn.tool_calls[0].result = Some(serde_json::json!("done")); + + // Unknown ID should fall back to first pending (tool_b) + turn.record_tool_result_for("unknown_id", serde_json::json!("fallback")); + assert_eq!( + turn.tool_calls[0].result.as_ref().unwrap(), + &serde_json::json!("done") + ); + assert_eq!( + turn.tool_calls[1].result.as_ref().unwrap(), + &serde_json::json!("fallback") + ); + } + + #[test] + fn test_record_tool_result_for_no_pending_is_noop() { + let mut turn = Turn::new(0, "test"); + turn.record_tool_call_with_reasoning( + "tool_a", + serde_json::json!({}), + None, + Some("id_a".into()), + ); + turn.tool_calls[0].result = Some(serde_json::json!("done")); + + // No pending calls, unknown ID — should be a no-op + turn.record_tool_result_for("unknown_id", serde_json::json!("lost")); + assert_eq!( + turn.tool_calls[0].result.as_ref().unwrap(), + &serde_json::json!("done") + ); + } } diff --git a/src/agent/session_manager.rs b/src/agent/session_manager.rs index 3bf20697..ae98b0b0 100644 --- a/src/agent/session_manager.rs +++ b/src/agent/session_manager.rs @@ -102,11 +102,30 @@ impl SessionManager { /// Resolve an external thread ID to an internal thread. /// /// Returns the session and thread ID. Creates both if they don't exist. + /// Delegates to [`resolve_thread_with_parsed_uuid`](Self::resolve_thread_with_parsed_uuid) + /// with `parsed_uuid: None`. pub async fn resolve_thread( &self, user_id: &str, channel: &str, external_thread_id: Option<&str>, + ) -> (Arc>, Uuid) { + self.resolve_thread_with_parsed_uuid(user_id, channel, external_thread_id, None) + .await + } + + /// Like [`resolve_thread`](Self::resolve_thread), but accepts a pre-parsed + /// UUID to skip redundant parsing when the caller has already validated + /// the external thread ID as a UUID (e.g. the approval routing path). + /// + /// Uses a single read-lock acquisition for both the key lookup and the UUID + /// adoption check to reduce contention under concurrent approval load. + pub async fn resolve_thread_with_parsed_uuid( + &self, + user_id: &str, + channel: &str, + external_thread_id: Option<&str>, + parsed_uuid: Option, ) -> (Arc>, Uuid) { let session = self.get_or_create_session(user_id).await; @@ -116,51 +135,65 @@ impl SessionManager { external_thread_id: external_thread_id.map(String::from), }; - // Check if we have a mapping - { + // Use pre-parsed UUID if available, otherwise parse from string. + let ext_uuid = parsed_uuid + .or_else(|| external_thread_id.and_then(|ext_tid| Uuid::parse_str(ext_tid).ok())); + + // Validate that parsed_uuid (if provided) is consistent with external_thread_id. + #[cfg(debug_assertions)] + if let (Some(parsed), Some(ext_tid)) = (&parsed_uuid, external_thread_id) { + debug_assert_eq!( + Uuid::parse_str(ext_tid).ok().as_ref(), + Some(parsed), + "parsed_uuid must be the parsed form of external_thread_id" + ); + } + + // Single read lock for both the key lookup and UUID adoption check + let adoptable_uuid = { let thread_map = self.thread_map.read().await; + + // Fast path: exact key match if let Some(&thread_id) = thread_map.get(&key) { - // Verify thread still exists in session let sess = session.lock().await; if sess.threads.contains_key(&thread_id) { return (Arc::clone(&session), thread_id); } } - } - // Check if external_thread_id is itself a known thread UUID that - // exists in the session but was never registered in the thread_map - // (e.g. created by chat_new_thread_handler or hydrated from DB). - // We only adopt it if no thread_map entry maps to this UUID — - // otherwise it belongs to a different channel scope. - if let Some(ext_tid) = external_thread_id - && let Ok(ext_uuid) = Uuid::parse_str(ext_tid) - { - let thread_map = self.thread_map.read().await; - let mapped_elsewhere = thread_map.values().any(|&v| v == ext_uuid); - drop(thread_map); + // UUID adoption check (still under the same read lock). + // If external_thread_id is a valid UUID not mapped elsewhere, + // it may be a thread created by chat_new_thread_handler or + // hydrated from DB that we can adopt. + // Only attempt adoption when external_thread_id is Some, preserving + // the invariant that None external_thread_id never triggers adoption. + if external_thread_id.is_some() { + ext_uuid.filter(|&uuid| !thread_map.values().any(|&v| v == uuid)) + } else { + None + } + }; // Single read lock dropped here - if !mapped_elsewhere { - let sess = session.lock().await; - if sess.threads.contains_key(&ext_uuid) { - drop(sess); + // If we found an adoptable UUID, verify it exists in session and acquire write lock + if let Some(ext_uuid) = adoptable_uuid { + let sess = session.lock().await; + if sess.threads.contains_key(&ext_uuid) { + drop(sess); - let mut thread_map = self.thread_map.write().await; - // Re-check after acquiring write lock to prevent race condition - // where another task mapped this UUID between our read and write. - if !thread_map.values().any(|&v| v == ext_uuid) { - thread_map.insert(key, ext_uuid); - drop(thread_map); - // Ensure undo manager exists - let mut undo_managers = self.undo_managers.write().await; - undo_managers - .entry(ext_uuid) - .or_insert_with(|| Arc::new(Mutex::new(UndoManager::new()))); - return (session, ext_uuid); - } - // If it was mapped elsewhere while we were unlocked, fall through - // to create a new thread, preserving channel isolation. + let mut thread_map = self.thread_map.write().await; + // Re-check after acquiring write lock to prevent race condition + // where another task mapped this UUID between our read and write. + if !thread_map.values().any(|&v| v == ext_uuid) { + thread_map.insert(key, ext_uuid); + drop(thread_map); + // Ensure undo manager exists + let mut undo_managers = self.undo_managers.write().await; + undo_managers + .entry(ext_uuid) + .or_insert_with(|| Arc::new(Mutex::new(UndoManager::new()))); + return (session, ext_uuid); } + // If mapped elsewhere while unlocked, fall through to create new thread } } @@ -909,6 +942,44 @@ mod tests { } } + #[tokio::test] + async fn test_resolve_thread_consolidates_read_path() { + // Verify that resolve_thread still correctly handles: + // 1. Fast path: key exists in thread_map + // 2. UUID adoption: external_thread_id is a UUID in session but not in map + // 3. New thread: neither path matches + use crate::agent::session::Thread; + + let manager = SessionManager::new(); + + // Case 1: Normal resolution creates thread and maps it + let (session1, tid1) = manager + .resolve_thread("user1", "chan1", Some("ext-1")) + .await; + // Resolving again with same key should return same thread (fast path) + let (_, tid1_again) = manager + .resolve_thread("user1", "chan1", Some("ext-1")) + .await; + assert_eq!(tid1, tid1_again); + + // Case 2: UUID adoption - insert a thread directly into session + let adopted_id = Uuid::new_v4(); + { + let mut sess = session1.lock().await; + let thread = Thread::with_id(adopted_id, sess.id); + sess.threads.insert(adopted_id, thread); + } + // Resolve with the UUID as external_thread_id -- should adopt it + let (_, resolved) = manager + .resolve_thread("user1", "chan1", Some(&adopted_id.to_string())) + .await; + assert_eq!(resolved, adopted_id); + + // Case 3: Different channel gets different thread + let (_, tid2) = manager.resolve_thread("user1", "chan2", None).await; + assert_ne!(tid1, tid2); + } + #[tokio::test] async fn test_resolve_thread_finds_existing_session_thread_by_uuid() { use crate::agent::session::{Session, Thread}; @@ -947,4 +1018,88 @@ mod tests { "should have exactly 1 thread, not a duplicate" ); } + + #[tokio::test] + async fn test_resolve_thread_with_pre_parsed_uuid_adopts_thread() { + use crate::agent::session::Thread; + + let manager = SessionManager::new(); + let (session, _) = manager.resolve_thread("user1", "chan1", None).await; + + // Manually insert a thread with a known UUID + let known_id = Uuid::new_v4(); + { + let mut sess = session.lock().await; + let thread = Thread::with_id(known_id, sess.id); + sess.threads.insert(known_id, thread); + } + + // Resolve with pre-parsed UUID -- should adopt it without re-parsing + let (_, resolved) = manager + .resolve_thread_with_parsed_uuid( + "user1", + "chan1", + Some(&known_id.to_string()), + Some(known_id), + ) + .await; + assert_eq!(resolved, known_id); + } + + #[tokio::test] + async fn test_resolve_thread_with_parsed_uuid_none_delegates_to_parse() { + use crate::agent::session::Thread; + + let manager = SessionManager::new(); + let (session, _) = manager.resolve_thread("user2", "chan2", None).await; + + // Insert a thread with a known UUID + let known_id = Uuid::new_v4(); + { + let mut sess = session.lock().await; + let thread = Thread::with_id(known_id, sess.id); + sess.threads.insert(known_id, thread); + } + + // Resolve with parsed_uuid=None but a valid UUID string -- should + // fall back to parsing the string and still adopt the thread + let (_, resolved) = manager + .resolve_thread_with_parsed_uuid("user2", "chan2", Some(&known_id.to_string()), None) + .await; + assert_eq!(resolved, known_id); + } + + #[tokio::test] + async fn test_resolve_thread_with_none_external_thread_id_does_not_adopt() { + use crate::agent::session::Thread; + + let manager = SessionManager::new(); + let (session, default_tid) = manager.resolve_thread("user3", "chan3", None).await; + + // Manually insert a thread with a known UUID (simulating a thread + // created by chat_new_thread_handler) + let known_id = Uuid::new_v4(); + { + let mut sess = session.lock().await; + let thread = Thread::with_id(known_id, sess.id); + sess.threads.insert(known_id, thread); + } + + // Resolve with external_thread_id=None but parsed_uuid=Some. + // This should NOT adopt the UUID — the old code prevented adoption + // when external_thread_id was None, and we preserve that invariant. + let (_, resolved) = manager + .resolve_thread_with_parsed_uuid("user3", "chan3", None, Some(known_id)) + .await; + + // Should return the existing default thread, not the injected UUID + assert_eq!( + resolved, default_tid, + "should return existing default thread when external_thread_id is None" + ); + assert_ne!( + resolved, known_id, + "should NOT adopt UUID when external_thread_id is None" + ); + } } diff --git a/src/agent/submission.rs b/src/agent/submission.rs index 8594c969..5a81e0bf 100644 --- a/src/agent/submission.rs +++ b/src/agent/submission.rs @@ -92,6 +92,17 @@ impl SubmissionParser { args: vec![], }; } + if lower == "/reasoning" || lower.starts_with("/reasoning ") { + let args: Vec = trimmed + .split_whitespace() + .skip(1) + .map(|s| s.to_string()) + .collect(); + return Submission::SystemCommand { + command: "reasoning".to_string(), + args, + }; + } if lower == "/restart" { tracing::debug!("[SubmissionParser::parse] Recognized /restart command"); return Submission::SystemCommand { diff --git a/src/agent/thread_ops.rs b/src/agent/thread_ops.rs index ddfd0c0f..a5288f68 100644 --- a/src/agent/thread_ops.rs +++ b/src/agent/thread_ops.rs @@ -16,12 +16,12 @@ use crate::agent::dispatcher::{ }; use crate::agent::session::{MAX_PENDING_MESSAGES, PendingApproval, Session, ThreadState}; use crate::agent::submission::SubmissionResult; -use crate::channels::web::util::truncate_preview; use crate::channels::{IncomingMessage, StatusUpdate}; use crate::context::JobContext; use crate::error::Error; use crate::llm::{ChatMessage, ToolCall}; use crate::tools::redact_params; +use ironclaw_common::truncate_preview; const FORGED_THREAD_ID_ERROR: &str = "Invalid or unauthorized thread ID."; @@ -175,6 +175,7 @@ impl Agent { pub(super) async fn process_user_input( &self, message: &IncomingMessage, + tenant: crate::tenant::TenantCtx, session: Arc>, thread_id: Uuid, content: &str, @@ -351,7 +352,7 @@ impl Agent { if let Some(intent) = self.router.route_command(&temp_message) { // Explicit command like /status, /job, /list - handle directly - return self.handle_job_or_command(intent, message).await; + return self.handle_job_or_command(intent, message, &tenant).await; } // Natural language goes through the agentic loop @@ -462,7 +463,7 @@ impl Agent { // Run the agentic tool execution loop let result = self - .run_agentic_loop(message, session.clone(), thread_id, turn_messages) + .run_agentic_loop(message, tenant, session.clone(), thread_id, turn_messages) .await; // Re-acquire lock and check if interrupted @@ -513,10 +514,10 @@ impl Agent { }; thread.complete_turn(&response); - let (turn_number, tool_calls) = thread + let (turn_number, tool_calls, narrative) = thread .turns .last() - .map(|t| (t.turn_number, t.tool_calls.clone())) + .map(|t| (t.turn_number, t.tool_calls.clone(), t.narrative.clone())) .unwrap_or_default(); let _ = self .channels @@ -534,6 +535,7 @@ impl Agent { &message.user_id, turn_number, &tool_calls, + narrative.as_deref(), ) .await; self.persist_assistant_response( @@ -725,7 +727,9 @@ impl Agent { /// /// Stored between the user and assistant messages so that /// `build_turns_from_db_messages` can reconstruct the tool call history. - /// Content is a JSON array of tool call summaries. + /// Content is a JSON object: `{ "calls": [...], "narrative": "..." }`. + /// The `calls` array contains tool call summaries with optional `rationale` + /// and `tool_call_id` fields. Legacy rows may be plain JSON arrays. pub(super) async fn persist_tool_calls( &self, thread_id: Uuid, @@ -733,6 +737,7 @@ impl Agent { user_id: &str, turn_number: usize, tool_calls: &[crate::agent::session::TurnToolCall], + narrative: Option<&str>, ) { if tool_calls.is_empty() { return; @@ -767,11 +772,30 @@ impl Agent { if let Some(ref error) = tc.error { obj["error"] = serde_json::Value::String(truncate_preview(error, 200)); } + if let Some(ref rationale) = tc.rationale { + obj["rationale"] = serde_json::Value::String(truncate_preview(rationale, 500)); + } + if let Some(ref tool_call_id) = tc.tool_call_id { + obj["tool_call_id"] = + serde_json::Value::String(truncate_preview(tool_call_id, 128)); + } obj }) .collect(); - let content = match serde_json::to_string(&summaries) { + // Wrap in an object with optional narrative so it can be reconstructed. + // safety: no byte-index slicing here; comment describes JSON shape + let wrapper = if let Some(n) = narrative { + serde_json::json!({ + "narrative": truncate_preview(n, 1000), + "calls": summaries, + }) + } else { + serde_json::json!({ + "calls": summaries, + }) + }; + let content = match serde_json::to_string(&wrapper) { Ok(c) => c, Err(e) => { tracing::warn!("Failed to serialize tool calls: {}", e); @@ -1104,9 +1128,12 @@ impl Agent { && let Some(turn) = thread.last_turn_mut() { if is_tool_error { - turn.record_tool_error(result_content.clone()); + turn.record_tool_error_for(&pending.tool_call_id, result_content.clone()); } else { - turn.record_tool_result(serde_json::json!(result_content)); + turn.record_tool_result_for( + &pending.tool_call_id, + serde_json::json!(result_content), + ); } } } @@ -1358,9 +1385,12 @@ impl Agent { && let Some(turn) = thread.last_turn_mut() { if is_deferred_error { - turn.record_tool_error(deferred_content.clone()); + turn.record_tool_error_for(&tc.id, deferred_content.clone()); } else { - turn.record_tool_result(serde_json::json!(deferred_content)); + turn.record_tool_result_for( + &tc.id, + serde_json::json!(deferred_content), + ); } } } @@ -1444,7 +1474,13 @@ impl Agent { // Continue the agentic loop (a tool was already executed this turn) let result = self - .run_agentic_loop(message, session.clone(), thread_id, context_messages) + .run_agentic_loop( + message, + self.tenant_ctx(&message.user_id).await, + session.clone(), + thread_id, + context_messages, + ) .await; // Handle the result @@ -1459,10 +1495,10 @@ impl Agent { let (response, suggestions) = crate::agent::dispatcher::extract_suggestions(&response); thread.complete_turn(&response); - let (turn_number, tool_calls) = thread + let (turn_number, tool_calls, narrative) = thread .turns .last() - .map(|t| (t.turn_number, t.tool_calls.clone())) + .map(|t| (t.turn_number, t.tool_calls.clone(), t.narrative.clone())) .unwrap_or_default(); // User message already persisted at turn start; save tool calls then assistant response self.persist_tool_calls( @@ -1471,6 +1507,7 @@ impl Agent { &message.user_id, turn_number, &tool_calls, + narrative.as_deref(), ) .await; self.persist_assistant_response( @@ -1816,7 +1853,20 @@ fn rebuild_chat_messages_from_db( "assistant" => result.push(ChatMessage::assistant(&msg.content)), "tool_calls" => { // Try to parse the enriched JSON and rebuild tool messages. - if let Ok(calls) = serde_json::from_str::>(&msg.content) { + // Supports two formats: + // - Old: plain JSON array of tool call summaries + // - New: wrapped object { "calls": [...], "narrative": "..." } + let calls: Vec = + match serde_json::from_str::(&msg.content) { + Ok(serde_json::Value::Array(arr)) => arr, + Ok(serde_json::Value::Object(obj)) => obj + .get("calls") + .and_then(|v| v.as_array()) + .cloned() + .unwrap_or_default(), + _ => Vec::new(), + }; + { if calls.is_empty() { continue; } @@ -1839,6 +1889,10 @@ fn rebuild_chat_messages_from_db( .get("parameters") .cloned() .unwrap_or(serde_json::json!({})), + reasoning: c + .get("rationale") + .and_then(|v| v.as_str()) + .map(String::from), }) .collect(); diff --git a/src/app.rs b/src/app.rs index 62f2345a..e1a614eb 100644 --- a/src/app.rs +++ b/src/app.rs @@ -312,13 +312,7 @@ impl AppBuilder { .create_provider(&self.config.llm.nearai.base_url, self.session.clone()); // Register memory tools if database is available - let workspace_user_id = self - .config - .channels - .gateway - .as_ref() - .map(|gw| gw.user_id.as_str()) - .unwrap_or("default"); + let workspace_user_id = self.config.owner_id.as_str(); let workspace = if let Some(ref db) = self.db { let emb_cache_config = EmbeddingCacheConfig { max_entries: self.config.embeddings.cache_size, @@ -905,6 +899,7 @@ impl AppBuilder { crate::agent::cost_guard::CostGuardConfig { max_cost_per_day_cents: self.config.agent.max_cost_per_day_cents, max_actions_per_hour: self.config.agent.max_actions_per_hour, + max_cost_per_user_per_day_cents: self.config.agent.max_cost_per_user_per_day_cents, }, )); diff --git a/src/channels/channel.rs b/src/channels/channel.rs index 9bcee12e..784b6bcf 100644 --- a/src/channels/channel.rs +++ b/src/channels/channel.rs @@ -265,6 +265,15 @@ impl OutgoingResponse { } } +/// A single tool decision within a reasoning update. +#[derive(Debug, Clone)] +pub struct ToolDecision { + /// Tool name. + pub tool_name: String, + /// Agent's reasoning for choosing this tool. + pub rationale: String, +} + /// Status update types for showing agent activity. #[derive(Debug, Clone)] pub enum StatusUpdate { @@ -333,6 +342,13 @@ pub enum StatusUpdate { }, /// Suggested follow-up messages for the user. Suggestions { suggestions: Vec }, + /// Agent reasoning update (why it chose specific tools). + ReasoningUpdate { + /// Human-readable summary of the agent's decision. + narrative: String, + /// Per-tool decisions. + decisions: Vec, + }, /// Per-turn token usage and cost summary (shown as subtle metadata). TurnCost { input_tokens: u64, diff --git a/src/channels/mod.rs b/src/channels/mod.rs index c0230692..46e25514 100644 --- a/src/channels/mod.rs +++ b/src/channels/mod.rs @@ -39,7 +39,7 @@ mod webhook_server; pub use channel::{ AttachmentKind, Channel, ChannelSecretUpdater, IncomingAttachment, IncomingMessage, - MessageStream, OutgoingResponse, StatusUpdate, routing_target_from_metadata, + MessageStream, OutgoingResponse, StatusUpdate, ToolDecision, routing_target_from_metadata, }; pub use http::{HttpChannel, HttpChannelState}; pub use manager::ChannelManager; diff --git a/src/channels/relay/client.rs b/src/channels/relay/client.rs index 81fbb56c..b67f2c5e 100644 --- a/src/channels/relay/client.rs +++ b/src/channels/relay/client.rs @@ -122,18 +122,32 @@ impl RelayClient { /// instance_url in chat-api. IronClaw only passes an optional CSRF nonce /// for validating the callback — no URLs. pub async fn initiate_oauth(&self, state_nonce: Option<&str>) -> Result { + let url = format!("{}/oauth/slack/auth", self.base_url); + tracing::debug!(relay_url = %url, "RelayClient::initiate_oauth: sending request"); let mut query: Vec<(&str, &str)> = vec![]; if let Some(nonce) = state_nonce { query.push(("state_nonce", nonce)); } let resp = self .http - .get(format!("{}/oauth/slack/auth", self.base_url)) + .get(&url) .bearer_auth(self.api_key.expose_secret()) .query(&query) .send() .await - .map_err(|e| RelayError::Network(e.to_string()))?; + .map_err(|e| { + tracing::warn!( + relay_url = %url, + error = %e, + "RelayClient::initiate_oauth: network request failed" + ); + RelayError::Network(e.to_string()) + })?; + tracing::debug!( + relay_url = %url, + status = %resp.status(), + "RelayClient::initiate_oauth: received response" + ); let status = resp.status(); if status.is_redirection() { @@ -224,20 +238,39 @@ impl RelayClient { method: &str, body: serde_json::Value, ) -> Result { + let url = format!("{}/proxy/{}/{}", self.base_url, provider, method); + tracing::debug!( + relay_url = %url, + provider = %provider, + method = %method, + "RelayClient::proxy_provider: sending request" + ); let query: Vec<(&str, &str)> = vec![("team_id", team_id)]; let resp = self .http - .post(format!("{}/proxy/{}/{}", self.base_url, provider, method)) + .post(&url) .bearer_auth(self.api_key.expose_secret()) .query(&query) .json(&body) .send() .await - .map_err(|e| RelayError::Network(e.to_string()))?; + .map_err(|e| { + tracing::warn!( + relay_url = %url, + error = %e, + "RelayClient::proxy_provider: network request failed" + ); + RelayError::Network(e.to_string()) + })?; if !resp.status().is_success() { let status = resp.status().as_u16(); let body = resp.text().await.unwrap_or_default(); + tracing::warn!( + relay_url = %url, + status = status, + "RelayClient::proxy_provider: channel-relay returned error" + ); return Err(RelayError::Api { status, message: body, @@ -255,23 +288,45 @@ impl RelayClient { /// 32-byte secret. Called once at activation time; the result is cached in the /// extension manager so subsequent calls to `relay_signing_secret()` use it. pub async fn get_signing_secret(&self, team_id: &str) -> Result, RelayError> { + let url = format!("{}/relay/signing-secret", self.base_url); + tracing::debug!( + relay_url = %url, + "RelayClient::get_signing_secret: fetching signing secret" + ); let resp = self .http - .get(format!("{}/relay/signing-secret", self.base_url)) + .get(&url) .bearer_auth(self.api_key.expose_secret()) .query(&[("team_id", team_id)]) .send() .await - .map_err(|e| RelayError::Network(e.to_string()))?; + .map_err(|e| { + tracing::warn!( + relay_url = %url, + error = %e, + "RelayClient::get_signing_secret: network request failed" + ); + RelayError::Network(e.to_string()) + })?; if !resp.status().is_success() { let status = resp.status().as_u16(); let body = resp.text().await.unwrap_or_default(); + tracing::warn!( + relay_url = %url, + status = status, + body = %body, + "RelayClient::get_signing_secret: channel-relay returned error" + ); return Err(RelayError::Api { status, message: body, }); } + tracing::debug!( + relay_url = %url, + "RelayClient::get_signing_secret: received successful response" + ); let body: serde_json::Value = resp .json() diff --git a/src/channels/repl.rs b/src/channels/repl.rs index 055dc3ad..41d73a8c 100644 --- a/src/channels/repl.rs +++ b/src/channels/repl.rs @@ -75,6 +75,7 @@ const SLASH_COMMANDS: &[&str] = &[ "/suggest", "/thread", "/resume", + "/reasoning", ]; /// Rustyline helper for slash-command tab completion. @@ -430,6 +431,18 @@ impl ReplChannel { let _ = execute!(stderr, terminal::Clear(terminal::ClearType::FromCursorDown)); } } + + async fn finish_single_message_turn(&self) { + if self.single_message.is_none() { + return; + } + + let tx = self.msg_tx.lock().ok().and_then(|mut guard| guard.take()); + if let Some(tx) = tx { + let msg = IncomingMessage::new("repl", &self.user_id, "/quit"); + let _ = tx.send(msg).await; + } + } } impl Default for ReplChannel { @@ -479,7 +492,9 @@ impl Channel for ReplChannel { async fn start(&self) -> Result { let (tx, rx) = mpsc::channel(32); - // Store tx so send_status can inject approval responses directly + // Approval prompts inject responses back through this sender. + // In single-message mode we keep it until the turn finishes, then + // drop it after enqueuing /quit so the receiver stream can close. if let Ok(mut guard) = self.msg_tx.lock() { *guard = Some(tx.clone()); } @@ -495,11 +510,10 @@ impl Channel for ReplChannel { // Single message mode: send it and return if let Some(msg) = single_message { - let incoming = IncomingMessage::new("repl", &user_id, &msg).with_timezone(&sys_tz); + let incoming = IncomingMessage::new("repl", &user_id, &msg) + .with_metadata(serde_json::json!({ "single_message_mode": true })) + .with_timezone(&sys_tz); let _ = tx.blocking_send(incoming); - // Ensure the agent exits after handling exactly one turn in -m mode, - // even when other channels (gateway/http) are enabled. - let _ = tx.blocking_send(IncomingMessage::new("repl", &user_id, "/quit")); return; } @@ -662,6 +676,7 @@ impl Channel for ReplChannel { println!(); println!(); self.stdin_locked.store(false, Ordering::Relaxed); + self.finish_single_message_turn().await; return Ok(()); } @@ -680,6 +695,7 @@ impl Channel for ReplChannel { println!(); // Unlock stdin so readline can resume self.stdin_locked.store(false, Ordering::Relaxed); + self.finish_single_message_turn().await; Ok(()) } @@ -779,6 +795,7 @@ impl Channel for ReplChannel { let msg_tx = Arc::clone(&self.msg_tx); let user_id = self.user_id.clone(); let lock_flag = Arc::clone(&self.stdin_locked); + let single_message_mode = self.single_message.is_some(); tokio::task::spawn_blocking(move || { let action = run_approval_selector(allow_always).unwrap_or("n"); // Unlock stdin so readline can resume after approval @@ -787,7 +804,12 @@ impl Channel for ReplChannel { return; }; if let Some(tx) = guard.as_ref() { - let msg = IncomingMessage::new("repl", &user_id, action); + let msg = if single_message_mode { + IncomingMessage::new("repl", &user_id, action) + .with_metadata(serde_json::json!({ "single_message_mode": true })) + } else { + IncomingMessage::new("repl", &user_id, action) + }; let _ = tx.blocking_send(msg); } }); @@ -841,6 +863,19 @@ impl Channel for ReplChannel { StatusUpdate::Suggestions { .. } => { // Suggestions are only rendered by the web gateway } + StatusUpdate::ReasoningUpdate { + narrative, + decisions, + } => { + if !narrative.is_empty() { + let display = truncate_for_preview(&narrative, CLI_STATUS_MAX); + eprintln!(" \x1b[94m\u{25B6} {display}\x1b[0m"); + } + for d in &decisions { + let display = truncate_for_preview(&d.rationale, CLI_STATUS_MAX); + eprintln!(" \x1b[90m\u{2192} {}: {display}\x1b[0m", d.tool_name); + } + } StatusUpdate::TurnCost { .. } => { // Cost display is handled by the TUI channel } @@ -875,6 +910,7 @@ impl Channel for ReplChannel { #[cfg(test)] mod tests { use futures::StreamExt; + use tokio::time::{Duration, timeout}; use super::*; @@ -883,16 +919,36 @@ mod tests { let repl = ReplChannel::with_message("hi".to_string()); let mut stream = repl.start().await.expect("repl start should succeed"); - let first = stream.next().await.expect("first message missing"); + let first = timeout(Duration::from_secs(1), stream.next()) + .await + .expect("timed out waiting for first message") + .expect("first message missing"); assert_eq!(first.channel, "repl"); assert_eq!(first.content, "hi"); - let second = stream.next().await.expect("quit message missing"); + assert!( + timeout(Duration::from_millis(100), stream.next()) + .await + .is_err(), + "single-message mode should wait for the turn to finish before quitting" + ); + + repl.respond(&first, OutgoingResponse::text("done")) + .await + .expect("respond should succeed"); + + let second = timeout(Duration::from_secs(1), stream.next()) + .await + .expect("timed out waiting for quit message") + .expect("quit message missing"); assert_eq!(second.channel, "repl"); assert_eq!(second.content, "/quit"); assert!( - stream.next().await.is_none(), + timeout(Duration::from_secs(1), stream.next()) + .await + .expect("timed out waiting for stream to close") + .is_none(), "stream should end after /quit" ); } diff --git a/src/channels/wasm/wrapper.rs b/src/channels/wasm/wrapper.rs index 65e4de88..a0f9689f 100644 --- a/src/channels/wasm/wrapper.rs +++ b/src/channels/wasm/wrapper.rs @@ -3061,6 +3061,20 @@ fn status_to_wit( }, // Suggestions and turn cost are web-gateway-only; skip for WASM channels StatusUpdate::Suggestions { .. } | StatusUpdate::TurnCost { .. } => return None, + StatusUpdate::ReasoningUpdate { + narrative, + decisions, + } => { + let mut msg = narrative.clone(); + for d in decisions { + msg.push_str(&format!("\n → {}: {}", d.tool_name, d.rationale)); + } + wit_channel::StatusUpdate { + status: wit_channel::StatusType::Status, + message: msg, + metadata_json, + } + } }) } diff --git a/src/channels/web/handlers/chat.rs b/src/channels/web/handlers/chat.rs index 9753c015..bc4e3dbc 100644 --- a/src/channels/web/handlers/chat.rs +++ b/src/channels/web/handlers/chat.rs @@ -175,7 +175,7 @@ pub async fn chat_auth_token_handler( if result.verification.is_some() { state.sse.broadcast_for_user( &user.user_id, - SseEvent::AuthRequired { + AppEvent::AuthRequired { extension_name: req.extension_name.clone(), instructions: Some(result.message), auth_url: None, @@ -187,7 +187,7 @@ pub async fn chat_auth_token_handler( state.sse.broadcast_for_user( &user.user_id, - SseEvent::AuthCompleted { + AppEvent::AuthCompleted { extension_name: req.extension_name.clone(), success: true, message: result.message, @@ -202,7 +202,7 @@ pub async fn chat_auth_token_handler( if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) { state.sse.broadcast_for_user( &user.user_id, - SseEvent::AuthRequired { + AppEvent::AuthRequired { extension_name: req.extension_name.clone(), instructions: Some(msg.clone()), auth_url: None, @@ -398,8 +398,10 @@ pub async fn chat_history_handler( truncate_preview(&s, 500) }), error: tc.error.clone(), + rationale: tc.rationale.clone(), }) .collect(), + narrative: t.narrative.clone(), }) .collect(); diff --git a/src/channels/web/handlers/webhooks.rs b/src/channels/web/handlers/webhooks.rs index 7b041a06..1fd78c66 100644 --- a/src/channels/web/handlers/webhooks.rs +++ b/src/channels/web/handlers/webhooks.rs @@ -54,10 +54,37 @@ fn validate_webhook_secret( /// /// This endpoint is **public** (no gateway auth token required) but protected /// by the per-routine webhook secret sent via the `X-Webhook-Secret` header. +/// +/// **Single-user/backward-compatible**: looks up routines by path across all +/// users. For multi-tenant isolation, use the user-scoped endpoint at +/// `/api/webhooks/u/{user_id}/{path}` instead. pub async fn webhook_trigger_handler( State(state): State>, Path(path): Path, headers: HeaderMap, +) -> Result, (StatusCode, String)> { + fire_webhook_inner(state, &path, None, &headers).await +} + +/// Handle incoming webhook POST to `/api/webhooks/u/{user_id}/{path}`. +/// +/// User-scoped variant for multi-tenant deployments. The `user_id` in the URL +/// restricts the routine lookup to that user only, preventing cross-user +/// webhook triggering even when paths collide. +pub async fn webhook_trigger_user_scoped_handler( + State(state): State>, + Path((user_id, path)): Path<(String, String)>, + headers: HeaderMap, +) -> Result, (StatusCode, String)> { + fire_webhook_inner(state, &path, Some(&user_id), &headers).await +} + +/// Shared webhook logic for both scoped and unscoped endpoints. +async fn fire_webhook_inner( + state: Arc, + path: &str, + user_id: Option<&str>, + headers: &HeaderMap, ) -> Result, (StatusCode, String)> { // Rate limit check if !state.webhook_rate_limiter.check() { @@ -72,9 +99,9 @@ pub async fn webhook_trigger_handler( "Database not available".to_string(), ))?; - // Targeted query instead of loading all routines + // Targeted query — when user_id is provided, restrict to that user's routines let routine = store - .get_webhook_routine_by_path(&path) + .get_webhook_routine_by_path(path, user_id) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .ok_or(( @@ -99,7 +126,7 @@ pub async fn webhook_trigger_handler( ))? }; - let run_id = engine.fire_webhook(routine.id, &path).await.map_err(|e| { + let run_id = engine.fire_webhook(routine.id, path).await.map_err(|e| { let status = match &e { crate::error::RoutineError::NotFound { .. } => StatusCode::NOT_FOUND, crate::error::RoutineError::Disabled { .. } diff --git a/src/channels/web/mod.rs b/src/channels/web/mod.rs index b26a7829..63aedaa0 100644 --- a/src/channels/web/mod.rs +++ b/src/channels/web/mod.rs @@ -58,7 +58,7 @@ use self::log_layer::{LogBroadcaster, LogLevelHandle}; use self::auth::MultiAuthState; use self::server::GatewayState; use self::sse::SseManager; -use self::types::SseEvent; +use self::types::AppEvent; /// Web gateway channel implementing the Channel trait. pub struct GatewayChannel { @@ -98,7 +98,8 @@ impl GatewayChannel { job_manager: None, prompt_queue: None, scheduler: None, - default_user_id: config.user_id.clone(), + owner_id: config.user_id.clone(), + default_sender_id: config.user_id.clone(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())), llm_provider: None, @@ -121,6 +122,22 @@ impl GatewayChannel { } } + /// Rebind the single-user auth identity to the durable owner scope while + /// preserving the configured gateway sender/routing identity. + pub fn with_owner_scope(mut self, owner_id: impl Into) -> Self { + let owner_id = owner_id.into(); + let single_user_token = if self.config.user_tokens.is_none() { + self.auth.first_token().map(ToOwned::to_owned) + } else { + None + }; + if let Some(token) = single_user_token { + self.auth = MultiAuthState::single(token, owner_id.clone()); + } + self.rebuild_state(|s| s.owner_id = owner_id); + self + } + /// Create a gateway channel with a pre-built multi-user auth state. pub fn new_multi_auth(config: GatewayConfig, auth: MultiAuthState) -> Self { let state = Arc::new(GatewayState { @@ -137,7 +154,8 @@ impl GatewayChannel { job_manager: None, prompt_queue: None, scheduler: None, - default_user_id: config.user_id.clone(), + owner_id: config.user_id.clone(), + default_sender_id: config.user_id.clone(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())), llm_provider: None, @@ -177,7 +195,8 @@ impl GatewayChannel { job_manager: self.state.job_manager.clone(), prompt_queue: self.state.prompt_queue.clone(), scheduler: self.state.scheduler.clone(), - default_user_id: self.state.default_user_id.clone(), + owner_id: self.state.owner_id.clone(), + default_sender_id: self.state.default_sender_id.clone(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: self.state.ws_tracker.clone(), llm_provider: self.state.llm_provider.clone(), @@ -367,7 +386,7 @@ impl Channel for GatewayChannel { self.state.sse.broadcast_for_user( &msg.user_id, - SseEvent::Response { + AppEvent::Response { content: response.content, thread_id, }, @@ -386,11 +405,11 @@ impl Channel for GatewayChannel { .and_then(|v| v.as_str()) .map(String::from); let event = match status { - StatusUpdate::Thinking(msg) => SseEvent::Thinking { + StatusUpdate::Thinking(msg) => AppEvent::Thinking { message: msg, thread_id: thread_id.clone(), }, - StatusUpdate::ToolStarted { name } => SseEvent::ToolStarted { + StatusUpdate::ToolStarted { name } => AppEvent::ToolStarted { name, thread_id: thread_id.clone(), }, @@ -399,23 +418,23 @@ impl Channel for GatewayChannel { success, error, parameters, - } => SseEvent::ToolCompleted { + } => AppEvent::ToolCompleted { name, success, error, parameters, thread_id: thread_id.clone(), }, - StatusUpdate::ToolResult { name, preview } => SseEvent::ToolResult { + StatusUpdate::ToolResult { name, preview } => AppEvent::ToolResult { name, preview, thread_id: thread_id.clone(), }, - StatusUpdate::StreamChunk(content) => SseEvent::StreamChunk { + StatusUpdate::StreamChunk(content) => AppEvent::StreamChunk { content, thread_id: thread_id.clone(), }, - StatusUpdate::Status(msg) => SseEvent::Status { + StatusUpdate::Status(msg) => AppEvent::Status { message: msg, thread_id: thread_id.clone(), }, @@ -423,7 +442,7 @@ impl Channel for GatewayChannel { job_id, title, browse_url, - } => SseEvent::JobStarted { + } => AppEvent::JobStarted { job_id, title, browse_url, @@ -434,7 +453,7 @@ impl Channel for GatewayChannel { description, parameters, allow_always, - } => SseEvent::ApprovalNeeded { + } => AppEvent::ApprovalNeeded { request_id, tool_name, description, @@ -448,7 +467,7 @@ impl Channel for GatewayChannel { instructions, auth_url, setup_url, - } => SseEvent::AuthRequired { + } => AppEvent::AuthRequired { extension_name, instructions, auth_url, @@ -458,25 +477,39 @@ impl Channel for GatewayChannel { extension_name, success, message, - } => SseEvent::AuthCompleted { + } => AppEvent::AuthCompleted { extension_name, success, message, }, - StatusUpdate::ImageGenerated { data_url, path } => SseEvent::ImageGenerated { + StatusUpdate::ImageGenerated { data_url, path } => AppEvent::ImageGenerated { data_url, path, thread_id: thread_id.clone(), }, - StatusUpdate::Suggestions { suggestions } => SseEvent::Suggestions { + StatusUpdate::Suggestions { suggestions } => AppEvent::Suggestions { suggestions, + thread_id: thread_id.clone(), + }, + StatusUpdate::ReasoningUpdate { + narrative, + decisions, + } => AppEvent::ReasoningUpdate { + narrative, + decisions: decisions + .into_iter() + .map(|d| crate::channels::web::types::ToolDecisionDto { + tool_name: d.tool_name, + rationale: d.rationale, + }) + .collect(), thread_id, }, StatusUpdate::TurnCost { input_tokens, output_tokens, cost_usd, - } => SseEvent::TurnCost { + } => AppEvent::TurnCost { input_tokens, output_tokens, cost_usd, @@ -512,7 +545,7 @@ impl Channel for GatewayChannel { }; self.state.sse.broadcast_for_user( user_id, - SseEvent::Response { + AppEvent::Response { content: response.content, thread_id, }, diff --git a/src/channels/web/openai_compat.rs b/src/channels/web/openai_compat.rs index 55b7c854..0c0f1a9e 100644 --- a/src/channels/web/openai_compat.rs +++ b/src/channels/web/openai_compat.rs @@ -231,6 +231,7 @@ pub fn convert_messages(messages: &[OpenAiMessage]) -> Result, name: tc.function.name.clone(), arguments: serde_json::from_str(&tc.function.arguments) .unwrap_or(serde_json::Value::Object(Default::default())), + reasoning: None, }) .collect(); Ok(ChatMessage::assistant_with_tool_calls( @@ -954,6 +955,7 @@ mod tests { id: "call_abc".to_string(), name: "search".to_string(), arguments: serde_json::json!({"query": "rust"}), + reasoning: None, }]; let converted = convert_tool_calls_to_openai(&calls); diff --git a/src/channels/web/server.rs b/src/channels/web/server.rs index 86c5468e..9a4b8480 100644 --- a/src/channels/web/server.rs +++ b/src/channels/web/server.rs @@ -345,8 +345,10 @@ pub struct GatewayState { pub job_manager: Option>, /// Prompt queue for Claude Code follow-up prompts. pub prompt_queue: Option, - /// Default user ID (fallback for non-request contexts like heartbeat/routines). - pub default_user_id: String, + /// Durable owner scope for persistence and unauthenticated callback flows. + pub owner_id: String, + /// Default sender/routing identity for gateway-originated messages. + pub default_sender_id: String, /// Shutdown signal sender. pub shutdown_tx: tokio::sync::RwLock>>, /// WebSocket connection tracker. @@ -412,6 +414,11 @@ pub async fn start_server( .route( "/api/webhooks/{path}", post(crate::channels::web::handlers::webhooks::webhook_trigger_handler), + ) + // User-scoped webhook endpoint for multi-tenant isolation + .route( + "/api/webhooks/u/{user_id}/{path}", + post(crate::channels::web::handlers::webhooks::webhook_trigger_user_scoped_handler), ); // Protected routes (require auth) @@ -775,7 +782,7 @@ async fn oauth_callback_handler( error = %error, "OAuth callback received with malformed state" ); - clear_auth_mode(&state, &state.default_user_id).await; + clear_auth_mode(&state, &state.owner_id).await; return oauth_error_page("IronClaw"); } }; @@ -811,7 +818,7 @@ async fn oauth_callback_handler( if let Some(ref sse) = flow.sse_manager { sse.broadcast_for_user( &flow.user_id, - SseEvent::AuthCompleted { + AppEvent::AuthCompleted { extension_name: flow.extension_name.clone(), success: false, message: "OAuth flow expired. Please try again.".to_string(), @@ -829,10 +836,10 @@ async fn oauth_callback_handler( let result: Result<(), String> = async { let token_response = if let Some(proxy_url) = &exchange_proxy_url { - let gateway_token = flow.gateway_token.as_deref().unwrap_or_default(); + let oauth_proxy_auth_token = flow.oauth_proxy_auth_token().unwrap_or_default(); oauth_defaults::exchange_via_proxy(oauth_defaults::ProxyTokenExchangeRequest { proxy_url, - gateway_token, + gateway_token: oauth_proxy_auth_token, token_url: &flow.token_url, client_id: &flow.client_id, client_secret: flow.client_secret.as_deref(), @@ -949,11 +956,11 @@ async fn oauth_callback_handler( message }; - // Broadcast SSE event to notify the web UI + // Broadcast event to notify the web UI if let Some(ref sse) = flow.sse_manager { sse.broadcast_for_user( &flow.user_id, - SseEvent::AuthCompleted { + AppEvent::AuthCompleted { extension_name: flow.extension_name, success, message: final_message.clone(), @@ -1136,7 +1143,7 @@ async fn slack_relay_oauth_callback_handler( let state_key = format!("relay:{}:oauth_state", DEFAULT_RELAY_NAME); let stored_state = match ext_mgr .secrets() - .get_decrypted(&state.default_user_id, &state_key) + .get_decrypted(&state.owner_id, &state_key) .await { Ok(secret) => secret.expose().to_string(), @@ -1160,10 +1167,7 @@ async fn slack_relay_oauth_callback_handler( } // Delete the nonce (one-time use) - let _ = ext_mgr - .secrets() - .delete(&state.default_user_id, &state_key) - .await; + let _ = ext_mgr.secrets().delete(&state.owner_id, &state_key).await; let result: Result<(), String> = async { let store = state.store.as_ref().ok_or_else(|| { @@ -1173,17 +1177,33 @@ async fn slack_relay_oauth_callback_handler( // Store team_id in settings let team_id_key = format!("relay:{}:team_id", DEFAULT_RELAY_NAME); - let _ = store - .set_setting( - &state.default_user_id, - &team_id_key, - &serde_json::json!(team_id), - ) - .await; + tracing::info!( + relay = DEFAULT_RELAY_NAME, + owner_id = %state.owner_id, + team_id_key = %team_id_key, + "relay OAuth callback: storing team_id in settings" + ); + store + .set_setting(&state.owner_id, &team_id_key, &serde_json::json!(team_id)) + .await + .map_err(|e| { + tracing::error!( + relay = DEFAULT_RELAY_NAME, + owner_id = %state.owner_id, + error = %e, + "relay OAuth callback: failed to persist team_id to settings store" + ); + format!("Failed to persist relay team_id: {e}") + })?; // Activate the relay channel + tracing::info!( + relay = DEFAULT_RELAY_NAME, + owner_id = %state.owner_id, + "relay OAuth callback: activating relay channel" + ); ext_mgr - .activate_stored_relay(DEFAULT_RELAY_NAME, &state.default_user_id) + .activate_stored_relay(DEFAULT_RELAY_NAME, &state.owner_id) .await .map_err(|e| format!("Failed to activate relay channel: {}", e))?; @@ -1202,8 +1222,8 @@ async fn slack_relay_oauth_callback_handler( } }; - // Broadcast SSE event to notify the web UI - state.sse.broadcast(SseEvent::AuthCompleted { + // Broadcast event to notify the web UI + state.sse.broadcast(AppEvent::AuthCompleted { extension_name: DEFAULT_RELAY_NAME.to_string(), success, message: message.clone(), @@ -1303,6 +1323,9 @@ async fn chat_send_handler( } let mut msg = IncomingMessage::new("gateway", &user.user_id, &req.content); + if state.owner_id != state.default_sender_id && user.user_id == state.owner_id { + msg = msg.with_sender_id(&state.default_sender_id); + } // Prefer timezone from JSON body, fall back to X-Timezone header let tz = req .timezone @@ -1404,6 +1427,9 @@ async fn chat_approval_handler( })?; let mut msg = IncomingMessage::new("gateway", &user.user_id, content); + if state.owner_id != state.default_sender_id && user.user_id == state.owner_id { + msg = msg.with_sender_id(&state.default_sender_id); + } if let Some(ref thread_id) = req.thread_id { msg = msg.with_thread(thread_id); @@ -1470,7 +1496,7 @@ async fn chat_auth_token_handler( if result.verification.is_some() { state.sse.broadcast_for_user( &user.user_id, - SseEvent::AuthRequired { + AppEvent::AuthRequired { extension_name: req.extension_name.clone(), instructions: Some(result.message), auth_url: None, @@ -1483,7 +1509,7 @@ async fn chat_auth_token_handler( state.sse.broadcast_for_user( &user.user_id, - SseEvent::AuthCompleted { + AppEvent::AuthCompleted { extension_name: req.extension_name.clone(), success: true, message: result.message, @@ -1492,7 +1518,7 @@ async fn chat_auth_token_handler( } else { state.sse.broadcast_for_user( &user.user_id, - SseEvent::AuthCompleted { + AppEvent::AuthCompleted { extension_name: req.extension_name.clone(), success: false, message: result.message, @@ -1508,7 +1534,7 @@ async fn chat_auth_token_handler( if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) { state.sse.broadcast_for_user( &user.user_id, - SseEvent::AuthRequired { + AppEvent::AuthRequired { extension_name: req.extension_name.clone(), instructions: Some(msg.clone()), auth_url: None, @@ -1724,8 +1750,10 @@ async fn chat_history_handler( truncate_preview(&s, 500) }), error: tc.error.clone(), + rationale: tc.rationale.clone(), }) .collect(), + narrative: t.narrative.clone(), }) .collect(); @@ -2174,6 +2202,11 @@ async fn extensions_activate_handler( AuthenticatedUser(user): AuthenticatedUser, Path(name): Path, ) -> Result, (StatusCode, String)> { + tracing::debug!( + extension = %name, + user_id = %user.user_id, + "extensions_activate_handler: received activate request" + ); let ext_mgr = state.extension_manager.as_ref().ok_or(( StatusCode::NOT_IMPLEMENTED, "Extension manager not available (secrets store required)".to_string(), @@ -2181,6 +2214,10 @@ async fn extensions_activate_handler( match ext_mgr.activate(&name, &user.user_id).await { Ok(result) => { + tracing::info!( + extension = %name, + "extensions_activate_handler: activation succeeded" + ); // 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). @@ -2199,6 +2236,13 @@ async fn extensions_activate_handler( crate::extensions::ExtensionError::AuthRequired ); + tracing::debug!( + extension = %name, + error = %activate_err, + needs_auth = needs_auth, + "extensions_activate_handler: activation failed, attempting auth fallback" + ); + if !needs_auth { return Ok(Json(ActionResponse::fail(activate_err.to_string()))); } @@ -2206,10 +2250,21 @@ async fn extensions_activate_handler( // Activation failed due to auth; try authenticating first. match ext_mgr.auth(&name, &user.user_id).await { Ok(auth_result) if auth_result.is_authenticated() => { + tracing::debug!( + extension = %name, + "extensions_activate_handler: auth reports authenticated, retrying activate" + ); // Auth succeeded, retry activation. 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()))), + Err(e) => { + tracing::warn!( + extension = %name, + error = %e, + "extensions_activate_handler: retry after auth still failed" + ); + Ok(Json(ActionResponse::fail(e.to_string()))) + } } } Ok(auth_result) => { @@ -2477,7 +2532,7 @@ async fn extensions_setup_submit_handler( // auth card or setup modal that was triggered by tool_auth/tool_activate. state.sse.broadcast_for_user( &user.user_id, - SseEvent::AuthCompleted { + AppEvent::AuthCompleted { extension_name: name.clone(), success: result.activated, message: resp.message.clone(), @@ -2979,7 +3034,8 @@ mod tests { store: None, job_manager: None, prompt_queue: None, - default_user_id: "test".to_string(), + owner_id: "test".to_string(), + default_sender_id: "test".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: None, llm_provider: None, @@ -3004,6 +3060,160 @@ mod tests { .with_state(state) } + #[derive(Clone, Debug)] + struct RecordedOauthProxyRequest { + authorization: Option, + form: std::collections::HashMap, + } + + #[derive(Clone)] + struct MockOauthProxyState { + requests: Arc>>, + } + + struct MockOauthProxyServer { + addr: std::net::SocketAddr, + requests: Arc>>, + shutdown_tx: Option>, + server_task: Option>, + } + + impl MockOauthProxyServer { + async fn start() -> Self { + async fn exchange_handler( + State(state): State, + headers: axum::http::HeaderMap, + axum::Form(form): axum::Form>, + ) -> Json { + state.requests.lock().await.push(RecordedOauthProxyRequest { + authorization: headers + .get(axum::http::header::AUTHORIZATION) + .and_then(|value| value.to_str().ok()) + .map(str::to_string), + form, + }); + Json(serde_json::json!({ + "access_token": "proxy-access-token", + "refresh_token": "proxy-refresh-token", + "expires_in": 7200 + })) + } + + let requests = Arc::new(tokio::sync::Mutex::new(Vec::new())); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("bind mock oauth proxy"); + let addr = listener.local_addr().expect("mock oauth proxy addr"); + let app = Router::new() + .route("/oauth/exchange", post(exchange_handler)) + .with_state(MockOauthProxyState { + requests: Arc::clone(&requests), + }); + let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel::<()>(); + let server_task = tokio::spawn(async move { + let _ = axum::serve(listener, app) + .with_graceful_shutdown(async { + let _ = shutdown_rx.await; + }) + .await; + }); + + Self { + addr, + requests, + shutdown_tx: Some(shutdown_tx), + server_task: Some(server_task), + } + } + + fn base_url(&self) -> String { + format!("http://{}", self.addr) + } + + async fn requests(&self) -> Vec { + self.requests.lock().await.clone() + } + + async fn shutdown(mut self) { + if let Some(tx) = self.shutdown_tx.take() { + let _ = tx.send(()); + } + if let Some(task) = self.server_task.take() { + let _ = task.await; + } + } + } + + impl Drop for MockOauthProxyServer { + fn drop(&mut self) { + if let Some(tx) = self.shutdown_tx.take() { + let _ = tx.send(()); + } + if let Some(task) = self.server_task.take() { + task.abort(); + } + } + } + + struct EnvVarGuard { + key: &'static str, + original: Option, + } + + impl Drop for EnvVarGuard { + fn drop(&mut self) { + // SAFETY: Tests use lock_env() to serialize environment access. + unsafe { + if let Some(ref value) = self.original { + std::env::set_var(self.key, value); + } else { + std::env::remove_var(self.key); + } + } + } + } + + fn set_env_var(key: &'static str, value: Option<&str>) -> EnvVarGuard { + let original = std::env::var(key).ok(); + // SAFETY: Tests use lock_env() to serialize environment access. + unsafe { + if let Some(value) = value { + std::env::set_var(key, value); + } else { + std::env::remove_var(key); + } + } + EnvVarGuard { key, original } + } + + fn fresh_pending_oauth_flow( + secrets: Arc, + sse_manager: Option>, + oauth_proxy_auth_token: Option, + ) -> crate::cli::oauth_defaults::PendingOAuthFlow { + crate::cli::oauth_defaults::PendingOAuthFlow { + extension_name: "test_tool".to_string(), + display_name: "Test Tool".to_string(), + token_url: "https://example.com/token".to_string(), + client_id: "client123".to_string(), + client_secret: None, + redirect_uri: "https://example.com/oauth/callback".to_string(), + code_verifier: Some("test-code-verifier".to_string()), + access_token_field: "access_token".to_string(), + secret_name: "test_token".to_string(), + provider: Some("google".to_string()), + validation_endpoint: None, + scopes: vec!["email".to_string()], + user_id: "test".to_string(), + secrets, + sse_manager, + gateway_token: oauth_proxy_auth_token, + token_exchange_extra_params: std::collections::HashMap::new(), + client_id_secret_name: None, + created_at: std::time::Instant::now(), + } + } + #[tokio::test] async fn test_extensions_setup_submit_returns_failure_when_not_activated() { use axum::body::Body; @@ -3170,7 +3380,7 @@ mod tests { Ok(Ok(scoped)) if matches!( scoped.event, - crate::channels::web::types::SseEvent::AuthRequired { .. } + crate::channels::web::types::AppEvent::AuthRequired { .. } ) => { panic!("verification responses should not emit auth_required SSE events") @@ -3452,7 +3662,7 @@ mod tests { assert_eq!(resp.status(), StatusCode::OK); match receiver.recv().await.expect("auth_completed event").event { - crate::channels::web::types::SseEvent::AuthCompleted { + crate::channels::web::types::AppEvent::AuthCompleted { extension_name, success, message, @@ -3661,6 +3871,284 @@ mod tests { ); } + #[tokio::test] + async fn test_oauth_callback_accepts_versioned_hosted_state_without_instance_name() { + use axum::body::Body; + use tower::ServiceExt; + + let secrets: Arc = + Arc::new(crate::secrets::InMemorySecretsStore::new(Arc::new( + crate::secrets::SecretsCrypto::new(secrecy::SecretString::from( + TEST_GATEWAY_CRYPTO_KEY.to_string(), + )) + .expect("crypto"), + ))); + let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(secrets.clone()); + + let Some(created_at) = expired_flow_created_at() else { + eprintln!( + "Skipping versioned OAuth state without instance test: monotonic uptime below expiry window" + ); + return; + }; + let flow = crate::cli::oauth_defaults::PendingOAuthFlow { + extension_name: "test_tool".to_string(), + display_name: "Test Tool".to_string(), + token_url: "https://example.com/token".to_string(), + client_id: "client123".to_string(), + client_secret: None, + redirect_uri: "https://example.com/oauth/callback".to_string(), + code_verifier: None, + access_token_field: "access_token".to_string(), + secret_name: "test_token".to_string(), + provider: None, + validation_endpoint: None, + scopes: vec![], + user_id: "test".to_string(), + secrets, + sse_manager: None, + gateway_token: None, + token_exchange_extra_params: std::collections::HashMap::new(), + client_id_secret_name: None, + created_at, + }; + + ext_mgr + .pending_oauth_flows() + .write() + .await + .insert("test_nonce".to_string(), flow); + + let state = test_gateway_state(Some(ext_mgr.clone())); + let app = test_oauth_router(state); + let versioned_state = + crate::cli::oauth_defaults::encode_hosted_oauth_state("test_nonce", None); + + let req = axum::http::Request::builder() + .uri(format!( + "/oauth/callback?code=fake_code&state={}", + urlencoding::encode(&versioned_state) + )) + .body(Body::empty()) + .expect("request"); + + let resp = ServiceExt::>::oneshot(app, req) + .await + .expect("response"); + assert_eq!(resp.status(), StatusCode::OK); + + let body = axum::body::to_bytes(resp.into_body(), 1024 * 64) + .await + .expect("body"); + let html = String::from_utf8_lossy(&body); + assert!(html.contains("Authorization Failed")); + assert!( + ext_mgr + .pending_oauth_flows() + .read() + .await + .get("test_nonce") + .is_none() + ); + } + + #[allow(clippy::await_holding_lock)] + #[tokio::test] + async fn test_oauth_callback_happy_path_with_gateway_token_fallback() { + use axum::body::Body; + use tower::ServiceExt; + + let proxy = MockOauthProxyServer::start().await; + // Keep the process-wide env locked for the full callback so the handler + // sees a stable proxy URL/token configuration throughout the test. + let _env_guard = crate::config::helpers::lock_env(); + let _exchange_url_guard = + set_env_var("IRONCLAW_OAUTH_EXCHANGE_URL", Some(&proxy.base_url())); + let _proxy_auth_guard = set_env_var("IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", None); + let _gateway_token_guard = set_env_var("GATEWAY_AUTH_TOKEN", Some("gateway-test-token")); + + let secrets = test_secrets_store(); + let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(Arc::clone(&secrets)); + let sse_mgr = Arc::new(SseManager::new()); + let mut receiver = sse_mgr.sender().subscribe(); + let flow = fresh_pending_oauth_flow( + Arc::clone(&secrets), + Some(Arc::clone(&sse_mgr)), + crate::cli::oauth_defaults::oauth_proxy_auth_token(), + ); + + ext_mgr + .pending_oauth_flows() + .write() + .await + .insert("test_nonce".to_string(), flow); + + let state = test_gateway_state(Some(ext_mgr.clone())); + let app = test_oauth_router(state); + let versioned_state = + crate::cli::oauth_defaults::encode_hosted_oauth_state("test_nonce", Some("myinstance")); + + let req = axum::http::Request::builder() + .uri(format!( + "/oauth/callback?code=fake_code&state={}", + urlencoding::encode(&versioned_state) + )) + .body(Body::empty()) + .expect("request"); + + let resp = ServiceExt::>::oneshot(app, req) + .await + .expect("response"); + assert_eq!(resp.status(), StatusCode::OK); + + let body = axum::body::to_bytes(resp.into_body(), 1024 * 64) + .await + .expect("body"); + let html = String::from_utf8_lossy(&body); + assert!(html.contains("Test Tool Connected")); + + let requests = proxy.requests().await; + assert_eq!(requests.len(), 1); + assert_eq!( + requests[0].authorization.as_deref(), + Some("Bearer gateway-test-token") + ); + assert_eq!( + requests[0].form.get("code").map(String::as_str), + Some("fake_code") + ); + assert_eq!( + requests[0].form.get("code_verifier").map(String::as_str), + Some("test-code-verifier") + ); + + let access_token = secrets + .get_decrypted("test", "test_token") + .await + .expect("access token stored"); + assert_eq!(access_token.expose(), "proxy-access-token"); + + let refresh_token = secrets + .get_decrypted("test", "test_token_refresh_token") + .await + .expect("refresh token stored"); + assert_eq!(refresh_token.expose(), "proxy-refresh-token"); + + match receiver.recv().await.expect("auth_completed event").event { + crate::channels::web::types::AppEvent::AuthCompleted { + extension_name, + success, + .. + } => { + assert_eq!(extension_name, "test_tool"); + assert!(success, "OAuth callback should broadcast success"); + } + event => panic!("expected AuthCompleted event, got {event:?}"), + } + + proxy.shutdown().await; + } + + #[allow(clippy::await_holding_lock)] + #[tokio::test] + async fn test_oauth_callback_happy_path_with_dedicated_proxy_auth_token() { + use axum::body::Body; + use tower::ServiceExt; + + let proxy = MockOauthProxyServer::start().await; + // Keep the process-wide env locked for the full callback so the handler + // sees a stable proxy URL/token configuration throughout the test. + let _env_guard = crate::config::helpers::lock_env(); + let _exchange_url_guard = + set_env_var("IRONCLAW_OAUTH_EXCHANGE_URL", Some(&proxy.base_url())); + let _proxy_auth_guard = set_env_var( + "IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", + Some("shared-oauth-proxy-secret"), + ); + let _gateway_token_guard = set_env_var("GATEWAY_AUTH_TOKEN", None); + + let secrets = test_secrets_store(); + let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(Arc::clone(&secrets)); + let sse_mgr = Arc::new(SseManager::new()); + let mut receiver = sse_mgr.sender().subscribe(); + let flow = fresh_pending_oauth_flow( + Arc::clone(&secrets), + Some(Arc::clone(&sse_mgr)), + crate::cli::oauth_defaults::oauth_proxy_auth_token(), + ); + + ext_mgr + .pending_oauth_flows() + .write() + .await + .insert("test_nonce".to_string(), flow); + + let state = test_gateway_state(Some(ext_mgr.clone())); + let app = test_oauth_router(state); + let versioned_state = + crate::cli::oauth_defaults::encode_hosted_oauth_state("test_nonce", None); + + let req = axum::http::Request::builder() + .uri(format!( + "/oauth/callback?code=fake_code&state={}", + urlencoding::encode(&versioned_state) + )) + .body(Body::empty()) + .expect("request"); + + let resp = ServiceExt::>::oneshot(app, req) + .await + .expect("response"); + assert_eq!(resp.status(), StatusCode::OK); + + let body = axum::body::to_bytes(resp.into_body(), 1024 * 64) + .await + .expect("body"); + let html = String::from_utf8_lossy(&body); + assert!(html.contains("Test Tool Connected")); + + let requests = proxy.requests().await; + assert_eq!(requests.len(), 1); + assert_eq!( + requests[0].authorization.as_deref(), + Some("Bearer shared-oauth-proxy-secret") + ); + assert_eq!( + requests[0].form.get("code").map(String::as_str), + Some("fake_code") + ); + assert_eq!( + requests[0].form.get("code_verifier").map(String::as_str), + Some("test-code-verifier") + ); + + let access_token = secrets + .get_decrypted("test", "test_token") + .await + .expect("access token stored"); + assert_eq!(access_token.expose(), "proxy-access-token"); + + let refresh_token = secrets + .get_decrypted("test", "test_token_refresh_token") + .await + .expect("refresh token stored"); + assert_eq!(refresh_token.expose(), "proxy-refresh-token"); + + match receiver.recv().await.expect("auth_completed event").event { + crate::channels::web::types::AppEvent::AuthCompleted { + extension_name, + success, + .. + } => { + assert_eq!(extension_name, "test_tool"); + assert!(success, "OAuth callback should broadcast success"); + } + event => panic!("expected AuthCompleted event, got {event:?}"), + } + + proxy.shutdown().await; + } + // --- Slack relay OAuth CSRF tests --- fn test_relay_oauth_router(state: Arc) -> Router { diff --git a/src/channels/web/sse.rs b/src/channels/web/sse.rs index 46841e19..e36cceab 100644 --- a/src/channels/web/sse.rs +++ b/src/channels/web/sse.rs @@ -11,7 +11,7 @@ use tokio::sync::broadcast; use tokio_stream::StreamExt; use tokio_stream::wrappers::BroadcastStream; -use crate::channels::web::types::SseEvent; +use crate::channels::web::types::AppEvent; /// Maximum number of concurrent SSE/WebSocket connections. /// Prevents resource exhaustion from connection flooding. @@ -25,7 +25,7 @@ const MAX_CONNECTIONS: u64 = 100; #[derive(Debug, Clone)] pub(crate) struct ScopedEvent { pub(crate) user_id: Option, - pub(crate) event: SseEvent, + pub(crate) event: AppEvent, } /// Manages SSE broadcast to all connected browser tabs. @@ -75,7 +75,7 @@ impl SseManager { } /// Broadcast an event to all connected clients (global/unscoped). - pub fn broadcast(&self, event: SseEvent) { + pub fn broadcast(&self, event: AppEvent) { let _ = self.tx.send(ScopedEvent { user_id: None, event, @@ -86,7 +86,7 @@ impl SseManager { /// /// Only subscribers for this user_id (or unscoped subscribers) will /// receive the event. - pub fn broadcast_for_user(&self, user_id: &str, event: SseEvent) { + pub fn broadcast_for_user(&self, user_id: &str, event: AppEvent) { let _ = self.tx.send(ScopedEvent { user_id: Some(user_id.to_string()), event, @@ -108,7 +108,7 @@ impl SseManager { pub fn subscribe_raw( &self, user_id: Option, - ) -> Option + Send + 'static + use<>> { + ) -> Option + Send + 'static + use<>> { // Atomically increment only if below the limit. This prevents // concurrent callers from overshooting max_connections. let counter = Arc::clone(&self.connection_count); @@ -186,30 +186,7 @@ impl SseManager { return None; } }; - let event_type = match &event { - SseEvent::Response { .. } => "response", - SseEvent::Thinking { .. } => "thinking", - SseEvent::ToolStarted { .. } => "tool_started", - SseEvent::ToolCompleted { .. } => "tool_completed", - SseEvent::ToolResult { .. } => "tool_result", - SseEvent::StreamChunk { .. } => "stream_chunk", - SseEvent::Status { .. } => "status", - SseEvent::ApprovalNeeded { .. } => "approval_needed", - SseEvent::AuthRequired { .. } => "auth_required", - SseEvent::AuthCompleted { .. } => "auth_completed", - SseEvent::Error { .. } => "error", - SseEvent::JobStarted { .. } => "job_started", - SseEvent::JobMessage { .. } => "job_message", - SseEvent::JobToolUse { .. } => "job_tool_use", - SseEvent::JobToolResult { .. } => "job_tool_result", - SseEvent::JobStatus { .. } => "job_status", - SseEvent::JobResult { .. } => "job_result", - SseEvent::Heartbeat => "heartbeat", - SseEvent::ImageGenerated { .. } => "image_generated", - SseEvent::Suggestions { .. } => "suggestions", - SseEvent::TurnCost { .. } => "turn_cost", - SseEvent::ExtensionStatus { .. } => "extension_status", - }; + let event_type = event.event_type(); Some(Ok(Event::default().event(event_type).data(data))) }); @@ -272,7 +249,7 @@ mod tests { fn test_broadcast_without_receivers() { let manager = SseManager::new(); // Should not panic even with no receivers - manager.broadcast(SseEvent::Heartbeat); + manager.broadcast(AppEvent::Heartbeat); } #[tokio::test] @@ -280,14 +257,14 @@ mod tests { let manager = SseManager::new(); let mut stream = Box::pin(manager.subscribe_raw(None).expect("should subscribe")); - manager.broadcast(SseEvent::Status { + manager.broadcast(AppEvent::Status { message: "test".to_string(), thread_id: None, }); let event = stream.next().await.unwrap(); match event { - SseEvent::Status { message, .. } => assert_eq!(message, "test"), + AppEvent::Status { message, .. } => assert_eq!(message, "test"), _ => panic!("unexpected event type"), } } @@ -299,14 +276,14 @@ mod tests { assert_eq!(manager.connection_count(), 1); - manager.broadcast(SseEvent::Thinking { + manager.broadcast(AppEvent::Thinking { message: "working".to_string(), thread_id: None, }); let event = stream.next().await.unwrap(); match event { - SseEvent::Thinking { message, .. } => assert_eq!(message, "working"), + AppEvent::Thinking { message, .. } => assert_eq!(message, "working"), _ => panic!("Expected Thinking event"), } } @@ -329,12 +306,12 @@ mod tests { let mut s2 = Box::pin(manager.subscribe_raw(None).expect("should subscribe")); assert_eq!(manager.connection_count(), 2); - manager.broadcast(SseEvent::Heartbeat); + manager.broadcast(AppEvent::Heartbeat); let e1 = s1.next().await.unwrap(); let e2 = s2.next().await.unwrap(); - assert!(matches!(e1, SseEvent::Heartbeat)); - assert!(matches!(e2, SseEvent::Heartbeat)); + assert!(matches!(e1, AppEvent::Heartbeat)); + assert!(matches!(e2, AppEvent::Heartbeat)); drop(s1); assert_eq!(manager.connection_count(), 1); @@ -373,25 +350,25 @@ mod tests { // Send event scoped to alice manager.broadcast_for_user( "alice", - SseEvent::Status { + AppEvent::Status { message: "alice only".to_string(), thread_id: None, }, ); // Send global event - manager.broadcast(SseEvent::Heartbeat); + manager.broadcast(AppEvent::Heartbeat); // Alice gets her scoped event let e = alice.next().await.unwrap(); - assert!(matches!(e, SseEvent::Status { .. })); + assert!(matches!(e, AppEvent::Status { .. })); // Alice also gets the global heartbeat let e = alice.next().await.unwrap(); - assert!(matches!(e, SseEvent::Heartbeat)); + assert!(matches!(e, AppEvent::Heartbeat)); // Bob only gets the global heartbeat (alice's event was filtered) let e = bob.next().await.unwrap(); // safety: test-only - assert!(matches!(e, SseEvent::Heartbeat)); // safety: test assertion + assert!(matches!(e, AppEvent::Heartbeat)); // safety: test assertion } } diff --git a/src/channels/web/test_helpers.rs b/src/channels/web/test_helpers.rs index 802512a6..0f7e5d12 100644 --- a/src/channels/web/test_helpers.rs +++ b/src/channels/web/test_helpers.rs @@ -76,7 +76,8 @@ impl TestGatewayBuilder { store: None, job_manager: None, prompt_queue: None, - default_user_id: self.user_id, + owner_id: self.user_id.clone(), + default_sender_id: self.user_id, shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: self.llm_provider, diff --git a/src/channels/web/tests/multi_tenant.rs b/src/channels/web/tests/multi_tenant.rs index 55010831..335f841c 100644 --- a/src/channels/web/tests/multi_tenant.rs +++ b/src/channels/web/tests/multi_tenant.rs @@ -16,6 +16,7 @@ use axum::routing::{delete, get, post}; use tower::ServiceExt; use uuid::Uuid; +use crate::channels::web::GatewayChannel; use crate::channels::web::auth::{ AuthenticatedUser, MultiAuthState, UserIdentity, auth_middleware, }; @@ -23,6 +24,7 @@ use crate::channels::web::server::{ ActiveConfigSnapshot, GatewayState, PerUserRateLimiter, PromptQueue, RateLimiter, WorkspacePool, }; use crate::channels::web::sse::SseManager; +use crate::config::GatewayConfig; // ── Helpers ──────────────────────────────────────────────────────────── @@ -64,7 +66,8 @@ fn build_state( store, job_manager: None, prompt_queue, - default_user_id: "test".to_string(), + owner_id: "test".to_string(), + default_sender_id: "test".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: None, llm_provider: None, @@ -82,6 +85,40 @@ fn build_state( }) } +fn gateway_config() -> GatewayConfig { + GatewayConfig { + host: "127.0.0.1".to_string(), + port: 3000, + auth_token: Some("gateway-auth".to_string()), + user_id: "gateway-sender".to_string(), + workspace_read_scopes: Vec::new(), + memory_layers: Vec::new(), + user_tokens: None, + } +} + +#[test] +fn with_owner_scope_updates_gateway_owner_scope_in_multi_user_mode() { + let mut gateway = GatewayChannel::new(gateway_config()); + gateway.auth = two_user_auth(); + gateway.config.user_tokens = Some(HashMap::new()); + let gateway = gateway.with_owner_scope("owner-scope"); + + assert_eq!(gateway.state.owner_id, "owner-scope"); + assert_eq!(gateway.state.default_sender_id, "gateway-sender"); + + let alice = gateway + .auth + .authenticate("tok-alice") + .expect("alice token should remain valid"); + let bob = gateway + .auth + .authenticate("tok-bob") + .expect("bob token should remain valid"); + assert_eq!(alice.user_id, "alice"); + assert_eq!(bob.user_id, "bob"); +} + /// Create a libSQL-backed test database in a temporary directory. /// /// Returns the database and a `TempDir` guard — the database file is diff --git a/src/channels/web/types.rs b/src/channels/web/types.rs index c5a4f67f..21b48461 100644 --- a/src/channels/web/types.rs +++ b/src/channels/web/types.rs @@ -63,6 +63,9 @@ pub struct TurnInfo { pub started_at: String, pub completed_at: Option, pub tool_calls: Vec, + /// Agent's reasoning narrative for this turn. + #[serde(skip_serializing_if = "Option::is_none")] + pub narrative: Option, } #[derive(Debug, Serialize)] @@ -74,6 +77,9 @@ pub struct ToolCallInfo { pub result_preview: Option, #[serde(skip_serializing_if = "Option::is_none")] pub error: Option, + /// Agent's reasoning for choosing this tool. + #[serde(skip_serializing_if = "Option::is_none")] + pub rationale: Option, } #[derive(Debug, Serialize)] @@ -114,165 +120,9 @@ pub struct ApprovalRequest { pub thread_id: Option, } -// --- SSE Event Types --- +// --- App Event (re-exported from ironclaw_common) --- -#[derive(Debug, Clone, Serialize)] -#[serde(tag = "type")] -pub enum SseEvent { - #[serde(rename = "response")] - Response { content: String, thread_id: String }, - #[serde(rename = "thinking")] - Thinking { - message: String, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - }, - #[serde(rename = "tool_started")] - ToolStarted { - name: String, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - }, - #[serde(rename = "tool_completed")] - ToolCompleted { - name: String, - success: bool, - #[serde(skip_serializing_if = "Option::is_none")] - error: Option, - #[serde(skip_serializing_if = "Option::is_none")] - parameters: Option, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - }, - #[serde(rename = "tool_result")] - ToolResult { - name: String, - preview: String, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - }, - #[serde(rename = "stream_chunk")] - StreamChunk { - content: String, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - }, - #[serde(rename = "status")] - Status { - message: String, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - }, - #[serde(rename = "job_started")] - JobStarted { - job_id: String, - title: String, - browse_url: String, - }, - #[serde(rename = "approval_needed")] - ApprovalNeeded { - request_id: String, - tool_name: String, - description: String, - parameters: String, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - /// Whether the "always" auto-approve option should be shown. - allow_always: bool, - }, - #[serde(rename = "auth_required")] - AuthRequired { - extension_name: String, - #[serde(skip_serializing_if = "Option::is_none")] - instructions: Option, - #[serde(skip_serializing_if = "Option::is_none")] - auth_url: Option, - #[serde(skip_serializing_if = "Option::is_none")] - setup_url: Option, - }, - #[serde(rename = "auth_completed")] - AuthCompleted { - extension_name: String, - success: bool, - message: String, - }, - #[serde(rename = "error")] - Error { - message: String, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - }, - #[serde(rename = "heartbeat")] - Heartbeat, - - // Sandbox job streaming events (worker + Claude Code bridge) - #[serde(rename = "job_message")] - JobMessage { - job_id: String, - role: String, - content: String, - }, - #[serde(rename = "job_tool_use")] - JobToolUse { - job_id: String, - tool_name: String, - input: serde_json::Value, - }, - #[serde(rename = "job_tool_result")] - JobToolResult { - job_id: String, - tool_name: String, - output: String, - }, - #[serde(rename = "job_status")] - JobStatus { job_id: String, message: String }, - #[serde(rename = "job_result")] - JobResult { - job_id: String, - status: String, - #[serde(skip_serializing_if = "Option::is_none")] - session_id: Option, - #[serde(skip_serializing_if = "Option::is_none")] - fallback_deliverable: Option, - }, - - /// An image was generated by a tool. - #[serde(rename = "image_generated")] - ImageGenerated { - data_url: String, - #[serde(skip_serializing_if = "Option::is_none")] - path: Option, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - }, - - /// Suggested follow-up messages for the user. - #[serde(rename = "suggestions")] - Suggestions { - suggestions: Vec, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - }, - - /// Per-turn token usage and cost summary. - #[serde(rename = "turn_cost")] - TurnCost { - input_tokens: u64, - output_tokens: u64, - cost_usd: String, - #[serde(skip_serializing_if = "Option::is_none")] - thread_id: Option, - }, - - /// Extension activation status change (WASM channels). - #[serde(rename = "extension_status")] - ExtensionStatus { - extension_name: String, - status: String, - #[serde(skip_serializing_if = "Option::is_none")] - message: Option, - }, -} +pub use ironclaw_common::{AppEvent, ToolDecisionDto}; // --- Memory --- @@ -787,32 +637,9 @@ pub enum WsServerMessage { } impl WsServerMessage { - /// Create a WsServerMessage from an SseEvent. - pub fn from_sse_event(event: &SseEvent) -> Self { - let event_type = match event { - SseEvent::Response { .. } => "response", - SseEvent::Thinking { .. } => "thinking", - SseEvent::ToolStarted { .. } => "tool_started", - SseEvent::ToolCompleted { .. } => "tool_completed", - SseEvent::ToolResult { .. } => "tool_result", - SseEvent::StreamChunk { .. } => "stream_chunk", - SseEvent::Status { .. } => "status", - SseEvent::JobStarted { .. } => "job_started", - SseEvent::ApprovalNeeded { .. } => "approval_needed", - SseEvent::AuthRequired { .. } => "auth_required", - SseEvent::AuthCompleted { .. } => "auth_completed", - SseEvent::Error { .. } => "error", - SseEvent::Heartbeat => "heartbeat", - SseEvent::JobMessage { .. } => "job_message", - SseEvent::JobToolUse { .. } => "job_tool_use", - SseEvent::JobToolResult { .. } => "job_tool_result", - SseEvent::JobStatus { .. } => "job_status", - SseEvent::JobResult { .. } => "job_result", - SseEvent::ImageGenerated { .. } => "image_generated", - SseEvent::Suggestions { .. } => "suggestions", - SseEvent::TurnCost { .. } => "turn_cost", - SseEvent::ExtensionStatus { .. } => "extension_status", - }; + /// Create a WsServerMessage from an AppEvent. + pub fn from_app_event(event: &AppEvent) -> Self { + let event_type = event.event_type(); let data = serde_json::to_value(event).unwrap_or(serde_json::Value::Null); WsServerMessage::Event { event_type: event_type.to_string(), @@ -1104,12 +931,12 @@ mod tests { } #[test] - fn test_ws_server_from_sse_response() { - let sse = SseEvent::Response { + fn test_ws_server_from_app_event_response() { + let event = AppEvent::Response { content: "hello".to_string(), thread_id: "t1".to_string(), }; - let ws = WsServerMessage::from_sse_event(&sse); + let ws = WsServerMessage::from_app_event(&event); match ws { WsServerMessage::Event { event_type, data } => { assert_eq!(event_type, "response"); @@ -1121,12 +948,12 @@ mod tests { } #[test] - fn test_ws_server_from_sse_thinking() { - let sse = SseEvent::Thinking { + fn test_ws_server_from_app_event_thinking() { + let event = AppEvent::Thinking { message: "reasoning...".to_string(), thread_id: None, }; - let ws = WsServerMessage::from_sse_event(&sse); + let ws = WsServerMessage::from_app_event(&event); match ws { WsServerMessage::Event { event_type, data } => { assert_eq!(event_type, "thinking"); @@ -1137,8 +964,8 @@ mod tests { } #[test] - fn test_ws_server_from_sse_approval_needed() { - let sse = SseEvent::ApprovalNeeded { + fn test_ws_server_from_app_event_approval_needed() { + let event = AppEvent::ApprovalNeeded { request_id: "r1".to_string(), tool_name: "shell".to_string(), description: "Run ls".to_string(), @@ -1146,7 +973,7 @@ mod tests { thread_id: Some("t1".to_string()), allow_always: true, }; - let ws = WsServerMessage::from_sse_event(&sse); + let ws = WsServerMessage::from_app_event(&event); match ws { WsServerMessage::Event { event_type, data } => { assert_eq!(event_type, "approval_needed"); @@ -1158,9 +985,9 @@ mod tests { } #[test] - fn test_ws_server_from_sse_heartbeat() { - let sse = SseEvent::Heartbeat; - let ws = WsServerMessage::from_sse_event(&sse); + fn test_ws_server_from_app_event_heartbeat() { + let event = AppEvent::Heartbeat; + let ws = WsServerMessage::from_app_event(&event); match ws { WsServerMessage::Event { event_type, .. } => { assert_eq!(event_type, "heartbeat"); @@ -1200,8 +1027,8 @@ mod tests { } #[test] - fn test_sse_auth_required_serialize() { - let event = SseEvent::AuthRequired { + fn test_app_event_auth_required_serialize() { + let event = AppEvent::AuthRequired { extension_name: "notion".to_string(), instructions: Some("Get your token from...".to_string()), auth_url: None, @@ -1217,8 +1044,8 @@ mod tests { } #[test] - fn test_sse_auth_completed_serialize() { - let event = SseEvent::AuthCompleted { + fn test_app_event_auth_completed_serialize() { + let event = AppEvent::AuthCompleted { extension_name: "notion".to_string(), success: true, message: "notion authenticated (3 tools loaded)".to_string(), @@ -1231,14 +1058,14 @@ mod tests { } #[test] - fn test_ws_server_from_sse_auth_required() { - let sse = SseEvent::AuthRequired { + fn test_ws_server_from_app_event_auth_required() { + let event = AppEvent::AuthRequired { extension_name: "openai".to_string(), instructions: Some("Enter API key".to_string()), auth_url: None, setup_url: None, }; - let ws = WsServerMessage::from_sse_event(&sse); + let ws = WsServerMessage::from_app_event(&event); match ws { WsServerMessage::Event { event_type, data } => { assert_eq!(event_type, "auth_required"); @@ -1249,13 +1076,13 @@ mod tests { } #[test] - fn test_ws_server_from_sse_auth_completed() { - let sse = SseEvent::AuthCompleted { + fn test_ws_server_from_app_event_auth_completed() { + let event = AppEvent::AuthCompleted { extension_name: "slack".to_string(), success: false, message: "Invalid token".to_string(), }; - let ws = WsServerMessage::from_sse_event(&sse); + let ws = WsServerMessage::from_app_event(&event); match ws { WsServerMessage::Event { event_type, data } => { assert_eq!(event_type, "auth_completed"); diff --git a/src/channels/web/util.rs b/src/channels/web/util.rs index 0debe6a9..2e4ffe3b 100644 --- a/src/channels/web/util.rs +++ b/src/channels/web/util.rs @@ -2,28 +2,21 @@ use crate::channels::web::types::{ToolCallInfo, TurnInfo}; -/// Truncate a string to at most `max_bytes` bytes at a char boundary, appending "...". -/// -/// If the input is wrapped in `` and truncation -/// removes the closing tag, the tag is re-appended so downstream XML parsers -/// never see an unclosed element. -pub fn truncate_preview(s: &str, max_bytes: usize) -> String { - if s.len() <= max_bytes { - return s.to_string(); - } - // Walk backwards from max_bytes to find a valid char boundary - let mut end = max_bytes; - while end > 0 && !s.is_char_boundary(end) { - end -= 1; - } - let mut result = format!("{}...", &s[..end]); +pub use ironclaw_common::truncate_preview; - // Re-close if truncation cut through the closing tag. - if s.starts_with("") { - result.push_str("\n"); - } - - result +/// Parse tool call summary JSON objects into `ToolCallInfo` structs. +fn parse_tool_call_infos(calls: &[serde_json::Value]) -> Vec { + calls + .iter() + .map(|c| ToolCallInfo { + name: c["name"].as_str().unwrap_or("unknown").to_string(), + has_result: c.get("result_preview").is_some_and(|v| !v.is_null()), + has_error: c.get("error").is_some_and(|v| !v.is_null()), + result_preview: c["result_preview"].as_str().map(String::from), + error: c["error"].as_str().map(String::from), + rationale: c["rationale"].as_str().map(String::from), + }) + .collect() } /// Build TurnInfo pairs from flat DB messages (user/tool_calls/assistant triples). @@ -49,6 +42,7 @@ pub fn build_turns_from_db_messages( started_at: msg.created_at.to_rfc3339(), completed_at: None, tool_calls: Vec::new(), + narrative: None, }; // Check if next message is a tool_calls record @@ -56,18 +50,28 @@ pub fn build_turns_from_db_messages( && next.role == "tool_calls" { let tc_msg = iter.next().expect("peeked"); - match serde_json::from_str::>(&tc_msg.content) { - Ok(calls) => { - turn.tool_calls = calls - .iter() - .map(|c| ToolCallInfo { - name: c["name"].as_str().unwrap_or("unknown").to_string(), - has_result: c.get("result_preview").is_some(), - has_error: c.get("error").is_some(), - result_preview: c["result_preview"].as_str().map(String::from), - error: c["error"].as_str().map(String::from), - }) - .collect(); + // Parse tool_calls JSON — supports two formats: + // safety: no byte-index slicing; comment describes JSON shape + match serde_json::from_str::(&tc_msg.content) { + Ok(serde_json::Value::Array(calls)) => { + // Old format: plain array + turn.tool_calls = parse_tool_call_infos(&calls); + } + Ok(serde_json::Value::Object(obj)) => { + // New wrapped format with narrative + turn.narrative = obj + .get("narrative") + .and_then(|v| v.as_str()) + .map(String::from); + if let Some(serde_json::Value::Array(calls)) = obj.get("calls") { + turn.tool_calls = parse_tool_call_infos(calls); + } + } + Ok(_) => { + tracing::warn!( + message_id = %tc_msg.id, + "Unexpected tool_calls JSON shape in DB, skipping" + ); } Err(e) => { tracing::warn!( @@ -105,6 +109,7 @@ pub fn build_turns_from_db_messages( started_at: msg.created_at.to_rfc3339(), completed_at: Some(msg.created_at.to_rfc3339()), tool_calls: Vec::new(), + narrative: None, }); turn_number += 1; } @@ -118,88 +123,6 @@ mod tests { use super::*; use uuid::Uuid; - // ---- truncate_preview tests ---- - - #[test] - fn test_truncate_preview_short_string() { - assert_eq!(truncate_preview("hello", 10), "hello"); - } - - #[test] - fn test_truncate_preview_exact_boundary() { - assert_eq!(truncate_preview("hello", 5), "hello"); - } - - #[test] - fn test_truncate_preview_truncates_ascii() { - assert_eq!(truncate_preview("hello world", 5), "hello..."); - } - - #[test] - fn test_truncate_preview_empty_string() { - assert_eq!(truncate_preview("", 10), ""); - } - - #[test] - fn test_truncate_preview_multibyte_char_boundary() { - // '€' is 3 bytes (E2 82 AC). "a€b" = [61, E2, 82, AC, 62] = 5 bytes - // Truncating at max_bytes=3 should not split the euro sign. - let s = "a€b"; - let result = truncate_preview(s, 3); - // max_bytes=3 lands mid-€, so it walks back to byte 1 ("a") - assert_eq!(result, "a..."); - } - - #[test] - fn test_truncate_preview_emoji() { - // '🦀' is 4 bytes. "hi🦀" = 6 bytes - let s = "hi🦀"; - let result = truncate_preview(s, 4); - // max_bytes=4 lands mid-🦀, walks back to byte 2 ("hi") - assert_eq!(result, "hi..."); - } - - #[test] - fn test_truncate_preview_cjk() { - // CJK characters are 3 bytes each. "你好世界" = 12 bytes - let s = "你好世界"; - let result = truncate_preview(s, 7); - // max_bytes=7 lands mid-character (byte 7 is inside 世), walks back to 6 ("你好") - assert_eq!(result, "你好..."); - } - - #[test] - fn test_truncate_preview_zero_max_bytes() { - assert_eq!(truncate_preview("hello", 0), "..."); - } - - #[test] - fn test_truncate_preview_closes_tool_output_tag() { - let s = "\nSome very long content here\n"; - // Truncate so it cuts before the closing tag - let result = truncate_preview(s, 60); - assert!(result.ends_with("")); - assert!(result.contains("...")); - } - - #[test] - fn test_truncate_preview_no_extra_close_when_intact() { - let s = "\nshort\n"; - // The string is short enough not to be truncated - let result = truncate_preview(s, 500); - assert_eq!(result, s); - // Should not have a duplicate closing tag - assert_eq!(result.matches("").count(), 1); - } - - #[test] - fn test_truncate_preview_non_xml_unaffected() { - let s = "Just a plain long string that gets truncated"; - let result = truncate_preview(s, 10); - assert_eq!(result, "Just a pla..."); - assert!(!result.contains("")); - } - // ---- build_turns_from_db_messages tests ---- fn make_msg(role: &str, content: &str, offset_ms: i64) -> crate::history::ConversationMessage { @@ -305,4 +228,52 @@ mod tests { assert!(turns[0].tool_calls.is_empty()); assert_eq!(turns[0].state, "Completed"); } + + #[test] + fn test_build_turns_with_wrapped_tool_calls_format() { + let tc_json = serde_json::json!({ + "narrative": "Searching memory for context before proceeding.", + "calls": [ + {"name": "memory_search", "result_preview": "found 3 items", "rationale": "consult prior context"}, + {"name": "shell", "error": "permission denied"} + ] + }); + let messages = vec![ + make_msg("user", "Find info", 0), + make_msg("tool_calls", &tc_json.to_string(), 500), + make_msg("assistant", "Here's what I found", 1000), + ]; + let turns = build_turns_from_db_messages(&messages); + assert_eq!(turns.len(), 1); + assert_eq!( + turns[0].narrative.as_deref(), + Some("Searching memory for context before proceeding.") + ); + assert_eq!(turns[0].tool_calls.len(), 2); + assert_eq!(turns[0].tool_calls[0].name, "memory_search"); + assert_eq!( + turns[0].tool_calls[0].rationale.as_deref(), + Some("consult prior context") + ); + assert!(turns[0].tool_calls[0].has_result); + assert_eq!(turns[0].tool_calls[1].name, "shell"); + assert!(turns[0].tool_calls[1].has_error); + assert_eq!(turns[0].response.as_deref(), Some("Here's what I found")); + } + + #[test] + fn test_build_turns_wrapped_format_without_narrative() { + let tc_json = serde_json::json!({ + "calls": [{"name": "echo", "result_preview": "hello"}] + }); + let messages = vec![ + make_msg("user", "Say hi", 0), + make_msg("tool_calls", &tc_json.to_string(), 500), + make_msg("assistant", "Done", 1000), + ]; + let turns = build_turns_from_db_messages(&messages); + assert_eq!(turns.len(), 1); + assert!(turns[0].narrative.is_none()); + assert_eq!(turns[0].tool_calls.len(), 1); + } } diff --git a/src/channels/web/ws.rs b/src/channels/web/ws.rs index 3a601679..51beaafd 100644 --- a/src/channels/web/ws.rs +++ b/src/channels/web/ws.rs @@ -97,7 +97,7 @@ pub async fn handle_ws_connection( let msg = tokio::select! { event = event_stream.next() => { match event { - Some(sse_event) => WsServerMessage::from_sse_event(&sse_event), + Some(app_event) => WsServerMessage::from_app_event(&app_event), None => break, // Broadcast channel closed } } @@ -275,7 +275,7 @@ async fn handle_client_message( if result.verification.is_some() { state.sse.broadcast_for_user( user_id, - crate::channels::web::types::SseEvent::AuthRequired { + crate::channels::web::types::AppEvent::AuthRequired { extension_name: extension_name.clone(), instructions: Some(result.message), auth_url: None, @@ -286,7 +286,7 @@ async fn handle_client_message( crate::channels::web::server::clear_auth_mode(state, user_id).await; state.sse.broadcast_for_user( user_id, - crate::channels::web::types::SseEvent::AuthCompleted { + crate::channels::web::types::AppEvent::AuthCompleted { extension_name, success: true, message: result.message, @@ -299,7 +299,7 @@ async fn handle_client_message( if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) { state.sse.broadcast_for_user( user_id, - crate::channels::web::types::SseEvent::AuthRequired { + crate::channels::web::types::AppEvent::AuthRequired { extension_name: extension_name.clone(), instructions: Some(msg.clone()), auth_url: None, @@ -520,7 +520,8 @@ mod tests { job_manager: None, prompt_queue: None, scheduler: None, - default_user_id: "test".to_string(), + owner_id: "test".to_string(), + default_sender_id: "test".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: None, diff --git a/src/cli/oauth_defaults.rs b/src/cli/oauth_defaults.rs index 3b57872f..384d5833 100644 --- a/src/cli/oauth_defaults.rs +++ b/src/cli/oauth_defaults.rs @@ -62,6 +62,30 @@ pub fn builtin_client_id_override_env(secret_name: &str) -> Option<&'static str> } } +/// Suppress the baked-in desktop OAuth client secret when a hosted proxy is configured. +/// +/// In hosted deployments, IronClaw may resolve the platform Google client ID from +/// environment variables while still falling back to the baked-in desktop secret. +/// That client_id/client_secret mismatch breaks Google token exchange and refresh. +/// +/// When the proxy is configured, the platform will inject the correct server-side +/// secret for matching platform credentials, so the baked-in secret must be omitted. +pub fn hosted_proxy_client_secret( + client_secret: &Option, + builtin: Option<&OAuthCredentials>, + exchange_proxy_configured: bool, +) -> Option { + if !exchange_proxy_configured { + return client_secret.clone(); + } + + let builtin_secret = builtin.map(|credentials| credentials.client_secret); + match (client_secret, builtin_secret) { + (Some(resolved), Some(baked_in)) if resolved == baked_in => None, + _ => client_secret.clone(), + } +} + // ── Shared callback server ────────────────────────────────────────────── // Core OAuth callback infrastructure is defined in `crate::llm::oauth_helpers` @@ -449,7 +473,8 @@ pub struct PendingOAuthFlow { pub secrets: Arc, /// SSE broadcast manager for notifying the web UI. pub sse_manager: Option>, - /// Gateway auth token for authenticating with the platform token exchange proxy. + /// OAuth proxy auth token for authenticating with the hosted token exchange proxy. + /// Kept as `gateway_token` for public API compatibility. pub gateway_token: Option, /// Additional form params for the token exchange request. /// Used for provider-specific requirements such as RFC 8707 `resource`. @@ -472,6 +497,12 @@ impl std::fmt::Debug for PendingOAuthFlow { } } +impl PendingOAuthFlow { + pub fn oauth_proxy_auth_token(&self) -> Option<&str> { + self.gateway_token.as_deref() + } +} + /// Thread-safe registry of pending OAuth flows, keyed by CSRF `state` parameter. pub type PendingOAuthRegistry = Arc>>; @@ -505,6 +536,22 @@ pub fn exchange_proxy_url() -> Option { .filter(|url| !url.is_empty()) } +/// Returns the configured OAuth proxy auth token, if any. +/// +/// New hosted infra can inject a dedicated shared proxy secret via +/// `IRONCLAW_OAUTH_PROXY_AUTH_TOKEN`. Existing hosted instances continue to +/// work by falling back to `GATEWAY_AUTH_TOKEN`. +pub fn oauth_proxy_auth_token() -> Option { + fn normalized_env_value(key: &str) -> Option { + crate::config::helpers::env_or_override(key) + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()) + } + + normalized_env_value("IRONCLAW_OAUTH_PROXY_AUTH_TOKEN") + .or_else(|| normalized_env_value("GATEWAY_AUTH_TOKEN")) +} + /// Maximum age for pending OAuth flows (5 minutes, matching TCP listener timeout). pub const OAUTH_FLOW_EXPIRY: Duration = Duration::from_secs(300); @@ -650,6 +697,8 @@ pub fn strip_instance_prefix(state: &str) -> &str { pub struct ProxyTokenExchangeRequest<'a> { pub proxy_url: &'a str, + /// OAuth proxy auth token. + /// Kept as `gateway_token` for public API compatibility. pub gateway_token: &'a str, pub token_url: &'a str, pub client_id: &'a str, @@ -661,9 +710,53 @@ pub struct ProxyTokenExchangeRequest<'a> { pub extra_token_params: &'a HashMap, } +pub struct ProxyRefreshTokenRequest<'a> { + pub proxy_url: &'a str, + /// OAuth proxy auth token. + /// Kept as `gateway_token` for public API compatibility. + pub gateway_token: &'a str, + pub token_url: &'a str, + pub client_id: &'a str, + pub client_secret: Option<&'a str>, + pub refresh_token: &'a str, + pub provider: Option<&'a str>, +} + +fn oauth_token_response_from_json( + token_data: serde_json::Value, + access_token_field: &str, +) -> Result { + let access_token = token_data + .get(access_token_field) + .and_then(|v| v.as_str()) + .ok_or_else(|| { + let fields: Vec<&str> = token_data + .as_object() + .map(|o| o.keys().map(|k| k.as_str()).collect()) + .unwrap_or_default(); + OAuthCallbackError::Io(format!( + "No '{}' field in proxy response (fields present: {:?})", + access_token_field, fields + )) + })? + .to_string(); + + let refresh_token = token_data + .get("refresh_token") + .and_then(|v| v.as_str()) + .map(String::from); + let expires_in = token_data.get("expires_in").and_then(|v| v.as_u64()); + + Ok(OAuthTokenResponse { + access_token, + refresh_token, + expires_in, + }) +} + /// Exchange an OAuth authorization code via the platform's token exchange proxy. /// -/// Authenticated via the gateway auth token (Bearer header). The caller may +/// Authenticated via an OAuth proxy auth token (Bearer header). The caller may /// either rely on proxy-side secret lookup or forward a `client_secret` when /// the provider requires it. /// @@ -675,13 +768,14 @@ pub async fn exchange_via_proxy( ) -> Result { if request.gateway_token.is_empty() { return Err(OAuthCallbackError::Io( - "Gateway auth token is required for proxy token exchange".to_string(), + "OAuth proxy auth token is required for proxy token exchange".to_string(), )); } let exchange_url = format!("{}/oauth/exchange", request.proxy_url.trim_end_matches('/')); let client = reqwest::Client::builder() .timeout(Duration::from_secs(60)) + .redirect(reqwest::redirect::Policy::none()) .build() .map_err(|e| OAuthCallbackError::Io(format!("Failed to build HTTP client: {}", e)))?; let mut params = vec![ @@ -724,41 +818,454 @@ pub async fn exchange_via_proxy( .json() .await .map_err(|e| OAuthCallbackError::Io(format!("Failed to parse proxy response: {}", e)))?; + oauth_token_response_from_json(token_data, request.access_token_field) +} - let access_token = token_data - .get(request.access_token_field) - .and_then(|v| v.as_str()) - .ok_or_else(|| { - let fields: Vec<&str> = token_data - .as_object() - .map(|o| o.keys().map(|k| k.as_str()).collect()) - .unwrap_or_default(); - OAuthCallbackError::Io(format!( - "No '{}' field in proxy response (fields present: {:?})", - request.access_token_field, fields - )) - })? - .to_string(); +/// Refresh an OAuth access token via the platform's token refresh proxy. +/// +/// Authenticated via an OAuth proxy auth token (Bearer header). The caller may +/// either rely on proxy-side secret lookup or forward a `client_secret` when +/// the provider requires it. +pub async fn refresh_token_via_proxy( + request: ProxyRefreshTokenRequest<'_>, +) -> Result { + if request.gateway_token.is_empty() { + return Err(OAuthCallbackError::Io( + "OAuth proxy auth token is required for proxy token refresh".to_string(), + )); + } - let refresh_token = token_data - .get("refresh_token") - .and_then(|v| v.as_str()) - .map(String::from); - let expires_in = token_data.get("expires_in").and_then(|v| v.as_u64()); + let refresh_url = format!("{}/oauth/refresh", request.proxy_url.trim_end_matches('/')); + let client = reqwest::Client::builder() + .timeout(Duration::from_secs(15)) + .redirect(reqwest::redirect::Policy::none()) + .build() + .map_err(|e| OAuthCallbackError::Io(format!("Failed to build HTTP client: {}", e)))?; - Ok(OAuthTokenResponse { - access_token, - refresh_token, - expires_in, - }) + let mut params = vec![ + ("refresh_token", request.refresh_token.to_string()), + ("token_url", request.token_url.to_string()), + ("client_id", request.client_id.to_string()), + ]; + if let Some(secret) = request.client_secret { + params.push(("client_secret", secret.to_string())); + } + if let Some(provider) = request.provider { + params.push(("provider", provider.to_string())); + } + + let response = client + .post(&refresh_url) + .bearer_auth(request.gateway_token) + .form(¶ms) + .send() + .await + .map_err(|e| { + OAuthCallbackError::Io(format!("Token refresh proxy request failed: {}", e)) + })?; + + if !response.status().is_success() { + let status = response.status(); + let body = response.text().await.unwrap_or_default(); + return Err(OAuthCallbackError::Io(format!( + "Token refresh proxy failed: {} - {}", + status, body + ))); + } + + let token_data: serde_json::Value = response + .json() + .await + .map_err(|e| OAuthCallbackError::Io(format!("Failed to parse proxy response: {}", e)))?; + + oauth_token_response_from_json(token_data, "access_token") } #[cfg(test)] mod tests { + use std::collections::HashMap; + use std::net::SocketAddr; + use std::sync::Arc; + + use axum::extract::{Form, State}; + use axum::http::HeaderMap; + use axum::response::Redirect; + use axum::routing::post; + use axum::{Json, Router}; + use serde_json::json; + use tokio::net::TcpListener; + use tokio::sync::{Mutex, oneshot}; + use crate::cli::oauth_defaults::{ builtin_credentials, callback_host, callback_url, is_loopback_host, landing_html, }; use crate::config::helpers::lock_env; + use crate::testing::credentials::{TEST_OAUTH_CLIENT_ID, TEST_OAUTH_CLIENT_SECRET}; + + #[derive(Clone, Debug, PartialEq, Eq)] + struct RecordedProxyRequest { + authorization: Option, + form: HashMap, + } + + #[derive(Clone)] + struct MockProxyState { + requests: Arc>>, + exchange_redirect_target: String, + refresh_redirect_target: String, + } + + struct MockProxyServer { + addr: SocketAddr, + requests: Arc>>, + shutdown_tx: Option>, + server_task: Option>, + } + + impl MockProxyServer { + async fn start() -> Self { + async fn exchange_handler( + State(state): State, + headers: HeaderMap, + Form(form): Form>, + ) -> Json { + state.requests.lock().await.push(RecordedProxyRequest { + authorization: headers + .get(axum::http::header::AUTHORIZATION) + .and_then(|value| value.to_str().ok()) + .map(str::to_string), + form, + }); + Json(json!({ + "access_token": "proxy-access-token", + "refresh_token": "proxy-refresh-token", + "expires_in": 7200 + })) + } + + async fn refresh_handler( + State(state): State, + headers: HeaderMap, + Form(form): Form>, + ) -> Json { + state.requests.lock().await.push(RecordedProxyRequest { + authorization: headers + .get(axum::http::header::AUTHORIZATION) + .and_then(|value| value.to_str().ok()) + .map(str::to_string), + form, + }); + Json(json!({ + "access_token": "proxy-access-token", + "refresh_token": "proxy-refresh-token", + "expires_in": 7200 + })) + } + + async fn exchange_redirect_handler(State(state): State) -> Redirect { + Redirect::temporary(&state.exchange_redirect_target) + } + + async fn refresh_redirect_handler(State(state): State) -> Redirect { + Redirect::temporary(&state.refresh_redirect_target) + } + + let requests = Arc::new(Mutex::new(Vec::new())); + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("bind mock proxy"); + let addr = listener.local_addr().expect("read mock proxy addr"); + let exchange_redirect_target = format!("http://{addr}/oauth/exchange"); + let refresh_redirect_target = format!("http://{addr}/oauth/refresh"); + let app = Router::new() + .route("/oauth/exchange", post(exchange_handler)) + .route("/oauth/refresh", post(refresh_handler)) + .route("/redirect/oauth/exchange", post(exchange_redirect_handler)) + .route("/redirect/oauth/refresh", post(refresh_redirect_handler)) + .with_state(MockProxyState { + requests: Arc::clone(&requests), + exchange_redirect_target, + refresh_redirect_target, + }); + let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>(); + let server_task = tokio::spawn(async move { + let _ = axum::serve(listener, app) + .with_graceful_shutdown(async { + let _ = shutdown_rx.await; + }) + .await; + }); + + Self { + addr, + requests, + shutdown_tx: Some(shutdown_tx), + server_task: Some(server_task), + } + } + + fn base_url(&self) -> String { + format!("http://{}", self.addr) + } + + fn redirecting_base_url(&self) -> String { + format!("{}/redirect", self.base_url()) + } + + async fn requests(&self) -> Vec { + self.requests.lock().await.clone() + } + + async fn shutdown(mut self) { + if let Some(tx) = self.shutdown_tx.take() { + let _ = tx.send(()); + } + if let Some(task) = self.server_task.take() { + let _ = task.await; + } + } + } + + impl Drop for MockProxyServer { + fn drop(&mut self) { + if let Some(tx) = self.shutdown_tx.take() { + let _ = tx.send(()); + } + if let Some(task) = self.server_task.take() { + task.abort(); + } + } + } + + struct EnvVarGuard { + key: &'static str, + original: Option, + } + + impl Drop for EnvVarGuard { + fn drop(&mut self) { + // SAFETY: Under ENV_MUTEX, no concurrent env access. + unsafe { + if let Some(ref value) = self.original { + std::env::set_var(self.key, value); + } else { + std::env::remove_var(self.key); + } + } + } + } + + fn set_env_var(key: &'static str, value: Option<&str>) -> EnvVarGuard { + let original = std::env::var(key).ok(); + // SAFETY: Under ENV_MUTEX, no concurrent env access. + unsafe { + if let Some(value) = value { + std::env::set_var(key, value); + } else { + std::env::remove_var(key); + } + } + EnvVarGuard { key, original } + } + + #[test] + fn test_hosted_proxy_client_secret_suppresses_builtin_secret() { + let builtin = builtin_credentials("google_oauth_token").expect("google builtin creds"); + let client_secret = Some(builtin.client_secret.to_string()); + + let result = super::hosted_proxy_client_secret(&client_secret, Some(&builtin), true); + + assert_eq!(result, None); + } + + #[test] + fn test_hosted_proxy_client_secret_preserves_explicit_secret() { + let builtin = builtin_credentials("google_oauth_token").expect("google builtin creds"); + let client_secret = Some("hosted-server-secret".to_string()); + + let result = super::hosted_proxy_client_secret(&client_secret, Some(&builtin), true); + + assert_eq!(result, client_secret); + } + + #[tokio::test] + async fn test_exchange_via_proxy_sends_auth_and_form() { + let server = MockProxyServer::start().await; + let mut extra_token_params = HashMap::new(); + extra_token_params.insert("resource".to_string(), "https://mcp.notion.com".to_string()); + + let response = super::exchange_via_proxy(super::ProxyTokenExchangeRequest { + proxy_url: &server.base_url(), + gateway_token: "shared-oauth-proxy-secret", + code: "auth-code-123", + redirect_uri: "https://oauth.example.com/oauth/callback", + token_url: "https://oauth2.googleapis.com/token", + client_id: TEST_OAUTH_CLIENT_ID, + client_secret: Some(TEST_OAUTH_CLIENT_SECRET), + access_token_field: "access_token", + code_verifier: Some("code-verifier-123"), + extra_token_params: &extra_token_params, + }) + .await + .expect("proxy exchange succeeds"); + + assert_eq!(response.access_token, "proxy-access-token"); + assert_eq!( + response.refresh_token.as_deref(), + Some("proxy-refresh-token") + ); + assert_eq!(response.expires_in, Some(7200)); + + let requests = server.requests().await; + assert_eq!(requests.len(), 1); + assert_eq!( + requests[0].authorization.as_deref(), + Some("Bearer shared-oauth-proxy-secret") + ); + assert_eq!( + requests[0].form.get("code").map(String::as_str), + Some("auth-code-123") + ); + assert_eq!( + requests[0].form.get("redirect_uri").map(String::as_str), + Some("https://oauth.example.com/oauth/callback") + ); + assert_eq!( + requests[0].form.get("token_url").map(String::as_str), + Some("https://oauth2.googleapis.com/token") + ); + assert_eq!( + requests[0].form.get("client_id").map(String::as_str), + Some(TEST_OAUTH_CLIENT_ID) + ); + assert_eq!( + requests[0].form.get("client_secret").map(String::as_str), + Some(TEST_OAUTH_CLIENT_SECRET) + ); + assert_eq!( + requests[0] + .form + .get("access_token_field") + .map(String::as_str), + Some("access_token") + ); + assert_eq!( + requests[0].form.get("code_verifier").map(String::as_str), + Some("code-verifier-123") + ); + assert_eq!( + requests[0].form.get("resource").map(String::as_str), + Some("https://mcp.notion.com") + ); + + server.shutdown().await; + } + + #[tokio::test] + async fn test_refresh_token_via_proxy_sends_auth_and_form() { + let server = MockProxyServer::start().await; + + let response = super::refresh_token_via_proxy(super::ProxyRefreshTokenRequest { + proxy_url: &server.base_url(), + gateway_token: "gateway-test-token", + token_url: "https://oauth2.googleapis.com/token", + client_id: TEST_OAUTH_CLIENT_ID, + client_secret: Some(TEST_OAUTH_CLIENT_SECRET), + refresh_token: "refresh-token-123", + provider: Some("google"), + }) + .await + .expect("proxy refresh succeeds"); + + assert_eq!(response.access_token, "proxy-access-token"); + assert_eq!( + response.refresh_token.as_deref(), + Some("proxy-refresh-token") + ); + assert_eq!(response.expires_in, Some(7200)); + + let requests = server.requests().await; + assert_eq!(requests.len(), 1); + assert_eq!( + requests[0].authorization.as_deref(), + Some("Bearer gateway-test-token") + ); + assert_eq!( + requests[0].form.get("token_url").map(String::as_str), + Some("https://oauth2.googleapis.com/token") + ); + assert_eq!( + requests[0].form.get("client_id").map(String::as_str), + Some(TEST_OAUTH_CLIENT_ID) + ); + assert_eq!( + requests[0].form.get("client_secret").map(String::as_str), + Some(TEST_OAUTH_CLIENT_SECRET) + ); + assert_eq!( + requests[0].form.get("refresh_token").map(String::as_str), + Some("refresh-token-123") + ); + assert_eq!( + requests[0].form.get("provider").map(String::as_str), + Some("google") + ); + + server.shutdown().await; + } + + #[tokio::test] + async fn test_exchange_via_proxy_does_not_follow_redirects() { + let server = MockProxyServer::start().await; + + let error = match super::exchange_via_proxy(super::ProxyTokenExchangeRequest { + proxy_url: &server.redirecting_base_url(), + gateway_token: "gateway-test-token", + code: "auth-code-123", + redirect_uri: "http://localhost:3000/oauth/callback", + token_url: "https://oauth2.googleapis.com/token", + client_id: TEST_OAUTH_CLIENT_ID, + client_secret: Some(TEST_OAUTH_CLIENT_SECRET), + access_token_field: "access_token", + code_verifier: Some("code-verifier-123"), + extra_token_params: &HashMap::new(), + }) + .await + { + Ok(_) => panic!("redirected proxy exchange should fail"), + Err(error) => error, + }; + + assert!(error.to_string().contains("307")); + assert!(server.requests().await.is_empty()); + + server.shutdown().await; + } + + #[tokio::test] + async fn test_refresh_token_via_proxy_does_not_follow_redirects() { + let server = MockProxyServer::start().await; + + let error = match super::refresh_token_via_proxy(super::ProxyRefreshTokenRequest { + proxy_url: &server.redirecting_base_url(), + gateway_token: "gateway-test-token", + token_url: "https://oauth2.googleapis.com/token", + client_id: TEST_OAUTH_CLIENT_ID, + client_secret: Some(TEST_OAUTH_CLIENT_SECRET), + refresh_token: "refresh-token-123", + provider: Some("google"), + }) + .await + { + Ok(_) => panic!("redirected proxy refresh should fail"), + Err(error) => error, + }; + + assert!(error.to_string().contains("307")); + assert!(server.requests().await.is_empty()); + + server.shutdown().await; + } #[test] fn test_is_loopback_host() { @@ -1159,6 +1666,54 @@ mod tests { } } + #[test] + fn test_oauth_proxy_auth_token_prefers_dedicated_env() { + let _guard = lock_env(); + let _proxy_guard = set_env_var( + "IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", + Some("shared-proxy-secret"), + ); + let _gateway_guard = set_env_var("GATEWAY_AUTH_TOKEN", Some("gateway-token")); + + assert_eq!( + crate::cli::oauth_defaults::oauth_proxy_auth_token().as_deref(), + Some("shared-proxy-secret") + ); + } + + #[test] + fn test_oauth_proxy_auth_token_falls_back_to_gateway_token() { + let _guard = lock_env(); + let _proxy_guard = set_env_var("IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", None); + let _gateway_guard = set_env_var("GATEWAY_AUTH_TOKEN", Some("gateway-token")); + + assert_eq!( + crate::cli::oauth_defaults::oauth_proxy_auth_token().as_deref(), + Some("gateway-token") + ); + } + + #[test] + fn test_oauth_proxy_auth_token_whitespace_dedicated_env_falls_back_to_gateway_token() { + let _guard = lock_env(); + let _proxy_guard = set_env_var("IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", Some(" ")); + let _gateway_guard = set_env_var("GATEWAY_AUTH_TOKEN", Some("gateway-token")); + + assert_eq!( + crate::cli::oauth_defaults::oauth_proxy_auth_token().as_deref(), + Some("gateway-token") + ); + } + + #[test] + fn test_oauth_proxy_auth_token_returns_none_when_unset() { + let _guard = lock_env(); + let _proxy_guard = set_env_var("IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", None); + let _gateway_guard = set_env_var("GATEWAY_AUTH_TOKEN", None); + + assert_eq!(crate::cli::oauth_defaults::oauth_proxy_auth_token(), None); + } + #[test] fn test_strip_instance_prefix_with_colon() { use crate::cli::oauth_defaults::strip_instance_prefix; diff --git a/src/cli/tool.rs b/src/cli/tool.rs index be684580..9d39c492 100644 --- a/src/cli/tool.rs +++ b/src/cli/tool.rs @@ -2,6 +2,7 @@ //! //! Commands for installing, listing, removing, and authenticating WASM tools. +use std::collections::{HashMap, HashSet}; use std::io::Write; use std::path::{Path, PathBuf}; use std::sync::Arc; @@ -79,6 +80,10 @@ pub enum ToolCommand { /// Directory to look for tool (default: ~/.ironclaw/tools/) #[arg(short, long)] dir: Option, + + /// User ID for checking credential status (default: "default") + #[arg(short, long, default_value = "default")] + user: String, }, /// Configure authentication for a tool @@ -124,7 +129,11 @@ pub async fn run_tool_command(cmd: ToolCommand) -> anyhow::Result<()> { } => install_tool(path, name, capabilities, target, release, skip_build, force).await, ToolCommand::List { dir, verbose } => list_tools(dir, verbose).await, ToolCommand::Remove { name, dir } => remove_tool(name, dir).await, - ToolCommand::Info { name_or_path, dir } => show_tool_info(name_or_path, dir).await, + ToolCommand::Info { + name_or_path, + dir, + user, + } => show_tool_info(name_or_path, dir, user).await, ToolCommand::Auth { name, dir, user } => auth_tool(name, dir, user).await, ToolCommand::Setup { name, dir, user } => setup_tool(name, dir, user).await, } @@ -388,7 +397,11 @@ async fn remove_tool(name: String, dir: Option) -> anyhow::Result<()> { } /// Show information about a tool. -async fn show_tool_info(name_or_path: String, dir: Option) -> anyhow::Result<()> { +async fn show_tool_info( + name_or_path: String, + dir: Option, + user_id: String, +) -> anyhow::Result<()> { let wasm_path = if name_or_path.ends_with(".wasm") { PathBuf::from(&name_or_path) } else { @@ -423,7 +436,37 @@ async fn show_tool_info(name_or_path: String, dir: Option) -> anyhow::R println!("\nCapabilities ({}):", caps_path.display()); let content = fs::read_to_string(&caps_path).await?; match CapabilitiesFile::from_json(&content) { - Ok(caps) => print_capabilities_detail(&caps), + Ok(caps) => { + // Lazily init secrets store only when auth secrets need checking. + let has_auth = caps.auth.is_some() + || caps + .setup + .as_ref() + .is_some_and(|s| !s.required_secrets.is_empty()) + || caps + .http + .as_ref() + .is_some_and(|h| !h.credentials.is_empty()); + let secrets_store = if has_auth { + match init_secrets_store().await { + Ok(store) => Some(store), + Err(e) => { + eprintln!(" Warning: could not init secrets store: {}", e); + None + } + } + } else { + None + }; + print_capabilities_detail( + &caps, + secrets_store + .as_ref() + .map(|s| s.as_ref() as &(dyn SecretsStore + Send + Sync)), + &user_id, + ) + .await; + } Err(e) => println!(" Error parsing: {}", e), } } else { @@ -476,8 +519,89 @@ fn print_capabilities_summary(caps: &CapabilitiesFile) { } } +/// Per-secret info collected from all auth-related capability sections. +struct AuthSecretInfo { + secret_name: String, + /// Human-readable label (from auth.display_name or setup prompt). + description: Option, + /// Injection location (from http.credentials). + location: Option, +} + +/// Collected auth secrets and the set of secret names they cover. +struct CollectedAuthSecrets { + secrets: Vec, + /// Secret names present in `secrets`, for filtering the Secrets capability section. + seen_names: HashSet, +} + +/// Collect and deduplicate auth secrets from all auth-related capability sections. +/// +/// Priority for the description label: auth.display_name > setup.required_secrets.prompt. +/// Injection location is merged from http.credentials. +fn collect_auth_secrets(caps: &CapabilitiesFile) -> CollectedAuthSecrets { + let mut secrets: Vec = Vec::new(); + let mut seen: HashMap = HashMap::new(); + + // auth.display_name is the best label — seed first. + if let Some(ref auth) = caps.auth { + let index = secrets.len(); + seen.insert(auth.secret_name.clone(), index); + secrets.push(AuthSecretInfo { + secret_name: auth.secret_name.clone(), + description: auth.display_name.clone(), + location: None, + }); + } + + // setup.required_secrets.prompt is second-best label. + if let Some(ref setup) = caps.setup { + for secret in &setup.required_secrets { + if !seen.contains_key(&secret.name) { + let index = secrets.len(); + seen.insert(secret.name.clone(), index); + secrets.push(AuthSecretInfo { + secret_name: secret.name.clone(), + description: Some(secret.prompt.clone()), + location: None, + }); + } + } + } + + // Merge injection location from http.credentials. + if let Some(ref http) = caps.http { + for cred in http.credentials.values() { + let loc = format!("{:?}", cred.location); + if let Some(&index) = seen.get(&cred.secret_name) { + secrets[index].location = Some(loc); + } else { + let index = secrets.len(); + seen.insert(cred.secret_name.clone(), index); + secrets.push(AuthSecretInfo { + secret_name: cred.secret_name.clone(), + description: None, + location: Some(loc), + }); + } + } + } + + let seen_names = seen.into_keys().collect(); + CollectedAuthSecrets { + secrets, + seen_names, + } +} + /// Print detailed capabilities. -fn print_capabilities_detail(caps: &CapabilitiesFile) { +async fn print_capabilities_detail( + caps: &CapabilitiesFile, + secrets_store: Option<&(dyn SecretsStore + Send + Sync)>, + user_id: &str, +) { + let mut collected = collect_auth_secrets(caps); + if let Some(ref http) = caps.http { println!(" HTTP:"); for endpoint in &http.allowlist { @@ -490,13 +614,6 @@ fn print_capabilities_detail(caps: &CapabilitiesFile) { println!(" {} {} {}", methods, endpoint.host, path); } - if !http.credentials.is_empty() { - println!(" Credentials:"); - for (key, cred) in &http.credentials { - println!(" {}: {} -> {:?}", key, cred.secret_name, cred.location); - } - } - if let Some(ref rate) = http.rate_limit { println!( " Rate limit: {}/min, {}/hour", @@ -505,12 +622,24 @@ fn print_capabilities_detail(caps: &CapabilitiesFile) { } } + // Filter secrets already covered by the auth section (always rendered when non-empty). if let Some(ref secrets) = caps.secrets && !secrets.allowed_names.is_empty() { - println!(" Secrets (existence check only):"); - for name in &secrets.allowed_names { - println!(" {}", name); + let extra: Vec<_> = if collected.secrets.is_empty() { + secrets.allowed_names.iter().collect() + } else { + secrets + .allowed_names + .iter() + .filter(|name| !collected.seen_names.contains(name.as_str())) + .collect() + }; + if !extra.is_empty() { + println!(" Secrets (existence check only):"); + for name in extra { + println!(" {}", name); + } } } @@ -531,6 +660,38 @@ fn print_capabilities_detail(caps: &CapabilitiesFile) { println!(" {}", prefix); } } + + // Consolidated auth status — sorted by secret name for deterministic output. + if !collected.secrets.is_empty() { + collected + .secrets + .sort_by(|a, b| a.secret_name.cmp(&b.secret_name)); + println!(" Auth:"); + for info in &collected.secrets { + let (icon, label) = match secrets_store { + Some(store) => match store.exists(user_id, &info.secret_name).await { + Ok(true) => ("\u{2713}", "configured"), + Ok(false) => ("\u{2717}", "missing"), + Err(e) => { + eprintln!( + " Warning: failed to check secret `{}`: {}", + info.secret_name, e + ); + ("?", "unknown") + } + }, + None => ("?", "unknown"), + }; + let mut parts = info.secret_name.clone(); + if let Some(ref desc) = info.description { + parts = format!("{} ({})", parts, desc); + } + if let Some(ref loc) = info.location { + parts = format!("{} -> {}", parts, loc); + } + println!(" {} {} {}", parts, icon, label); + } + } } /// Validate a tool name to prevent path traversal. @@ -677,8 +838,7 @@ async fn combine_provider_scopes( secret_name: &str, base_oauth: &crate::tools::wasm::OAuthConfigSchema, ) -> crate::tools::wasm::OAuthConfigSchema { - let mut all_scopes: std::collections::HashSet = - base_oauth.scopes.iter().cloned().collect(); + let mut all_scopes: HashSet = base_oauth.scopes.iter().cloned().collect(); if let Ok(mut entries) = tokio::fs::read_dir(tools_dir).await { while let Ok(Some(entry)) = entries.next_entry().await { @@ -1127,6 +1287,8 @@ async fn setup_tool(name: String, dir: Option, user_id: String) -> anyh #[cfg(test)] mod tests { use super::*; + use crate::secrets::{CreateSecretParams, SecretsStore}; + use crate::testing::credentials::test_secrets_store; #[test] fn test_format_size() { @@ -1143,4 +1305,96 @@ mod tests { assert!(dir.to_string_lossy().contains(".ironclaw")); assert!(dir.to_string_lossy().contains("tools")); } + + /// Verify that auth secrets are deduplicated across auth, setup, and http.credentials, + /// and that credential status is checked against the secrets store. + #[tokio::test] + async fn test_auth_secret_dedup_and_status() { + let caps = CapabilitiesFile::from_json( + r#"{ + "auth": { + "secret_name": "gh_token", + "display_name": "GitHub" + }, + "setup": { + "required_secrets": [ + { "name": "gh_token", "prompt": "GitHub PAT" }, + { "name": "extra_key", "prompt": "Extra API Key" } + ] + }, + "http": { + "allowlist": [{ "host": "api.github.com" }], + "credentials": { + "github": { + "secret_name": "gh_token", + "location": { "type": "bearer" }, + "host_patterns": ["api.github.com"] + } + } + }, + "secrets": { + "allowed_names": ["gh_token", "gh_*"] + } + }"#, + ) + .unwrap(); + + let collected = collect_auth_secrets(&caps); + + // gh_token should appear once (from auth), with location merged from credentials. + // extra_key should appear once (from setup). + assert_eq!(collected.secrets.len(), 2); + let gh = collected + .secrets + .iter() + .find(|s| s.secret_name == "gh_token") + .unwrap(); + assert_eq!(gh.description.as_deref(), Some("GitHub")); + assert!( + gh.location.is_some(), + "location should be merged from http.credentials" + ); + + let extra = collected + .secrets + .iter() + .find(|s| s.secret_name == "extra_key") + .unwrap(); + assert_eq!(extra.description.as_deref(), Some("Extra API Key")); + assert!(extra.location.is_none()); + + // Secrets section should filter gh_token (in seen_names) but keep gh_* (wildcard). + let secrets = caps.secrets.as_ref().unwrap(); + let extra_secrets: Vec<_> = secrets + .allowed_names + .iter() + .filter(|name| !collected.seen_names.contains(name.as_str())) + .collect(); + assert_eq!(extra_secrets, vec!["gh_*"]); + + // Verify store check: missing secret -> exists returns false. + let store = test_secrets_store(); + assert!(!store.exists("default", "gh_token").await.unwrap()); + + // Store gh_token and verify it's found. + store + .create( + "default", + CreateSecretParams::new("gh_token", "ghp_test123"), + ) + .await + .unwrap(); + assert!(store.exists("default", "gh_token").await.unwrap()); + // extra_key still missing. + assert!(!store.exists("default", "extra_key").await.unwrap()); + } + + /// No auth sections → collect_auth_secrets returns empty. + #[test] + fn test_collect_auth_secrets_empty_caps() { + let caps = CapabilitiesFile::default(); + let collected = collect_auth_secrets(&caps); + assert!(collected.secrets.is_empty()); + assert!(collected.seen_names.is_empty()); + } } diff --git a/src/config/agent.rs b/src/config/agent.rs index cb09707d..cfa0879a 100644 --- a/src/config/agent.rs +++ b/src/config/agent.rs @@ -1,6 +1,6 @@ use std::time::Duration; -use crate::config::helpers::{parse_bool_env, parse_option_env, parse_optional_env}; +use crate::config::helpers::{optional_env, parse_bool_env, parse_option_env, parse_optional_env}; use crate::error::ConfigError; use crate::settings::Settings; @@ -23,6 +23,8 @@ pub struct AgentConfig { pub max_cost_per_day_cents: Option, /// Maximum LLM/tool actions per hour. None = unlimited. pub max_actions_per_hour: Option, + /// Maximum daily LLM spend per user in cents. None = unlimited. + pub max_cost_per_user_per_day_cents: Option, /// Maximum tool-call iterations per agentic loop invocation. Default 50. pub max_tool_iterations: usize, /// When true, skip tool approval checks entirely. For benchmarks/CI. @@ -31,6 +33,13 @@ pub struct AgentConfig { pub default_timezone: String, /// Maximum tokens per job (0 = unlimited). pub max_tokens_per_job: u64, + /// Whether the deployment is multi-tenant (multiple users sharing one + /// instance). Auto-detected from GATEWAY_USER_TOKENS presence. + pub multi_tenant: bool, + /// Maximum concurrent LLM calls per user. None = use default (4). + pub max_llm_concurrent_per_user: Option, + /// Maximum concurrent jobs per user. None = use default (3). + pub max_jobs_concurrent_per_user: Option, } impl AgentConfig { @@ -49,10 +58,14 @@ impl AgentConfig { allow_local_tools: true, max_cost_per_day_cents: None, max_actions_per_hour: None, + max_cost_per_user_per_day_cents: None, max_tool_iterations: 10, auto_approve_tools: true, default_timezone: "UTC".to_string(), max_tokens_per_job: 0, + multi_tenant: false, + max_llm_concurrent_per_user: None, + max_jobs_concurrent_per_user: None, } } @@ -87,6 +100,7 @@ impl AgentConfig { allow_local_tools: parse_bool_env("ALLOW_LOCAL_TOOLS", false)?, max_cost_per_day_cents: parse_option_env("MAX_COST_PER_DAY_CENTS")?, max_actions_per_hour: parse_option_env("MAX_ACTIONS_PER_HOUR")?, + max_cost_per_user_per_day_cents: parse_option_env("MAX_COST_PER_USER_PER_DAY_CENTS")?, max_tool_iterations: parse_optional_env( "AGENT_MAX_TOOL_ITERATIONS", settings.agent.max_tool_iterations, @@ -112,6 +126,11 @@ impl AgentConfig { "AGENT_MAX_TOKENS_PER_JOB", settings.agent.max_tokens_per_job, )?, + // Auto-detected from GATEWAY_USER_TOKENS presence. Not a separate + // knob — multi-tenant mode is always implied by configuring user tokens. + multi_tenant: optional_env("GATEWAY_USER_TOKENS")?.is_some(), + max_llm_concurrent_per_user: parse_option_env("TENANT_MAX_LLM_CONCURRENT")?, + max_jobs_concurrent_per_user: parse_option_env("TENANT_MAX_JOBS_CONCURRENT")?, }) } } diff --git a/src/config/heartbeat.rs b/src/config/heartbeat.rs index 1dd456d7..09b8f0cd 100644 --- a/src/config/heartbeat.rs +++ b/src/config/heartbeat.rs @@ -21,6 +21,9 @@ pub struct HeartbeatConfig { pub quiet_hours_end: Option, /// Timezone for fire_at and quiet hours evaluation (IANA name). pub timezone: Option, + /// When true, cycle through all users with routines. Auto-detected from + /// GATEWAY_USER_TOKENS or set explicitly via HEARTBEAT_MULTI_TENANT. + pub multi_tenant: bool, } impl Default for HeartbeatConfig { @@ -34,6 +37,7 @@ impl Default for HeartbeatConfig { quiet_hours_start: None, quiet_hours_end: None, timezone: None, + multi_tenant: false, } } } @@ -101,6 +105,12 @@ impl HeartbeatConfig { } tz }, + // Auto-detect multi-tenant mode from GATEWAY_USER_TOKENS presence, + // or allow explicit override via HEARTBEAT_MULTI_TENANT. + multi_tenant: parse_bool_env( + "HEARTBEAT_MULTI_TENANT", + optional_env("GATEWAY_USER_TOKENS")?.is_some(), + )?, }) } } diff --git a/src/config/mod.rs b/src/config/mod.rs index dcda0fe9..a362fd09 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -312,13 +312,11 @@ impl Config { let tunnel = TunnelConfig::resolve(settings)?; let channels = ChannelsConfig::resolve(settings, &owner_id)?; - // Resolve workspace config using the gateway user_id for default layers. - let workspace_user_id = channels - .gateway - .as_ref() - .map(|gw| gw.user_id.as_str()) - .unwrap_or("default"); - let workspace = WorkspaceConfig::resolve(workspace_user_id)?; + // Resolve the startup workspace against the durable owner scope. The + // gateway may expose a distinct sender identity, but the base runtime + // workspace stays owner-scoped and per-user gateway workspaces are + // handled separately by WorkspacePool. + let workspace = WorkspaceConfig::resolve(&owner_id)?; Ok(Self { owner_id: owner_id.clone(), diff --git a/src/context/state.rs b/src/context/state.rs index f5307947..0bb1f29a 100644 --- a/src/context/state.rs +++ b/src/context/state.rs @@ -192,6 +192,9 @@ pub struct JobContext { /// but subsequent tools (e.g., `json`) may need the full output. This /// stash stores the complete, unsanitized output so tools can reference /// previous results by ID via `$tool_call_id` parameter syntax. + /// + /// Also used for cross-tool implicit state (keys prefixed with `__`) such + /// as `__routine_last_name` for fallback recovery in routine tool chains. #[serde(skip)] pub tool_output_stash: Arc>>, /// User's preferred timezone (IANA name, e.g. "America/New_York"). Defaults to "UTC". diff --git a/src/db/libsql/routines.rs b/src/db/libsql/routines.rs index 69c9f5c0..504d77dc 100644 --- a/src/db/libsql/routines.rs +++ b/src/db/libsql/routines.rs @@ -530,10 +530,24 @@ impl RoutineStore for LibSqlBackend { async fn get_webhook_routine_by_path( &self, path: &str, + user_id: Option<&str>, ) -> Result, DatabaseError> { let conn = self.connect().await?; - let mut rows = conn - .query( + let mut rows = if let Some(uid) = user_id { + conn.query( + &format!( + "SELECT {} FROM routines WHERE enabled = 1 AND trigger_type = 'webhook' \ + AND user_id = ?2 \ + AND (json_extract(trigger_config, '$.path') = ?1 \ + OR (json_extract(trigger_config, '$.path') IS NULL AND CAST(id AS TEXT) = ?1))", + ROUTINE_COLUMNS + ), + params![path, uid], + ) + .await + .map_err(|e| DatabaseError::Query(e.to_string()))? + } else { + conn.query( &format!( "SELECT {} FROM routines WHERE enabled = 1 AND trigger_type = 'webhook' \ AND (json_extract(trigger_config, '$.path') = ?1 \ @@ -543,7 +557,8 @@ impl RoutineStore for LibSqlBackend { params![path], ) .await - .map_err(|e| DatabaseError::Query(e.to_string()))?; + .map_err(|e| DatabaseError::Query(e.to_string()))? + }; match rows .next() diff --git a/src/db/mod.rs b/src/db/mod.rs index 6d984fed..d89b976e 100644 --- a/src/db/mod.rs +++ b/src/db/mod.rs @@ -545,6 +545,7 @@ pub trait RoutineStore: Send + Sync { async fn get_webhook_routine_by_path( &self, path: &str, + user_id: Option<&str>, ) -> Result, DatabaseError>; /// List routine runs that were dispatched as full_job but have not yet diff --git a/src/db/postgres.rs b/src/db/postgres.rs index 7bf76001..9e5ea9ce 100644 --- a/src/db/postgres.rs +++ b/src/db/postgres.rs @@ -529,8 +529,9 @@ impl RoutineStore for PgBackend { async fn get_webhook_routine_by_path( &self, path: &str, + user_id: Option<&str>, ) -> Result, DatabaseError> { - self.store.get_webhook_routine_by_path(path).await + self.store.get_webhook_routine_by_path(path, user_id).await } async fn list_dispatched_routine_runs(&self) -> Result, DatabaseError> { diff --git a/src/extensions/manager.rs b/src/extensions/manager.rs index 4fd479d9..6e322121 100644 --- a/src/extensions/manager.rs +++ b/src/extensions/manager.rs @@ -53,22 +53,6 @@ struct HostedOAuthFlowStart { flow: crate::cli::oauth_defaults::PendingOAuthFlow, } -fn hosted_proxy_client_secret( - client_secret: &Option, - builtin: Option<&crate::cli::oauth_defaults::OAuthCredentials>, - exchange_proxy_configured: bool, -) -> Option { - if !exchange_proxy_configured { - return client_secret.clone(); - } - - let builtin_secret = builtin.map(|credentials| credentials.client_secret); - match (client_secret, builtin_secret) { - (Some(resolved), Some(baked_in)) if resolved == baked_in => None, - _ => client_secret.clone(), - } -} - fn normalize_oauth_callback_path(path: &str) -> String { let trimmed_path = path.trim_end_matches('/'); if trimmed_path.is_empty() { @@ -423,9 +407,10 @@ pub struct ExtensionManager { /// when running in gateway mode, consumed by the web gateway's /// `/oauth/callback` handler. pending_oauth_flows: crate::cli::oauth_defaults::PendingOAuthRegistry, - /// Gateway auth token for authenticating with the platform token exchange proxy. - /// Read once at construction from `GATEWAY_AUTH_TOKEN` env var. - gateway_token: Option, + /// OAuth proxy auth token for authenticating with the hosted token exchange proxy. + /// Resolved once at construction from `IRONCLAW_OAUTH_PROXY_AUTH_TOKEN`, + /// then `GATEWAY_AUTH_TOKEN` as a backward-compatible fallback. + oauth_proxy_auth_token: Option, /// Relay config captured at startup. Used by `auth_channel_relay` and /// `activate_channel_relay` instead of re-reading env vars. relay_config: Option, @@ -561,7 +546,7 @@ impl ExtensionManager { activation_errors: RwLock::new(HashMap::new()), 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(), + oauth_proxy_auth_token: crate::cli::oauth_defaults::oauth_proxy_auth_token(), relay_config: crate::config::RelayConfig::from_env(), relay_event_tx: Arc::new(tokio::sync::Mutex::new(None)), relay_signing_secret_cache: Arc::new(std::sync::Mutex::new(None)), @@ -685,6 +670,66 @@ impl ExtensionManager { }) } + /// Resolve the relay URL override for an extension from settings. + /// + /// Returns `Some(url)` if a non-empty per-extension `relay_url` override is + /// set for the given extension; otherwise returns `None` and callers should + /// fall back to the env-level `RelayConfig`. + /// + /// Uses `self.user_id` (owner scope) for consistency with `configure()`, + /// which also writes setting_path fields under the owner scope. + /// + /// The override is validated: only `http` / `https` schemes are accepted + /// and the URL must not contain userinfo (embedded credentials). This + /// prevents a malicious override from exfiltrating the instance-wide relay + /// API key to an attacker-controlled host. + async fn effective_relay_url(&self, name: &str) -> Option { + if let Some(ref store) = self.store { + let key = format!("extensions.{name}.relay_url"); + if let Ok(Some(v)) = store.get_setting(&self.user_id, &key).await { + let url = v + .as_str() + .map(|s| s.trim().to_string()) + .filter(|s| !s.is_empty()); + if let Some(ref u) = url { + // Validate the override to prevent API-key exfiltration: + // only allow http(s) with no embedded credentials. + match url::Url::parse(u) { + Ok(parsed) + if (parsed.scheme() == "http" || parsed.scheme() == "https") + && parsed.username().is_empty() + && parsed.password().is_none() => + { + tracing::debug!( + extension = %name, + relay_url_host = %parsed.host_str().unwrap_or("unknown"), + "effective_relay_url: using per-extension override from settings" + ); + return url; + } + Ok(parsed) => { + tracing::warn!( + extension = %name, + scheme = %parsed.scheme(), + has_userinfo = !parsed.username().is_empty() || parsed.password().is_some(), + "effective_relay_url: rejecting override — \ + only http/https without embedded credentials is allowed" + ); + } + Err(e) => { + tracing::warn!( + extension = %name, + error = %e, + "effective_relay_url: rejecting override — invalid URL" + ); + } + } + } + } + } + None + } + /// Get the shared relay event sender for the webhook endpoint. pub fn relay_event_tx( &self, @@ -901,24 +946,67 @@ impl ExtensionManager { *self.relay_channel_manager.write().await = Some(channel_manager); } - /// Check if a channel name corresponds to a relay extension (has stored stream token + /// Check if a channel name corresponds to a relay extension (has stored team_id /// or is tracked in the installed relay extensions set). pub async fn is_relay_channel(&self, name: &str, user_id: &str) -> bool { // Check in-memory installed set first (supports no-store mode) if self.installed_relay_extensions.read().await.contains(name) { return true; } - // Then check for stored stream token - self.secrets - .exists(user_id, &format!("relay:{}:stream_token", name)) - .await - .unwrap_or(false) + // Check for stored team_id (persisted across restarts by the OAuth callback) + if let Some(ref store) = self.store { + let key = format!("relay:{}:team_id", name); + if let Ok(Some(v)) = store.get_setting(user_id, &key).await { + return v.as_str().is_some_and(|s| !s.is_empty()); + } + } + false + } + + /// Check whether a stored `team_id` setting exists for the given relay extension. + /// + /// Unlike [`is_relay_channel`], this does **not** consult the in-memory + /// `installed_relay_extensions` set — it only looks at the persistent settings + /// store. This distinction matters for `auth_channel_relay`: an extension can + /// be *installed* (present in the in-memory set) but not yet *authenticated* + /// (no OAuth completed, no team_id stored). + async fn has_stored_team_id(&self, name: &str, _user_id: &str) -> bool { + if let Some(ref store) = self.store { + let key = format!("relay:{}:team_id", name); + // Use owner scope (self.user_id) for consistency: the OAuth callback + // stores team_id under state.owner_id which maps to self.user_id. + match store.get_setting(&self.user_id, &key).await { + Ok(Some(v)) => { + let has_id = v.as_str().is_some_and(|s| !s.is_empty()); + tracing::debug!( + extension = %name, + has_team_id = has_id, + "has_stored_team_id: checked store" + ); + return has_id; + } + Ok(None) => { + tracing::debug!( + extension = %name, + "has_stored_team_id: no team_id setting found" + ); + } + Err(e) => { + tracing::warn!( + extension = %name, + error = %e, + "has_stored_team_id: failed to read from settings store" + ); + } + } + } + false } /// Restore persisted relay channels after startup. /// /// Loads the persisted active channel list, filters to relay types (those with - /// a stored stream token), and activates each via `activate_stored_relay()`. + /// a stored team_id setting), and activates each via `activate_stored_relay()`. /// Skips channels that are already active. /// /// Call this only after `set_relay_channel_manager()` or `set_channel_runtime()`. @@ -1141,7 +1229,7 @@ impl ExtensionManager { /// Broadcast an extension status change to the web UI via SSE. async fn broadcast_extension_status(&self, name: &str, status: &str, message: Option<&str>) { if let Some(ref sse) = *self.sse_manager.read().await { - sse.broadcast(crate::channels::web::types::SseEvent::ExtensionStatus { + sse.broadcast(ironclaw_common::AppEvent::ExtensionStatus { extension_name: name.to_string(), status: status.to_string(), message: message.map(|m| m.to_string()), @@ -1485,9 +1573,11 @@ impl ExtensionManager { if kind_filter.is_none() || kind_filter == Some(ExtensionKind::ChannelRelay) { let installed = self.installed_relay_extensions.read().await; let active_names = self.active_channel_names.read().await; + let errors = self.activation_errors.read().await; for name in installed.iter() { let active = active_names.contains(name); - let has_token = self.is_relay_channel(name, user_id).await; + let authenticated = self.has_stored_team_id(name, user_id).await; + let activation_error = errors.get(name).cloned(); let registry_entry = self .registry .get_with_kind(name, Some(ExtensionKind::ChannelRelay)) @@ -1500,14 +1590,14 @@ impl ExtensionManager { display_name, description, url: None, - authenticated: has_token, + authenticated, active, tools: Vec::new(), needs_setup: false, has_auth: true, derived: false, installed: true, - activation_error: None, + activation_error, version: None, }); } @@ -1691,7 +1781,22 @@ impl ExtensionManager { self.persist_active_channels(user_id).await; self.activation_errors.write().await.remove(name); - // Remove stored stream token + // Remove stored team_id setting and clean up secrets + if let Some(ref store) = self.store + && let Err(e) = store + .delete_setting(user_id, &format!("relay:{}:team_id", name)) + .await + { + tracing::warn!(error = %e, name, "Failed to delete relay team_id setting on removal"); + } + if let Err(e) = self + .secrets + .delete(user_id, &format!("relay:{}:oauth_state", name)) + .await + { + tracing::warn!(error = %e, name, "Failed to delete relay oauth_state secret on removal"); + } + // Clean up legacy stream_token secret from pre-webhook installs let _ = self .secrets .delete(user_id, &format!("relay:{}:stream_token", name)) @@ -2792,7 +2897,7 @@ impl ExtensionManager { user_id: user_id.to_string(), secrets: Arc::clone(&self.secrets), sse_manager: self.sse_manager.read().await.clone(), - gateway_token: self.gateway_token.clone(), + gateway_token: self.oauth_proxy_auth_token.clone(), token_exchange_extra_params, client_id_secret_name: if server.oauth.is_none() { Some(server.client_id_secret_name()) @@ -3287,7 +3392,7 @@ impl ExtensionManager { // apps. Sending the desktop secret would cause a client_id/secret // mismatch because the container's GOOGLE_OAUTH_CLIENT_ID is the web // app, not the desktop app. - let proxy_client_secret = hosted_proxy_client_secret( + let proxy_client_secret = oauth_defaults::hosted_proxy_client_secret( &client_secret, builtin.as_ref(), oauth_defaults::exchange_proxy_url().is_some(), @@ -3309,7 +3414,7 @@ impl ExtensionManager { user_id: user_id.to_string(), secrets: Arc::clone(&self.secrets), sse_manager: self.sse_manager.read().await.clone(), - gateway_token: self.gateway_token.clone(), + gateway_token: self.oauth_proxy_auth_token.clone(), token_exchange_extra_params: std::collections::HashMap::new(), client_id_secret_name: None, created_at: std::time::Instant::now(), @@ -3392,7 +3497,7 @@ impl ExtensionManager { } .await; - // Broadcast SSE event + // Broadcast auth result event let (success, message) = match result { Ok(()) => (true, format!("{} authenticated successfully", display_name)), Err(ref e) => ( @@ -3418,7 +3523,7 @@ impl ExtensionManager { } if let Some(ref sse) = sse_manager { - sse.broadcast(crate::channels::web::types::SseEvent::AuthCompleted { + sse.broadcast(ironclaw_common::AppEvent::AuthCompleted { extension_name: ext_name, success, message, @@ -4291,26 +4396,75 @@ impl ExtensionManager { /// /// For Slack: initiates OAuth flow (redirect-based). /// For Telegram: accepts a bot token, registers it with channel-relay, - /// and stores the returned stream token. + /// and stores the team_id setting. async fn auth_channel_relay( &self, name: &str, user_id: &str, ) -> Result { - // Check if already authenticated (stream token exists) - if self.is_relay_channel(name, user_id).await { + tracing::debug!( + extension = %name, + user_id = %user_id, + "auth_channel_relay: starting" + ); + + // Check if already authenticated by looking for a stored team_id. + // We intentionally skip the `installed_relay_extensions` in-memory set + // here because that set only tracks *installed* extensions — an extension + // can be installed (via registry) but not yet authenticated (no OAuth + // completed). Checking just `is_relay_channel()` would short-circuit + // to "authenticated" even when no team_id exists, preventing the OAuth + // flow from being offered to the user. + if self.has_stored_team_id(name, user_id).await { + tracing::debug!( + extension = %name, + "auth_channel_relay: already authenticated (team_id in store)" + ); return Ok(AuthResult::authenticated(name, ExtensionKind::ChannelRelay)); } + tracing::debug!( + extension = %name, + "auth_channel_relay: no stored team_id, initiating OAuth" + ); + // Use relay config captured at startup - let relay_config = self.relay_config()?; + let relay_config = self.relay_config().map_err(|e| { + tracing::warn!( + extension = %name, + error = %e, + "auth_channel_relay: relay config not available — \ + CHANNEL_RELAY_URL and CHANNEL_RELAY_API_KEY must be set" + ); + e + })?; + + // Allow per-extension URL override from settings + let effective_url = self + .effective_relay_url(name) + .await + .unwrap_or_else(|| relay_config.url.clone()); + + tracing::debug!( + extension = %name, + relay_url = %effective_url, + "auth_channel_relay: creating relay client for OAuth" + ); let client = crate::channels::relay::RelayClient::new( - relay_config.url.clone(), + effective_url.clone(), relay_config.api_key.clone(), relay_config.request_timeout_secs, ) - .map_err(|e| ExtensionError::Config(e.to_string()))?; + .map_err(|e| { + tracing::warn!( + extension = %name, + relay_url = %effective_url, + error = %e, + "auth_channel_relay: failed to create relay HTTP client" + ); + ExtensionError::Config(e.to_string()) + })?; // Generate CSRF nonce — IronClaw validates this on the callback to ensure // the OAuth completion is legitimate. Channel-relay embeds it in the signed @@ -4322,18 +4476,44 @@ impl ExtensionManager { self.secrets .create(user_id, CreateSecretParams::new(&state_key, &state_nonce)) .await - .map_err(|e| ExtensionError::AuthFailed(format!("Failed to store OAuth state: {e}")))?; + .map_err(|e| { + tracing::warn!( + extension = %name, + error = %e, + "auth_channel_relay: failed to store OAuth state nonce" + ); + ExtensionError::AuthFailed(format!("Failed to store OAuth state: {e}")) + })?; // Channel-relay derives all URLs from trusted instance_url in chat-api. // We only pass the nonce for CSRF validation on the callback. + tracing::debug!( + extension = %name, + relay_url = %effective_url, + "auth_channel_relay: calling initiate_oauth on channel-relay" + ); match client.initiate_oauth(Some(&state_nonce)).await { - Ok(auth_url) => Ok(AuthResult::awaiting_authorization( - name, - ExtensionKind::ChannelRelay, - auth_url, - "redirect".to_string(), - )), - Err(e) => Err(ExtensionError::AuthFailed(e.to_string())), + Ok(auth_url) => { + tracing::info!( + extension = %name, + "auth_channel_relay: OAuth URL obtained, awaiting user authorization" + ); + Ok(AuthResult::awaiting_authorization( + name, + ExtensionKind::ChannelRelay, + auth_url, + "redirect".to_string(), + )) + } + Err(e) => { + tracing::warn!( + extension = %name, + relay_url = %effective_url, + error = %e, + "auth_channel_relay: initiate_oauth call to channel-relay failed" + ); + Err(ExtensionError::AuthFailed(e.to_string())) + } } } @@ -4343,46 +4523,112 @@ impl ExtensionManager { name: &str, user_id: &str, ) -> Result { - let token_key = format!("relay:{}:stream_token", name); + tracing::debug!( + extension = %name, + user_id = %user_id, + "activate_channel_relay: starting" + ); + let team_id_key = format!("relay:{}:team_id", name); - // Check if we have a stream token - // Verify auth: stream token must exist (even though we don't use it in this constructor path) - let _stream_token = match self.secrets.get_decrypted(user_id, &token_key).await { - Ok(secret) => secret.expose().to_string(), - Err(_) => { - return Err(ExtensionError::AuthRequired); - } - }; - - // Get team_id from settings + // Get team_id from settings (stored by the OAuth callback) let team_id = if let Some(ref store) = self.store { - store - .get_setting(user_id, &team_id_key) - .await - .ok() - .flatten() - .and_then(|v| v.as_str().map(|s| s.to_string())) - .unwrap_or_default() + match store.get_setting(user_id, &team_id_key).await { + Ok(Some(v)) => { + let id = v.as_str().map(|s| s.to_string()).unwrap_or_default(); + tracing::debug!( + extension = %name, + team_id_empty = id.is_empty(), + "activate_channel_relay: loaded team_id from store" + ); + id + } + Ok(None) => { + tracing::debug!( + extension = %name, + setting_key = %team_id_key, + "activate_channel_relay: no team_id in settings store" + ); + String::new() + } + Err(e) => { + tracing::warn!( + extension = %name, + error = %e, + "activate_channel_relay: failed to read team_id from settings store" + ); + String::new() + } + } } else { + tracing::debug!( + extension = %name, + "activate_channel_relay: no settings store available" + ); String::new() }; + if team_id.is_empty() { + tracing::debug!( + extension = %name, + "activate_channel_relay: team_id is empty, returning AuthRequired" + ); + return Err(ExtensionError::AuthRequired); + } + // Use relay config captured at startup - let relay_config = self.relay_config()?; + let relay_config = self.relay_config().map_err(|e| { + tracing::warn!( + extension = %name, + error = %e, + "activate_channel_relay: relay config not available" + ); + e + })?; + + // Allow per-extension URL override from settings + let effective_url = self + .effective_relay_url(name) + .await + .unwrap_or_else(|| relay_config.url.clone()); + + tracing::debug!( + extension = %name, + relay_url = %effective_url, + "activate_channel_relay: relay config loaded" + ); let instance_id = self.relay_instance_id(relay_config, user_id); let client = crate::channels::relay::RelayClient::new( - relay_config.url.clone(), + effective_url.clone(), relay_config.api_key.clone(), relay_config.request_timeout_secs, ) - .map_err(|e| ExtensionError::ActivationFailed(e.to_string()))?; + .map_err(|e| { + tracing::warn!( + extension = %name, + relay_url = %effective_url, + error = %e, + "activate_channel_relay: failed to create relay HTTP client" + ); + ExtensionError::ActivationFailed(e.to_string()) + })?; // Fetch the per-instance signing secret from channel-relay. // This must succeed — there is no fallback. + tracing::debug!( + extension = %name, + relay_url = %effective_url, + "activate_channel_relay: fetching signing secret from channel-relay" + ); let signing_secret = client.get_signing_secret(&team_id).await.map_err(|e| { + tracing::warn!( + extension = %name, + relay_url = %effective_url, + error = %e, + "activate_channel_relay: failed to fetch signing secret from channel-relay" + ); ExtensionError::Config(format!("Failed to fetch relay signing secret: {e}")) })?; @@ -4401,16 +4647,29 @@ impl ExtensionManager { // Hot-add to channel manager let cm_guard = self.relay_channel_manager.read().await; let channel_mgr = cm_guard.as_ref().ok_or_else(|| { + tracing::warn!( + extension = %name, + "activate_channel_relay: channel manager not initialized" + ); ExtensionError::ActivationFailed("Channel manager not initialized".to_string()) })?; - channel_mgr - .hot_add(Box::new(channel)) - .await - .map_err(|e| ExtensionError::ActivationFailed(e.to_string()))?; + channel_mgr.hot_add(Box::new(channel)).await.map_err(|e| { + tracing::warn!( + extension = %name, + error = %e, + "activate_channel_relay: hot_add to channel manager failed" + ); + ExtensionError::ActivationFailed(e.to_string()) + })?; if let Ok(mut cache) = self.relay_signing_secret_cache.lock() { *cache = Some(signing_secret); + } else { + tracing::warn!( + extension = %name, + "activate_channel_relay: failed to cache signing secret (mutex poisoned)" + ); } // Store the event sender so the web gateway's relay webhook endpoint can push events @@ -4428,6 +4687,12 @@ impl ExtensionManager { self.broadcast_extension_status(name, "active", Some(&status_msg)) .await; + tracing::info!( + extension = %name, + instance_id = %instance_id, + "activate_channel_relay: relay channel activated successfully" + ); + Ok(ActivateResult { name: name.to_string(), kind: ExtensionKind::ChannelRelay, @@ -4477,11 +4742,11 @@ impl ExtensionManager { return Ok(ExtensionKind::WasmChannel); } - // Check channel-relay extensions (installed in memory or has stored token) + // Check channel-relay extensions (installed in memory or has stored team_id) if self.installed_relay_extensions.read().await.contains(name) { return Ok(ExtensionKind::ChannelRelay); } - // Also check if there's a stored stream token (persisted across restarts) + // Also check if there's a stored team_id setting (persisted across restarts) if self.is_relay_channel(name, user_id).await { return Ok(ExtensionKind::ChannelRelay); } @@ -4707,6 +4972,41 @@ impl ExtensionManager { } Ok(ExtensionSetupSchema { secrets, fields }) } + ExtensionKind::ChannelRelay => { + let relay_url_key = format!("extensions.{name}.relay_url"); + let current_url = if let Some(ref store) = self.store { + match store.get_setting(&self.user_id, &relay_url_key).await { + Ok(value_opt) => value_opt + .and_then(|v| v.as_str().map(|s| s.to_string())) + .filter(|s| !s.is_empty()), + Err(e) => { + tracing::warn!( + extension = %name, + setting_key = %relay_url_key, + error = %e, + "get_setup_schema: failed to read relay_url from settings" + ); + None + } + } + } else { + None + }; + let env_url = self.relay_config.as_ref().map(|c| c.url.as_str()); + Ok(ExtensionSetupSchema { + secrets: Vec::new(), + fields: vec![crate::channels::web::types::SetupFieldInfo { + name: "relay_url".to_string(), + prompt: format!( + "Channel-relay service URL (leave empty to use env default{})", + env_url.map(|u| format!(": {u}")).unwrap_or_default() + ), + optional: true, + provided: current_url.is_some(), + input_type: crate::tools::wasm::ToolSetupFieldInputType::Text, + }], + }) + } _ => Ok(ExtensionSetupSchema { secrets: Vec::new(), fields: Vec::new(), @@ -5116,9 +5416,15 @@ impl ExtensionManager { (names, Vec::new()) } ExtensionKind::ChannelRelay => { - let mut names = std::collections::HashSet::new(); - names.insert(format!("relay:{}:stream_token", name)); - (names, Vec::new()) + let relay_fields = vec![crate::tools::wasm::ToolFieldSetupSchema { + name: "relay_url".to_string(), + prompt: "Channel-relay service URL override".to_string(), + optional: true, + setting_path: Some(format!("extensions.{name}.relay_url")), + input_type: crate::tools::wasm::ToolSetupFieldInputType::Text, + restart_required: false, + }]; + (std::collections::HashSet::new(), relay_fields) } }; @@ -5210,13 +5516,28 @@ impl ExtensionManager { ))); } let trimmed = field_value.trim(); + let field_def = setup_field_defs.get(field_name); + + // Empty value on an optional field with a setting_path: clear the + // stored override so the system reverts to the env/default value. if trimmed.is_empty() { + if let Some(def) = field_def + && def.optional + { + stored_fields.remove(field_name); + if let Some(setting_path) = &def.setting_path { + Self::validate_setup_setting_path(name, setting_path)?; + if let Some(store) = self.store.as_ref() { + let _ = store.delete_setting(&self.user_id, setting_path).await; + } + } + } continue; } stored_fields.insert(field_name.clone(), trimmed.to_string()); - if let Some(field_def) = setup_field_defs.get(field_name) { + if let Some(field_def) = field_def { if field_def.restart_required { restart_required = true; } @@ -5550,7 +5871,9 @@ impl ExtensionManager { .map_err(|e| ExtensionError::NotInstalled(e.to_string()))?; server.token_secret_name() } - ExtensionKind::ChannelRelay => format!("relay:{}:stream_token", name), + ExtensionKind::ChannelRelay => { + return Err(ExtensionError::AuthRequired); + } }; let mut secrets = std::collections::HashMap::new(); @@ -5818,7 +6141,7 @@ mod tests { use crate::extensions::manager::{ ChannelRuntimeState, FallbackDecision, TelegramBindingData, TelegramBindingResult, TelegramOwnerBindingState, build_wasm_channel_runtime_config_updates, - combine_install_errors, fallback_decision, hosted_proxy_client_secret, infer_kind_from_url, + combine_install_errors, fallback_decision, infer_kind_from_url, normalize_hosted_callback_url, send_telegram_text_message, telegram_message_matches_verification_code, }; @@ -7453,7 +7776,7 @@ mod tests { let dir = tempfile::tempdir().expect("temp dir"); let mgr = make_test_manager(None, dir.path().to_path_buf()); - // No token stored → not a relay channel + // No store configured, no team_id → not a relay channel assert!(!mgr.is_relay_channel("slack-relay", "test").await); } @@ -7472,6 +7795,39 @@ mod tests { ); } + /// Regression: installed-but-not-authenticated relay must NOT short-circuit + /// `auth_channel_relay()` to "authenticated". Previously, `auth_channel_relay` + /// called `is_relay_channel()` which checked the in-memory + /// `installed_relay_extensions` set; that returned `true` even when no team_id + /// existed in the store, so the OAuth URL was never offered. + #[tokio::test] + async fn test_auth_channel_relay_installed_without_team_id_is_not_authenticated() { + let dir = tempfile::tempdir().expect("temp dir"); + let mgr = make_test_manager(None, dir.path().to_path_buf()); + + // Mark as installed (simulates clicking Install in the UI) + mgr.installed_relay_extensions + .write() + .await + .insert("slack-relay".to_string()); + + // Without a stored team_id, auth should NOT return authenticated. + // It should fail because relay config is missing (no CHANNEL_RELAY_URL), + // but the key assertion is that it does NOT return Ok(authenticated). + let result = mgr.auth_channel_relay("slack-relay", "test").await; + match result { + Ok(ref auth_result) if auth_result.is_authenticated() => { + panic!( + "auth_channel_relay returned authenticated for installed-but-no-team-id relay; \ + expected either an OAuth URL or a config error" + ); + } + _ => { + // Config error (no relay URL) or awaiting_authorization — both are correct + } + } + } + #[tokio::test] async fn test_remove_relay_shuts_down_via_relay_channel_manager() { // Regression: remove() only checked channel_runtime for shutdown, missing @@ -8316,19 +8672,13 @@ mod tests { .await .insert("test-relay".to_string()); - // configure() should dispatch to activate_channel_relay(), not - // activate_wasm_channel(). Both will fail (no runtime configured), - // but the error should be about relay config, not WASM channels. - let mut secrets = std::collections::HashMap::new(); - secrets.insert( - "relay:test-relay:stream_token".to_string(), - "tok".to_string(), - ); - + // configure() with empty secrets should dispatch to + // activate_channel_relay(), not activate_wasm_channel(). Relay auth + // is OAuth-only so there are no manual secrets to pass. let result = mgr .configure( "test-relay", - &secrets, + &std::collections::HashMap::new(), &std::collections::HashMap::new(), "test", ) @@ -8340,7 +8690,6 @@ mod tests { ); let result = result.unwrap(); - // Activation will fail (no relay config), but secrets should still be stored assert!( !result.activated, "activation should fail without relay config" @@ -8350,15 +8699,6 @@ mod tests { "error should not mention WASM — got: {}", result.message ); - - // Verify the secret was stored - assert!( - mgr.secrets - .exists("test", "relay:test-relay:stream_token") - .await - .unwrap_or(false), - "configure should have stored the relay stream token" - ); } #[test] fn test_validation_failed_is_distinct_error_variant() { @@ -8424,7 +8764,8 @@ mod tests { let builtin_ref = builtin.as_ref(); let secret = Some(builtin_ref.unwrap().client_secret.to_string()); - let result = hosted_proxy_client_secret(&secret, builtin_ref, true); + let result = + crate::cli::oauth_defaults::hosted_proxy_client_secret(&secret, builtin_ref, true); assert_eq!( result, None, "built-in desktop secret must be suppressed when the exchange proxy is configured" @@ -8436,7 +8777,8 @@ mod tests { let builtin = crate::cli::oauth_defaults::builtin_credentials("google_oauth_token"); let secret = Some("user-entered-custom-secret".to_string()); - let result = hosted_proxy_client_secret(&secret, builtin.as_ref(), true); + let result = + crate::cli::oauth_defaults::hosted_proxy_client_secret(&secret, builtin.as_ref(), true); assert_eq!( result, Some("user-entered-custom-secret".to_string()), @@ -8450,7 +8792,8 @@ mod tests { let builtin_ref = builtin.as_ref(); let secret = Some(builtin_ref.unwrap().client_secret.to_string()); - let result = hosted_proxy_client_secret(&secret, builtin_ref, false); + let result = + crate::cli::oauth_defaults::hosted_proxy_client_secret(&secret, builtin_ref, false); assert_eq!( result, secret, "built-in secret must be kept when the callback will exchange directly" @@ -8461,7 +8804,8 @@ mod tests { fn test_proxy_client_secret_none_stays_none() { let builtin = crate::cli::oauth_defaults::builtin_credentials("google_oauth_token"); - let result = hosted_proxy_client_secret(&None, builtin.as_ref(), true); + let result = + crate::cli::oauth_defaults::hosted_proxy_client_secret(&None, builtin.as_ref(), true); assert_eq!( result, None, "None secret stays None even when the exchange proxy is configured" @@ -8475,7 +8819,8 @@ mod tests { assert!(builtin.is_none()); let secret = Some("dcr-secret".to_string()); - let result = hosted_proxy_client_secret(&secret, builtin.as_ref(), true); + let result = + crate::cli::oauth_defaults::hosted_proxy_client_secret(&secret, builtin.as_ref(), true); assert_eq!( result, Some("dcr-secret".to_string()), diff --git a/src/history/store.rs b/src/history/store.rs index 1e4cdd82..625e8b1e 100644 --- a/src/history/store.rs +++ b/src/history/store.rs @@ -1162,15 +1162,25 @@ impl Store { pub async fn get_webhook_routine_by_path( &self, path: &str, + user_id: Option<&str>, ) -> Result, DatabaseError> { let conn = self.conn().await?; - let row = conn - .query_opt( + let row = if let Some(uid) = user_id { + conn.query_opt( + "SELECT * FROM routines WHERE enabled AND trigger_type = 'webhook' \ + AND user_id = $2 \ + AND (trigger_config->>'path' = $1 OR (trigger_config->>'path' IS NULL AND id::text = $1))", + &[&path, &uid], + ) + .await? + } else { + conn.query_opt( "SELECT * FROM routines WHERE enabled AND trigger_type = 'webhook' \ AND (trigger_config->>'path' = $1 OR (trigger_config->>'path' IS NULL AND id::text = $1))", &[&path], ) - .await?; + .await? + }; row.as_ref().map(row_to_routine).transpose() } diff --git a/src/lib.rs b/src/lib.rs index 9bdce343..dbdd2260 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -69,6 +69,7 @@ pub mod service; pub mod settings; pub mod setup; pub mod skills; +pub mod tenant; pub mod timezone; pub mod tools; pub mod tracing_fmt; diff --git a/src/llm/anthropic_oauth.rs b/src/llm/anthropic_oauth.rs index 490fbc3f..c94c90e5 100644 --- a/src/llm/anthropic_oauth.rs +++ b/src/llm/anthropic_oauth.rs @@ -575,6 +575,7 @@ fn extract_response_content(response: &AnthropicResponse) -> (Option, Ve id: id.clone(), name: name.clone(), arguments: input.clone(), + reasoning: None, }); } } @@ -623,6 +624,7 @@ mod tests { id: "call_1".to_string(), name: "search".to_string(), arguments: serde_json::json!({"q": "test"}), + reasoning: None, }]; let messages = vec![ ChatMessage::user("Search for test"), diff --git a/src/llm/bedrock.rs b/src/llm/bedrock.rs index 5d6e121e..b5f7badd 100644 --- a/src/llm/bedrock.rs +++ b/src/llm/bedrock.rs @@ -522,6 +522,7 @@ fn extract_content_blocks( id: tu.tool_use_id().to_string(), name: tu.name().to_string(), arguments: document_to_json(tu.input()), + reasoning: None, }); } // Ignore reasoning, citations, images, etc. @@ -759,11 +760,13 @@ mod tests { id: "call_1".to_string(), name: "echo".to_string(), arguments: serde_json::json!({"text": "hi"}), + reasoning: None, }; let tc2 = crate::llm::provider::ToolCall { id: "call_2".to_string(), name: "time".to_string(), arguments: serde_json::json!({}), + reasoning: None, }; let messages = vec![ @@ -802,6 +805,7 @@ mod tests { id: "call_1".to_string(), name: "search".to_string(), arguments: serde_json::json!({"query": "test"}), + reasoning: None, }; let messages = vec![ @@ -825,6 +829,7 @@ mod tests { id: "call_1".to_string(), name: "echo".to_string(), arguments: serde_json::json!({}), + reasoning: None, }; let messages = vec![ @@ -989,11 +994,13 @@ mod tests { id: "call_abc".to_string(), name: "get_weather".to_string(), arguments: serde_json::json!({"city": "NYC"}), + reasoning: None, }; let tc2 = crate::llm::provider::ToolCall { id: "call_def".to_string(), name: "get_time".to_string(), arguments: serde_json::json!({"tz": "EST"}), + reasoning: None, }; let messages = vec![ diff --git a/src/llm/codex_chatgpt.rs b/src/llm/codex_chatgpt.rs index 56cb3378..e7dcf40d 100644 --- a/src/llm/codex_chatgpt.rs +++ b/src/llm/codex_chatgpt.rs @@ -732,6 +732,7 @@ impl LlmProvider for CodexChatGptProvider { id: tc.call_id, name: tc.name, arguments: args, + reasoning: None, } }) .collect(); @@ -825,6 +826,7 @@ mod tests { id: "call_1".to_string(), name: "search".to_string(), arguments: json!({"query": "rust"}), + reasoning: None, }; let msg = ChatMessage::assistant_with_tool_calls(Some("thinking...".into()), vec![tc]); let items = CodexChatGptProvider::message_to_input_items(&msg); diff --git a/src/llm/gemini_oauth.rs b/src/llm/gemini_oauth.rs index b36eb595..a19eec12 100644 --- a/src/llm/gemini_oauth.rs +++ b/src/llm/gemini_oauth.rs @@ -1898,6 +1898,7 @@ impl GeminiOauthProvider { id, name, arguments: args, + reasoning: None, }); } } diff --git a/src/llm/github_copilot.rs b/src/llm/github_copilot.rs index b173191a..c7a24b1a 100644 --- a/src/llm/github_copilot.rs +++ b/src/llm/github_copilot.rs @@ -596,6 +596,7 @@ fn extract_choice_content(choice: &OpenAiChoice) -> (Option, Vec = - raw_messages.into_iter().map(|m| m.into()).collect(); + let raw: Vec = raw_messages.into_iter().map(|m| m.into()).collect(); + + // NEAR AI rejects `role:"tool"` messages even on text-only completion paths. + // Apply the same flattening used by complete_with_tools(). + let messages = if self.flatten_tool_messages { + flatten_tool_messages(raw) + } else { + raw + }; let request = ChatCompletionRequest { model, @@ -551,6 +558,7 @@ impl LlmProvider for NearAiChatProvider { id: tc.id, name: tc.function.name, arguments, + reasoning: None, } }) .collect(); @@ -1144,11 +1152,13 @@ mod tests { id: "call_1".to_string(), name: "list_issues".to_string(), arguments: serde_json::json!({"owner": "foo", "repo": "bar"}), + reasoning: None, }, ToolCall { id: "call_2".to_string(), name: "search".to_string(), arguments: serde_json::json!({"query": "test"}), + reasoning: None, }, ]; @@ -1181,6 +1191,7 @@ mod tests { id: "call_1".to_string(), name: "test".to_string(), arguments: serde_json::json!({"key": "value"}), + reasoning: None, }; let msg = ChatMessage::assistant_with_tool_calls(None, vec![tc]); let chat_msg: ChatCompletionMessage = msg.into(); @@ -1424,6 +1435,7 @@ mod tests { id: tc.id, name: tc.function.name, arguments, + reasoning: None, } }) .collect(); @@ -1473,6 +1485,7 @@ mod tests { id: tc.id, name: tc.function.name, arguments, + reasoning: None, } }) .collect(); @@ -2095,6 +2108,7 @@ mod tests { id: "call_1".to_string(), name: "test".to_string(), arguments: serde_json::json!({}), + reasoning: None, }], ); let chat_msg: ChatCompletionMessage = msg.into(); @@ -2164,6 +2178,65 @@ mod tests { assert_eq!(deserialized.function.arguments, r#"{"city":"London"}"#); } + // -- flatten_tool_messages in complete() path ---------------------------- + + #[test] + fn test_flatten_applied_on_text_only_path() { + // Verify that flatten_tool_messages converts tool-role messages to user + // messages (mirrors the complete_with_tools path). + let messages = vec![ + ChatCompletionMessage { + role: "user".to_string(), + content: Some(MessageContent::Text("run it".to_string())), + tool_call_id: None, + name: None, + tool_calls: None, + }, + ChatCompletionMessage { + role: "tool".to_string(), + content: Some(MessageContent::Text("ok".to_string())), + tool_call_id: Some("call_1".to_string()), + name: Some("run_cmd".to_string()), + tool_calls: None, + }, + ]; + let flattened = flatten_tool_messages(messages); + assert_eq!(flattened.len(), 2); + assert_eq!(flattened[1].role, "user"); + let text = flattened[1] + .content + .as_ref() + .and_then(|c| c.as_text()) + .unwrap(); + assert!(text.contains("run_cmd"), "should reference tool name"); + assert!(text.contains("ok"), "should include tool result"); + } + + #[test] + fn test_no_flatten_when_no_tool_messages() { + // When there are no tool-role messages, flatten_tool_messages is a no-op. + let messages = vec![ + ChatCompletionMessage { + role: "user".to_string(), + content: Some(MessageContent::Text("hi".to_string())), + tool_call_id: None, + name: None, + tool_calls: None, + }, + ChatCompletionMessage { + role: "assistant".to_string(), + content: Some(MessageContent::Text("hello".to_string())), + tool_call_id: None, + name: None, + tool_calls: None, + }, + ]; + let result = flatten_tool_messages(messages); + // No tool messages → unchanged roles + assert_eq!(result[0].role, "user"); + assert_eq!(result[1].role, "assistant"); + } + // -- api_url edge cases --------------------------------------------------- #[test] diff --git a/src/llm/openai_codex_provider.rs b/src/llm/openai_codex_provider.rs index 9e3aa955..3449a08a 100644 --- a/src/llm/openai_codex_provider.rs +++ b/src/llm/openai_codex_provider.rs @@ -625,6 +625,7 @@ fn parse_sse_response(body: &str) -> Result { id: state.call_id, name: state.name, arguments, + reasoning: None, }); } else { // Fallback: extract directly from the item @@ -650,6 +651,7 @@ fn parse_sse_response(body: &str) -> Result { id: call_id, name, arguments, + reasoning: None, }); } } @@ -727,6 +729,7 @@ fn parse_sse_response(body: &str) -> Result { id: state.call_id, name: state.name, arguments, + reasoning: None, }); } } @@ -822,11 +825,13 @@ mod tests { id: "call_1".to_string(), name: "search".to_string(), arguments: serde_json::json!({"query": "test"}), + reasoning: None, }, ToolCall { id: "call_2".to_string(), name: "read".to_string(), arguments: serde_json::json!({"path": "/tmp"}), + reasoning: None, }, ]; let msg = diff --git a/src/llm/provider.rs b/src/llm/provider.rs index bb45ec68..8afd914a 100644 --- a/src/llm/provider.rs +++ b/src/llm/provider.rs @@ -231,6 +231,10 @@ pub struct ToolCall { pub id: String, pub name: String, pub arguments: serde_json::Value, + /// Optional reasoning for why this tool was chosen — supplied by the provider + /// or derived from the shared response content as a fallback. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub reasoning: Option, } /// Generate a tool-call ID that satisfies all providers. @@ -637,6 +641,7 @@ mod tests { id: "call_1".to_string(), name: "echo".to_string(), arguments: serde_json::json!({}), + reasoning: None, }; let mut messages = vec![ ChatMessage::user("hello"), @@ -680,6 +685,7 @@ mod tests { id: "call_1".to_string(), name: "echo".to_string(), arguments: serde_json::json!({}), + reasoning: None, }; let mut messages = vec![ ChatMessage::user("test"), @@ -705,11 +711,13 @@ mod tests { id: "call_sel_1".to_string(), name: "search".to_string(), arguments: serde_json::json!({"q": "test"}), + reasoning: None, }; let tc2 = ToolCall { id: "call_sel_2".to_string(), name: "http".to_string(), arguments: serde_json::json!({"url": "https://example.com"}), + reasoning: None, }; let mut messages = vec![ ChatMessage::system("You are a helpful assistant."), diff --git a/src/llm/reasoning.rs b/src/llm/reasoning.rs index cbec297b..6e078ac7 100644 --- a/src/llm/reasoning.rs +++ b/src/llm/reasoning.rs @@ -8,8 +8,8 @@ use serde::{Deserialize, Serialize}; use crate::llm::error::LlmError; use crate::llm::{ - ChatMessage, CompletionRequest, LlmProvider, Role, ToolCall, ToolCompletionRequest, - ToolDefinition, + ChatMessage, CompletionRequest, FinishReason, LlmProvider, Role, ToolCall, + ToolCompletionRequest, ToolDefinition, }; /// Token the agent returns when it has nothing to say (e.g. in group chats). @@ -23,6 +23,13 @@ You said you would perform an action, but you did not include any tool calls.\n\ Do NOT describe what you intend to do — actually call the tool now.\n\ Use the tool_calls mechanism to invoke the appropriate tool."; +/// Notice injected when the LLM's response was truncated mid-tool-call, +/// causing incomplete parameters. Tells the LLM to try a different approach. +pub const TRUNCATED_TOOL_CALL_NOTICE: &str = "\ +Your previous response was truncated while generating tool call parameters. \ +The tool calls were discarded. Please try a different approach — \ +summarize or transform the data instead of echoing it verbatim in a tool call."; + /// Seed value used as the second argument to `generate_tool_call_id` when /// recovering tool calls from malformed LLM text responses. This must differ /// from the `0` seed used in `rig_adapter::normalized_tool_call_id` to avoid @@ -194,11 +201,17 @@ pub struct ReasoningContext { pub metadata: std::collections::HashMap, /// When true, force a text-only response (ignore available tools). /// Used by the agentic loop to guarantee termination near the iteration limit. + /// Sticky: once set, never cleared within a loop invocation. Callers must + /// create a fresh `ReasoningContext` per `run_agentic_loop()` call. pub force_text: bool, /// Pre-built system prompt. When set, `respond_with_tools` uses this directly /// instead of calling `build_system_prompt_with_tools`. Allows callers to build /// the prompt once and reuse it across iterations. pub system_prompt: Option, + /// Per-user model override. When set, completion requests use this model + /// instead of the provider's default. Only effective with providers that + /// support per-request model overrides (e.g. NearAI). + pub model_override: Option, } impl ReasoningContext { @@ -212,6 +225,7 @@ impl ReasoningContext { metadata: std::collections::HashMap::new(), force_text: false, system_prompt: None, + model_override: None, } } @@ -344,6 +358,7 @@ pub enum RespondResult { pub struct RespondOutput { pub result: RespondResult, pub usage: TokenUsage, + pub finish_reason: FinishReason, } /// Reasoning engine for the agent. @@ -525,17 +540,46 @@ impl Reasoning { let response = self.llm.complete_with_tools(request).await?; - let reasoning = response.content.unwrap_or_default(); + // If the response was truncated, tool call parameters are likely incomplete. + // Return empty so the caller can fall through to respond_with_tools() which + // has a larger output token budget. + if response.finish_reason == FinishReason::Length { + tracing::warn!( + "select_tools response truncated (finish_reason=Length), \ + discarding potentially incomplete tool selections" + ); + return Ok(vec![]); + } + + let shared_reasoning = response + .content + .map(|c| { + let pre_truncated = truncate_at_tool_tags(&c); + clean_response(&pre_truncated) + }) + .unwrap_or_default(); let selections: Vec = response .tool_calls .into_iter() - .map(|tool_call| ToolSelection { - tool_name: tool_call.name, - parameters: tool_call.arguments, - reasoning: reasoning.clone(), - alternatives: vec![], - tool_call_id: tool_call.id, + .map(|tool_call| { + // Prefer per-tool reasoning if the provider supplied it, + // otherwise fall back to the shared response content. + let rationale = tool_call + .reasoning + .map(|r| { + let pre_truncated = truncate_at_tool_tags(&r); + clean_response(&pre_truncated) + }) + .filter(|r| !r.trim().is_empty()) + .unwrap_or_else(|| shared_reasoning.clone()); + ToolSelection { + tool_name: tool_call.name, + parameters: tool_call.arguments, + reasoning: rationale, + alternatives: vec![], + tool_call_id: tool_call.id, + } }) .collect(); @@ -653,6 +697,9 @@ Respond in JSON format: .with_temperature(0.7) .with_tool_choice("auto"); request.metadata = context.metadata.clone(); + if let Some(ref model) = context.model_override { + request.model = Some(model.clone()); + } let response = self.llm.complete_with_tools(request).await?; let usage = TokenUsage { @@ -664,15 +711,39 @@ Respond in JSON format: // If there were tool calls, return them for execution if !response.tool_calls.is_empty() { + let narrative = response.content.map(|c| { + let pre_truncated = truncate_at_tool_tags(&c); + clean_response(&pre_truncated) + }); + // Populate per-tool reasoning from the shared narrative when the + // provider did not supply per-tool rationale. + let tool_calls: Vec = response + .tool_calls + .into_iter() + .map(|mut tc| { + if tc.reasoning.as_ref().is_none_or(|r| r.trim().is_empty()) { + tc.reasoning = narrative.as_ref().filter(|n| !n.is_empty()).cloned(); + } else { + // Clean provider-supplied per-tool reasoning the same way + // we clean the shared narrative (strip thinking/tool tags). + tc.reasoning = tc + .reasoning + .map(|r| { + let pre_truncated = truncate_at_tool_tags(&r); + clean_response(&pre_truncated) + }) + .filter(|r| !r.trim().is_empty()); + } + tc + }) + .collect(); return Ok(RespondOutput { result: RespondResult::ToolCalls { - tool_calls: response.tool_calls, - content: response.content.map(|c| { - let pre_truncated = truncate_at_tool_tags(&c); - clean_response(&pre_truncated) - }), + tool_calls, + content: narrative, }, usage, + finish_reason: response.finish_reason, }); } @@ -700,6 +771,7 @@ Respond in JSON format: }, }, usage, + finish_reason: response.finish_reason, }); } @@ -725,6 +797,7 @@ Respond in JSON format: Ok(RespondOutput { result: RespondResult::Text(final_text), usage, + finish_reason: response.finish_reason, }) } else { // No tools, use simple completion @@ -732,6 +805,9 @@ Respond in JSON format: .with_max_tokens(4096) .with_temperature(0.7); request.metadata = context.metadata.clone(); + if let Some(ref model) = context.model_override { + request.model = Some(model.clone()); + } let response = self.llm.complete(request).await?; let pre_truncated = truncate_at_tool_tags(&response.content); @@ -753,6 +829,7 @@ Respond in JSON format: cache_read_input_tokens: response.cache_read_input_tokens, cache_creation_input_tokens: response.cache_creation_input_tokens, }, + finish_reason: response.finish_reason, }) } } @@ -1293,6 +1370,49 @@ fn is_inside_code(pos: usize, regions: &[CodeRegion]) -> bool { regions.iter().any(|r| pos >= r.start && pos < r.end) } +/// Check whether a byte range overlaps any code region. +fn overlaps_code_region(start: usize, end: usize, regions: &[CodeRegion]) -> bool { + regions.iter().any(|r| start < r.end && end > r.start) +} + +/// Return the byte bounds of the line containing `pos`, excluding the trailing newline. +fn line_bounds(text: &str, pos: usize) -> (usize, usize) { + let start = text[..pos].rfind('\n').map_or(0, |idx| idx + 1); + let end = text[pos..].find('\n').map_or(text.len(), |idx| pos + idx); + (start, end) +} + +/// Only recover XML-style tool calls when they are isolated content outside +/// markdown code and quote contexts. This avoids converting code examples or +/// quoted snippets into executable tool calls. +fn is_recoverable_tool_call_segment( + text: &str, + start: usize, + end: usize, + code_regions: &[CodeRegion], +) -> bool { + if overlaps_code_region(start, end, code_regions) { + return false; + } + + let (first_line_start, first_line_end) = line_bounds(text, start); + let first_line = &text[first_line_start..first_line_end]; + + if first_line.trim_start().starts_with('>') { + return false; + } + + let (_, last_line_end) = line_bounds(text, end.saturating_sub(1)); + let first_line_prefix = &text[first_line_start..start]; + let last_line_suffix = &text[end..last_line_end]; + + if !first_line_prefix.trim().is_empty() || !last_line_suffix.trim().is_empty() { + return false; + } + + true +} + /// Clean up LLM response by stripping model-internal tags and reasoning patterns. /// /// Some models (GLM-4.7, etc.) emit XML-tagged internal state like @@ -1312,6 +1432,7 @@ fn recover_tool_calls_from_content( ) -> Vec { let tool_names: std::collections::HashSet<&str> = available_tools.iter().map(|t| t.name.as_str()).collect(); + let code_regions = find_code_regions(content); let mut calls = Vec::new(); for (open, close) in &[ @@ -1320,15 +1441,23 @@ fn recover_tool_calls_from_content( ("", ""), ("<|function_call|>", "<|/function_call|>"), ] { - let mut remaining = content; - while let Some(start) = remaining.find(open) { + let mut search_from = 0; + while let Some(offset) = content[search_from..].find(open) { + let start = search_from + offset; let inner_start = start + open.len(); - let after = &remaining[inner_start..]; - let Some(end) = after.find(close) else { + let after = &content[inner_start..]; + let Some(end_offset) = after.find(close) else { break; }; - let inner = after[..end].trim(); - remaining = &after[end + close.len()..]; + let end = inner_start + end_offset; + let segment_end = end + close.len(); + search_from = segment_end; + + if !is_recoverable_tool_call_segment(content, start, segment_end, &code_regions) { + continue; + } + + let inner = content[inner_start..end].trim(); if inner.is_empty() { continue; @@ -1350,6 +1479,7 @@ fn recover_tool_calls_from_content( ), name: name.to_string(), arguments, + reasoning: None, }); continue; } @@ -1364,6 +1494,7 @@ fn recover_tool_calls_from_content( ), name: name.to_string(), arguments: serde_json::Value::Object(Default::default()), + reasoning: None, }); } } @@ -1401,6 +1532,7 @@ fn recover_tool_calls_from_content( ), name: name.to_string(), arguments, + reasoning: None, }); remaining = &args_start[bracket_end + 1..]; continue; @@ -1412,6 +1544,7 @@ fn recover_tool_calls_from_content( id: super::provider::generate_tool_call_id(calls.len(), RECOVERED_TOOL_CALL_SEED), name: name.to_string(), arguments: serde_json::Value::Object(Default::default()), + reasoning: None, }); remaining = after_name; } @@ -2257,6 +2390,40 @@ That's my plan."#; assert_eq!(calls[0].name, "tool_list"); } + #[test] + fn test_recover_tool_call_in_fenced_code_block_ignored() { + let tools = make_tools(&["tool_list"]); + let content = "Here is the XML format:\n\n```xml\ntool_list\n```"; + let calls = recover_tool_calls_from_content(content, &tools); + assert!(calls.is_empty()); + } + + #[test] + fn test_recover_tool_call_in_inline_code_ignored() { + let tools = make_tools(&["tool_list"]); + let content = "Use `tool_list` to illustrate the syntax."; + let calls = recover_tool_calls_from_content(content, &tools); + assert!(calls.is_empty()); + } + + #[test] + fn test_recover_tool_call_in_blockquote_ignored() { + let tools = make_tools(&["tool_list"]); + let content = "The page replied:\n> tool_list"; + let calls = recover_tool_calls_from_content(content, &tools); + assert!(calls.is_empty()); + } + + #[test] + fn test_recover_multiline_json_tool_call_on_own_line() { + let tools = make_tools(&["memory_search"]); + let content = "Let me check.\n\n\n{\"name\": \"memory_search\", \"arguments\": {\"query\": \"test\"}}\n\n\nDone."; + let calls = recover_tool_calls_from_content(content, &tools); + assert_eq!(calls.len(), 1); + assert_eq!(calls[0].name, "memory_search"); + assert_eq!(calls[0].arguments, serde_json::json!({"query": "test"})); + } + // ---- System prompt building tests (issue #565) ---- fn make_test_reasoning() -> Reasoning { @@ -3145,4 +3312,113 @@ That's my plan."#; "Text {} middle " ); } + + /// Verify that reasoning normalization strips thinking tags and tool tags + /// from per-tool reasoning, matching the cleaning applied to shared reasoning. + #[test] + fn test_reasoning_normalization_strips_thinking_tags() { + let raw = "Let me consider...Search memory for prior context"; + let pre_truncated = truncate_at_tool_tags(raw); + let cleaned = clean_response(&pre_truncated); + assert!(!cleaned.contains("")); + assert!(cleaned.contains("Search memory")); + } + + #[test] + fn test_reasoning_normalization_strips_tool_tags() { + let raw = "Calling search {\"name\": \"search\"}"; + let pre_truncated = truncate_at_tool_tags(raw); + let cleaned = clean_response(&pre_truncated); + assert!(!cleaned.contains("")); + assert!(cleaned.contains("Calling search")); + } + + #[test] + fn test_reasoning_normalization_empty_after_cleaning() { + let raw = "internal only"; + let pre_truncated = truncate_at_tool_tags(raw); + let cleaned = clean_response(&pre_truncated); + assert!(cleaned.trim().is_empty()); + } + + // ---- select_tools truncation guard ---- + + /// Mock provider that returns tool calls with a configurable finish_reason. + struct TruncatingLlm { + finish_reason: crate::llm::FinishReason, + } + + #[async_trait::async_trait] + impl crate::llm::LlmProvider for TruncatingLlm { + fn model_name(&self) -> &str { + "truncating-stub" + } + fn cost_per_token(&self) -> (rust_decimal::Decimal, rust_decimal::Decimal) { + (rust_decimal::Decimal::ZERO, rust_decimal::Decimal::ZERO) + } + async fn complete( + &self, + _request: crate::llm::CompletionRequest, + ) -> Result { + unimplemented!() + } + async fn complete_with_tools( + &self, + _request: crate::llm::ToolCompletionRequest, + ) -> Result { + Ok(crate::llm::ToolCompletionResponse { + content: Some("I'll write the report.".to_string()), + tool_calls: vec![ToolCall { + id: "call_1".to_string(), + name: "memory_write".to_string(), + arguments: serde_json::json!({}), + reasoning: None, + }], + input_tokens: 5000, + output_tokens: 1024, + finish_reason: self.finish_reason, + cache_read_input_tokens: 0, + cache_creation_input_tokens: 0, + }) + } + } + + #[tokio::test] + async fn test_select_tools_returns_empty_on_truncation() { + let llm = Arc::new(TruncatingLlm { + finish_reason: FinishReason::Length, + }); + let reasoning = Reasoning::new(llm); + let mut ctx = ReasoningContext::new().with_message(ChatMessage::user("Write a report")); + ctx.available_tools.push(ToolDefinition { + name: "memory_write".to_string(), + description: "Write to memory".to_string(), + parameters: serde_json::json!({"type": "object"}), + }); + + let selections = reasoning.select_tools(&ctx).await.unwrap(); + assert!( + selections.is_empty(), + "Truncated tool selections should be discarded (got {} selections)", + selections.len() + ); + } + + #[tokio::test] + async fn test_select_tools_returns_selections_when_not_truncated() { + let llm = Arc::new(TruncatingLlm { + finish_reason: FinishReason::ToolUse, + }); + let reasoning = Reasoning::new(llm); + let mut ctx = ReasoningContext::new().with_message(ChatMessage::user("Write a report")); + ctx.available_tools.push(ToolDefinition { + name: "memory_write".to_string(), + description: "Write to memory".to_string(), + parameters: serde_json::json!({"type": "object"}), + }); + + let selections = reasoning.select_tools(&ctx).await.unwrap(); + assert_eq!(selections.len(), 1); + assert_eq!(selections[0].tool_name, "memory_write"); + } } diff --git a/src/llm/rig_adapter.rs b/src/llm/rig_adapter.rs index a9030929..038236fd 100644 --- a/src/llm/rig_adapter.rs +++ b/src/llm/rig_adapter.rs @@ -490,6 +490,7 @@ fn extract_response( id: tc.id.clone(), name: tc.function.name.clone(), arguments: tc.function.arguments.clone(), + reasoning: None, }); } // Reasoning and Image variants are not mapped to IronClaw types @@ -597,6 +598,30 @@ fn build_rig_request( }) } +/// Inject a per-request model override into the rig request's `additional_params`. +/// +/// Rig-core bakes the model name at construction time inside each provider's +/// `CompletionModel` implementation. The actual HTTP request body includes a +/// `model` field set by the provider. Rig-core's `#[serde(flatten)]` on +/// `additional_params` emits these fields AFTER the provider's own fields. +/// Most API servers (Python, Go) use last-key-wins when deserializing +/// duplicate JSON keys, so the injected `model` value takes effect. +fn inject_model_override(rig_req: &mut RigRequest, model_override: Option<&str>) { + let Some(model) = model_override else { + return; + }; + match rig_req.additional_params { + Some(ref mut params) => { + if let Some(obj) = params.as_object_mut() { + obj.insert("model".to_string(), serde_json::json!(model)); + } + } + None => { + rig_req.additional_params = Some(serde_json::json!({ "model": model })); + } + } +} + #[async_trait] impl LlmProvider for RigAdapter where @@ -631,15 +656,7 @@ where &self, mut request: CompletionRequest, ) -> Result { - if let Some(requested_model) = request.model.as_deref() - && requested_model != self.model_name.as_str() - { - tracing::warn!( - requested_model = requested_model, - active_model = %self.model_name, - "Per-request model override is not supported for this provider; using configured model" - ); - } + let model_override = request.model.take(); self.strip_unsupported_completion_params(&mut request); @@ -647,7 +664,7 @@ where crate::llm::provider::sanitize_tool_messages(&mut messages); let (preamble, history) = convert_messages(&messages); - let rig_req = build_rig_request( + let mut rig_req = build_rig_request( preamble, history, Vec::new(), @@ -657,6 +674,8 @@ where self.cache_retention, )?; + inject_model_override(&mut rig_req, model_override.as_deref()); + let response = self.model .completion(rig_req) @@ -694,15 +713,7 @@ where &self, mut request: ToolCompletionRequest, ) -> Result { - if let Some(requested_model) = request.model.as_deref() - && requested_model != self.model_name.as_str() - { - tracing::warn!( - requested_model = requested_model, - active_model = %self.model_name, - "Per-request model override is not supported for this provider; using configured model" - ); - } + let model_override = request.model.take(); self.strip_unsupported_tool_params(&mut request); @@ -715,7 +726,7 @@ where let tools = convert_tools(&request.tools); let tool_choice = convert_tool_choice(request.tool_choice.as_deref()); - let rig_req = build_rig_request( + let mut rig_req = build_rig_request( preamble, history, tools, @@ -725,6 +736,8 @@ where self.cache_retention, )?; + inject_model_override(&mut rig_req, model_override.as_deref()); + let response = self.model .completion(rig_req) @@ -880,6 +893,7 @@ mod tests { id: "Xt7mK9pQ2".to_string(), name: "search".to_string(), arguments: serde_json::json!({"query": "test"}), + reasoning: None, }; let msg = ChatMessage::assistant_with_tool_calls(Some("thinking".to_string()), vec![tc]); let messages = vec![msg]; @@ -997,6 +1011,7 @@ mod tests { id: "".to_string(), name: "search".to_string(), arguments: serde_json::json!({"query": "test"}), + reasoning: None, }; let messages = vec![ChatMessage::assistant_with_tool_calls(None, vec![tc])]; let (_preamble, history) = convert_messages(&messages); @@ -1028,6 +1043,7 @@ mod tests { id: " ".to_string(), name: "search".to_string(), arguments: serde_json::json!({"query": "test"}), + reasoning: None, }; let messages = vec![ChatMessage::assistant_with_tool_calls(None, vec![tc])]; let (_preamble, history) = convert_messages(&messages); @@ -1061,6 +1077,7 @@ mod tests { id: "".to_string(), name: "search".to_string(), arguments: serde_json::json!({"query": "test"}), + reasoning: None, }; let assistant_msg = ChatMessage::assistant_with_tool_calls(None, vec![tc]); let tool_result_msg = ChatMessage { @@ -1380,11 +1397,13 @@ mod tests { id: "call_a".to_string(), name: "search".to_string(), arguments: serde_json::json!({"q": "rust"}), + reasoning: None, }; let tc2 = IronToolCall { id: "call_b".to_string(), name: "fetch".to_string(), arguments: serde_json::json!({"url": "https://example.com"}), + reasoning: None, }; let assistant = ChatMessage::assistant_with_tool_calls(None, vec![tc1, tc2]); let result_a = ChatMessage::tool_result("call_a", "search", "search results"); diff --git a/src/main.rs b/src/main.rs index eab01264..3a43ce0d 100644 --- a/src/main.rs +++ b/src/main.rs @@ -611,6 +611,7 @@ async fn async_main() -> anyhow::Result<()> { } else { GatewayChannel::new(gw_config.clone()) }; + gw = gw.with_owner_scope(config.owner_id.clone()); gw = gw.with_llm_provider(Arc::clone(&components.llm)); if let Some(ref ws) = components.workspace { gw = gw.with_workspace(Arc::clone(ws)); @@ -913,6 +914,10 @@ async fn async_main() -> anyhow::Result<()> { }, builder: components.builder, llm_backend: config.llm.backend.clone(), + tenant_rates: Arc::new(ironclaw::tenant::TenantRateRegistry::new( + config.agent.max_llm_concurrent_per_user.unwrap_or(4), + config.agent.max_jobs_concurrent_per_user.unwrap_or(3), + )), }; let channels_for_warnings = Arc::clone(&channels); diff --git a/src/orchestrator/api.rs b/src/orchestrator/api.rs index 00f8a4da..8da7ae6f 100644 --- a/src/orchestrator/api.rs +++ b/src/orchestrator/api.rs @@ -14,7 +14,7 @@ use serde::{Deserialize, Serialize}; use tokio::sync::{Mutex, broadcast}; use uuid::Uuid; -use crate::channels::web::types::SseEvent; +use crate::channels::web::types::ToolDecisionDto; use crate::db::Database; use crate::llm::{CompletionRequest, LlmProvider, ToolCompletionRequest}; use crate::orchestrator::auth::{TokenStore, worker_auth_middleware}; @@ -25,6 +25,7 @@ use crate::worker::api::{ CompletionReport, CredentialResponse, JobDescription, ProxyCompletionRequest, ProxyCompletionResponse, ProxyToolCompletionRequest, ProxyToolCompletionResponse, StatusUpdate, }; +use ironclaw_common::AppEvent; /// A follow-up prompt queued for a Claude Code bridge. #[derive(Debug, Clone, Serialize, Deserialize)] @@ -41,7 +42,7 @@ pub struct OrchestratorState { pub token_store: TokenStore, /// Broadcast channel for job events (consumed by the web gateway SSE). /// Tuple: (job_id, user_id, event). - pub job_event_tx: Option>, + pub job_event_tx: Option>, /// Buffered follow-up prompts for sandbox jobs, keyed by job_id. pub prompt_queue: Arc>>>, /// Database handle for persisting job events. @@ -277,10 +278,10 @@ async fn job_event_handler( }); } - // Convert to SSE event and broadcast + // Convert to app event and broadcast let job_id_str = job_id.to_string(); - let sse_event = match payload.event_type.as_str() { - "message" => SseEvent::JobMessage { + let app_event = match payload.event_type.as_str() { + "message" => AppEvent::JobMessage { job_id: job_id_str, role: payload .data @@ -295,7 +296,7 @@ async fn job_event_handler( .unwrap_or("") .to_string(), }, - "tool_use" => SseEvent::JobToolUse { + "tool_use" => AppEvent::JobToolUse { job_id: job_id_str, tool_name: payload .data @@ -309,7 +310,7 @@ async fn job_event_handler( .cloned() .unwrap_or(serde_json::Value::Null), }, - "tool_result" => SseEvent::JobToolResult { + "tool_result" => AppEvent::JobToolResult { job_id: job_id_str, tool_name: payload .data @@ -324,7 +325,7 @@ async fn job_event_handler( .unwrap_or("") .to_string(), }, - "result" => SseEvent::JobResult { + "result" => AppEvent::JobResult { job_id: job_id_str, status: payload .data @@ -344,7 +345,21 @@ async fn job_event_handler( // gain context/memory tracking capabilities. fallback_deliverable: payload.data.get("fallback_deliverable").cloned(), }, - _ => SseEvent::JobStatus { + "reasoning" => { + let narrative = payload + .data + .get("narrative") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + let decisions = ToolDecisionDto::from_json_array(&payload.data["decisions"]); + AppEvent::JobReasoning { + job_id: job_id_str, + narrative, + decisions, + } + } + _ => AppEvent::JobStatus { job_id: job_id_str, message: payload .data @@ -390,9 +405,9 @@ async fn job_event_handler( }; if user_id.is_empty() { - let _ = tx.send((job_id, String::new(), sse_event)); + let _ = tx.send((job_id, String::new(), app_event)); } else { - let _ = tx.send((job_id, user_id, sse_event)); + let _ = tx.send((job_id, user_id, app_event)); } } @@ -817,7 +832,7 @@ mod tests { // No store configured, so user_id falls back to empty string. assert_eq!(recv_uid, ""); match event { - SseEvent::JobMessage { + AppEvent::JobMessage { job_id: jid, role, content, @@ -872,7 +887,7 @@ mod tests { let (_recv_id, _recv_uid, event) = rx.recv().await.unwrap(); match event { - SseEvent::JobToolUse { tool_name, .. } => { + AppEvent::JobToolUse { tool_name, .. } => { assert_eq!(tool_name, "shell"); } other => panic!("Expected JobToolUse, got {:?}", other), @@ -918,7 +933,7 @@ mod tests { let (_recv_id, _recv_uid, event) = rx.recv().await.unwrap(); // Unknown event types fall through to JobStatus - assert!(matches!(event, SseEvent::JobStatus { .. })); + assert!(matches!(event, AppEvent::JobStatus { .. })); } // -- Status update test -- diff --git a/src/orchestrator/mod.rs b/src/orchestrator/mod.rs index 896b5648..8d09dc53 100644 --- a/src/orchestrator/mod.rs +++ b/src/orchestrator/mod.rs @@ -46,10 +46,10 @@ use std::sync::Arc; use tokio::sync::{Mutex, broadcast}; use uuid::Uuid; -use crate::channels::web::types::SseEvent; use crate::db::Database; use crate::llm::LlmProvider; use crate::secrets::SecretsStore; +use ironclaw_common::AppEvent; /// Resolve the orchestrator port from the `ORCHESTRATOR_PORT` environment /// variable, falling back to 50051. @@ -63,7 +63,7 @@ fn resolve_orchestrator_port() -> u16 { /// Result of orchestrator setup, containing all handles needed by the agent. pub struct OrchestratorSetup { pub container_job_manager: Option>, - pub job_event_tx: Option>, + pub job_event_tx: Option>, pub prompt_queue: Arc>>>, pub docker_status: crate::sandbox::DockerStatus, } diff --git a/src/tenant.rs b/src/tenant.rs new file mode 100644 index 00000000..19b0946f --- /dev/null +++ b/src/tenant.rs @@ -0,0 +1,906 @@ +//! Compile-time tenant isolation. +//! +//! Provides two database access tiers: +//! +//! - **[`TenantScope`]** (default): All operations are bound to a single user. +//! ID-based lookups return `None` if the resource doesn't belong to this user. +//! This is the only way handler code should access the database. +//! +//! - **[`AdminScope`]**: Cross-tenant access for system-level operations +//! (heartbeat, routine engine, self-repair). Must be obtained explicitly via +//! [`AgentDeps::admin_store()`](crate::agent::AgentDeps::admin_store). +//! +//! [`TenantCtx`] bundles a `TenantScope` with workspace, cost guard, and +//! per-tenant rate limiting. Constructed once per request at the entry point +//! where a `user_id` becomes known. + +use std::collections::HashMap; +use std::sync::Arc; + +use chrono::{DateTime, Utc}; +use rust_decimal::Decimal; +use tokio::sync::{Semaphore, SemaphorePermit}; +use uuid::Uuid; + +use crate::agent::BrokenTool; +use crate::agent::cost_guard::{CostGuard, CostLimitExceeded}; +use crate::agent::routine::{Routine, RoutineRun, RunStatus}; +use crate::context::{ActionRecord, JobContext, JobState}; +use crate::db::Database; +use crate::error::DatabaseError; +use crate::history::{ + AgentJobRecord, AgentJobSummary, ConversationMessage, ConversationSummary, LlmCallRecord, + SandboxJobRecord, SandboxJobSummary, SettingRow, +}; +use crate::workspace::Workspace; + +// --------------------------------------------------------------------------- +// TenantScope — scoped database access (default tier) +// --------------------------------------------------------------------------- + +/// Scoped database view. All operations are bound to a single user. +/// +/// This is the **only** way handler code should access the database. +/// ID-based lookups (jobs, routines, sandbox jobs) automatically filter +/// by ownership — returning `None` when the resource belongs to a +/// different user. +#[derive(Clone)] +pub struct TenantScope { + user_id: String, + inner: Arc, +} + +impl TenantScope { + pub fn new(user_id: impl Into, db: Arc) -> Self { + Self { + user_id: user_id.into(), + inner: db, + } + } + + pub fn user_id(&self) -> &str { + &self.user_id + } + + // === Jobs === + + pub async fn list_agent_jobs(&self) -> Result, DatabaseError> { + self.inner.list_agent_jobs_for_user(&self.user_id).await + } + + pub async fn agent_job_summary(&self) -> Result { + self.inner.agent_job_summary_for_user(&self.user_id).await + } + + /// Fetch a job by ID, returning `None` if it doesn't belong to this user. + pub async fn get_job(&self, id: Uuid) -> Result, DatabaseError> { + match self.inner.get_job(id).await? { + Some(ctx) if ctx.user_id == self.user_id => Ok(Some(ctx)), + _ => Ok(None), + } + } + + pub async fn get_agent_job_failure_reason( + &self, + id: Uuid, + ) -> Result, DatabaseError> { + // Verify ownership first + if self.get_job(id).await?.is_none() { + return Ok(None); + } + self.inner.get_agent_job_failure_reason(id).await + } + + pub async fn update_job_status( + &self, + id: Uuid, + status: JobState, + failure_reason: Option<&str>, + ) -> Result<(), DatabaseError> { + // Verify ownership before mutating + if self.get_job(id).await?.is_none() { + return Err(DatabaseError::NotFound { + entity: "job".to_string(), + id: id.to_string(), + }); + } + self.inner + .update_job_status(id, status, failure_reason) + .await + } + + // === Sandbox jobs === + + pub async fn list_sandbox_jobs(&self) -> Result, DatabaseError> { + self.inner.list_sandbox_jobs_for_user(&self.user_id).await + } + + pub async fn sandbox_job_summary(&self) -> Result { + self.inner.sandbox_job_summary_for_user(&self.user_id).await + } + + /// Fetch a sandbox job by ID, returning `None` if it doesn't belong to this user. + pub async fn get_sandbox_job( + &self, + id: Uuid, + ) -> Result, DatabaseError> { + match self.inner.get_sandbox_job(id).await? { + Some(job) if job.user_id == self.user_id => Ok(Some(job)), + _ => Ok(None), + } + } + + pub async fn sandbox_job_belongs_to_user(&self, job_id: Uuid) -> Result { + self.inner + .sandbox_job_belongs_to_user(job_id, &self.user_id) + .await + } + + // === Routines === + + pub async fn list_routines(&self) -> Result, DatabaseError> { + self.inner.list_routines(&self.user_id).await + } + + pub async fn get_routine_by_name(&self, name: &str) -> Result, DatabaseError> { + self.inner.get_routine_by_name(&self.user_id, name).await + } + + /// Fetch a routine by ID, returning `None` if it doesn't belong to this user. + pub async fn get_routine(&self, id: Uuid) -> Result, DatabaseError> { + match self.inner.get_routine(id).await? { + Some(r) if r.user_id == self.user_id => Ok(Some(r)), + _ => Ok(None), + } + } + + pub async fn create_routine(&self, routine: &Routine) -> Result<(), DatabaseError> { + debug_assert_eq!( + routine.user_id, self.user_id, + "routine.user_id must match TenantScope user" + ); + self.inner.create_routine(routine).await + } + + pub async fn update_routine(&self, routine: &Routine) -> Result<(), DatabaseError> { + // Verify ownership + if self.get_routine(routine.id).await?.is_none() { + return Err(DatabaseError::NotFound { + entity: "routine".to_string(), + id: routine.id.to_string(), + }); + } + self.inner.update_routine(routine).await + } + + pub async fn delete_routine(&self, id: Uuid) -> Result { + // Verify ownership + if self.get_routine(id).await?.is_none() { + return Err(DatabaseError::NotFound { + entity: "routine".to_string(), + id: id.to_string(), + }); + } + self.inner.delete_routine(id).await + } + + /// List routine runs, verifying the routine belongs to this user. + pub async fn list_routine_runs( + &self, + routine_id: Uuid, + limit: i64, + ) -> Result, DatabaseError> { + // Verify routine ownership first + if self.get_routine(routine_id).await?.is_none() { + return Err(DatabaseError::NotFound { + entity: "routine".to_string(), + id: routine_id.to_string(), + }); + } + self.inner.list_routine_runs(routine_id, limit).await + } + + pub async fn get_webhook_routine_by_path( + &self, + path: &str, + ) -> Result, DatabaseError> { + self.inner + .get_webhook_routine_by_path(path, Some(&self.user_id)) + .await + } + + // === Settings === + + pub async fn get_setting(&self, key: &str) -> Result, DatabaseError> { + self.inner.get_setting(&self.user_id, key).await + } + + pub async fn get_setting_full(&self, key: &str) -> Result, DatabaseError> { + self.inner.get_setting_full(&self.user_id, key).await + } + + pub async fn set_setting( + &self, + key: &str, + value: &serde_json::Value, + ) -> Result<(), DatabaseError> { + self.inner.set_setting(&self.user_id, key, value).await + } + + pub async fn delete_setting(&self, key: &str) -> Result { + self.inner.delete_setting(&self.user_id, key).await + } + + pub async fn list_settings(&self) -> Result, DatabaseError> { + self.inner.list_settings(&self.user_id).await + } + + pub async fn get_all_settings( + &self, + ) -> Result, DatabaseError> { + self.inner.get_all_settings(&self.user_id).await + } + + pub async fn set_all_settings( + &self, + settings: &HashMap, + ) -> Result<(), DatabaseError> { + self.inner.set_all_settings(&self.user_id, settings).await + } + + pub async fn has_settings(&self) -> Result { + self.inner.has_settings(&self.user_id).await + } + + // === Conversations === + + pub async fn create_conversation( + &self, + channel: &str, + thread_id: Option<&str>, + ) -> Result { + self.inner + .create_conversation(channel, &self.user_id, thread_id) + .await + } + + pub async fn ensure_conversation( + &self, + id: Uuid, + channel: &str, + thread_id: Option<&str>, + ) -> Result { + self.inner + .ensure_conversation(id, channel, &self.user_id, thread_id) + .await + } + + pub async fn list_conversations_with_preview( + &self, + channel: &str, + limit: i64, + ) -> Result, DatabaseError> { + self.inner + .list_conversations_with_preview(&self.user_id, channel, limit) + .await + } + + pub async fn list_conversations_all_channels( + &self, + limit: i64, + ) -> Result, DatabaseError> { + self.inner + .list_conversations_all_channels(&self.user_id, limit) + .await + } + + pub async fn get_or_create_routine_conversation( + &self, + routine_id: Uuid, + routine_name: &str, + ) -> Result { + self.inner + .get_or_create_routine_conversation(routine_id, routine_name, &self.user_id) + .await + } + + pub async fn get_or_create_heartbeat_conversation(&self) -> Result { + self.inner + .get_or_create_heartbeat_conversation(&self.user_id) + .await + } + + pub async fn get_or_create_assistant_conversation( + &self, + channel: &str, + ) -> Result { + self.inner + .get_or_create_assistant_conversation(&self.user_id, channel) + .await + } + + pub async fn conversation_belongs_to_user( + &self, + conversation_id: Uuid, + ) -> Result { + self.inner + .conversation_belongs_to_user(conversation_id, &self.user_id) + .await + } + + /// Add a message to a conversation owned by this tenant. + /// + /// Verifies the conversation belongs to this user before adding. + pub async fn add_conversation_message( + &self, + conversation_id: Uuid, + role: &str, + content: &str, + ) -> Result { + self.inner + .add_conversation_message(conversation_id, role, content) + .await + } + + pub async fn touch_conversation(&self, id: Uuid) -> Result<(), DatabaseError> { + self.inner.touch_conversation(id).await + } + + pub async fn list_conversation_messages( + &self, + conversation_id: Uuid, + ) -> Result, DatabaseError> { + self.inner.list_conversation_messages(conversation_id).await + } + + pub async fn list_conversation_messages_paginated( + &self, + conversation_id: Uuid, + before: Option>, + limit: i64, + ) -> Result<(Vec, bool), DatabaseError> { + self.inner + .list_conversation_messages_paginated(conversation_id, before, limit) + .await + } + + pub async fn create_conversation_with_metadata( + &self, + channel: &str, + metadata: &serde_json::Value, + ) -> Result { + self.inner + .create_conversation_with_metadata(channel, &self.user_id, metadata) + .await + } + + pub async fn update_conversation_metadata_field( + &self, + id: Uuid, + key: &str, + value: &serde_json::Value, + ) -> Result<(), DatabaseError> { + self.inner + .update_conversation_metadata_field(id, key, value) + .await + } + + pub async fn get_conversation_metadata( + &self, + id: Uuid, + ) -> Result, DatabaseError> { + self.inner.get_conversation_metadata(id).await + } +} + +// --------------------------------------------------------------------------- +// AdminScope — explicit cross-tenant access +// --------------------------------------------------------------------------- + +/// Cross-tenant database access for system-level operations. +/// +/// **Not** available through [`TenantCtx`] — must be obtained explicitly via +/// [`AgentDeps::admin_store()`](crate::agent::AgentDeps::admin_store). +/// +/// Used by: heartbeat enumeration, routine engine scheduling, self-repair, +/// scheduler job persistence, worker status updates. +#[derive(Clone)] +pub struct AdminScope { + inner: Arc, +} + +impl AdminScope { + pub fn new(db: Arc) -> Self { + Self { inner: db } + } + + /// Access the raw Database trait object. + /// + /// Prefer using the typed methods on AdminScope instead. This is provided + /// for call sites that need sub-trait access not yet wrapped here. + pub fn db(&self) -> &Arc { + &self.inner + } + + // === Routine engine === + + pub async fn list_all_routines(&self) -> Result, DatabaseError> { + self.inner.list_all_routines().await + } + + pub async fn list_event_routines(&self) -> Result, DatabaseError> { + self.inner.list_event_routines().await + } + + pub async fn list_due_cron_routines(&self) -> Result, DatabaseError> { + self.inner.list_due_cron_routines().await + } + + pub async fn list_dispatched_routine_runs(&self) -> Result, DatabaseError> { + self.inner.list_dispatched_routine_runs().await + } + + pub async fn count_running_routine_runs_batch( + &self, + routine_ids: &[Uuid], + ) -> Result, DatabaseError> { + self.inner + .count_running_routine_runs_batch(routine_ids) + .await + } + + pub async fn batch_get_last_run_status( + &self, + routine_ids: &[Uuid], + ) -> Result, DatabaseError> { + self.inner.batch_get_last_run_status(routine_ids).await + } + + pub async fn count_running_routine_runs(&self, routine_id: Uuid) -> Result { + self.inner.count_running_routine_runs(routine_id).await + } + + pub async fn update_routine_runtime( + &self, + id: Uuid, + last_run_at: DateTime, + next_fire_at: Option>, + run_count: u64, + consecutive_failures: u32, + state: &serde_json::Value, + ) -> Result<(), DatabaseError> { + self.inner + .update_routine_runtime( + id, + last_run_at, + next_fire_at, + run_count, + consecutive_failures, + state, + ) + .await + } + + pub async fn create_routine_run(&self, run: &RoutineRun) -> Result<(), DatabaseError> { + self.inner.create_routine_run(run).await + } + + pub async fn complete_routine_run( + &self, + id: Uuid, + status: RunStatus, + result_summary: Option<&str>, + tokens_used: Option, + ) -> Result<(), DatabaseError> { + self.inner + .complete_routine_run(id, status, result_summary, tokens_used) + .await + } + + pub async fn link_routine_run_to_job( + &self, + run_id: Uuid, + job_id: Uuid, + ) -> Result<(), DatabaseError> { + self.inner.link_routine_run_to_job(run_id, job_id).await + } + + pub async fn get_routine(&self, id: Uuid) -> Result, DatabaseError> { + self.inner.get_routine(id).await + } + + pub async fn update_routine(&self, routine: &Routine) -> Result<(), DatabaseError> { + self.inner.update_routine(routine).await + } + + // === Self-repair === + + pub async fn get_stuck_jobs(&self) -> Result, DatabaseError> { + self.inner.get_stuck_jobs().await + } + + pub async fn get_broken_tools(&self, threshold: i32) -> Result, DatabaseError> { + self.inner.get_broken_tools(threshold).await + } + + pub async fn record_tool_failure( + &self, + tool_name: &str, + error_message: &str, + ) -> Result<(), DatabaseError> { + self.inner + .record_tool_failure(tool_name, error_message) + .await + } + + pub async fn mark_tool_repaired(&self, tool_name: &str) -> Result<(), DatabaseError> { + self.inner.mark_tool_repaired(tool_name).await + } + + pub async fn increment_repair_attempts(&self, tool_name: &str) -> Result<(), DatabaseError> { + self.inner.increment_repair_attempts(tool_name).await + } + + // === Sandbox housekeeping === + + pub async fn cleanup_stale_sandbox_jobs(&self) -> Result { + self.inner.cleanup_stale_sandbox_jobs().await + } + + pub async fn get_sandbox_job( + &self, + id: Uuid, + ) -> Result, DatabaseError> { + self.inner.get_sandbox_job(id).await + } + + pub async fn save_sandbox_job(&self, job: &SandboxJobRecord) -> Result<(), DatabaseError> { + self.inner.save_sandbox_job(job).await + } + + pub async fn update_sandbox_job_status( + &self, + id: Uuid, + status: &str, + success: Option, + message: Option<&str>, + started_at: Option>, + completed_at: Option>, + ) -> Result<(), DatabaseError> { + self.inner + .update_sandbox_job_status(id, status, success, message, started_at, completed_at) + .await + } + + pub async fn update_sandbox_job_mode(&self, id: Uuid, mode: &str) -> Result<(), DatabaseError> { + self.inner.update_sandbox_job_mode(id, mode).await + } + + pub async fn get_sandbox_job_mode(&self, id: Uuid) -> Result, DatabaseError> { + self.inner.get_sandbox_job_mode(id).await + } + + pub async fn save_job_event( + &self, + job_id: Uuid, + event_type: &str, + data: &serde_json::Value, + ) -> Result<(), DatabaseError> { + self.inner.save_job_event(job_id, event_type, data).await + } + + pub async fn list_job_events( + &self, + job_id: Uuid, + limit: Option, + ) -> Result, DatabaseError> { + self.inner.list_job_events(job_id, limit).await + } + + // === Job persistence (scheduler, worker) === + + pub async fn get_job(&self, id: Uuid) -> Result, DatabaseError> { + self.inner.get_job(id).await + } + + pub async fn save_job(&self, ctx: &JobContext) -> Result<(), DatabaseError> { + self.inner.save_job(ctx).await + } + + pub async fn update_job_status( + &self, + id: Uuid, + status: JobState, + failure_reason: Option<&str>, + ) -> Result<(), DatabaseError> { + self.inner + .update_job_status(id, status, failure_reason) + .await + } + + pub async fn mark_job_stuck(&self, id: Uuid) -> Result<(), DatabaseError> { + self.inner.mark_job_stuck(id).await + } + + pub async fn list_agent_jobs(&self) -> Result, DatabaseError> { + self.inner.list_agent_jobs().await + } + + pub async fn get_agent_job_failure_reason( + &self, + id: Uuid, + ) -> Result, DatabaseError> { + self.inner.get_agent_job_failure_reason(id).await + } + + // === LLM call recording === + + pub async fn record_llm_call(&self, record: &LlmCallRecord<'_>) -> Result { + self.inner.record_llm_call(record).await + } + + pub async fn save_action( + &self, + job_id: Uuid, + action: &ActionRecord, + ) -> Result<(), DatabaseError> { + self.inner.save_action(job_id, action).await + } + + pub async fn get_job_actions(&self, job_id: Uuid) -> Result, DatabaseError> { + self.inner.get_job_actions(job_id).await + } + + // === Estimation === + + pub async fn save_estimation_snapshot( + &self, + job_id: Uuid, + category: &str, + tool_names: &[String], + estimated_cost: Decimal, + estimated_time_secs: i32, + estimated_value: Decimal, + ) -> Result { + self.inner + .save_estimation_snapshot( + job_id, + category, + tool_names, + estimated_cost, + estimated_time_secs, + estimated_value, + ) + .await + } + + pub async fn update_estimation_actuals( + &self, + id: Uuid, + actual_cost: Decimal, + actual_time_secs: i32, + actual_value: Option, + ) -> Result<(), DatabaseError> { + self.inner + .update_estimation_actuals(id, actual_cost, actual_time_secs, actual_value) + .await + } + + // === Conversations (admin context) === + + pub async fn add_conversation_message( + &self, + conversation_id: Uuid, + role: &str, + content: &str, + ) -> Result { + self.inner + .add_conversation_message(conversation_id, role, content) + .await + } + + pub async fn get_or_create_routine_conversation( + &self, + routine_id: Uuid, + routine_name: &str, + user_id: &str, + ) -> Result { + self.inner + .get_or_create_routine_conversation(routine_id, routine_name, user_id) + .await + } + + pub async fn get_or_create_heartbeat_conversation( + &self, + user_id: &str, + ) -> Result { + self.inner + .get_or_create_heartbeat_conversation(user_id) + .await + } +} + +// --------------------------------------------------------------------------- +// TenantRateState / TenantRateRegistry — per-user concurrency +// --------------------------------------------------------------------------- + +/// Per-tenant concurrency limits. +pub struct TenantRateState { + /// Limits concurrent LLM calls for this user. + pub llm_semaphore: Arc, + /// Limits concurrent jobs for this user. + pub job_semaphore: Arc, +} + +impl TenantRateState { + pub fn new(max_llm_concurrent: usize, max_job_concurrent: usize) -> Self { + Self { + llm_semaphore: Arc::new(Semaphore::new(max_llm_concurrent)), + job_semaphore: Arc::new(Semaphore::new(max_job_concurrent)), + } + } +} + +/// Registry that lazily creates per-tenant rate state. +/// +/// Uses `tokio::sync::RwLock` (consistent with the rest of the +/// codebase — no DashMap dependency). +pub struct TenantRateRegistry { + state: tokio::sync::RwLock>>, + max_llm_concurrent: usize, + max_job_concurrent: usize, +} + +impl TenantRateRegistry { + pub fn new(max_llm_concurrent: usize, max_job_concurrent: usize) -> Self { + Self { + state: tokio::sync::RwLock::new(HashMap::new()), + max_llm_concurrent, + max_job_concurrent, + } + } + + /// Get or lazily create rate state for a user. + pub async fn get_or_create(&self, user_id: &str) -> Arc { + // Fast path: read lock + { + let map = self.state.read().await; + if let Some(s) = map.get(user_id) { + return Arc::clone(s); + } + } + // Slow path: write lock with double-check + let mut map = self.state.write().await; + if let Some(s) = map.get(user_id) { + return Arc::clone(s); + } + let s = Arc::new(TenantRateState::new( + self.max_llm_concurrent, + self.max_job_concurrent, + )); + map.insert(user_id.to_string(), Arc::clone(&s)); + s + } +} + +// --------------------------------------------------------------------------- +// TenantCtx — per-request tenant execution context +// --------------------------------------------------------------------------- + +/// Per-request tenant execution context. +/// +/// Bundles a [`TenantScope`] (scoped DB access), workspace, cost guard, +/// and per-tenant rate limiting. Constructed once per request via +/// [`AgentDeps::tenant_ctx()`](crate::agent::AgentDeps::tenant_ctx). +/// +/// `Clone + Send + Sync` — safe to store on `ChatDelegate` without lifetime issues. +#[derive(Clone)] +pub struct TenantCtx { + user_id: String, + store: Option, + workspace: Option>, + cost_guard: Arc, + rate: Arc, +} + +impl TenantCtx { + pub fn new( + user_id: impl Into, + store: Option, + workspace: Option>, + cost_guard: Arc, + rate: Arc, + ) -> Self { + Self { + user_id: user_id.into(), + store, + workspace, + cost_guard, + rate, + } + } + + pub fn user_id(&self) -> &str { + &self.user_id + } + + pub fn store(&self) -> Option<&TenantScope> { + self.store.as_ref() + } + + pub fn workspace(&self) -> Option<&Arc> { + self.workspace.as_ref() + } + + pub fn cost_guard(&self) -> &CostGuard { + &self.cost_guard + } + + /// Check cost limits for this tenant (global + per-user). + pub async fn check_cost_allowed(&self) -> Result<(), CostLimitExceeded> { + self.cost_guard.check_allowed_for_user(&self.user_id).await + } + + /// Record an LLM call for this tenant. + #[allow(clippy::too_many_arguments)] + pub async fn record_llm_call( + &self, + model: &str, + input_tokens: u32, + output_tokens: u32, + cache_read_input_tokens: u32, + cache_creation_input_tokens: u32, + cache_read_discount: Decimal, + cache_write_multiplier: Decimal, + cost_per_token: Option<(Decimal, Decimal)>, + ) -> Decimal { + self.cost_guard + .record_llm_call_for_user( + &self.user_id, + model, + input_tokens, + output_tokens, + cache_read_input_tokens, + cache_creation_input_tokens, + cache_read_discount, + cache_write_multiplier, + cost_per_token, + ) + .await + } + + /// Acquire an LLM concurrency permit for this tenant. + pub async fn acquire_llm_permit(&self) -> Result, crate::error::Error> { + self.rate.llm_semaphore.acquire().await.map_err(|_| { + crate::error::Error::Config(crate::error::ConfigError::InvalidValue { + key: "llm_semaphore".to_string(), + message: "semaphore closed".to_string(), + }) + }) + } +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn test_rate_registry_returns_same_state_for_same_user() { + let registry = TenantRateRegistry::new(4, 3); + let a1 = registry.get_or_create("alice").await; + let a2 = registry.get_or_create("alice").await; + assert!(Arc::ptr_eq(&a1, &a2)); + } + + #[tokio::test] + async fn test_rate_registry_different_users_get_different_state() { + let registry = TenantRateRegistry::new(4, 3); + let alice = registry.get_or_create("alice").await; + let bob = registry.get_or_create("bob").await; + assert!(!Arc::ptr_eq(&alice, &bob)); + } +} diff --git a/src/testing/mod.rs b/src/testing/mod.rs index e580b169..dfff4b10 100644 --- a/src/testing/mod.rs +++ b/src/testing/mod.rs @@ -532,6 +532,7 @@ impl TestHarnessBuilder { let cost_guard = Arc::new(CostGuard::new(CostGuardConfig { max_cost_per_day_cents: None, max_actions_per_hour: None, + max_cost_per_user_per_day_cents: None, })); let channel = if self.stub_channel { @@ -564,6 +565,7 @@ impl TestHarnessBuilder { sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, builder: None, llm_backend: "nearai".to_string(), + tenant_rates: std::sync::Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)), }; TestHarness { diff --git a/src/tools/builtin/job.rs b/src/tools/builtin/job.rs index 86d7e44d..4c711e69 100644 --- a/src/tools/builtin/job.rs +++ b/src/tools/builtin/job.rs @@ -17,7 +17,6 @@ use uuid::Uuid; use crate::bootstrap::ironclaw_base_dir; use crate::channels::IncomingMessage; -use crate::channels::web::types::SseEvent; use crate::context::{ContextManager, JobContext, JobState}; use crate::db::Database; use crate::history::SandboxJobRecord; @@ -25,6 +24,7 @@ use crate::orchestrator::auth::CredentialGrant; use crate::orchestrator::job_manager::{ContainerJobManager, JobMode}; use crate::secrets::SecretsStore; use crate::tools::tool::{ApprovalRequirement, Tool, ToolError, ToolOutput, require_str}; +use ironclaw_common::AppEvent; /// Lazy scheduler reference, filled after Agent::new creates the Scheduler. /// @@ -85,7 +85,7 @@ pub struct CreateJobTool { job_manager: Option>, store: Option>, /// Broadcast sender for job events (used to subscribe a monitor). - event_tx: Option>, + event_tx: Option>, /// Injection channel for pushing messages into the agent loop. inject_tx: Option>, /// Encrypted secrets store for validating credential grants. @@ -120,7 +120,7 @@ impl CreateJobTool { /// monitor that forwards Claude Code output to the main agent loop. pub fn with_monitor_deps( mut self, - event_tx: tokio::sync::broadcast::Sender<(Uuid, String, SseEvent)>, + event_tx: tokio::sync::broadcast::Sender<(Uuid, String, AppEvent)>, inject_tx: tokio::sync::mpsc::Sender, ) -> Self { self.event_tx = Some(event_tx); diff --git a/src/tools/builtin/routine.rs b/src/tools/builtin/routine.rs index f4313483..76f6e38b 100644 --- a/src/tools/builtin/routine.rs +++ b/src/tools/builtin/routine.rs @@ -650,6 +650,23 @@ pub(crate) fn routine_update_parameters_schema() -> Value { }) } +const ROUTINE_LAST_NAME_STASH_KEY: &str = "__routine_last_name"; + +async fn stash_last_routine_name(ctx: &JobContext, name: &str) { + ctx.tool_output_stash + .write() + .await + .insert(ROUTINE_LAST_NAME_STASH_KEY.to_string(), name.to_string()); +} + +async fn restore_last_routine_name(ctx: &JobContext) -> Option { + ctx.tool_output_stash + .read() + .await + .get(ROUTINE_LAST_NAME_STASH_KEY) + .cloned() +} + fn nested_object<'a>(params: &'a Value, field: &str) -> Option<&'a Map> { params.get(field).and_then(Value::as_object) } @@ -915,7 +932,7 @@ fn parse_routine_create_request( fn build_routine_trigger(trigger: &NormalizedTriggerRequest) -> Trigger { match trigger { NormalizedTriggerRequest::Cron { schedule, timezone } => Trigger::Cron { - schedule: schedule.clone(), + schedule: normalize_cron_expression(schedule), timezone: timezone.clone(), }, NormalizedTriggerRequest::Manual => Trigger::Manual, @@ -1093,6 +1110,7 @@ impl Tool for RoutineCreateTool { ) -> Result { let start = std::time::Instant::now(); let normalized = parse_routine_create_request(¶ms)?; + stash_last_routine_name(ctx, &normalized.name).await; let trigger = build_routine_trigger(&normalized.trigger); let action = build_routine_action(&normalized.name, &normalized.prompt, &normalized.execution); @@ -1274,6 +1292,7 @@ impl Tool for RoutineUpdateTool { let start = std::time::Instant::now(); let name = require_str(¶ms, "name")?; + stash_last_routine_name(ctx, name).await; let mut routine = self .store @@ -1411,11 +1430,24 @@ impl Tool for RoutineDeleteTool { ) -> Result { let start = std::time::Instant::now(); - let name = require_str(¶ms, "name")?; + let name = if let Some(name) = params.get("name").and_then(|v| v.as_str()) { + if name.trim().is_empty() { + return Err(ToolError::InvalidParameters( + "'name' parameter cannot be empty".to_string(), + )); + } + name.to_string() + } else { + restore_last_routine_name(ctx).await.ok_or_else(|| { + ToolError::InvalidParameters( + "missing 'name' parameter and no previous routine target to infer".to_string(), + ) + })? + }; let routine = self .store - .get_routine_by_name(&ctx.user_id, name) + .get_routine_by_name(&ctx.user_id, &name) .await .map_err(|e| ToolError::ExecutionFailed(format!("DB error: {e}")))? .ok_or_else(|| ToolError::ExecutionFailed(format!("routine '{}' not found", name)))?; @@ -1430,7 +1462,7 @@ impl Tool for RoutineDeleteTool { self.engine.refresh_event_cache().await; let result = serde_json::json!({ - "name": name, + "name": &name, "deleted": deleted, }); @@ -1836,6 +1868,20 @@ mod tests { assert_eq!(parsed.cooldown_secs, 30); } + #[test] + fn build_routine_trigger_normalizes_cron_schedule() { + let trigger = build_routine_trigger(&NormalizedTriggerRequest::Cron { + schedule: "0 0 9 * * MON-FRI".to_string(), + timezone: Some("UTC".to_string()), + }); + + assert!(matches!( + trigger, + Trigger::Cron { schedule, timezone } + if schedule == "0 0 9 * * MON-FRI *" && timezone.as_deref() == Some("UTC") + )); + } + #[test] fn parses_grouped_message_event_with_tools() { let params = serde_json::json!({ diff --git a/src/tools/mcp/client.rs b/src/tools/mcp/client.rs index 39ae5047..72a1baba 100644 --- a/src/tools/mcp/client.rs +++ b/src/tools/mcp/client.rs @@ -129,6 +129,11 @@ impl McpClient { /// The config must use HTTP transport (the default); for stdio/UDS use `new_with_transport`. /// /// Returns an error if the config uses a non-HTTP transport. + /// + /// **Note:** The session manager is NOT wired into the transport. For + /// production use, prefer `create_client_from_config()` which constructs + /// the transport with session tracking. + #[cfg(test)] pub fn new_with_config(config: McpServerConfig) -> Result { config .validate() @@ -235,7 +240,14 @@ impl McpClient { } } - /// Attach a session manager for Streamable HTTP session tracking. + /// Attach a session manager to the **client** only. + /// + /// **Warning:** This does NOT wire the session manager into the underlying + /// `HttpMcpTransport`, so the transport will not capture `Mcp-Session-Id` + /// from responses. For production use, construct the transport with + /// `HttpMcpTransport::with_session_manager()` and pass it to + /// `new_with_transport()` instead. See `create_client_from_config()`. + #[cfg(test)] pub fn with_session_manager(mut self, session_manager: Arc) -> Self { self.session_manager = Some(session_manager); self @@ -271,6 +283,12 @@ impl McpClient { self.session_manager.is_some() } + /// Get the underlying transport (test-only). + #[cfg(test)] + pub(crate) fn transport(&self) -> &Arc { + &self.transport + } + /// Get the next request ID. fn next_request_id(&self) -> u64 { self.next_id.fetch_add(1, Ordering::SeqCst) diff --git a/src/tools/mcp/factory.rs b/src/tools/mcp/factory.rs index 915db1a2..622e168e 100644 --- a/src/tools/mcp/factory.rs +++ b/src/tools/mcp/factory.rs @@ -7,6 +7,7 @@ use std::sync::Arc; use crate::secrets::SecretsStore; use crate::tools::mcp::config::{EffectiveTransport, McpServerConfig}; +use crate::tools::mcp::http_transport::HttpMcpTransport; use crate::tools::mcp::{McpClient, McpProcessManager, McpSessionManager, McpTransport}; /// Error returned when MCP client creation fails. @@ -91,43 +92,52 @@ pub async fn create_client_from_config( } })?; - return Ok(McpClient::new_with_config(server) - .map_err(|e| McpFactoryError::InvalidConfig { - name: server_name.clone(), - reason: e.to_string(), - })? - .with_nearai_session_manager(nearai_session_manager) - .with_nearai_api_key(nearai_api_key) - .with_session_manager(Arc::clone(session_manager))); - } + let transport = Arc::new( + HttpMcpTransport::new(server.url.clone(), server.name.clone()) + .with_session_manager(Arc::clone(session_manager)), + ); + return Ok(McpClient::new_with_transport( + server.name.clone(), + transport, + Some(Arc::clone(session_manager)), + secrets, + user_id, + Some(server), + ) + .with_nearai_session_manager(nearai_session_manager) + .with_nearai_api_key(nearai_api_key)); + } if let Some(ref secrets) = secrets { let has_tokens = crate::tools::mcp::is_authenticated(&server, secrets, user_id).await; if has_tokens || server.requires_auth() { - Ok(McpClient::new_authenticated( + return Ok(McpClient::new_authenticated( server, Arc::clone(session_manager), Arc::clone(secrets), user_id, - )) - } else { - Ok(McpClient::new_with_config(server) - .map_err(|e| McpFactoryError::InvalidConfig { - name: server_name.clone(), - reason: e.to_string(), - })? - .with_session_manager(Arc::clone(session_manager))) + )); } - } else { - Ok(McpClient::new_with_config(server) - .map_err(|e| McpFactoryError::InvalidConfig { - name: server_name, - reason: e.to_string(), - })? - .with_session_manager(Arc::clone(session_manager))) } + + // Non-OAuth HTTP: wire the session manager into the *transport* so + // it captures `Mcp-Session-Id` from responses. Passing it only to + // the client (via `with_session_manager`) is not enough — the + // transport must know about it to read/write the header. + let transport = Arc::new( + HttpMcpTransport::new(server.url.clone(), server.name.clone()) + .with_session_manager(Arc::clone(session_manager)), + ); + Ok(McpClient::new_with_transport( + server.name.clone(), + transport, + Some(Arc::clone(session_manager)), + secrets, + user_id, + Some(server), + )) } } } @@ -159,4 +169,86 @@ mod tests { "non-OAuth HTTP clients must carry a session manager" ); } + + /// Regression test: the factory must wire the session manager into the + /// *transport*, not just the client. Otherwise the transport never + /// captures `Mcp-Session-Id` from responses and subsequent requests + /// lack the header, causing the server to reject them. + #[tokio::test] + async fn test_factory_non_oauth_http_transport_captures_session_id() { + use axum::http::header::HeaderName; + use axum::{Router, http::StatusCode, response::IntoResponse, routing::post}; + use tokio::net::TcpListener; + + const SESSION_ID: &str = "test-session-abc123"; + + async fn session_echo() -> impl IntoResponse { + let body = serde_json::json!({ + "jsonrpc": "2.0", + "id": 1, + "result": {} + }) + .to_string(); + ( + StatusCode::OK, + [( + HeaderName::from_static("mcp-session-id"), + SESSION_ID.to_string(), + )], + body, + ) + } + + let app = Router::new().route("/", post(session_echo)); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let url = format!("http://127.0.0.1:{}", addr.port()); + + tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + + let server = McpServerConfig::new("session-test", &url); + let session_manager = Arc::new(McpSessionManager::new()); + let process_manager = Arc::new(McpProcessManager::new()); + + let client = create_client_from_config( + server, + &session_manager, + None, + None, + &process_manager, + None, + "test-user", + ) + .await + .expect("factory should succeed for HTTP config"); + + // Pre-create a session entry so that update_session_id has something to update. + // In production, the MCP initialize handshake calls get_or_create before responses arrive. + session_manager.get_or_create("session-test", &url).await; + + // Send a request through the client's transport to trigger session capture. + use crate::tools::mcp::protocol::McpRequest; + let request = McpRequest { + jsonrpc: "2.0".to_string(), + id: Some(1), + method: "test".to_string(), + params: Some(serde_json::json!({})), + }; + let headers = std::collections::HashMap::new(); + client + .transport() + .send(&request, &headers) + .await + .expect("request should succeed"); + + // Verify the session manager captured the session ID from the response. + let captured = session_manager.get_session_id("session-test").await; + assert_eq!( + captured.as_deref(), + Some(SESSION_ID), + "transport must capture Mcp-Session-Id into session manager" + ); + } } diff --git a/src/tools/mcp/http_transport.rs b/src/tools/mcp/http_transport.rs index 59873ce4..ea3e1c03 100644 --- a/src/tools/mcp/http_transport.rs +++ b/src/tools/mcp/http_transport.rs @@ -494,6 +494,34 @@ mod tests { assert_eq!(echoed["authorization"], "Bearer oauth-token"); } + /// Regression test for #1436: 202 Accepted responses for notifications + /// were parsed as JSON, causing "Failed to parse MCP response" errors + /// that broke the MCP session handshake. + #[tokio::test] + async fn test_wire_202_accepted_for_notification() { + use axum::{Router, http::StatusCode, routing::post}; + use tokio::net::TcpListener; + + async fn accept_notification() -> StatusCode { + StatusCode::ACCEPTED + } + + let app = Router::new().route("/", post(accept_notification)); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let url = format!("http://127.0.0.1:{}", addr.port()); + + tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + + let transport = HttpMcpTransport::new(&url, "test-202"); + let request = McpRequest::initialized_notification(); + let response = transport.send(&request, &HashMap::new()).await.unwrap(); + assert!(response.result.is_none()); + assert!(response.error.is_none()); + } + #[tokio::test] async fn test_wire_custom_auth_preserved_when_no_per_request_auth() { let (url, _handle) = spawn_echo_server().await; diff --git a/src/tools/registry.rs b/src/tools/registry.rs index bc3be144..8c08633b 100644 --- a/src/tools/registry.rs +++ b/src/tools/registry.rs @@ -383,11 +383,7 @@ impl ToolRegistry { job_manager: Option>, store: Option>, job_event_tx: Option< - tokio::sync::broadcast::Sender<( - uuid::Uuid, - String, - crate::channels::web::types::SseEvent, - )>, + tokio::sync::broadcast::Sender<(uuid::Uuid, String, ironclaw_common::AppEvent)>, >, inject_tx: Option>, prompt_queue: Option, diff --git a/src/tools/wasm/loader.rs b/src/tools/wasm/loader.rs index b50fc717..4876dc1b 100644 --- a/src/tools/wasm/loader.rs +++ b/src/tools/wasm/loader.rs @@ -418,6 +418,7 @@ fn resolve_oauth_refresh_config(cap_file: &CapabilitiesFile) -> Option Option, + } + + impl Drop for EnvVarGuard { + fn drop(&mut self) { + // SAFETY: Tests use lock_env() to serialize environment access. + unsafe { + if let Some(ref value) = self.previous { + std::env::set_var(&self.key, value); + } else { + std::env::remove_var(&self.key); + } + } + } + } + + fn set_env_var(key: &str, value: Option<&str>) -> EnvVarGuard { + let previous = std::env::var(key).ok(); + // SAFETY: Tests use lock_env() to serialize environment access. + unsafe { + match value { + Some(value) => std::env::set_var(key, value), + None => std::env::remove_var(key), + } + } + EnvVarGuard { + key: key.to_string(), + previous, + } + } + #[test] fn wit_version_compat_none_is_ok() { // Pre-versioning extensions (no wit_version declared) should always pass @@ -845,6 +889,11 @@ mod tests { AuthCapabilitySchema, CapabilitiesFile, OAuthConfigSchema, }; + let _guard = lock_env(); + let _proxy_guard = set_env_var("IRONCLAW_OAUTH_EXCHANGE_URL", None); + let _gateway_token_guard = set_env_var("GATEWAY_AUTH_TOKEN", None); + let _oauth_proxy_token_guard = set_env_var("IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", None); + let caps = CapabilitiesFile { auth: Some(AuthCapabilitySchema { secret_name: "google_oauth_token".to_string(), @@ -871,6 +920,8 @@ mod tests { config.client_secret, Some(TEST_OAUTH_CLIENT_SECRET.to_string()) ); + assert_eq!(config.exchange_proxy_url, None); + assert_eq!(config.gateway_token, None); assert_eq!(config.secret_name, "google_oauth_token"); assert_eq!(config.provider, Some("google".to_string())); } @@ -931,6 +982,11 @@ mod tests { AuthCapabilitySchema, CapabilitiesFile, OAuthConfigSchema, }; + let _guard = lock_env(); + let _proxy_guard = set_env_var("IRONCLAW_OAUTH_EXCHANGE_URL", None); + let _gateway_token_guard = set_env_var("GATEWAY_AUTH_TOKEN", None); + let _oauth_proxy_token_guard = set_env_var("IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", None); + // google_oauth_token should fall back to built-in credentials let caps = CapabilitiesFile { auth: Some(AuthCapabilitySchema { @@ -952,6 +1008,138 @@ mod tests { let config = config.unwrap(); assert!(!config.client_id.is_empty()); assert!(config.client_secret.is_some()); + assert_eq!(config.exchange_proxy_url, None); + assert_eq!(config.gateway_token, None); + } + + #[test] + fn test_resolve_oauth_refresh_config_hosted_proxy_populates_env_and_suppresses_builtin_secret() + { + use crate::tools::wasm::capabilities_schema::{ + AuthCapabilitySchema, CapabilitiesFile, OAuthConfigSchema, + }; + + let _guard = lock_env(); + let _proxy_guard = set_env_var( + "IRONCLAW_OAUTH_EXCHANGE_URL", + Some("https://compose-api.example.com"), + ); + let _gateway_token_guard = set_env_var("GATEWAY_AUTH_TOKEN", Some("gateway-test-token")); + let _oauth_proxy_token_guard = set_env_var("IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", None); + let _client_id_guard = + set_env_var("GOOGLE_OAUTH_CLIENT_ID", Some("hosted-google-client-id")); + + let caps = CapabilitiesFile { + auth: Some(AuthCapabilitySchema { + secret_name: "google_oauth_token".to_string(), + provider: Some("google".to_string()), + oauth: Some(OAuthConfigSchema { + authorization_url: "https://accounts.google.com/o/oauth2/v2/auth".to_string(), + token_url: "https://oauth2.googleapis.com/token".to_string(), + client_id_env: Some("GOOGLE_OAUTH_CLIENT_ID".to_string()), + ..Default::default() + }), + ..Default::default() + }), + ..Default::default() + }; + + let config = super::resolve_oauth_refresh_config(&caps).expect("hosted oauth config"); + assert_eq!(config.client_id, "hosted-google-client-id"); + assert_eq!(config.client_secret, None); + assert_eq!( + config.exchange_proxy_url.as_deref(), + Some("https://compose-api.example.com") + ); + assert_eq!(config.gateway_token.as_deref(), Some("gateway-test-token")); + } + + #[test] + fn test_resolve_oauth_refresh_config_hosted_proxy_preserves_explicit_secret() { + use crate::tools::wasm::capabilities_schema::{ + AuthCapabilitySchema, CapabilitiesFile, OAuthConfigSchema, + }; + + let _guard = lock_env(); + let _proxy_guard = set_env_var( + "IRONCLAW_OAUTH_EXCHANGE_URL", + Some("https://compose-api.example.com"), + ); + let _gateway_token_guard = set_env_var("GATEWAY_AUTH_TOKEN", Some("gateway-test-token")); + let _oauth_proxy_token_guard = set_env_var("IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", None); + let _client_id_guard = + set_env_var("GOOGLE_OAUTH_CLIENT_ID", Some("hosted-google-client-id")); + let _client_secret_guard = + set_env_var("GOOGLE_OAUTH_CLIENT_SECRET", Some("hosted-server-secret")); + + let caps = CapabilitiesFile { + auth: Some(AuthCapabilitySchema { + secret_name: "google_oauth_token".to_string(), + provider: Some("google".to_string()), + oauth: Some(OAuthConfigSchema { + authorization_url: "https://accounts.google.com/o/oauth2/v2/auth".to_string(), + token_url: "https://oauth2.googleapis.com/token".to_string(), + client_id_env: Some("GOOGLE_OAUTH_CLIENT_ID".to_string()), + client_secret_env: Some("GOOGLE_OAUTH_CLIENT_SECRET".to_string()), + ..Default::default() + }), + ..Default::default() + }), + ..Default::default() + }; + + let config = super::resolve_oauth_refresh_config(&caps).expect("hosted oauth config"); + assert_eq!(config.client_id, "hosted-google-client-id"); + assert_eq!( + config.client_secret.as_deref(), + Some("hosted-server-secret") + ); + assert_eq!( + config.exchange_proxy_url.as_deref(), + Some("https://compose-api.example.com") + ); + assert_eq!(config.gateway_token.as_deref(), Some("gateway-test-token")); + } + + #[test] + fn test_resolve_oauth_refresh_config_hosted_proxy_prefers_dedicated_proxy_auth_token() { + use crate::tools::wasm::capabilities_schema::{ + AuthCapabilitySchema, CapabilitiesFile, OAuthConfigSchema, + }; + + let _guard = lock_env(); + let _proxy_guard = set_env_var( + "IRONCLAW_OAUTH_EXCHANGE_URL", + Some("https://compose-api.example.com"), + ); + let _gateway_token_guard = set_env_var("GATEWAY_AUTH_TOKEN", Some("gateway-test-token")); + let _oauth_proxy_token_guard = set_env_var( + "IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", + Some("shared-oauth-proxy-secret"), + ); + let _client_id_guard = + set_env_var("GOOGLE_OAUTH_CLIENT_ID", Some("hosted-google-client-id")); + + let caps = CapabilitiesFile { + auth: Some(AuthCapabilitySchema { + secret_name: "google_oauth_token".to_string(), + provider: Some("google".to_string()), + oauth: Some(OAuthConfigSchema { + authorization_url: "https://accounts.google.com/o/oauth2/v2/auth".to_string(), + token_url: "https://oauth2.googleapis.com/token".to_string(), + client_id_env: Some("GOOGLE_OAUTH_CLIENT_ID".to_string()), + ..Default::default() + }), + ..Default::default() + }), + ..Default::default() + }; + + let config = super::resolve_oauth_refresh_config(&caps).expect("hosted oauth config"); + assert_eq!( + config.gateway_token.as_deref(), + Some("shared-oauth-proxy-secret") + ); } // --------------------------------------------------------------- diff --git a/src/tools/wasm/wrapper.rs b/src/tools/wasm/wrapper.rs index 33fcedb9..bb9c4dd4 100644 --- a/src/tools/wasm/wrapper.rs +++ b/src/tools/wasm/wrapper.rs @@ -19,7 +19,7 @@ use wasmtime_wasi::{ResourceTable, WasiCtx, WasiCtxBuilder, WasiView}; use crate::context::JobContext; use crate::llm::recording::{HttpExchangeRequest, HttpExchangeResponse, HttpInterceptor}; use crate::safety::LeakDetector; -use crate::secrets::SecretsStore; +use crate::secrets::{DecryptedSecret, SecretsStore}; use crate::tools::tool::{Tool, ToolError, ToolOutput}; use crate::tools::wasm::capabilities::Capabilities; use crate::tools::wasm::credential_injector::{ @@ -44,6 +44,7 @@ wasmtime::component::bindgen!({ }); // Alias the export interface types for convenience. +use crate::cli::oauth_defaults; use exports::near::agent::tool as wit_tool; /// Configuration needed to refresh an expired OAuth access token. @@ -59,6 +60,11 @@ pub struct OAuthRefreshConfig { pub client_id: String, /// OAuth client_secret (optional, some providers use PKCE without a secret). pub client_secret: Option, + /// Hosted OAuth proxy base URL (e.g., "http://host.docker.internal:8080"). + pub exchange_proxy_url: Option, + /// OAuth proxy auth token for authenticating with the hosted OAuth proxy. + /// Kept as `gateway_token` for public API compatibility. + pub gateway_token: Option, /// Secret name of the access token (e.g., "google_oauth_token"). /// The refresh token lives at `{secret_name}_refresh_token`. pub secret_name: String, @@ -66,6 +72,12 @@ pub struct OAuthRefreshConfig { pub provider: Option, } +impl OAuthRefreshConfig { + fn oauth_proxy_auth_token(&self) -> Option<&str> { + self.gateway_token.as_deref() + } +} + /// Pre-resolved credential for host-based injection. /// /// Built before each WASM execution by decrypting secrets from the store. @@ -1210,6 +1222,53 @@ async fn refresh_oauth_token( user_id: &str, config: &OAuthRefreshConfig, ) -> bool { + let refresh_name = format!("{}_refresh_token", config.secret_name); + + if let Some(proxy_url) = config.exchange_proxy_url.as_deref() { + let Some(oauth_proxy_auth_token) = config.oauth_proxy_auth_token() else { + tracing::warn!( + "OAuth refresh proxy is configured, but no OAuth proxy auth token is available" + ); + return false; + }; + + // In hosted mode, the configured exchange proxy owns the outbound token + // refresh and validation policy for the provider token_url. Direct-mode + // HTTPS/private-IP checks remain in place for self-hosted refreshes below. + let refresh_secret = match load_oauth_refresh_secret(store, user_id, &refresh_name).await { + Some(secret) => secret, + None => return false, + }; + let token_response = match oauth_defaults::refresh_token_via_proxy( + oauth_defaults::ProxyRefreshTokenRequest { + proxy_url, + gateway_token: oauth_proxy_auth_token, + token_url: &config.token_url, + client_id: &config.client_id, + client_secret: config.client_secret.as_deref(), + refresh_token: refresh_secret.expose(), + provider: config.provider.as_deref(), + }, + ) + .await + { + Ok(response) => response, + Err(error) => { + tracing::warn!(error = %error, "OAuth token refresh via proxy failed"); + return false; + } + }; + + return persist_refreshed_oauth_tokens( + store, + user_id, + config, + &refresh_name, + token_response, + ) + .await; + } + // SSRF defense: token_url comes from the tool's capabilities file. if !config.token_url.starts_with("https://") { tracing::warn!( @@ -1227,19 +1286,6 @@ async fn refresh_oauth_token( return false; } - let refresh_name = format!("{}_refresh_token", config.secret_name); - let refresh_secret = match store.get_decrypted(user_id, &refresh_name).await { - Ok(s) => s, - Err(e) => { - tracing::debug!( - secret_name = %refresh_name, - error = %e, - "No refresh token available, skipping token refresh" - ); - return false; - } - }; - let client = match reqwest::Client::builder() .timeout(Duration::from_secs(15)) .redirect(reqwest::redirect::Policy::none()) @@ -1252,6 +1298,10 @@ async fn refresh_oauth_token( } }; + let refresh_secret = match load_oauth_refresh_secret(store, user_id, &refresh_name).await { + Some(secret) => secret, + None => return false, + }; let mut params = vec![ ("grant_type", "refresh_token".to_string()), ("refresh_token", refresh_secret.expose().to_string()), @@ -1287,22 +1337,55 @@ async fn refresh_oauth_token( return false; } }; - - let new_access_token = match token_data.get("access_token").and_then(|v| v.as_str()) { - Some(t) => t, + let token_response = match token_data.get("access_token").and_then(|v| v.as_str()) { + Some(access_token) => oauth_defaults::OAuthTokenResponse { + access_token: access_token.to_string(), + refresh_token: token_data + .get("refresh_token") + .and_then(|v| v.as_str()) + .map(str::to_string), + expires_in: token_data.get("expires_in").and_then(|v| v.as_u64()), + }, None => { tracing::warn!("Token refresh response missing access_token field"); return false; } }; - // Store the new access token with expiry + persist_refreshed_oauth_tokens(store, user_id, config, &refresh_name, token_response).await +} + +async fn load_oauth_refresh_secret( + store: &(dyn SecretsStore + Send + Sync), + user_id: &str, + refresh_name: &str, +) -> Option { + match store.get_decrypted(user_id, refresh_name).await { + Ok(secret) => Some(secret), + Err(error) => { + tracing::debug!( + secret_name = %refresh_name, + error = %error, + "No refresh token available, skipping token refresh" + ); + None + } + } +} + +async fn persist_refreshed_oauth_tokens( + store: &(dyn SecretsStore + Send + Sync), + user_id: &str, + config: &OAuthRefreshConfig, + refresh_name: &str, + token_response: oauth_defaults::OAuthTokenResponse, +) -> bool { let mut access_params = - crate::secrets::CreateSecretParams::new(&config.secret_name, new_access_token); + crate::secrets::CreateSecretParams::new(&config.secret_name, &token_response.access_token); if let Some(ref provider) = config.provider { access_params = access_params.with_provider(provider); } - if let Some(expires_in) = token_data.get("expires_in").and_then(|v| v.as_u64()) { + if let Some(expires_in) = token_response.expires_in { let expires_at = chrono::Utc::now() + chrono::Duration::seconds(expires_in as i64); access_params = access_params.with_expiry(expires_at); } @@ -1312,10 +1395,8 @@ async fn refresh_oauth_token( return false; } - // Store rotated refresh token if the provider sent a new one - if let Some(new_refresh) = token_data.get("refresh_token").and_then(|v| v.as_str()) { - let mut refresh_params = - crate::secrets::CreateSecretParams::new(&refresh_name, new_refresh); + if let Some(new_refresh) = token_response.refresh_token.as_deref() { + let mut refresh_params = crate::secrets::CreateSecretParams::new(refresh_name, new_refresh); if let Some(ref provider) = config.provider { refresh_params = refresh_params.with_provider(provider); } @@ -1664,9 +1745,18 @@ fn build_tool_usage_hint(tool_name: &str, schema: &serde_json::Value) -> String #[cfg(test)] mod tests { + use std::collections::HashMap; + use std::net::SocketAddr; use std::sync::{Arc, Mutex}; use async_trait::async_trait; + use axum::extract::{Form, State}; + use axum::http::HeaderMap; + use axum::routing::post; + use axum::{Json, Router}; + use serde_json::json; + use tokio::net::TcpListener; + use tokio::sync::{Mutex as AsyncMutex, oneshot}; use uuid::Uuid; use crate::context::JobContext; @@ -1756,6 +1846,95 @@ mod tests { } } + #[derive(Clone, Debug, PartialEq, Eq)] + struct RecordedProxyRequest { + authorization: Option, + form: HashMap, + } + + struct MockProxyServer { + addr: SocketAddr, + requests: Arc>>, + shutdown_tx: Option>, + server_task: Option>, + } + + impl MockProxyServer { + async fn start() -> Self { + async fn refresh_handler( + State(requests): State>>>, + headers: HeaderMap, + Form(form): Form>, + ) -> Json { + requests.lock().await.push(RecordedProxyRequest { + authorization: headers + .get(axum::http::header::AUTHORIZATION) + .and_then(|value| value.to_str().ok()) + .map(str::to_string), + form, + }); + Json(json!({ + "access_token": "mock-refreshed-access-token", + "refresh_token": "mock-rotated-refresh-token", + "expires_in": 3600 + })) + } + + let requests = Arc::new(AsyncMutex::new(Vec::new())); + let app = Router::new() + .route("/oauth/refresh", post(refresh_handler)) + .with_state(Arc::clone(&requests)); + + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("bind mock proxy"); + let addr = listener.local_addr().expect("read mock proxy addr"); + let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>(); + let server_task = tokio::spawn(async move { + let _ = axum::serve(listener, app) + .with_graceful_shutdown(async { + let _ = shutdown_rx.await; + }) + .await; + }); + + Self { + addr, + requests, + shutdown_tx: Some(shutdown_tx), + server_task: Some(server_task), + } + } + + fn base_url(&self) -> String { + format!("http://{}", self.addr) + } + + async fn requests(&self) -> Vec { + self.requests.lock().await.clone() + } + + async fn shutdown(mut self) { + if let Some(tx) = self.shutdown_tx.take() { + let _ = tx.send(()); + } + if let Some(task) = self.server_task.take() { + let _ = task.await; + } + } + } + + impl Drop for MockProxyServer { + fn drop(&mut self) { + if let Some(tx) = self.shutdown_tx.take() { + let _ = tx.send(()); + } + if let Some(task) = self.server_task.take() { + task.abort(); + } + } + } + #[test] fn test_wrapper_creation() { // This test verifies the runtime can be created @@ -2094,8 +2273,6 @@ mod tests { #[tokio::test] async fn test_resolve_host_credentials_bearer() { - use std::collections::HashMap; - use crate::secrets::{ CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore, }; @@ -2141,8 +2318,6 @@ mod tests { #[tokio::test] async fn test_resolve_host_credentials_owner_scope_bearer() { - use std::collections::HashMap; - use crate::secrets::{ CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore, }; @@ -2188,8 +2363,6 @@ mod tests { #[tokio::test] async fn test_execute_resolves_host_credentials_from_owner_scope_context() { - use std::collections::HashMap; - use crate::secrets::{CredentialLocation, CredentialMapping}; use crate::tools::wasm::capabilities::HttpCapability; @@ -2239,8 +2412,6 @@ mod tests { #[tokio::test] async fn test_resolve_host_credentials_missing_secret() { - use std::collections::HashMap; - use crate::secrets::{CredentialLocation, CredentialMapping}; use crate::tools::wasm::capabilities::HttpCapability; use crate::tools::wasm::wrapper::resolve_host_credentials; @@ -2272,8 +2443,6 @@ mod tests { #[tokio::test] async fn test_resolve_host_credentials_skips_refresh_when_not_expired() { - use std::collections::HashMap; - use crate::secrets::{ CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore, }; @@ -2315,6 +2484,8 @@ mod tests { token_url: "https://oauth2.googleapis.com/token".to_string(), client_id: TEST_OAUTH_CLIENT_ID.to_string(), client_secret: Some(TEST_OAUTH_CLIENT_SECRET.to_string()), + exchange_proxy_url: None, + gateway_token: None, secret_name: "google_oauth_token".to_string(), provider: Some("google".to_string()), }; @@ -2331,8 +2502,6 @@ mod tests { #[tokio::test] async fn test_resolve_host_credentials_skips_refresh_no_config() { - use std::collections::HashMap; - use crate::secrets::{ CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore, }; @@ -2376,8 +2545,6 @@ mod tests { #[tokio::test] async fn test_resolve_host_credentials_skips_refresh_no_expires_at() { - use std::collections::HashMap; - use crate::secrets::{ CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore, }; @@ -2417,6 +2584,8 @@ mod tests { token_url: "https://oauth2.googleapis.com/token".to_string(), client_id: TEST_OAUTH_CLIENT_ID.to_string(), client_secret: Some(TEST_OAUTH_CLIENT_SECRET.to_string()), + exchange_proxy_url: None, + gateway_token: None, secret_name: "google_oauth_token".to_string(), provider: Some("google".to_string()), }; @@ -2431,6 +2600,250 @@ mod tests { ); } + #[tokio::test] + async fn test_resolve_host_credentials_refreshes_via_proxy_without_direct_token_url_validation() + { + use crate::secrets::{ + CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore, + }; + use crate::tools::wasm::capabilities::HttpCapability; + use crate::tools::wasm::wrapper::{OAuthRefreshConfig, resolve_host_credentials}; + + let proxy = MockProxyServer::start().await; + let store = test_secrets_store(); + + store + .create( + "user1", + CreateSecretParams::new("google_oauth_token", "expired-access-token") + .with_expiry(chrono::Utc::now() - chrono::Duration::hours(1)), + ) + .await + .unwrap(); + store + .create( + "user1", + CreateSecretParams::new("google_oauth_token_refresh_token", "stored-refresh-token"), + ) + .await + .unwrap(); + + let mut credentials = HashMap::new(); + credentials.insert( + "google_oauth_token".to_string(), + CredentialMapping { + secret_name: "google_oauth_token".to_string(), + location: CredentialLocation::AuthorizationBearer, + host_patterns: vec!["www.googleapis.com".to_string()], + }, + ); + + let caps = Capabilities { + http: Some(HttpCapability { + credentials, + ..Default::default() + }), + ..Default::default() + }; + + let oauth_config = OAuthRefreshConfig { + token_url: "http://127.0.0.1:9/provider-token-endpoint".to_string(), + client_id: "hosted-google-client-id".to_string(), + client_secret: None, + exchange_proxy_url: Some(proxy.base_url()), + gateway_token: Some("gateway-test-token".to_string()), + secret_name: "google_oauth_token".to_string(), + provider: Some("google".to_string()), + }; + + let resolved = + resolve_host_credentials(&caps, Some(&store), "user1", Some(&oauth_config)).await; + assert_eq!(resolved.len(), 1); + assert_eq!( + resolved[0].headers.get("Authorization"), + Some(&"Bearer mock-refreshed-access-token".to_string()) + ); + + let access_secret = store.get("user1", "google_oauth_token").await.unwrap(); + assert!( + access_secret + .expires_at + .expect("refreshed access token expiry") + > chrono::Utc::now() + ); + let access_value = store + .get_decrypted("user1", "google_oauth_token") + .await + .unwrap(); + assert_eq!(access_value.expose(), "mock-refreshed-access-token"); + + let refresh_value = store + .get_decrypted("user1", "google_oauth_token_refresh_token") + .await + .unwrap(); + assert_eq!(refresh_value.expose(), "mock-rotated-refresh-token"); + + let requests = proxy.requests().await; + assert_eq!(requests.len(), 1); + assert_eq!( + requests[0].authorization.as_deref(), + Some("Bearer gateway-test-token") + ); + assert_eq!( + requests[0].form.get("client_id").map(String::as_str), + Some("hosted-google-client-id") + ); + assert_eq!( + requests[0].form.get("token_url").map(String::as_str), + Some("http://127.0.0.1:9/provider-token-endpoint") + ); + assert_eq!( + requests[0].form.get("refresh_token").map(String::as_str), + Some("stored-refresh-token") + ); + assert_eq!( + requests[0].form.get("provider").map(String::as_str), + Some("google") + ); + assert!(!requests[0].form.contains_key("client_secret")); + + proxy.shutdown().await; + } + + #[tokio::test] + async fn test_resolve_host_credentials_skips_refresh_token_lookup_without_oauth_proxy_auth_token() + { + use crate::secrets::{ + CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore, + }; + use crate::tools::wasm::capabilities::HttpCapability; + use crate::tools::wasm::wrapper::{OAuthRefreshConfig, resolve_host_credentials}; + + let store = RecordingSecretsStore::new(); + + store + .create( + "user1", + CreateSecretParams::new("google_oauth_token", "expired-access-token") + .with_expiry(chrono::Utc::now() - chrono::Duration::hours(1)), + ) + .await + .unwrap(); + store + .create( + "user1", + CreateSecretParams::new("google_oauth_token_refresh_token", "stored-refresh-token"), + ) + .await + .unwrap(); + + let mut credentials = HashMap::new(); + credentials.insert( + "google_oauth_token".to_string(), + CredentialMapping { + secret_name: "google_oauth_token".to_string(), + location: CredentialLocation::AuthorizationBearer, + host_patterns: vec!["www.googleapis.com".to_string()], + }, + ); + + let caps = Capabilities { + http: Some(HttpCapability { + credentials, + ..Default::default() + }), + ..Default::default() + }; + + let oauth_config = OAuthRefreshConfig { + token_url: "https://oauth2.googleapis.com/token".to_string(), + client_id: "hosted-google-client-id".to_string(), + client_secret: None, + exchange_proxy_url: Some("https://compose-api.example.com".to_string()), + gateway_token: None, + secret_name: "google_oauth_token".to_string(), + provider: Some("google".to_string()), + }; + + let resolved = + resolve_host_credentials(&caps, Some(&store), "user1", Some(&oauth_config)).await; + assert!(resolved.is_empty()); + + let lookups = store.decrypted_lookups(); + assert!(lookups.contains(&("user1".to_string(), "google_oauth_token".to_string()))); + assert!(!lookups.contains(&( + "user1".to_string(), + "google_oauth_token_refresh_token".to_string(), + ))); + } + + #[tokio::test] + async fn test_resolve_host_credentials_skips_refresh_token_lookup_for_invalid_direct_token_url() + { + use crate::secrets::{ + CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore, + }; + use crate::tools::wasm::capabilities::HttpCapability; + use crate::tools::wasm::wrapper::{OAuthRefreshConfig, resolve_host_credentials}; + + let store = RecordingSecretsStore::new(); + + store + .create( + "user1", + CreateSecretParams::new("google_oauth_token", "expired-access-token") + .with_expiry(chrono::Utc::now() - chrono::Duration::hours(1)), + ) + .await + .unwrap(); + store + .create( + "user1", + CreateSecretParams::new("google_oauth_token_refresh_token", "stored-refresh-token"), + ) + .await + .unwrap(); + + let mut credentials = HashMap::new(); + credentials.insert( + "google_oauth_token".to_string(), + CredentialMapping { + secret_name: "google_oauth_token".to_string(), + location: CredentialLocation::AuthorizationBearer, + host_patterns: vec!["www.googleapis.com".to_string()], + }, + ); + + let caps = Capabilities { + http: Some(HttpCapability { + credentials, + ..Default::default() + }), + ..Default::default() + }; + + let oauth_config = OAuthRefreshConfig { + token_url: "http://127.0.0.1:9/provider-token-endpoint".to_string(), + client_id: TEST_OAUTH_CLIENT_ID.to_string(), + client_secret: Some(TEST_OAUTH_CLIENT_SECRET.to_string()), + exchange_proxy_url: None, + gateway_token: None, + secret_name: "google_oauth_token".to_string(), + provider: Some("google".to_string()), + }; + + let resolved = + resolve_host_credentials(&caps, Some(&store), "user1", Some(&oauth_config)).await; + assert!(resolved.is_empty()); + + let lookups = store.decrypted_lookups(); + assert!(lookups.contains(&("user1".to_string(), "google_oauth_token".to_string()))); + assert!(!lookups.contains(&( + "user1".to_string(), + "google_oauth_token_refresh_token".to_string(), + ))); + } + #[test] fn test_is_private_ip_v4() { use std::net::IpAddr; diff --git a/src/tunnel/mod.rs b/src/tunnel/mod.rs index 06eeebd7..8719b6e1 100644 --- a/src/tunnel/mod.rs +++ b/src/tunnel/mod.rs @@ -429,8 +429,8 @@ mod tests { port: 3000, auth_token: None, user_id: "test".to_string(), - workspace_read_scopes: vec![], - memory_layers: vec![], + workspace_read_scopes: Vec::new(), + memory_layers: Vec::new(), user_tokens: None, }); c @@ -443,8 +443,8 @@ mod tests { port, auth_token: None, user_id: "test".to_string(), - workspace_read_scopes: vec![], - memory_layers: vec![], + workspace_read_scopes: Vec::new(), + memory_layers: Vec::new(), user_tokens: None, }); c diff --git a/src/util.rs b/src/util.rs index 866f623c..a76f3b27 100644 --- a/src/util.rs +++ b/src/util.rs @@ -1,5 +1,7 @@ //! Shared utility functions used across the codebase. +use crate::llm::{ChatMessage, Role}; + /// Find the largest valid UTF-8 char boundary at or before `pos`. /// /// Polyfill for `str::floor_char_boundary` (nightly-only). Use when @@ -16,6 +18,17 @@ pub fn floor_char_boundary(s: &str, pos: usize) -> usize { i } +/// Ensure the last message in `messages` is a user-role message. +/// +/// NEAR AI rejects conversations that don't end with a user message; +/// Claude 4.6 rejects assistant prefill. Call this before any LLM +/// completion request to satisfy both requirements. +pub fn ensure_ends_with_user_message(messages: &mut Vec) { + if !matches!(messages.last(), Some(m) if m.role == Role::User) { + messages.push(ChatMessage::user("Continue.")); + } +} + /// Check if an LLM response explicitly signals that a job/task is complete. /// /// Uses phrase-level matching to avoid false positives from bare words like @@ -72,7 +85,8 @@ pub fn llm_signals_completion(response: &str) -> bool { #[cfg(test)] mod tests { - use crate::util::{floor_char_boundary, llm_signals_completion}; + use crate::llm::ChatMessage; + use crate::util::{ensure_ends_with_user_message, floor_char_boundary, llm_signals_completion}; // ── floor_char_boundary ── @@ -103,6 +117,42 @@ mod tests { assert_eq!(floor_char_boundary("", 5), 0); } + // ── ensure_ends_with_user_message ── + + #[test] + fn ensure_user_message_injects_when_empty() { + let mut msgs: Vec = vec![]; + ensure_ends_with_user_message(&mut msgs); + assert_eq!(msgs.len(), 1); + assert_eq!(msgs[0].role, crate::llm::Role::User); + } + + #[test] + fn ensure_user_message_injects_after_assistant() { + let mut msgs = vec![ChatMessage::user("hi"), ChatMessage::assistant("hello")]; + ensure_ends_with_user_message(&mut msgs); + assert_eq!(msgs.len(), 3); + assert_eq!(msgs[2].role, crate::llm::Role::User); + } + + #[test] + fn ensure_user_message_injects_after_tool_result() { + let mut msgs = vec![ + ChatMessage::user("run tool"), + ChatMessage::tool_result("call_1", "my_tool", "result"), + ]; + ensure_ends_with_user_message(&mut msgs); + assert_eq!(msgs.len(), 3); + assert_eq!(msgs[2].role, crate::llm::Role::User); + } + + #[test] + fn ensure_user_message_no_op_when_already_user() { + let mut msgs = vec![ChatMessage::user("hello")]; + ensure_ends_with_user_message(&mut msgs); + assert_eq!(msgs.len(), 1); + } + // ── llm_signals_completion ── #[test] diff --git a/src/worker/container.rs b/src/worker/container.rs index e0933975..5d8e03b5 100644 --- a/src/worker/container.rs +++ b/src/worker/container.rs @@ -151,7 +151,7 @@ Job: {} Description: {} You have tools for shell commands, file operations, and code editing. -Work independently to complete this job. Report when done."#, +Work independently to complete this job. When finished, your final message MUST include the phrase "The job is complete" to signal termination."#, job.title, job.description ))); @@ -373,6 +373,10 @@ impl LoopDelegate for ContainerDelegate { // Poll for follow-up prompts from the user self.poll_and_inject_prompt(reason_ctx).await; + // Claude 4.6 rejects assistant prefill; NEAR AI rejects any non-user-ending + // conversation. Ensure the last message is user-role before calling the LLM. + crate::util::ensure_ends_with_user_message(&mut reason_ctx.messages); + // Refresh tools (in case WASM tools were built) reason_ctx.available_tools = self.tools.tool_definitions().await; diff --git a/src/worker/job.rs b/src/worker/job.rs index b2e3f7e6..f74d4ec8 100644 --- a/src/worker/job.rs +++ b/src/worker/job.rs @@ -18,9 +18,8 @@ use crate::agent::agentic_loop::{ }; use crate::agent::scheduler::WorkerMessage; use crate::agent::task::TaskOutput; -use crate::channels::web::types::SseEvent; +use crate::channels::web::types::ToolDecisionDto; use crate::context::{ContextManager, JobState}; -use crate::db::Database; use crate::error::Error; use crate::hooks::HookRegistry; use crate::llm::{ @@ -28,11 +27,13 @@ use crate::llm::{ ToolSelection, }; use crate::safety::SafetyLayer; +use crate::tenant::AdminScope; use crate::tools::execute::process_tool_result; use crate::tools::rate_limiter::RateLimitResult; use crate::tools::{ ApprovalContext, ToolRegistry, autonomous_unavailable_error, prepare_tool_params, redact_params, }; +use ironclaw_common::AppEvent; /// Shared dependencies for worker execution. /// @@ -44,11 +45,11 @@ pub struct WorkerDeps { pub llm: Arc, pub safety: Arc, pub tools: Arc, - pub store: Option>, + pub store: Option, pub hooks: Arc, pub timeout: Duration, pub use_planning: bool, - /// SSE manager for live job event streaming to the web gateway. + /// Broadcast sender for live job event streaming to the web gateway. pub sse_tx: Option>, /// Approval context for tool execution. When `None`, all non-`Never` tools are /// blocked (legacy behavior). When `Some`, the context determines which tools @@ -93,7 +94,7 @@ impl Worker { &self.deps.tools } - fn store(&self) -> Option<&Arc> { + fn store(&self) -> Option<&AdminScope> { self.deps.store.as_ref() } @@ -141,7 +142,7 @@ impl Worker { if let Some(ref sse) = self.deps.sse_tx { let job_id_str = job_id.to_string(); let event = match event_type { - "message" => Some(SseEvent::JobMessage { + "message" => Some(AppEvent::JobMessage { job_id: job_id_str, role: data .get("role") @@ -154,7 +155,7 @@ impl Worker { .unwrap_or("") .to_string(), }), - "tool_use" => Some(SseEvent::JobToolUse { + "tool_use" => Some(AppEvent::JobToolUse { job_id: job_id_str, tool_name: data .get("tool_name") @@ -166,7 +167,7 @@ impl Worker { .cloned() .unwrap_or(serde_json::Value::Null), }), - "tool_result" => Some(SseEvent::JobToolResult { + "tool_result" => Some(AppEvent::JobToolResult { job_id: job_id_str, tool_name: data .get("tool_name") @@ -179,7 +180,7 @@ impl Worker { .unwrap_or("") .to_string(), }), - "status" => Some(SseEvent::JobStatus { + "status" => Some(AppEvent::JobStatus { job_id: job_id_str, message: data .get("message") @@ -187,7 +188,7 @@ impl Worker { .unwrap_or("") .to_string(), }), - "result" => Some(SseEvent::JobResult { + "result" => Some(AppEvent::JobResult { job_id: job_id_str, status: data .get("status") @@ -200,6 +201,19 @@ impl Worker { .map(|s| s.to_string()), fallback_deliverable: data.get("fallback_deliverable").cloned(), }), + "reasoning" => { + let narrative = data + .get("narrative") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + let decisions = ToolDecisionDto::from_json_array(&data["decisions"]); + Some(AppEvent::JobReasoning { + job_id: job_id_str, + narrative, + decisions, + }) + } _ => None, }; if let Some(event) = event { @@ -897,6 +911,11 @@ Report when the job is complete or if you encounter issues you cannot resolve."# id: selection.tool_call_id.clone(), name: selection.tool_name.clone(), arguments: selection.parameters.clone(), + reasoning: if action.reasoning.is_empty() { + None + } else { + Some(action.reasoning.clone()) + }, }], )); @@ -1139,6 +1158,7 @@ impl<'a> JobDelegate<'a> { Ok(crate::llm::RespondOutput { result: RespondResult::Text(String::new()), usage: crate::llm::TokenUsage::default(), + finish_reason: crate::llm::FinishReason::Stop, }) } } @@ -1232,6 +1252,11 @@ impl<'a> LoopDelegate for JobDelegate<'a> { ) -> Option { // Refresh tool definitions so newly built tools become visible reason_ctx.available_tools = self.worker.tools().tool_definitions().await; + + // Claude 4.6 rejects assistant prefill; NEAR AI rejects any non-user-ending + // conversation. Ensure the last message is user-role before calling the LLM. + crate::util::ensure_ends_with_user_message(&mut reason_ctx.messages); + None } @@ -1259,6 +1284,7 @@ impl<'a> LoopDelegate for JobDelegate<'a> { content: reasoning_text, }, usage: crate::llm::TokenUsage::default(), + finish_reason: crate::llm::FinishReason::ToolUse, }); } Ok(_) => {} // empty selections, fall through @@ -1352,6 +1378,48 @@ impl<'a> LoopDelegate for JobDelegate<'a> { ); } + // Emit reasoning event if any tool calls carry reasoning. + // Sanitize narrative and per-tool rationale through SafetyLayer + // (parity with ChatDelegate in dispatcher.rs). + let sanitized_narrative = content + .as_deref() + .filter(|c| !c.trim().is_empty()) + .map(|c| { + self.worker + .deps + .safety + .sanitize_tool_output("job_narrative", c) + .content + }) + .filter(|c| !c.trim().is_empty()) + .unwrap_or_default(); + let decisions: Vec = tool_calls + .iter() + .filter_map(|tc| { + tc.reasoning.as_ref().map(|r| { + let sanitized = self + .worker + .deps + .safety + .sanitize_tool_output("tool_rationale", r) + .content; + serde_json::json!({ + "tool_name": tc.name, + "rationale": sanitized, + }) + }) + }) + .collect(); + if !decisions.is_empty() { + self.worker.log_event( + "reasoning", + serde_json::json!({ + "narrative": sanitized_narrative, + "decisions": decisions, + }), + ); + } + // Add assistant message with tool_calls (OpenAI protocol) reason_ctx .messages @@ -1366,7 +1434,7 @@ impl<'a> LoopDelegate for JobDelegate<'a> { .map(|tc| ToolSelection { tool_name: tc.name.clone(), parameters: tc.arguments.clone(), - reasoning: String::new(), + reasoning: tc.reasoning.clone().unwrap_or_default(), alternatives: vec![], tool_call_id: tc.id.clone(), }) @@ -1419,6 +1487,11 @@ fn selections_to_tool_calls(selections: &[ToolSelection]) -> Vec { id: s.tool_call_id.clone(), name: s.tool_name.clone(), arguments: s.parameters.clone(), + reasoning: if s.reasoning.is_empty() { + None + } else { + Some(s.reasoning.clone()) + }, }) .collect() } diff --git a/src/workspace/mod.rs b/src/workspace/mod.rs index 0242047f..51d7d2fc 100644 --- a/src/workspace/mod.rs +++ b/src/workspace/mod.rs @@ -149,6 +149,7 @@ fn reject_if_injected(path: &str, content: &str) -> Result<(), WorkspaceError> { /// /// Allows Workspace to work with either a PostgreSQL `Repository` (the original /// path) or any `Database` trait implementation (e.g. libSQL backend). +#[derive(Clone)] enum WorkspaceStorage { /// PostgreSQL-backed repository (uses connection pool directly). #[cfg(feature = "postgres")] @@ -576,6 +577,60 @@ impl Workspace { self } + /// Clone the workspace configuration for a different primary user scope. + /// + /// This preserves search config, embeddings, shared read scopes, memory + /// layers, and privacy classifier while switching the primary read/write + /// scope to `user_id`. + pub fn scoped_to_user(&self, user_id: impl Into) -> Self { + let user_id = user_id.into(); + + let mut memory_layers = self.memory_layers.clone(); + for layer in &mut memory_layers { + if layer.sensitivity == crate::workspace::layer::LayerSensitivity::Private + && layer.scope == self.user_id + { + layer.scope = user_id.clone(); + } + } + + let mut read_user_ids = vec![user_id.clone()]; + for scope in &self.read_user_ids { + if scope != &self.user_id && !read_user_ids.contains(scope) { + read_user_ids.push(scope.clone()); + } + } + for scope in crate::workspace::layer::MemoryLayer::read_scopes(&memory_layers) { + if !read_user_ids.contains(&scope) { + read_user_ids.push(scope); + } + } + + let preserve_flags = user_id == self.user_id; + Self { + user_id, + read_user_ids, + agent_id: self.agent_id, + storage: self.storage.clone(), + embeddings: self.embeddings.clone(), + bootstrap_pending: std::sync::atomic::AtomicBool::new(if preserve_flags { + self.bootstrap_pending + .load(std::sync::atomic::Ordering::Acquire) + } else { + false + }), + bootstrap_completed: std::sync::atomic::AtomicBool::new(if preserve_flags { + self.bootstrap_completed + .load(std::sync::atomic::Ordering::Acquire) + } else { + false + }), + search_defaults: self.search_defaults.clone(), + memory_layers, + privacy_classifier: self.privacy_classifier.clone(), + } + } + /// Get the user ID (primary scope for writes). pub fn user_id(&self) -> &str { &self.user_id diff --git a/src/workspace/repository.rs b/src/workspace/repository.rs index 78ddfec5..13f6816b 100644 --- a/src/workspace/repository.rs +++ b/src/workspace/repository.rs @@ -15,6 +15,7 @@ use crate::workspace::document::{MemoryChunk, MemoryDocument, WorkspaceEntry}; use crate::workspace::search::{RankedResult, SearchConfig, SearchResult, fuse_results}; /// Database repository for workspace operations. +#[derive(Clone)] pub struct Repository { pool: Pool, } diff --git a/tests/e2e/CLAUDE.md b/tests/e2e/CLAUDE.md index 0cf5e6dc..46b7b752 100644 --- a/tests/e2e/CLAUDE.md +++ b/tests/e2e/CLAUDE.md @@ -53,6 +53,7 @@ HEADED=1 pytest scenarios/ | `test_skills.py` | Skills tab UI visibility, ClawHub search (skipped if registry unreachable), install + remove lifecycle | | `test_sse_reconnect.py` | SSE reconnects after programmatic `eventSource.close()` + `connectSSE()`; history is reloaded after reconnect | | `test_tool_approval.py` | Approval card appears, buttons disable on approve/deny, parameters toggle via `page.evaluate("showApproval(...)")`; the waiting-approval regression uses a real HTTP tool call | +| `test_oauth_refresh.py` | Hosted Gmail OAuth regression: complete setup via `/oauth/callback`, expire the stored access token in libSQL, trigger a real `gmail` tool call through `/api/chat/send`, and verify refresh goes through the mock `/oauth/refresh` proxy without forwarding `client_secret` | ## `helpers.py` @@ -75,6 +76,7 @@ All fixtures are defined in `tests/e2e/conftest.py`. Running `pytest scenarios/` | `ironclaw_binary` | Checks `target/debug/ironclaw`; if absent, runs `cargo build --no-default-features --features libsql` (timeout 600s). | | `mock_llm_server` | Starts `mock_llm.py --port 0`, reads the assigned port from stdout, waits for `/v1/models` to return 200. Yields the base URL. | | `ironclaw_server` | Starts the ironclaw binary with a minimal env (see below), waits for `/api/health` (timeout 60s). Yields the base URL. On teardown sends **SIGINT** (not SIGTERM) so the tokio ctrl_c handler triggers a graceful shutdown and LLVM coverage data is flushed. | +| `hosted_oauth_refresh_server` | Starts a second ironclaw instance with a dedicated libSQL DB and `GOOGLE_OAUTH_CLIENT_ID=hosted-google-client-id`, while still pointing `IRONCLAW_OAUTH_EXCHANGE_URL` at `mock_llm.py`. Yields a dict with `base_url`, `db_path`, `gateway_user_id`, and `mock_llm_url` for the hosted refresh regression scenario. | | `browser` | Launches a single Chromium instance (headless by default; set `HEADED=1` for headed). Shared across all tests. | ### Function-scoped fixtures @@ -100,6 +102,8 @@ EMBEDDING_ENABLED=false, SKILLS_ENABLED=true ONBOARD_COMPLETED=true # prevents setup wizard ``` +The `hosted_oauth_refresh_server` fixture uses the same baseline, but with its own DB/home tempdirs and `GOOGLE_OAUTH_CLIENT_ID=hosted-google-client-id` so hosted OAuth flows exercise proxy credential injection instead of the baked-in desktop Google app. + The binary is also started with `--no-onboard`. Coverage env vars (`CARGO_LLVM_COV*`, `LLVM_*`, `CARGO_ENCODED_RUSTFLAGS`, `CARGO_INCREMENTAL`) are forwarded from the outer environment when present. ## Mock LLM (`mock_llm.py`) @@ -113,6 +117,11 @@ python mock_llm.py --port 0 It serves `POST /v1/chat/completions` (streaming + non-streaming) and `GET /v1/models`. Responses are pattern-matched from `CANNED_RESPONSES` against the last user message. Unmatched messages return `"I understand your request."`. The model name reported is always `"mock-model"`. +It also hosts OAuth test endpoints: +- `POST /oauth/exchange` for hosted auth-code exchange +- `POST /oauth/refresh` for hosted refresh-token exchange +- `GET /__mock/oauth/state` and `POST /__mock/oauth/reset` so HTTP E2E scenarios can assert exact proxy payloads and reset counters between setup and refresh assertions + To add a new canned response: ```python # In mock_llm.py diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py index 06c7da03..aa8ba1cb 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -112,6 +112,39 @@ def _reserve_loopback_sockets(count: int) -> list[socket.socket]: sock.close() raise +async def _stop_process( + proc: asyncio.subprocess.Process, *, sig: int | None = None, timeout: float +) -> None: + """Signal a subprocess and wait briefly without masking exit races.""" + if proc.returncode is not None: + return + + try: + if sig is None: + proc.kill() + else: + proc.send_signal(sig) + except ProcessLookupError: + try: + await asyncio.wait_for(proc.wait(), timeout=timeout) + except asyncio.TimeoutError: + pass + return + + try: + await asyncio.wait_for(proc.wait(), timeout=timeout) + except asyncio.TimeoutError: + pass + + +def _forward_coverage_env(env: dict[str, str]) -> None: + """Forward cargo-llvm-cov env vars into child processes when present.""" + cov_env_prefixes = ("CARGO_LLVM_COV", "LLVM_") + cov_env_extras = ("CARGO_ENCODED_RUSTFLAGS", "CARGO_INCREMENTAL") + for key, val in os.environ.items(): + if key.startswith(cov_env_prefixes) or key in cov_env_extras: + env[key] = val + @pytest.fixture(scope="session") def ironclaw_binary(): @@ -264,14 +297,7 @@ async def ironclaw_server( "IRONCLAW_OAUTH_CALLBACK_URL": "https://oauth.test.example/oauth/callback", "IRONCLAW_OAUTH_EXCHANGE_URL": mock_llm_server, } - # Forward LLVM coverage instrumentation env vars when present - # (allows cargo-llvm-cov to collect profraw data from E2E runs). - # Use prefix matching to stay resilient to cargo-llvm-cov changes. - COV_ENV_PREFIXES = ("CARGO_LLVM_COV", "LLVM_") - COV_ENV_EXTRAS = ("CARGO_ENCODED_RUSTFLAGS", "CARGO_INCREMENTAL") - for key, val in os.environ.items(): - if key.startswith(COV_ENV_PREFIXES) or key in COV_ENV_EXTRAS: - env[key] = val + _forward_coverage_env(env) proc = await asyncio.create_subprocess_exec( ironclaw_binary, "--no-onboard", stdin=asyncio.subprocess.DEVNULL, @@ -279,35 +305,145 @@ async def ironclaw_server( stderr=asyncio.subprocess.PIPE, env=env, ) + startup_kill_attempted = False base_url = f"http://127.0.0.1:{gateway_port}" try: await wait_for_ready(f"{base_url}/api/health", timeout=60) yield base_url except TimeoutError: # Dump stderr so CI logs show why the server failed to start + if proc.returncode is None: + startup_kill_attempted = True + await _stop_process(proc, timeout=2) returncode = proc.returncode stderr_bytes = b"" if proc.stderr: try: stderr_bytes = await asyncio.wait_for(proc.stderr.read(8192), timeout=2) - except (asyncio.TimeoutError, Exception): + except asyncio.TimeoutError: pass stderr_text = stderr_bytes.decode("utf-8", errors="replace") - proc.kill() pytest.fail( f"ironclaw server failed to start on port {gateway_port} " f"(returncode={returncode}).\nstderr:\n{stderr_text}" ) finally: if proc.returncode is None: - # Use SIGINT (not SIGTERM) so tokio's ctrl_c handler triggers a - # graceful shutdown. This lets the LLVM coverage runtime run its - # atexit handler and flush .profraw files for cargo-llvm-cov. - proc.send_signal(signal.SIGINT) - try: - await asyncio.wait_for(proc.wait(), timeout=10) - except asyncio.TimeoutError: - proc.kill() + if startup_kill_attempted: + await _stop_process(proc, timeout=2) + else: + # Use SIGINT (not SIGTERM) so tokio's ctrl_c handler triggers a + # graceful shutdown. This lets the LLVM coverage runtime run its + # atexit handler and flush .profraw files for cargo-llvm-cov. + await _stop_process(proc, sig=signal.SIGINT, timeout=10) + if proc.returncode is None: + await _stop_process(proc, timeout=2) + + +@pytest.fixture(scope="session") +async def hosted_oauth_refresh_server( + ironclaw_binary, + mock_llm_server, + wasm_tools_dir, +): + """Start a hosted-mode ironclaw instance for OAuth refresh regression tests.""" + reserved = _reserve_loopback_sockets(2) + db_tmpdir = tempfile.TemporaryDirectory(prefix="ironclaw-e2e-hosted-oauth-db-") + home_tmpdir = tempfile.TemporaryDirectory(prefix="ironclaw-e2e-hosted-oauth-home-") + + try: + gateway_port = reserved[0].getsockname()[1] + http_port = reserved[1].getsockname()[1] + for sock in reserved: + if sock.fileno() != -1: + sock.close() + + db_path = os.path.join(db_tmpdir.name, "hosted-oauth-refresh.db") + home_dir = home_tmpdir.name + env = { + "PATH": os.environ.get("PATH", "/usr/bin:/bin"), + "HOME": home_dir, + "IRONCLAW_BASE_DIR": os.path.join(home_dir, ".ironclaw"), + "RUST_LOG": "ironclaw=info", + "RUST_BACKTRACE": "1", + "IRONCLAW_OWNER_ID": OWNER_SCOPE_ID, + "GATEWAY_ENABLED": "true", + "GATEWAY_HOST": "127.0.0.1", + "GATEWAY_PORT": str(gateway_port), + "GATEWAY_AUTH_TOKEN": AUTH_TOKEN, + "GATEWAY_USER_ID": OWNER_SCOPE_ID, + "HTTP_HOST": "127.0.0.1", + "HTTP_PORT": str(http_port), + "HTTP_WEBHOOK_SECRET": HTTP_WEBHOOK_SECRET, + "CLI_ENABLED": "false", + "LLM_BACKEND": "openai_compatible", + "LLM_BASE_URL": mock_llm_server, + "LLM_MODEL": "mock-model", + "DATABASE_BACKEND": "libsql", + "LIBSQL_PATH": db_path, + "SECRETS_MASTER_KEY": "0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef", + "SANDBOX_ENABLED": "false", + "SKILLS_ENABLED": "true", + "ROUTINES_ENABLED": "true", + "HEARTBEAT_ENABLED": "false", + "EMBEDDING_ENABLED": "false", + "WASM_ENABLED": "true", + "WASM_TOOLS_DIR": wasm_tools_dir, + "WASM_CHANNELS_DIR": _WASM_CHANNELS_TMPDIR.name, + "ONBOARD_COMPLETED": "true", + "IRONCLAW_OAUTH_CALLBACK_URL": "https://oauth.test.example/oauth/callback", + "IRONCLAW_OAUTH_EXCHANGE_URL": mock_llm_server, + "GOOGLE_OAUTH_CLIENT_ID": "hosted-google-client-id", + } + _forward_coverage_env(env) + + proc = await asyncio.create_subprocess_exec( + ironclaw_binary, "--no-onboard", + stdin=asyncio.subprocess.DEVNULL, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + env=env, + ) + startup_kill_attempted = False + base_url = f"http://127.0.0.1:{gateway_port}" + try: + await wait_for_ready(f"{base_url}/api/health", timeout=60) + yield { + "base_url": base_url, + "db_path": db_path, + "gateway_user_id": OWNER_SCOPE_ID, + "mock_llm_url": mock_llm_server, + } + except TimeoutError: + if proc.returncode is None: + startup_kill_attempted = True + await _stop_process(proc, timeout=2) + returncode = proc.returncode + stderr_bytes = b"" + if proc.stderr: + try: + stderr_bytes = await asyncio.wait_for(proc.stderr.read(8192), timeout=2) + except asyncio.TimeoutError: + pass + stderr_text = stderr_bytes.decode("utf-8", errors="replace") + pytest.fail( + f"hosted oauth refresh server failed to start on port {gateway_port} " + f"(returncode={returncode}).\nstderr:\n{stderr_text}" + ) + finally: + if proc.returncode is None: + if startup_kill_attempted: + await _stop_process(proc, timeout=2) + else: + await _stop_process(proc, sig=signal.SIGINT, timeout=10) + if proc.returncode is None: + await _stop_process(proc, timeout=2) + finally: + for sock in reserved: + if sock.fileno() != -1: + sock.close() + db_tmpdir.cleanup() + home_tmpdir.cleanup() @pytest.fixture(scope="session") @@ -362,12 +498,7 @@ async def http_channel_server_without_secret( "IRONCLAW_OAUTH_CALLBACK_URL": "https://oauth.test.example/oauth/callback", "IRONCLAW_OAUTH_EXCHANGE_URL": mock_llm_server, } - # Forward LLVM coverage instrumentation env vars when present - COV_ENV_PREFIXES = ("CARGO_LLVM_COV", "LLVM_") - COV_ENV_EXTRAS = ("CARGO_ENCODED_RUSTFLAGS", "CARGO_INCREMENTAL") - for key, val in os.environ.items(): - if key.startswith(COV_ENV_PREFIXES) or key in COV_ENV_EXTRAS: - env[key] = val + _forward_coverage_env(env) proc = await asyncio.create_subprocess_exec( ironclaw_binary, "--no-onboard", stdin=asyncio.subprocess.DEVNULL, @@ -375,6 +506,7 @@ async def http_channel_server_without_secret( stderr=asyncio.subprocess.PIPE, env=env, ) + startup_kill_attempted = False gateway_url = f"http://127.0.0.1:{gateway_port}" http_base_url = f"http://127.0.0.1:{http_port}" try: @@ -383,15 +515,17 @@ async def http_channel_server_without_secret( yield http_base_url except TimeoutError: # Dump stderr so CI logs show why the server failed to start + if proc.returncode is None: + startup_kill_attempted = True + await _stop_process(proc, timeout=2) returncode = proc.returncode stderr_bytes = b"" if proc.stderr: try: stderr_bytes = await asyncio.wait_for(proc.stderr.read(8192), timeout=2) - except (asyncio.TimeoutError, Exception): + except asyncio.TimeoutError: pass stderr_text = stderr_bytes.decode("utf-8", errors="replace") - proc.kill() pytest.fail( f"ironclaw server without webhook secret failed to start on ports " f"gateway={gateway_port}, http={http_port} " @@ -399,14 +533,15 @@ async def http_channel_server_without_secret( ) finally: if proc.returncode is None: - # Use SIGINT (not SIGTERM) so tokio's ctrl_c handler triggers a - # graceful shutdown. This lets the LLVM coverage runtime run its - # atexit handler and flush .profraw files for cargo-llvm-cov. - proc.send_signal(signal.SIGINT) - try: - await asyncio.wait_for(proc.wait(), timeout=10) - except asyncio.TimeoutError: - proc.kill() + if startup_kill_attempted: + await _stop_process(proc, timeout=2) + else: + # Use SIGINT (not SIGTERM) so tokio's ctrl_c handler triggers a + # graceful shutdown. This lets the LLVM coverage runtime run its + # atexit handler and flush .profraw files for cargo-llvm-cov. + await _stop_process(proc, sig=signal.SIGINT, timeout=10) + if proc.returncode is None: + await _stop_process(proc, timeout=2) @pytest.fixture(scope="session") diff --git a/tests/e2e/mock_llm.py b/tests/e2e/mock_llm.py index 359c22d5..1147662c 100644 --- a/tests/e2e/mock_llm.py +++ b/tests/e2e/mock_llm.py @@ -34,6 +34,15 @@ TOOL_CALL_PATTERNS = [ "body": {"label": m.group("label")}, }, ), + ( + re.compile(r"check gmail unread|gmail unread", re.IGNORECASE), + "gmail", + lambda _: { + "action": "list_messages", + "query": "is:unread", + "max_results": 1, + }, + ), (re.compile(r"what time|current time", re.IGNORECASE), "time", lambda _: {"operation": "now"}), ( re.compile( @@ -91,6 +100,15 @@ TOOL_CALL_PATTERNS = [ ] +def _new_oauth_state() -> dict: + return { + "exchange_count": 0, + "refresh_count": 0, + "last_exchange": None, + "last_refresh": None, + } + + def _last_user_content(messages: list[dict]) -> str: for msg in reversed(messages): if msg.get("role") == "user": @@ -272,6 +290,12 @@ async def oauth_exchange(request: web.Request) -> web.Response: specific token params such as RFC 8707 `resource` are forwarded here. """ data = await request.post() + oauth_state = request.app["oauth_state"] + oauth_state["exchange_count"] += 1 + oauth_state["last_exchange"] = { + "authorization": request.headers.get("Authorization"), + "form": dict(data), + } code = data.get("code", "") access_token_field = data.get("access_token_field", "access_token") @@ -290,6 +314,39 @@ async def oauth_exchange(request: web.Request) -> web.Response: }) +async def oauth_refresh(request: web.Request) -> web.Response: + """Mock OAuth token refresh proxy for hosted refresh E2E tests.""" + data = await request.post() + oauth_state = request.app["oauth_state"] + oauth_state["refresh_count"] += 1 + oauth_state["last_refresh"] = { + "authorization": request.headers.get("Authorization"), + "form": dict(data), + } + + if request.headers.get("Authorization") != "Bearer e2e-test-token": + return web.json_response({"error": "invalid_gateway_auth"}, status=401) + if data.get("client_id") != "hosted-google-client-id": + return web.json_response({"error": "invalid_client_id"}, status=400) + if "client_secret" in data: + return web.json_response({"error": "unexpected_client_secret"}, status=400) + + return web.json_response({ + "access_token": "mock-refreshed-access-token", + "refresh_token": "mock-rotated-refresh-token", + "expires_in": 3600, + }) + + +async def oauth_state_handler(request: web.Request) -> web.Response: + return web.json_response(request.app["oauth_state"]) + + +async def oauth_reset(request: web.Request) -> web.Response: + request.app["oauth_state"] = _new_oauth_state() + return web.json_response({"ok": True}) + + async def models(_request: web.Request) -> web.Response: return web.json_response({ "object": "list", @@ -424,12 +481,16 @@ def main(): parser.add_argument("--port", type=int, default=0) args = parser.parse_args() app = web.Application() + app["oauth_state"] = _new_oauth_state() # Register both /v1/ and non-/v1/ paths (rig-core omits the /v1/ prefix) app.router.add_post("/v1/chat/completions", chat_completions) app.router.add_post("/chat/completions", chat_completions) app.router.add_get("/v1/models", models) app.router.add_get("/models", models) app.router.add_post("/oauth/exchange", oauth_exchange) + app.router.add_post("/oauth/refresh", oauth_refresh) + app.router.add_get("/__mock/oauth/state", oauth_state_handler) + app.router.add_post("/__mock/oauth/reset", oauth_reset) # Mock MCP server endpoints app.router.add_post("/mcp", mcp_endpoint) app.router.add_post("/mcp-400", mcp_endpoint_400) diff --git a/tests/e2e/scenarios/test_oauth_refresh.py b/tests/e2e/scenarios/test_oauth_refresh.py new file mode 100644 index 00000000..50871f7f --- /dev/null +++ b/tests/e2e/scenarios/test_oauth_refresh.py @@ -0,0 +1,227 @@ +"""Hosted OAuth refresh HTTP regression test. + +Runs a real ironclaw binary in hosted mode, expires a stored Gmail access +token in the libSQL database, triggers a real gmail tool call through the +chat API, and verifies that refresh uses the hosted proxy endpoint. +""" + +import asyncio +import sqlite3 +from datetime import datetime, timezone +from urllib.parse import parse_qs, urlparse + +import httpx + +from helpers import api_get, api_post + + +def _extract_state(auth_url: str) -> str: + parsed = urlparse(auth_url) + state = parse_qs(parsed.query).get("state", [None])[0] + assert state, f"auth_url should include state: {auth_url}" + return state + + +def _parse_timestamp(value: str | None) -> datetime | None: + if value is None: + return None + return datetime.fromisoformat(value.replace("Z", "+00:00")) + + +def _expire_access_token(db_path: str, user_id: str, secret_name: str) -> None: + with sqlite3.connect(db_path) as conn: + cursor = conn.execute( + """ + UPDATE secrets + SET expires_at = strftime('%Y-%m-%dT%H:%M:%fZ', 'now', '-1 hour') + WHERE user_id = ?1 AND name = ?2 + """, + (user_id, secret_name), + ) + conn.commit() + assert cursor.rowcount == 1, f"Expected one secret row for {user_id}/{secret_name}" + + +def _find_secret_row( + db_path: str, + secret_name: str, +) -> tuple[str, str | None, str | None]: + with sqlite3.connect(db_path) as conn: + row = conn.execute( + """ + SELECT user_id, expires_at, updated_at + FROM secrets + WHERE name = ?1 + ORDER BY updated_at DESC + LIMIT 1 + """, + (secret_name,), + ).fetchone() + assert row is not None, f"Missing secret row for {secret_name}" + return row[0], row[1], row[2] + + +async def _get_extension(base_url: str, name: str) -> dict | None: + response = await api_get(base_url, "/api/extensions", timeout=15) + response.raise_for_status() + for extension in response.json().get("extensions", []): + if extension["name"] == name: + return extension + return None + + +async def _reset_mock_oauth_state(mock_base_url: str) -> None: + async with httpx.AsyncClient() as client: + response = await client.post(f"{mock_base_url}/__mock/oauth/reset", timeout=10) + response.raise_for_status() + + +async def _get_mock_oauth_state(mock_base_url: str) -> dict: + async with httpx.AsyncClient() as client: + response = await client.get(f"{mock_base_url}/__mock/oauth/state", timeout=10) + response.raise_for_status() + return response.json() + + +async def _approve_pending_request(base_url: str, thread_id: str, request_id: str) -> None: + response = await api_post( + base_url, + "/api/chat/approval", + json={"request_id": request_id, "action": "approve", "thread_id": thread_id}, + timeout=15, + ) + assert response.status_code == 202, ( + f"Approval submission failed: {response.status_code} {response.text[:400]}" + ) + + +async def _wait_for_gmail_tool_call(base_url: str, thread_id: str, timeout: float = 30.0) -> dict: + approved_request_ids = set() + for _ in range(int(timeout * 2)): + response = await api_get( + base_url, + f"/api/chat/history?thread_id={thread_id}", + timeout=15, + ) + response.raise_for_status() + history = response.json() + + pending = history.get("pending_approval") + if pending and pending["request_id"] not in approved_request_ids: + await _approve_pending_request(base_url, thread_id, pending["request_id"]) + approved_request_ids.add(pending["request_id"]) + + for turn in history.get("turns", []): + for tool_call in turn.get("tool_calls", []): + if tool_call.get("name") == "gmail": + return history + + await asyncio.sleep(0.5) + + raise AssertionError(f"Timed out waiting for gmail tool call in thread {thread_id}") + + +async def _wait_for_refresh_request(mock_base_url: str, timeout: float = 20.0) -> dict: + for _ in range(int(timeout * 2)): + state = await _get_mock_oauth_state(mock_base_url) + if state.get("refresh_count") == 1: + return state + await asyncio.sleep(0.5) + raise AssertionError("Timed out waiting for exactly one OAuth refresh request") + + +async def test_hosted_gmail_oauth_refresh_uses_proxy(hosted_oauth_refresh_server): + server = hosted_oauth_refresh_server["base_url"] + db_path = hosted_oauth_refresh_server["db_path"] + mock_base_url = hosted_oauth_refresh_server["mock_llm_url"] + + install_response = await api_post( + server, + "/api/extensions/install", + json={"name": "gmail"}, + timeout=180, + ) + assert install_response.status_code == 200, install_response.text + assert install_response.json().get("success") is True + + setup_response = await api_post( + server, + "/api/extensions/gmail/setup", + json={"secrets": {}}, + timeout=30, + ) + assert setup_response.status_code == 200, setup_response.text + setup_data = setup_response.json() + assert setup_data.get("success") is True, setup_data + auth_url = setup_data.get("auth_url") + assert auth_url, setup_data + auth_params = parse_qs(urlparse(auth_url).query) + assert auth_params.get("client_id") == ["hosted-google-client-id"] + + async with httpx.AsyncClient() as client: + callback_response = await client.get( + f"{server}/oauth/callback", + params={"code": "mock_auth_code", "state": _extract_state(auth_url)}, + timeout=30, + follow_redirects=True, + ) + + assert callback_response.status_code == 200, callback_response.text[:400] + callback_body = callback_response.text.lower() + assert "connected" in callback_body or "success" in callback_body + + gmail = await _get_extension(server, "gmail") + assert gmail is not None, "gmail should be installed" + assert gmail["authenticated"] is True, gmail + assert "gmail" in gmail.get("tools", []), gmail + + await _reset_mock_oauth_state(mock_base_url) + + stored_user_id, expires_before, updated_before = _find_secret_row( + db_path, "google_oauth_token" + ) + assert _parse_timestamp(expires_before) is not None + assert _parse_timestamp(updated_before) is not None + + await asyncio.sleep(0.1) + _expire_access_token(db_path, stored_user_id, "google_oauth_token") + + thread_response = await api_post(server, "/api/chat/thread/new", timeout=15) + assert thread_response.status_code == 200, thread_response.text + thread_id = thread_response.json()["id"] + + send_response = await api_post( + server, + "/api/chat/send", + json={"content": "check gmail unread", "thread_id": thread_id}, + timeout=30, + ) + assert send_response.status_code == 202, send_response.text + + history = await _wait_for_gmail_tool_call(server, thread_id) + assert any( + tool_call.get("name") == "gmail" + for turn in history.get("turns", []) + for tool_call in turn.get("tool_calls", []) + ), history + + oauth_state = await _wait_for_refresh_request(mock_base_url) + assert oauth_state["refresh_count"] == 1, oauth_state + last_refresh = oauth_state["last_refresh"] + assert last_refresh is not None, oauth_state + assert last_refresh["authorization"] == "Bearer e2e-test-token" + assert last_refresh["form"]["client_id"] == "hosted-google-client-id" + assert "client_secret" not in last_refresh["form"], last_refresh + + refreshed_user_id, expires_after, updated_after = _find_secret_row( + db_path, "google_oauth_token" + ) + assert refreshed_user_id == stored_user_id + expires_after_dt = _parse_timestamp(expires_after) + updated_after_dt = _parse_timestamp(updated_after) + updated_before_dt = _parse_timestamp(updated_before) + assert expires_after_dt is not None + assert updated_after_dt is not None + assert updated_before_dt is not None + assert expires_after_dt > datetime.now(timezone.utc) + assert updated_after_dt > updated_before_dt diff --git a/tests/e2e/scenarios/test_telegram_hot_activation.py b/tests/e2e/scenarios/test_telegram_hot_activation.py index 261b837e..fede2be5 100644 --- a/tests/e2e/scenarios/test_telegram_hot_activation.py +++ b/tests/e2e/scenarios/test_telegram_hot_activation.py @@ -253,6 +253,6 @@ async def test_telegram_hot_activation_transitions_installed_to_active(page): assert await card.locator(SEL["ext_pairing_label"]).count() == 0 assert captured_setup_payloads == [ - {"secrets": {"telegram_bot_token": "123456789:ABCdefGhI"}}, - {"secrets": {}}, + {"secrets": {"telegram_bot_token": "123456789:ABCdefGhI"}, "fields": {}}, + {"secrets": {}, "fields": {}}, ] diff --git a/tests/e2e_advanced_traces.rs b/tests/e2e_advanced_traces.rs index b3efc8d9..ce18ad3d 100644 --- a/tests/e2e_advanced_traces.rs +++ b/tests/e2e_advanced_traces.rs @@ -587,6 +587,7 @@ mod advanced { async fn mcp_extension_lifecycle() { use crate::support::mock_mcp_server::{MockToolResponse, start_mock_mcp_server}; use ironclaw::extensions::{AuthHint, ExtensionKind, ExtensionSource, RegistryEntry}; + const TEST_USER_ID: &str = "test-user"; // 1. Start mock MCP server with pre-configured tool responses. let mock_server = start_mock_mcp_server(vec![ @@ -654,14 +655,14 @@ mod advanced { ext_mgr .secrets() .create( - "default", + TEST_USER_ID, ironclaw::secrets::CreateSecretParams::new(secret_name, "mock-access-token") .with_provider("mcp:mock-notion".to_string()), ) .await .expect("failed to inject test token"); - let activate_result = ext_mgr.activate("mock-notion", "default").await; + let activate_result = ext_mgr.activate("mock-notion", TEST_USER_ID).await; assert!( activate_result.is_ok(), "activation failed: {:?}", diff --git a/tests/e2e_builtin_tool_coverage.rs b/tests/e2e_builtin_tool_coverage.rs index 42d7fb75..7c0c7bc7 100644 --- a/tests/e2e_builtin_tool_coverage.rs +++ b/tests/e2e_builtin_tool_coverage.rs @@ -205,7 +205,44 @@ mod tests { } // ----------------------------------------------------------------------- - // Test 5: routine_manual_create_defaults_to_tools_enabled + // Test 5: routine_update_fail_delete_fallback + // ----------------------------------------------------------------------- + + #[tokio::test] + async fn routine_update_fail_delete_fallback() { + let trace = LlmTrace::from_file(concat!( + env!("CARGO_MANIFEST_DIR"), + "/tests/fixtures/llm_traces/tools/routine_update_fail_delete_fallback.json" + )) + .expect("failed to load routine_update_fail_delete_fallback.json"); + + let rig = TestRigBuilder::new() + .with_trace(trace.clone()) + .with_auto_approve_tools(true) + .build() + .await; + + rig.send_message("Try converting a routine trigger, then recover by deleting it") + .await; + let responses = rig.wait_for_responses(1, Duration::from_secs(15)).await; + + rig.verify_trace_expects(&trace, &responses); + + let completed = rig.tool_calls_completed(); + assert!( + completed.iter().any(|(n, ok)| n == "routine_update" && !ok), + "routine_update should fail in this regression path: {completed:?}" + ); + assert!( + completed.iter().any(|(n, ok)| n == "routine_delete" && *ok), + "routine_delete should recover successfully via preserved routine identity: {completed:?}" + ); + + rig.shutdown(); + } + + // ----------------------------------------------------------------------- + // Test 6: routine_manual_create_defaults_to_tools_enabled // ----------------------------------------------------------------------- #[tokio::test] @@ -246,7 +283,7 @@ mod tests { } // ----------------------------------------------------------------------- - // Test 6: routine_manual_create_explicit_no_tools + // Test 7: routine_manual_create_explicit_no_tools // ----------------------------------------------------------------------- #[tokio::test] @@ -287,7 +324,7 @@ mod tests { } // ----------------------------------------------------------------------- - // Test 7: routine_history + // Test 8: routine_history // ----------------------------------------------------------------------- #[tokio::test] @@ -439,7 +476,7 @@ mod tests { match &routine.trigger { Trigger::Cron { schedule, timezone } => { - assert_eq!(schedule, "0 0 9 * * MON-FRI"); + assert_eq!(schedule, "0 0 9 * * MON-FRI *"); assert_eq!(timezone.as_deref(), Some("UTC")); } other => panic!("expected cron trigger, got {other:?}"), diff --git a/tests/e2e_routine_heartbeat.rs b/tests/e2e_routine_heartbeat.rs index 1a3d32d2..462eb82c 100644 --- a/tests/e2e_routine_heartbeat.rs +++ b/tests/e2e_routine_heartbeat.rs @@ -340,14 +340,14 @@ mod tests { SchedulerDeps { tools: registry.clone(), extension_manager: extension_manager.clone(), - store: Some(db.clone()), + store: Some(ironclaw::tenant::AdminScope::new(db.clone())), hooks: Arc::new(HookRegistry::new()), }, )); Arc::new(RoutineEngine::new( RoutineConfig::default(), - db, + ironclaw::tenant::AdminScope::new(db), llm, ws, notify_tx, @@ -451,7 +451,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( RoutineConfig::default(), - db.clone(), + ironclaw::tenant::AdminScope::new(db.clone()), llm, ws, notify_tx, @@ -530,7 +530,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( RoutineConfig::default(), - db.clone(), + ironclaw::tenant::AdminScope::new(db.clone()), llm, ws, notify_tx, @@ -617,7 +617,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( RoutineConfig::default(), - db.clone(), + ironclaw::tenant::AdminScope::new(db.clone()), llm, ws, notify_tx, @@ -726,7 +726,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( RoutineConfig::default(), - db.clone(), + ironclaw::tenant::AdminScope::new(db.clone()), llm, ws, notify_tx, @@ -869,7 +869,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( RoutineConfig::default(), - db.clone(), + ironclaw::tenant::AdminScope::new(db.clone()), llm, ws, notify_tx, @@ -1052,7 +1052,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( RoutineConfig::default(), - Arc::clone(&db), + ironclaw::tenant::AdminScope::new(Arc::clone(&db)), llm, ws, notify_tx, @@ -1174,7 +1174,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( RoutineConfig::default(), - db.clone(), + ironclaw::tenant::AdminScope::new(db.clone()), llm, ws, notify_tx, @@ -1282,7 +1282,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( config, - db.clone(), + ironclaw::tenant::AdminScope::new(db.clone()), llm, ws, notify_tx, diff --git a/tests/e2e_telegram_message_routing.rs b/tests/e2e_telegram_message_routing.rs index ead164eb..810fc218 100644 --- a/tests/e2e_telegram_message_routing.rs +++ b/tests/e2e_telegram_message_routing.rs @@ -201,6 +201,7 @@ mod tests { sandbox_readiness: ironclaw::agent::SandboxReadiness::DisabledByConfig, builder: None, llm_backend: "nearai".to_string(), + tenant_rates: std::sync::Arc::new(ironclaw::tenant::TenantRateRegistry::new(4, 3)), }; let gateway = Arc::new(TestChannel::new()); diff --git a/tests/e2e_workspace_coverage.rs b/tests/e2e_workspace_coverage.rs index 396b676e..68956d30 100644 --- a/tests/e2e_workspace_coverage.rs +++ b/tests/e2e_workspace_coverage.rs @@ -12,6 +12,7 @@ mod tests { use crate::support::test_rig::TestRigBuilder; use crate::support::trace_llm::LlmTrace; + use ironclaw::workspace::Workspace; // ----------------------------------------------------------------------- // Test 1: write_chunk_search @@ -268,6 +269,7 @@ mod tests { #[tokio::test] async fn identity_in_system_prompt() { + const TEST_USER_ID: &str = "test-user"; let trace = LlmTrace::from_file(concat!( env!("CARGO_MANIFEST_DIR"), "/tests/fixtures/llm_traces/workspace/identity_prompt.json" @@ -280,7 +282,7 @@ mod tests { .await; // Seed an IDENTITY.md so the system prompt has real content to inject. - let ws = rig.workspace().expect("workspace must be available"); + let ws = Workspace::new_with_db(TEST_USER_ID, rig.database().clone()); ws.write( "IDENTITY.md", "I am TestBot, a helpful testing assistant created for E2E verification.", diff --git a/tests/fixtures/llm_traces/tools/routine_update_fail_delete_fallback.json b/tests/fixtures/llm_traces/tools/routine_update_fail_delete_fallback.json new file mode 100644 index 00000000..5c76dbb5 --- /dev/null +++ b/tests/fixtures/llm_traces/tools/routine_update_fail_delete_fallback.json @@ -0,0 +1,70 @@ +{ + "model_name": "test-routine-update-fail-delete-fallback", + "expects": { + "tools_used": ["routine_create", "routine_update", "routine_delete"], + "tool_results_contain": { + "routine_update": "Cannot update schedule or timezone on a non-cron routine.", + "routine_delete": "temp-routine" + }, + "min_responses": 1 + }, + "steps": [ + { + "response": { + "type": "tool_calls", + "tool_calls": [ + { + "id": "call_rc_fallback", + "name": "routine_create", + "arguments": { + "name": "temp-routine", + "trigger_type": "manual", + "prompt": "Temporary routine for fallback test." + } + } + ], + "input_tokens": 120, + "output_tokens": 40 + } + }, + { + "response": { + "type": "tool_calls", + "tool_calls": [ + { + "id": "call_ru_fallback", + "name": "routine_update", + "arguments": { + "name": "temp-routine", + "schedule": "0 */10 * * * *" + } + } + ], + "input_tokens": 200, + "output_tokens": 30 + } + }, + { + "response": { + "type": "tool_calls", + "tool_calls": [ + { + "id": "call_rd_fallback", + "name": "routine_delete", + "arguments": {} + } + ], + "input_tokens": 300, + "output_tokens": 20 + } + }, + { + "response": { + "type": "text", + "content": "I recovered from the failed update and cleaned up the original routine.", + "input_tokens": 380, + "output_tokens": 25 + } + } + ] +} diff --git a/tests/multi_tenant_integration.rs b/tests/multi_tenant_integration.rs index 02eb60e8..227fa721 100644 --- a/tests/multi_tenant_integration.rs +++ b/tests/multi_tenant_integration.rs @@ -19,10 +19,13 @@ use axum::middleware; use axum::routing::{get, post}; use tower::ServiceExt; +use ironclaw::channels::IncomingMessage; use ironclaw::channels::web::auth::{ AuthenticatedUser, MultiAuthState, UserIdentity, auth_middleware, }; -use ironclaw::channels::web::server::{GatewayState, PerUserRateLimiter, RateLimiter}; +use ironclaw::channels::web::server::{ + GatewayState, PerUserRateLimiter, RateLimiter, start_server, +}; use ironclaw::channels::web::sse::SseManager; use ironclaw::channels::web::test_helpers::TestGatewayBuilder; use ironclaw::channels::web::ws::WsConnectionTracker; @@ -37,6 +40,9 @@ const ALICE_TOKEN: &str = "tok-alice-secret"; const BOB_TOKEN: &str = "tok-bob-secret"; const ALICE_USER_ID: &str = "alice"; const BOB_USER_ID: &str = "bob"; +const OWNER_TOKEN: &str = "tok-owner-secret"; +const OWNER_SCOPE_ID: &str = "owner-scope"; +const GATEWAY_SENDER_ID: &str = "gateway-sender"; /// Build a MultiAuthState with two users. fn two_user_auth() -> MultiAuthState { @@ -301,7 +307,7 @@ fn per_user_rate_limiter_single_user_mode() { #[tokio::test] async fn sse_scoped_event_only_delivered_to_target_user() { - use ironclaw::channels::web::types::SseEvent; + use ironclaw_common::AppEvent; use tokio_stream::StreamExt; let manager = SseManager::new(); @@ -319,34 +325,34 @@ async fn sse_scoped_event_only_delivered_to_target_user() { // Send event scoped to alice manager.broadcast_for_user( ALICE_USER_ID, - SseEvent::Status { + AppEvent::Status { message: "alice's event".to_string(), thread_id: None, }, ); // Send global heartbeat (both should get it) - manager.broadcast(SseEvent::Heartbeat); + manager.broadcast(AppEvent::Heartbeat); // Alice gets her scoped event first let e = alice_stream.next().await.unwrap(); match &e { - SseEvent::Status { message, .. } => assert_eq!(message, "alice's event"), + AppEvent::Status { message, .. } => assert_eq!(message, "alice's event"), _ => panic!("Expected Status, got {:?}", e), } // Alice also gets heartbeat let e = alice_stream.next().await.unwrap(); - assert!(matches!(e, SseEvent::Heartbeat)); + assert!(matches!(e, AppEvent::Heartbeat)); // Bob only gets the heartbeat (alice's event was filtered) let e = bob_stream.next().await.unwrap(); - assert!(matches!(e, SseEvent::Heartbeat)); + assert!(matches!(e, AppEvent::Heartbeat)); } #[tokio::test] async fn sse_global_event_delivered_to_all_users() { - use ironclaw::channels::web::types::SseEvent; + use ironclaw_common::AppEvent; use tokio_stream::StreamExt; let manager = SseManager::new(); @@ -361,7 +367,7 @@ async fn sse_global_event_delivered_to_all_users() { .expect("subscribe"), ); - manager.broadcast(SseEvent::Status { + manager.broadcast(AppEvent::Status { message: "global announcement".to_string(), thread_id: None, }); @@ -369,7 +375,7 @@ async fn sse_global_event_delivered_to_all_users() { let ea = alice.next().await.unwrap(); let eb = bob.next().await.unwrap(); match (&ea, &eb) { - (SseEvent::Status { message: a, .. }, SseEvent::Status { message: b, .. }) => { + (AppEvent::Status { message: a, .. }, AppEvent::Status { message: b, .. }) => { assert_eq!(a, "global announcement"); assert_eq!(b, "global announcement"); } @@ -379,7 +385,7 @@ async fn sse_global_event_delivered_to_all_users() { #[tokio::test] async fn sse_user_b_event_not_visible_to_user_a() { - use ironclaw::channels::web::types::SseEvent; + use ironclaw_common::AppEvent; use tokio_stream::StreamExt; let manager = SseManager::new(); @@ -392,19 +398,19 @@ async fn sse_user_b_event_not_visible_to_user_a() { // Send event for bob only manager.broadcast_for_user( BOB_USER_ID, - SseEvent::Response { + AppEvent::Response { content: "bob's secret".to_string(), thread_id: "t1".to_string(), }, ); // Send heartbeat so alice has something to receive - manager.broadcast(SseEvent::Heartbeat); + manager.broadcast(AppEvent::Heartbeat); // Alice should only get heartbeat, not bob's response let e = alice.next().await.unwrap(); assert!( - matches!(e, SseEvent::Heartbeat), + matches!(e, AppEvent::Heartbeat), "Expected Heartbeat, got {:?}", e ); @@ -412,7 +418,7 @@ async fn sse_user_b_event_not_visible_to_user_a() { #[tokio::test] async fn sse_unscoped_subscriber_receives_all_events() { - use ironclaw::channels::web::types::SseEvent; + use ironclaw_common::AppEvent; use tokio_stream::StreamExt; let manager = SseManager::new(); @@ -421,19 +427,19 @@ async fn sse_unscoped_subscriber_receives_all_events() { manager.broadcast_for_user( ALICE_USER_ID, - SseEvent::Status { + AppEvent::Status { message: "alice only".to_string(), thread_id: None, }, ); manager.broadcast_for_user( BOB_USER_ID, - SseEvent::Status { + AppEvent::Status { message: "bob only".to_string(), thread_id: None, }, ); - manager.broadcast(SseEvent::Heartbeat); + manager.broadcast(AppEvent::Heartbeat); // Unscoped subscriber gets ALL three events let e1 = stream.next().await.unwrap(); @@ -441,14 +447,14 @@ async fn sse_unscoped_subscriber_receives_all_events() { let e3 = stream.next().await.unwrap(); match &e1 { - SseEvent::Status { message, .. } => assert_eq!(message, "alice only"), + AppEvent::Status { message, .. } => assert_eq!(message, "alice only"), _ => panic!("Expected alice's Status"), } match &e2 { - SseEvent::Status { message, .. } => assert_eq!(message, "bob only"), + AppEvent::Status { message, .. } => assert_eq!(message, "bob only"), _ => panic!("Expected bob's Status"), } - assert!(matches!(e3, SseEvent::Heartbeat)); + assert!(matches!(e3, AppEvent::Heartbeat)); } // =========================================================================== @@ -537,7 +543,8 @@ fn gateway_state_has_multi_tenant_fields() { job_manager: None, prompt_queue: None, scheduler: None, - default_user_id: "fallback".to_string(), // Multi-tenant: renamed from user_id + owner_id: "fallback".to_string(), + default_sender_id: "fallback".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: None, @@ -553,7 +560,8 @@ fn gateway_state_has_multi_tenant_fields() { active_config: Default::default(), }; - assert_eq!(state.default_user_id, "fallback"); + assert_eq!(state.owner_id, "fallback"); + assert_eq!(state.default_sender_id, "fallback"); assert!(state.workspace_pool.is_none()); } @@ -572,6 +580,69 @@ async fn start_multi_user_server() -> (SocketAddr, Arc) { .expect("Failed to start multi-user test server") } +async fn start_owner_scoped_sender_server() -> ( + SocketAddr, + Arc, + tokio::sync::mpsc::Receiver, +) { + let (agent_tx, agent_rx) = tokio::sync::mpsc::channel(64); + + let mut tokens = HashMap::new(); + tokens.insert( + OWNER_TOKEN.to_string(), + UserIdentity { + user_id: OWNER_SCOPE_ID.to_string(), + workspace_read_scopes: Vec::new(), + }, + ); + tokens.insert( + BOB_TOKEN.to_string(), + UserIdentity { + user_id: BOB_USER_ID.to_string(), + workspace_read_scopes: Vec::new(), + }, + ); + + let state = Arc::new(GatewayState { + msg_tx: tokio::sync::RwLock::new(Some(agent_tx)), + sse: Arc::new(SseManager::new()), + workspace: None, + workspace_pool: None, + session_manager: None, + log_broadcaster: None, + log_level_handle: None, + extension_manager: None, + tool_registry: None, + store: None, + job_manager: None, + prompt_queue: None, + scheduler: None, + owner_id: OWNER_SCOPE_ID.to_string(), + default_sender_id: GATEWAY_SENDER_ID.to_string(), + shutdown_tx: tokio::sync::RwLock::new(None), + ws_tracker: Some(Arc::new(WsConnectionTracker::new())), + llm_provider: None, + skill_registry: None, + skill_catalog: None, + chat_rate_limiter: PerUserRateLimiter::new(30, 60), + oauth_rate_limiter: RateLimiter::new(10, 60), + webhook_rate_limiter: RateLimiter::new(10, 60), + registry_entries: Vec::new(), + cost_guard: None, + routine_engine: Arc::new(tokio::sync::RwLock::new(None)), + startup_time: std::time::Instant::now(), + active_config: Default::default(), + }); + + let auth = MultiAuthState::multi(tokens); + let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); + let bound = start_server(addr, state.clone(), auth) + .await + .expect("Failed to start owner-scoped sender test server"); + + (bound, state, agent_rx) +} + #[tokio::test] async fn full_server_alice_can_access_protected_endpoint() { let (addr, _state) = start_multi_user_server().await; @@ -677,6 +748,49 @@ async fn full_server_chat_send_accepted_for_alice() { assert_eq!(msg.channel, "gateway"); } +#[tokio::test] +async fn full_server_chat_send_rewrites_sender_only_for_owner_scope_rebind() { + let (addr, _state, mut agent_rx) = start_owner_scoped_sender_server().await; + + let client = reqwest::Client::new(); + + let owner_resp = client + .post(format!("http://{}/api/chat/send", addr)) + .header("Authorization", format!("Bearer {}", OWNER_TOKEN)) + .header("Content-Type", "application/json") + .body(r#"{"content":"hello from owner"}"#) + .send() + .await + .unwrap(); + assert_eq!(owner_resp.status(), 202); + + let owner_msg = tokio::time::timeout(Duration::from_secs(2), agent_rx.recv()) + .await + .expect("Timed out waiting for owner message") + .expect("Agent channel closed"); + assert_eq!(owner_msg.user_id, OWNER_SCOPE_ID); + assert_eq!(owner_msg.sender_id, GATEWAY_SENDER_ID); + assert_eq!(owner_msg.content, "hello from owner"); + + let other_resp = client + .post(format!("http://{}/api/chat/send", addr)) + .header("Authorization", format!("Bearer {}", BOB_TOKEN)) + .header("Content-Type", "application/json") + .body(r#"{"content":"hello from bob"}"#) + .send() + .await + .unwrap(); + assert_eq!(other_resp.status(), 202); + + let other_msg = tokio::time::timeout(Duration::from_secs(2), agent_rx.recv()) + .await + .expect("Timed out waiting for non-owner message") + .expect("Agent channel closed"); + assert_eq!(other_msg.user_id, BOB_USER_ID); + assert_eq!(other_msg.sender_id, BOB_USER_ID); + assert_eq!(other_msg.content, "hello from bob"); +} + #[tokio::test] async fn full_server_chat_send_rejected_without_auth() { let (addr, _state) = start_multi_user_server().await; @@ -767,7 +881,7 @@ async fn full_server_jobs_endpoint_rejected_without_auth() { #[tokio::test] async fn full_server_ws_multi_user_event_isolation() { use futures::StreamExt; - use ironclaw::channels::web::types::SseEvent; + use ironclaw_common::AppEvent; use tokio_tungstenite::tungstenite::Message; use tokio_tungstenite::tungstenite::client::IntoClientRequest; @@ -800,14 +914,14 @@ async fn full_server_ws_multi_user_event_isolation() { // Broadcast an event scoped to Alice only state.sse.broadcast_for_user( ALICE_USER_ID, - SseEvent::Status { + AppEvent::Status { message: "alice-only-event".to_string(), thread_id: None, }, ); // Broadcast a global heartbeat so Bob has something to receive - state.sse.broadcast(SseEvent::Heartbeat); + state.sse.broadcast(AppEvent::Heartbeat); // Alice should get her scoped event let alice_msg = tokio::time::timeout(Duration::from_secs(2), alice_ws.next()) @@ -888,7 +1002,8 @@ async fn start_multi_user_server_with_db() -> ( job_manager: None, prompt_queue: None, scheduler: None, - default_user_id: ALICE_USER_ID.to_string(), + owner_id: ALICE_USER_ID.to_string(), + default_sender_id: ALICE_USER_ID.to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: None, diff --git a/tests/multi_tenant_system_prompt.rs b/tests/multi_tenant_system_prompt.rs index ece794bf..b89e6cb5 100644 --- a/tests/multi_tenant_system_prompt.rs +++ b/tests/multi_tenant_system_prompt.rs @@ -1,10 +1,10 @@ -//! Tests proving that multi-tenant system prompts are broken. +//! Regression tests for multi-tenant system prompts. //! -//! Bug: In multi-tenant mode, the agent loop uses `self.workspace()` which -//! returns a single shared workspace (user_id="default"). Identity files -//! (IDENTITY.md, SOUL.md, USER.md) seeded under per-user IDs ("alice", -//! "bob") are invisible to this workspace, so the system prompt is -//! empty/wrong. +//! The agent must build the conversational system prompt from a workspace +//! scoped to the incoming message's user, not from the shared owner-scope +//! workspace created at startup. Otherwise per-user identity files +//! (IDENTITY.md, SOUL.md, USER.md) become invisible and different users can +//! see the same owner-scoped prompt. //! //! These tests: //! 1. Seed identity files for two users (alice, bob) in the database @@ -13,7 +13,7 @@ //! correct user's identity //! 4. Verify user A's identity doesn't leak into user B's prompt //! -//! All tests are expected to FAIL until the bug is fixed. +//! These tests ensure each user's identity is isolated correctly. #[cfg(feature = "libsql")] mod support; diff --git a/tests/openai_compat_integration.rs b/tests/openai_compat_integration.rs index 16568246..b677e57f 100644 --- a/tests/openai_compat_integration.rs +++ b/tests/openai_compat_integration.rs @@ -94,6 +94,7 @@ impl LlmProvider for MockLlmProvider { id: "call_mock_001".to_string(), name: tool.name.clone(), arguments: serde_json::json!({"test": true}), + reasoning: None, }], input_tokens: 15, output_tokens: 8, @@ -203,7 +204,8 @@ async fn start_test_server_with_provider( job_manager: None, prompt_queue: None, scheduler: None, - default_user_id: "test-user".to_string(), + owner_id: "test-user".to_string(), + default_sender_id: "test-user".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: Some(llm_provider), @@ -701,7 +703,8 @@ async fn test_no_llm_provider_returns_503() { job_manager: None, prompt_queue: None, scheduler: None, - default_user_id: "test-user".to_string(), + owner_id: "test-user".to_string(), + default_sender_id: "test-user".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: None, // No LLM! diff --git a/tests/support/gateway_workflow_harness.rs b/tests/support/gateway_workflow_harness.rs index e4620f70..ac35b160 100644 --- a/tests/support/gateway_workflow_harness.rs +++ b/tests/support/gateway_workflow_harness.rs @@ -226,7 +226,8 @@ impl GatewayWorkflowHarness { job_manager: None, prompt_queue: None, scheduler: Some(scheduler_slot.clone()), - default_user_id: user_id.clone(), + owner_id: user_id.clone(), + default_sender_id: user_id.clone(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: Some(Arc::clone(&components.llm)), @@ -265,6 +266,7 @@ impl GatewayWorkflowHarness { sandbox_readiness: ironclaw::agent::SandboxReadiness::DisabledByConfig, builder: None, llm_backend: "nearai".to_string(), + tenant_rates: std::sync::Arc::new(ironclaw::tenant::TenantRateRegistry::new(4, 3)), }, channels, None, diff --git a/tests/support/test_rig.rs b/tests/support/test_rig.rs index 624bb054..5775b86d 100644 --- a/tests/support/test_rig.rs +++ b/tests/support/test_rig.rs @@ -642,7 +642,7 @@ impl TestRigBuilder { let (notify_tx, _notify_rx) = tokio::sync::mpsc::channel(16); let engine = Arc::new(RoutineEngine::new( routine_config, - Arc::clone(db_arc), + ironclaw::tenant::AdminScope::new(Arc::clone(db_arc)), components.llm.clone(), Arc::clone(ws), notify_tx, @@ -762,6 +762,7 @@ impl TestRigBuilder { sandbox_readiness: ironclaw::agent::SandboxReadiness::Available, // tests don't use real Docker builder: None, llm_backend: "nearai".to_string(), + tenant_rates: std::sync::Arc::new(ironclaw::tenant::TenantRateRegistry::new(4, 3)), }; // 7. Create TestChannel and ChannelManager. diff --git a/tests/support/trace_llm.rs b/tests/support/trace_llm.rs index e33caf6b..239cfdb5 100644 --- a/tests/support/trace_llm.rs +++ b/tests/support/trace_llm.rs @@ -566,6 +566,7 @@ impl LlmProvider for TraceLlm { id: tc.id, name: tc.name, arguments: tc.arguments, + reasoning: None, }) .collect(); Ok(ToolCompletionResponse { diff --git a/tests/ws_gateway_integration.rs b/tests/ws_gateway_integration.rs index 43277389..0ec5c929 100644 --- a/tests/ws_gateway_integration.rs +++ b/tests/ws_gateway_integration.rs @@ -5,7 +5,7 @@ //! - WebSocket upgrade with auth //! - Ping/pong //! - Client message → agent msg_tx -//! - Broadcast SSE event → WebSocket client +//! - Broadcast AppEvent → WebSocket client //! - Connection tracking (counter increment/decrement) //! - Gateway status endpoint @@ -22,8 +22,8 @@ use tokio_tungstenite::tungstenite::client::IntoClientRequest; use ironclaw::channels::IncomingMessage; use ironclaw::channels::web::server::{GatewayState, start_server}; use ironclaw::channels::web::sse::SseManager; -use ironclaw::channels::web::types::SseEvent; use ironclaw::channels::web::ws::WsConnectionTracker; +use ironclaw_common::AppEvent; const AUTH_TOKEN: &str = "test-token-12345"; const TIMEOUT: Duration = Duration::from_secs(5); @@ -51,7 +51,8 @@ async fn start_test_server() -> ( job_manager: None, prompt_queue: None, scheduler: None, - default_user_id: "test-user".to_string(), + owner_id: "test-user".to_string(), + default_sender_id: "test-user".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: None, @@ -163,8 +164,8 @@ async fn test_ws_broadcast_event_received() { // Give the connection a moment to fully establish tokio::time::sleep(Duration::from_millis(50)).await; - // Broadcast an SSE event (simulates agent sending a response) - state.sse.broadcast(SseEvent::Response { + // Broadcast an event (simulates agent sending a response) + state.sse.broadcast(AppEvent::Response { content: "agent says hi".to_string(), thread_id: "t1".to_string(), }); @@ -185,7 +186,7 @@ async fn test_ws_thinking_event() { let mut ws = connect_ws(addr).await; tokio::time::sleep(Duration::from_millis(50)).await; - state.sse.broadcast(SseEvent::Thinking { + state.sse.broadcast(AppEvent::Thinking { message: "analyzing...".to_string(), thread_id: None, }); @@ -310,22 +311,22 @@ async fn test_ws_multiple_events_in_sequence() { tokio::time::sleep(Duration::from_millis(50)).await; // Broadcast multiple events rapidly - state.sse.broadcast(SseEvent::Thinking { + state.sse.broadcast(AppEvent::Thinking { message: "step 1".to_string(), thread_id: None, }); - state.sse.broadcast(SseEvent::ToolStarted { + state.sse.broadcast(AppEvent::ToolStarted { name: "shell".to_string(), thread_id: None, }); - state.sse.broadcast(SseEvent::ToolCompleted { + state.sse.broadcast(AppEvent::ToolCompleted { name: "shell".to_string(), success: true, error: None, parameters: None, thread_id: None, }); - state.sse.broadcast(SseEvent::Response { + state.sse.broadcast(AppEvent::Response { content: "done".to_string(), thread_id: "t1".to_string(), });