Compare commits

..
Author SHA1 Message Date
Claude 47f80ddc22 style: fix formatting
https://claude.ai/code/session_01GG37cPSH8vfi9nriukuorf
2026-03-23 13:12:21 +00:00
ZakiandClaude Opus 4.6 d2098f1030 test(agent): strengthen regression test for missing thread error handling
Replace shallow assertion-only test with one that exercises the actual
match-based error detection pattern used in process_approval()'s
rejection and state-setting paths.

Addresses Gemini review feedback on #1579.

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-23 05:56:20 -07:00
ZakiandClaude Opus 4.6 79a2c5d9dd fix(agent): return errors when approval thread disappears (#1487)
Replace silent if-let-Some patterns with explicit match arms that log
errors and return error responses when threads are not found during
approval processing. Critical state mutations (complete turn, clear
approval, set Processing, await approval) return errors. Auxiliary
operations (record tool result) log errors but continue since the tool
already executed.

Closes #1487

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-22 23:26:59 -07:00
176 changed files with 2905 additions and 19366 deletions
+1 -1
View File
@@ -54,7 +54,7 @@ jobs:
- group: features - group: features
files: "tests/e2e/scenarios/test_skills.py tests/e2e/scenarios/test_tool_approval.py tests/e2e/scenarios/test_webhook.py" files: "tests/e2e/scenarios/test_skills.py tests/e2e/scenarios/test_tool_approval.py tests/e2e/scenarios/test_webhook.py"
- group: extensions - group: extensions
files: "tests/e2e/scenarios/test_extensions.py tests/e2e/scenarios/test_extension_oauth.py tests/e2e/scenarios/test_oauth_url_parameters.py tests/e2e/scenarios/test_telegram_token_validation.py tests/e2e/scenarios/test_telegram_hot_activation.py tests/e2e/scenarios/test_wasm_lifecycle.py tests/e2e/scenarios/test_tool_execution.py tests/e2e/scenarios/test_pairing.py tests/e2e/scenarios/test_mcp_auth_flow.py tests/e2e/scenarios/test_oauth_credential_fallback.py tests/e2e/scenarios/test_routine_oauth_credential_injection.py" files: "tests/e2e/scenarios/test_extensions.py tests/e2e/scenarios/test_extension_oauth.py tests/e2e/scenarios/test_telegram_token_validation.py tests/e2e/scenarios/test_telegram_hot_activation.py tests/e2e/scenarios/test_wasm_lifecycle.py tests/e2e/scenarios/test_tool_execution.py tests/e2e/scenarios/test_pairing.py tests/e2e/scenarios/test_mcp_auth_flow.py tests/e2e/scenarios/test_oauth_credential_fallback.py tests/e2e/scenarios/test_routine_oauth_credential_injection.py"
- group: routines - group: routines
files: "tests/e2e/scenarios/test_owner_scope.py tests/e2e/scenarios/test_routine_event_batch.py" files: "tests/e2e/scenarios/test_owner_scope.py tests/e2e/scenarios/test_routine_event_batch.py"
steps: steps:
+5 -19
View File
@@ -12,7 +12,6 @@ jobs:
tests: tests:
name: Tests (${{ matrix.name }}) name: Tests (${{ matrix.name }})
runs-on: ubuntu-latest runs-on: ubuntu-latest
timeout-minutes: 45
strategy: strategy:
fail-fast: false fail-fast: false
matrix: matrix:
@@ -41,14 +40,11 @@ jobs:
- name: Build WASM channels (for integration tests) - name: Build WASM channels (for integration tests)
run: ./scripts/build-wasm-extensions.sh --channels run: ./scripts/build-wasm-extensions.sh --channels
- name: Run Tests - name: Run Tests
run: | run: cargo test ${{ matrix.flags }} -- --nocapture
timeout --signal=INT --kill-after=30s 40m \
cargo test ${{ matrix.flags }} -- --nocapture
heavy-integration-tests: heavy-integration-tests:
name: Heavy Integration Tests name: Heavy Integration Tests
runs-on: ubuntu-latest runs-on: ubuntu-latest
timeout-minutes: 20
steps: steps:
- name: Checkout repository - name: Checkout repository
uses: actions/checkout@v6 uses: actions/checkout@v6
@@ -62,13 +58,9 @@ jobs:
- name: Build Telegram WASM channel - name: Build Telegram WASM channel
run: cargo build --manifest-path channels-src/telegram/Cargo.toml --target wasm32-wasip2 --release run: cargo build --manifest-path channels-src/telegram/Cargo.toml --target wasm32-wasip2 --release
- name: Run thread scheduling integration tests - name: Run thread scheduling integration tests
run: | run: cargo test --no-default-features --features libsql,integration --test e2e_thread_scheduling -- --nocapture
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 - name: Run Telegram thread-scope regression test
run: | run: cargo test --features integration --test telegram_auth_integration test_private_messages_use_chat_id_as_thread_scope -- --exact
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: telegram-tests:
name: Telegram Channel Tests name: Telegram Channel Tests
@@ -76,7 +68,6 @@ jobs:
github.event_name != 'pull_request' || github.event_name != 'pull_request' ||
github.base_ref != 'staging' github.base_ref != 'staging'
runs-on: ubuntu-latest runs-on: ubuntu-latest
timeout-minutes: 15
steps: steps:
- name: Checkout repository - name: Checkout repository
uses: actions/checkout@v6 uses: actions/checkout@v6
@@ -84,9 +75,7 @@ jobs:
uses: dtolnay/rust-toolchain@stable uses: dtolnay/rust-toolchain@stable
- uses: Swatinem/rust-cache@v2 - uses: Swatinem/rust-cache@v2
- name: Run Telegram Channel Tests - name: Run Telegram Channel Tests
run: | run: cargo test --manifest-path channels-src/telegram/Cargo.toml -- --nocapture
timeout --signal=INT --kill-after=30s 10m \
cargo test --manifest-path channels-src/telegram/Cargo.toml -- --nocapture
windows-build: windows-build:
name: Windows Build (${{ matrix.name }}) name: Windows Build (${{ matrix.name }})
@@ -121,7 +110,6 @@ jobs:
github.event_name != 'pull_request' || github.event_name != 'pull_request' ||
github.base_ref != 'staging' github.base_ref != 'staging'
runs-on: ubuntu-latest runs-on: ubuntu-latest
timeout-minutes: 30
steps: steps:
- name: Checkout repository - name: Checkout repository
uses: actions/checkout@v6 uses: actions/checkout@v6
@@ -137,9 +125,7 @@ jobs:
- name: Build all WASM extensions against current WIT - name: Build all WASM extensions against current WIT
run: ./scripts/build-wasm-extensions.sh run: ./scripts/build-wasm-extensions.sh
- name: Instantiation test (host linker compatibility) - name: Instantiation test (host linker compatibility)
run: | run: cargo test --all-features wit_compat -- --nocapture
timeout --signal=INT --kill-after=30s 20m \
cargo test --all-features wit_compat -- --nocapture
bench-compile: bench-compile:
name: Benchmark Compilation name: Benchmark Compilation
-132
View File
@@ -7,138 +7,6 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
## [Unreleased] ## [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 ## [0.19.0](https://github.com/nearai/ironclaw/compare/v0.18.0...v0.19.0) - 2026-03-17
### Added ### Added
Generated
+14 -25
View File
@@ -157,7 +157,7 @@ version = "1.1.5"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "40c48f72fd53cd289104fc64099abca73db4166ad86ea0b4341abe65af83dadc" checksum = "40c48f72fd53cd289104fc64099abca73db4166ad86ea0b4341abe65af83dadc"
dependencies = [ dependencies = [
"windows-sys 0.60.2", "windows-sys 0.61.2",
] ]
[[package]] [[package]]
@@ -168,7 +168,7 @@ checksum = "291e6a250ff86cd4a820112fb8898808a366d8f9f58ce16d1f538353ad55747d"
dependencies = [ dependencies = [
"anstyle", "anstyle",
"once_cell_polyfill", "once_cell_polyfill",
"windows-sys 0.60.2", "windows-sys 0.61.2",
] ]
[[package]] [[package]]
@@ -2136,7 +2136,7 @@ dependencies = [
"libc", "libc",
"option-ext", "option-ext",
"redox_users 0.5.2", "redox_users 0.5.2",
"windows-sys 0.59.0", "windows-sys 0.61.2",
] ]
[[package]] [[package]]
@@ -2323,7 +2323,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb"
dependencies = [ dependencies = [
"libc", "libc",
"windows-sys 0.59.0", "windows-sys 0.61.2",
] ]
[[package]] [[package]]
@@ -3390,7 +3390,7 @@ dependencies = [
[[package]] [[package]]
name = "ironclaw" name = "ironclaw"
version = "0.22.0" version = "0.19.0"
dependencies = [ dependencies = [
"aes-gcm", "aes-gcm",
"aho-corasick", "aho-corasick",
@@ -3428,7 +3428,6 @@ dependencies = [
"hyper-util", "hyper-util",
"iana-time-zone", "iana-time-zone",
"insta", "insta",
"ironclaw_common",
"ironclaw_safety", "ironclaw_safety",
"json5", "json5",
"libsql", "libsql",
@@ -3482,22 +3481,13 @@ dependencies = [
"wasmparser 0.220.1", "wasmparser 0.220.1",
"wasmtime", "wasmtime",
"wasmtime-wasi", "wasmtime-wasi",
"webpki-roots 0.26.11",
"zbus", "zbus",
"zip", "zip",
] ]
[[package]]
name = "ironclaw_common"
version = "0.1.0"
dependencies = [
"serde",
"serde_json",
]
[[package]] [[package]]
name = "ironclaw_safety" name = "ironclaw_safety"
version = "0.2.0" version = "0.1.0"
dependencies = [ dependencies = [
"aho-corasick", "aho-corasick",
"regex", "regex",
@@ -4144,7 +4134,7 @@ version = "0.50.3"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5" checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5"
dependencies = [ dependencies = [
"windows-sys 0.59.0", "windows-sys 0.61.2",
] ]
[[package]] [[package]]
@@ -5482,7 +5472,7 @@ dependencies = [
"errno", "errno",
"libc", "libc",
"linux-raw-sys 0.12.1", "linux-raw-sys 0.12.1",
"windows-sys 0.59.0", "windows-sys 0.61.2",
] ]
[[package]] [[package]]
@@ -6164,7 +6154,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3a766e1110788c36f4fa1c2b71b387a7815aa65f88ce0229841826633d93723e" checksum = "3a766e1110788c36f4fa1c2b71b387a7815aa65f88ce0229841826633d93723e"
dependencies = [ dependencies = [
"libc", "libc",
"windows-sys 0.60.2", "windows-sys 0.61.2",
] ]
[[package]] [[package]]
@@ -6364,9 +6354,9 @@ checksum = "55937e1799185b12863d447f42597ed69d9928686b8d88a1df17376a097d8369"
[[package]] [[package]]
name = "tar" name = "tar"
version = "0.4.45" version = "0.4.44"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "22692a6476a21fa75fdfc11d452fda482af402c008cdbaf3476414e122040973" checksum = "1d863878d212c87a19c1a610eb53bb01fe12951c0501cf5a0d65f724914a667a"
dependencies = [ dependencies = [
"filetime", "filetime",
"libc", "libc",
@@ -6389,7 +6379,7 @@ dependencies = [
"getrandom 0.4.2", "getrandom 0.4.2",
"once_cell", "once_cell",
"rustix 1.1.4", "rustix 1.1.4",
"windows-sys 0.59.0", "windows-sys 0.61.2",
] ]
[[package]] [[package]]
@@ -6992,7 +6982,6 @@ dependencies = [
"futures-util", "futures-util",
"http 1.4.0", "http 1.4.0",
"http-body 1.0.1", "http-body 1.0.1",
"http-body-util",
"iri-string", "iri-string",
"pin-project-lite", "pin-project-lite",
"tower 0.5.3", "tower 0.5.3",
@@ -7190,7 +7179,7 @@ checksum = "f2f6fb2847f6742cd76af783a2a2c49e9375d0a111c7bef6f71cd9e738c72d6e"
dependencies = [ dependencies = [
"memoffset", "memoffset",
"tempfile", "tempfile",
"windows-sys 0.60.2", "windows-sys 0.61.2",
] ]
[[package]] [[package]]
@@ -8040,7 +8029,7 @@ version = "0.1.11"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22"
dependencies = [ dependencies = [
"windows-sys 0.48.0", "windows-sys 0.61.2",
] ]
[[package]] [[package]]
+4 -10
View File
@@ -1,5 +1,5 @@
[workspace] [workspace]
members = [".", "crates/ironclaw_common", "crates/ironclaw_safety"] members = [".", "crates/ironclaw_safety"]
exclude = [ exclude = [
"channels-src/discord", "channels-src/discord",
"channels-src/telegram", "channels-src/telegram",
@@ -20,7 +20,7 @@ exclude = [
[package] [package]
name = "ironclaw" name = "ironclaw"
version = "0.22.0" version = "0.19.0"
edition = "2024" edition = "2024"
rust-version = "1.92" rust-version = "1.92"
description = "Secure personal AI assistant that protects your data and expands its capabilities on the fly" description = "Secure personal AI assistant that protects your data and expands its capabilities on the fly"
@@ -57,7 +57,6 @@ refinery = { version = "0.8", features = ["tokio-postgres"], optional = true }
tokio-postgres-rustls = { version = "0.13", optional = true } tokio-postgres-rustls = { version = "0.13", optional = true }
rustls = { version = "0.23", optional = true, default-features = false } rustls = { version = "0.23", optional = true, default-features = false }
rustls-native-certs = { version = "0.8", optional = true } rustls-native-certs = { version = "0.8", optional = true }
webpki-roots = { version = "0.26", optional = true }
# Database - libSQL/Turso (optional embedded database) # Database - libSQL/Turso (optional embedded database)
libsql = { version = "0.6", optional = true, default-features = false, features = ["core", "replication", "remote", "tls"] } libsql = { version = "0.6", optional = true, default-features = false, features = ["core", "replication", "remote", "tls"] }
@@ -96,16 +95,13 @@ termimad = "0.34"
# Channel integrations # Channel integrations
axum = { version = "0.8", features = ["ws"] } axum = { version = "0.8", features = ["ws"] }
tower = "0.5" tower = "0.5"
tower-http = { version = "0.6", features = ["trace", "cors", "set-header", "catch-panic"] } tower-http = { version = "0.6", features = ["trace", "cors", "set-header"] }
# Cron scheduling for routines # Cron scheduling for routines
cron = "0.13" cron = "0.13"
# Shared types
ironclaw_common = { path = "crates/ironclaw_common", version = "0.1.0" }
# Safety/sanitization # Safety/sanitization
ironclaw_safety = { path = "crates/ironclaw_safety", version = "0.2.0" } ironclaw_safety = { path = "crates/ironclaw_safety", version = "0.1.0" }
regex = "1" regex = "1"
aho-corasick = "1" aho-corasick = "1"
@@ -220,7 +216,6 @@ postgres = [
"dep:tokio-postgres-rustls", "dep:tokio-postgres-rustls",
"dep:rustls", "dep:rustls",
"dep:rustls-native-certs", "dep:rustls-native-certs",
"dep:webpki-roots",
"dep:postgres-types", "dep:postgres-types",
"dep:refinery", "dep:refinery",
"dep:pgvector", "dep:pgvector",
@@ -232,7 +227,6 @@ libsql = ["dep:libsql"]
integration = [] integration = []
html-to-markdown = ["dep:html-to-markdown-rs", "dep:readabilityrs"] html-to-markdown = ["dep:html-to-markdown-rs", "dep:readabilityrs"]
bedrock = ["dep:aws-config", "dep:aws-sdk-bedrockruntime", "dep:aws-smithy-types"] bedrock = ["dep:aws-config", "dep:aws-sdk-bedrockruntime", "dep:aws-smithy-types"]
demo = []
import = ["dep:json5", "libsql"] import = ["dep:json5", "libsql"]
[[test]] [[test]]
+9 -39
View File
@@ -1,78 +1,49 @@
# Multi-stage Dockerfile for the IronClaw agent (cloud deployment). # Multi-stage Dockerfile for the IronClaw agent (cloud deployment).
# #
# Uses cargo-chef for dependency caching — only rebuilds deps when
# Cargo.toml/Cargo.lock change, not on every source edit.
#
# Build: # Build:
# docker build --platform linux/amd64 -t ironclaw:latest . # docker build --platform linux/amd64 -t ironclaw:latest .
# #
# Run: # Run:
# docker run --env-file .env -p 3000:3000 ironclaw:latest # docker run --env-file .env -p 3000:3000 ironclaw:latest
# Stage 1: Install cargo-chef # Stage 1: Build
FROM rust:1.92-slim-bookworm AS chef FROM rust:1.92-slim-bookworm AS builder
RUN apt-get update && apt-get install -y --no-install-recommends \ RUN apt-get update && apt-get install -y --no-install-recommends \
pkg-config libssl-dev cmake gcc g++ \ pkg-config libssl-dev cmake gcc g++ \
&& rm -rf /var/lib/apt/lists/* \ && rm -rf /var/lib/apt/lists/* \
&& rustup target add wasm32-wasip2 \ && rustup target add wasm32-wasip2 \
&& cargo install cargo-chef wasm-tools && cargo install wasm-tools
WORKDIR /app WORKDIR /app
# Stage 2: Generate the dependency recipe (changes only when Cargo.toml/lock change) # Copy manifests first for layer caching
FROM chef AS planner
COPY Cargo.toml Cargo.lock ./ COPY Cargo.toml Cargo.lock ./
COPY crates/ crates/ COPY crates/ crates/
# Copy source, build script, tests, and supporting directories
COPY build.rs build.rs COPY build.rs build.rs
COPY src/ src/ COPY src/ src/
COPY tests/ tests/ COPY tests/ tests/
COPY benches/ benches/
COPY migrations/ migrations/ COPY migrations/ migrations/
COPY registry/ registry/ COPY registry/ registry/
COPY channels-src/ channels-src/ COPY channels-src/ channels-src/
COPY wit/ wit/ COPY wit/ wit/
COPY providers.json providers.json COPY providers.json providers.json
# [[bench]] entries in Cargo.toml require bench sources to exist for cargo to parse the manifest
RUN cargo chef prepare --recipe-path recipe.json
# Stage 3: Build dependencies (cached unless Cargo.toml/lock change)
FROM chef AS deps
COPY --from=planner /app/recipe.json recipe.json
RUN cargo chef cook --release --recipe-path recipe.json
# Stage 4: Build the actual binary (only recompiles ironclaw source)
FROM deps AS builder
COPY Cargo.toml Cargo.lock ./
COPY crates/ crates/
COPY build.rs build.rs
COPY src/ src/
COPY tests/ tests/
COPY benches/ benches/ COPY benches/ benches/
COPY migrations/ migrations/
COPY registry/ registry/
COPY channels-src/ channels-src/
COPY wit/ wit/
COPY providers.json providers.json
COPY skills/ skills/ RUN cargo build --release --bin ironclaw
RUN cargo build --release --features demo --bin ironclaw # Stage 2: Runtime
# Stage 5: Runtime
FROM debian:bookworm-slim FROM debian:bookworm-slim
RUN apt-get update && apt-get install -y --no-install-recommends \ RUN apt-get update && apt-get install -y --no-install-recommends \
ca-certificates libssl3 \ ca-certificates libssl3 \
&& update-ca-certificates \
&& rm -rf /var/lib/apt/lists/* && rm -rf /var/lib/apt/lists/*
COPY --from=builder /app/target/release/ironclaw /usr/local/bin/ironclaw COPY --from=builder /app/target/release/ironclaw /usr/local/bin/ironclaw
COPY --from=builder /app/migrations /app/migrations COPY --from=builder /app/migrations /app/migrations
COPY --from=builder /app/skills /app/skills
# Non-root user # Non-root user
RUN useradd -m -u 1000 -s /bin/bash ironclaw RUN useradd -m -u 1000 -s /bin/bash ironclaw
@@ -81,6 +52,5 @@ USER ironclaw
EXPOSE 3000 EXPOSE 3000
ENV RUST_LOG=ironclaw=info ENV RUST_LOG=ironclaw=info
ENV SKILLS_DIR=/app/skills
ENTRYPOINT ["ironclaw"] ENTRYPOINT ["ironclaw"]
+1 -1
View File
@@ -161,7 +161,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
| `config` | ✅ | ✅ | - | Read/write config plus validate/path helpers | | `config` | ✅ | ✅ | - | Read/write config plus validate/path helpers |
| `backup` | ✅ | ❌ | P3 | Create/verify local backup archives | | `backup` | ✅ | ❌ | P3 | Create/verify local backup archives |
| `channels` | ✅ | 🚧 | P2 | `list` implemented; `enable`/`disable`/`status` deferred pending config source unification | | `channels` | ✅ | 🚧 | P2 | `list` implemented; `enable`/`disable`/`status` deferred pending config source unification |
| `models` | ✅ | 🚧 | P1 | `models list [<provider>]` (`--verbose`, `--json`; fetches live model list when provider specified), `models status` (`--json`), `models set <model>`, `models set-provider <provider> [--model model]` (alias normalization, config.toml + .env persistence). Remaining: `set` doesn't validate model against live list. | | `models` | ✅ | 🚧 | - | Model selector in TUI |
| `status` | ✅ | ✅ | - | System status (enriched session details) | | `status` | ✅ | ✅ | - | System status (enriched session details) |
| `agents` | ✅ | ❌ | P3 | Multi-agent management | | `agents` | ✅ | ❌ | P3 | Multi-agent management |
| `sessions` | ✅ | ❌ | P3 | Session listing (shows subagent models) | | `sessions` | ✅ | ❌ | P3 | Session listing (shows subagent models) |
+1 -1
View File
@@ -21,7 +21,7 @@
}, },
{ {
"name": "feishu_app_secret", "name": "feishu_app_secret",
"prompt": "Enter your Feishu/Lark App Secret (from your app settings at open.feishu.cn)", "prompt": "Enter your Feishu/Lark App Secret",
"optional": false "optional": false
}, },
{ {
-17
View File
@@ -1,17 +0,0 @@
[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 <[email protected]>"]
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"
-393
View File
@@ -1,393 +0,0 @@
//! 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<Self> {
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<String>,
},
#[serde(rename = "tool_started")]
ToolStarted {
name: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "tool_completed")]
ToolCompleted {
name: String,
success: bool,
#[serde(skip_serializing_if = "Option::is_none")]
error: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
parameters: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "tool_result")]
ToolResult {
name: String,
preview: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "stream_chunk")]
StreamChunk {
content: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "status")]
Status {
message: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[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<String>,
/// 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<String>,
#[serde(skip_serializing_if = "Option::is_none")]
auth_url: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
setup_url: Option<String>,
},
#[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<String>,
},
#[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<String>,
#[serde(skip_serializing_if = "Option::is_none")]
fallback_deliverable: Option<serde_json::Value>,
},
/// An image was generated by a tool.
#[serde(rename = "image_generated")]
ImageGenerated {
data_url: String,
#[serde(skip_serializing_if = "Option::is_none")]
path: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
/// Suggested follow-up messages for the user.
#[serde(rename = "suggestions")]
Suggestions {
suggestions: Vec<String>,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
/// 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<String>,
},
/// 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<String>,
},
/// Agent reasoning update (why it chose specific tools).
#[serde(rename = "reasoning_update")]
ReasoningUpdate {
narrative: String,
decisions: Vec<ToolDecisionDto>,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
/// Reasoning update for a sandbox job.
#[serde(rename = "job_reasoning")]
JobReasoning {
job_id: String,
narrative: String,
decisions: Vec<ToolDecisionDto>,
},
}
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<AppEvent> = 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");
}
}
-7
View File
@@ -1,7 +0,0 @@
//! Shared types and utilities for the IronClaw workspace.
mod event;
mod util;
pub use event::{AppEvent, ToolDecisionDto};
pub use util::truncate_preview;
-100
View File
@@ -1,100 +0,0 @@
//! Shared utility functions.
/// Truncate a string to at most `max_bytes` bytes at a char boundary, appending "...".
///
/// If the input is wrapped in `<tool_output ...>...</tool_output>` 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 <tool_output> if truncation cut through the closing tag.
if s.starts_with("<tool_output") && !result.ends_with("</tool_output>") {
result.push_str("\n</tool_output>");
}
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 = "<tool_output name=\"search\">\nSome very long content here\n</tool_output>";
let result = truncate_preview(s, 60);
assert!(result.ends_with("</tool_output>"));
assert!(result.contains("..."));
}
#[test]
fn test_truncate_preview_no_extra_close_when_intact() {
let s = "<tool_output name=\"echo\">\nshort\n</tool_output>";
let result = truncate_preview(s, 500);
assert_eq!(result, s);
assert_eq!(result.matches("</tool_output>").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("</tool_output>"));
}
}
+2 -1
View File
@@ -1,6 +1,6 @@
[package] [package]
name = "ironclaw_safety" name = "ironclaw_safety"
version = "0.2.0" version = "0.1.0"
edition = "2024" edition = "2024"
rust-version = "1.92" rust-version = "1.92"
description = "Prompt injection defense, input validation, secret leak detection, and safety policy enforcement" description = "Prompt injection defense, input validation, secret leak detection, and safety policy enforcement"
@@ -8,6 +8,7 @@ authors = ["NEAR AI <[email protected]>"]
license = "MIT OR Apache-2.0" license = "MIT OR Apache-2.0"
homepage = "https://github.com/nearai/ironclaw" homepage = "https://github.com/nearai/ironclaw"
repository = "https://github.com/nearai/ironclaw" repository = "https://github.com/nearai/ironclaw"
publish = false
[package.metadata.dist] [package.metadata.dist]
dist = false dist = false
-571
View File
@@ -1,571 +0,0 @@
# User Management API
DB-backed user management for multi-tenant IronClaw deployments. Covers admin user CRUD, per-user secrets provisioning, self-service profile, API token management, and usage reporting.
## Authentication
All endpoints require `Authorization: Bearer <token>`. Tokens are either:
- **Env-var tokens** — configured via `GATEWAY_AUTH_TOKEN` (single-user) at startup
- **DB-backed tokens** — created via `POST /api/tokens` or `POST /api/admin/users`
DB tokens are SHA-256 hashed at rest; plaintext is returned exactly once at creation time.
Auth is cached in a bounded LRU (1024 entries, 60s TTL). Suspending a user or revoking a token may take up to 60s to take effect.
## Roles
| Role | Scope |
|------|-------|
| `admin` | Full access to all endpoints |
| `member` | Self-service profile + own token management only |
Endpoints marked **Admin** return `403 Forbidden` for `member` role.
---
## Admin: Users
### POST /api/admin/users
Create a new user. Returns the user record and a one-time plaintext API token.
**Auth:** Admin
**Request body:**
```json
{
"display_name": "Alice Smith",
"email": "[email protected]",
"role": "member"
}
```
| Field | Type | Required | Default | Notes |
|-------|------|----------|---------|-------|
| `display_name` | string | yes | | |
| `email` | string | no | `null` | Must be unique if provided |
| `role` | string | no | `"member"` | `"admin"` or `"member"` |
**Response:** `200 OK`
```json
{
"id": "550e8400-e29b-41d4-a716-446655440000",
"email": "[email protected]",
"display_name": "Alice Smith",
"status": "active",
"role": "member",
"token": "a1b2c3d4e5f6...64-char hex...",
"created_at": "2026-03-25T12:00:00+00:00",
"created_by": "admin-user-id"
}
```
The `token` field is the plaintext API token. It is shown **only once** — store it securely.
**Errors:** `400` (missing display_name, invalid role), `403` (not admin), `503` (no database)
---
### GET /api/admin/users
List all users.
**Auth:** Admin
**Response:** `200 OK`
```json
{
"users": [
{
"id": "550e8400-...",
"email": "[email protected]",
"display_name": "Alice Smith",
"status": "active",
"role": "member",
"created_at": "2026-03-25T12:00:00+00:00",
"updated_at": "2026-03-25T12:00:00+00:00",
"last_login_at": "2026-03-25T14:30:00+00:00",
"created_by": "admin-user-id"
}
]
}
```
---
### GET /api/admin/users/{id}
Get a single user by ID.
**Auth:** Admin
**Response:** `200 OK`
```json
{
"id": "550e8400-...",
"email": "[email protected]",
"display_name": "Alice Smith",
"status": "active",
"role": "member",
"created_at": "2026-03-25T12:00:00+00:00",
"updated_at": "2026-03-25T12:00:00+00:00",
"last_login_at": "2026-03-25T14:30:00+00:00",
"created_by": "admin-user-id",
"metadata": {}
}
```
**Errors:** `404` (user not found), `403` (not admin)
---
### PATCH /api/admin/users/{id}
Update a user's display name and/or metadata. Omitted fields are left unchanged.
**Auth:** Admin
**Request body:**
```json
{
"display_name": "Alice Johnson",
"metadata": {"department": "engineering"}
}
```
| Field | Type | Required | Notes |
|-------|------|----------|-------|
| `display_name` | string | no | |
| `metadata` | object | no | Replaces entire metadata object (merge patch) |
**Response:** `200 OK` — returns the full updated user record (same shape as GET detail, without `last_login_at`/`created_by`).
**Errors:** `404` (user not found), `403` (not admin)
---
### POST /api/admin/users/{id}/suspend
Suspend a user. Suspended users cannot authenticate (DB auth checks user status).
**Auth:** Admin
**Response:** `200 OK`
```json
{
"id": "550e8400-...",
"status": "suspended"
}
```
**Errors:** `404` (user not found), `403` (not admin)
---
### POST /api/admin/users/{id}/activate
Re-activate a suspended user.
**Auth:** Admin
**Response:** `200 OK`
```json
{
"id": "550e8400-...",
"status": "active"
}
```
**Errors:** `404` (user not found), `403` (not admin)
---
### DELETE /api/admin/users/{id}
Permanently delete a user and all associated data (tokens, jobs, conversations, memory, routines, settings, secrets).
**Auth:** Admin
**Response:** `200 OK`
```json
{
"id": "550e8400-...",
"deleted": true
}
```
**Errors:** `404` (user not found), `403` (not admin)
**Cascade:** Deletes from `api_tokens`, `agent_jobs`, `conversations`, `memory_documents`, `routines`, `secrets`, `settings`, `wasm_tools`, and related tables. On PostgreSQL this uses FK cascades; on libSQL it uses explicit deletes.
---
## Admin: Per-User Secrets
Provision secrets on behalf of individual users. The primary use case is an application backend (acting as admin) that configures per-user credentials so each user's IronClaw agent can call back to external services.
Secrets are encrypted at rest with AES-256-GCM using a per-secret HKDF-derived key. Plaintext values are **never returned** by any endpoint — they can only be used by the agent's tool system at runtime.
### PUT /api/admin/users/{user_id}/secrets/{name}
Create or update a secret for the specified user. If a secret with the same name already exists, it is overwritten.
**Auth:** Admin
**Path parameters:**
| Param | Type | Notes |
|-------|------|-------|
| `user_id` | string | The user's ID |
| `name` | string | Secret name (normalized to lowercase) |
**Request body:**
```json
{
"value": "sk-live-abc123...",
"provider": "my-app-backend",
"expires_in_days": 90
}
```
| Field | Type | Required | Notes |
|-------|------|----------|-------|
| `value` | string | yes | The secret value (encrypted at rest, never returned) |
| `provider` | string | no | Tag for grouping (e.g. `"stripe"`, `"my-app"`) |
| `expires_in_days` | integer | no | Auto-expire after N days; `null` = never |
**Response:** `200 OK`
```json
{
"user_id": "550e8400-...",
"name": "my_app_callback_token",
"status": "created"
}
```
**Errors:** `400` (missing value), `403` (not admin), `503` (secrets store not available)
**Example — application backend provisioning a callback token:**
```bash
# Admin creates a user
curl -X POST https://ironclaw.example.com/api/admin/users \
-H "Authorization: Bearer $ADMIN_TOKEN" \
-d '{"display_name": "Alice", "role": "member"}'
# Response includes: {"id": "alice-uuid", "token": "alice-bearer-token", ...}
# Admin provisions a per-user callback secret
curl -X PUT https://ironclaw.example.com/api/admin/users/alice-uuid/secrets/app_callback_token \
-H "Authorization: Bearer $ADMIN_TOKEN" \
-d '{"value": "per-user-jwt-for-alice", "provider": "my-app"}'
# Now Alice's IronClaw agent can use the "app_callback_token" secret
# when calling tools that need to authenticate back to the app backend.
```
---
### GET /api/admin/users/{user_id}/secrets
List a user's secrets. Returns names and providers only — **never values or hashes**.
**Auth:** Admin
**Response:** `200 OK`
```json
{
"user_id": "550e8400-...",
"secrets": [
{"name": "app_callback_token", "provider": "my-app"},
{"name": "openai_api_key", "provider": "openai"}
]
}
```
---
### DELETE /api/admin/users/{user_id}/secrets/{name}
Delete a specific secret for a user.
**Auth:** Admin
**Response:** `200 OK`
```json
{
"user_id": "550e8400-...",
"name": "app_callback_token",
"deleted": true
}
```
**Errors:** `404` (secret not found), `403` (not admin), `503` (secrets store not available)
---
## Admin: Usage
### GET /api/admin/usage
Per-user LLM usage statistics aggregated from `llm_calls` via `agent_jobs.user_id`.
**Auth:** Admin
**Query parameters:**
| Param | Type | Default | Notes |
|-------|------|---------|-------|
| `user_id` | string | all users | Filter to a single user |
| `period` | string | `"day"` | `"day"` (24h), `"week"` (7d), or `"month"` (30d) |
**Response:** `200 OK`
```json
{
"period": "week",
"since": "2026-03-18T12:00:00+00:00",
"usage": [
{
"user_id": "alice-id",
"model": "claude-sonnet-4-5-20250514",
"call_count": 42,
"input_tokens": 150000,
"output_tokens": 30000,
"total_cost": "1.23"
}
]
}
```
---
## Self-Service: Profile
### GET /api/profile
Get the authenticated user's own profile.
**Auth:** Any authenticated user
**Response:** `200 OK`
```json
{
"id": "550e8400-...",
"email": "[email protected]",
"display_name": "Alice Smith",
"status": "active",
"role": "member",
"created_at": "2026-03-25T12:00:00+00:00",
"last_login_at": "2026-03-25T14:30:00+00:00"
}
```
---
### PATCH /api/profile
Update the authenticated user's own display name and/or metadata.
**Auth:** Any authenticated user
**Request body:**
```json
{
"display_name": "Alice Johnson",
"metadata": {"theme": "dark"}
}
```
**Response:** `200 OK`
```json
{
"id": "550e8400-...",
"display_name": "Alice Johnson",
"updated": true
}
```
---
## Self-Service: Tokens
### POST /api/tokens
Create a new API token for the authenticated user. Admins can optionally create tokens for other users by including `user_id`.
**Auth:** Any authenticated user
**Request body:**
```json
{
"name": "CI pipeline",
"expires_in_days": 90,
"user_id": "other-user-id"
}
```
| Field | Type | Required | Notes |
|-------|------|----------|-------|
| `name` | string | yes | Human-readable label |
| `expires_in_days` | integer | no | `null` = never expires |
| `user_id` | string | no | Admin-only; create token for another user |
**Response:** `200 OK`
```json
{
"token": "a1b2c3d4...64-char hex...",
"id": "token-uuid",
"name": "CI pipeline",
"token_prefix": "a1b2c3d4",
"expires_at": "2026-06-23T12:00:00+00:00",
"created_at": "2026-03-25T12:00:00+00:00"
}
```
The `token` field is shown **only once**.
---
### GET /api/tokens
List the authenticated user's tokens. Token hashes are never returned.
**Auth:** Any authenticated user
**Response:** `200 OK`
```json
{
"tokens": [
{
"id": "token-uuid",
"name": "CI pipeline",
"token_prefix": "a1b2c3d4",
"expires_at": "2026-06-23T12:00:00+00:00",
"last_used_at": "2026-03-25T14:00:00+00:00",
"created_at": "2026-03-25T12:00:00+00:00",
"revoked_at": null
}
]
}
```
---
### DELETE /api/tokens/{id}
Revoke one of the authenticated user's tokens. Users can only revoke their own tokens.
**Auth:** Any authenticated user
**Path:** `id` — UUID of the token to revoke
**Response:** `200 OK`
```json
{
"status": "revoked",
"id": "token-uuid"
}
```
**Errors:** `400` (invalid UUID), `404` (token not found or belongs to another user)
---
## Error Format
All error responses return a plain text body with the error message and the corresponding HTTP status code:
| Code | Meaning |
|------|---------|
| `400` | Bad request (missing fields, invalid input) |
| `401` | Missing or invalid bearer token |
| `403` | Authenticated but insufficient role (member accessing admin endpoint) |
| `404` | Resource not found |
| `503` | Database or secrets store not available |
| `500` | Internal server error |
---
## Security Model
### Secrets Encryption
- **Algorithm:** AES-256-GCM with per-secret HKDF-SHA256 derived keys
- **Master key:** 32+ bytes, resolved from `SECRETS_MASTER_KEY` env var or OS keychain
- **Storage format:** `nonce (12B) || ciphertext || tag (16B)` in `encrypted_value` column
- **Per-secret salt:** 32 random bytes stored alongside the ciphertext
- **Zero-exposure:** Plaintext never appears in logs, debug output, API responses, or LLM conversations
### Auth Cache
- Bounded LRU cache (1024 entries max)
- 60-second TTL per entry
- Suspending a user or revoking a token takes up to 60s to propagate
---
## Database Schema
### users
| Column | Type (PG / libSQL) | Notes |
|--------|--------------------|-------|
| `id` | `UUID` / `TEXT` | Primary key, UUID v4 |
| `email` | `TEXT UNIQUE` | Nullable |
| `display_name` | `TEXT NOT NULL` | |
| `status` | `TEXT NOT NULL` | `"active"` or `"suspended"` |
| `role` | `TEXT NOT NULL` | `"admin"` or `"member"` |
| `created_at` | `TIMESTAMPTZ` / `TEXT` | |
| `updated_at` | `TIMESTAMPTZ` / `TEXT` | |
| `last_login_at` | `TIMESTAMPTZ` / `TEXT` | Nullable |
| `created_by` | `TEXT` | Nullable, references `users.id` |
| `metadata` | `JSONB` / `TEXT` | Default `{}` |
### api_tokens
| Column | Type (PG / libSQL) | Notes |
|--------|--------------------|-------|
| `id` | `UUID` / `TEXT` | Primary key |
| `user_id` | `TEXT NOT NULL` | FK to `users.id` (PG cascades; libSQL explicit cleanup) |
| `token_hash` | `BYTEA` / `BLOB` | SHA-256 of hex-encoded plaintext |
| `token_prefix` | `TEXT NOT NULL` | First 8 chars for identification |
| `name` | `TEXT NOT NULL` | Human-readable label |
| `expires_at` | `TIMESTAMPTZ` / `TEXT` | Nullable |
| `last_used_at` | `TIMESTAMPTZ` / `TEXT` | Nullable |
| `created_at` | `TIMESTAMPTZ` / `TEXT` | |
| `revoked_at` | `TIMESTAMPTZ` / `TEXT` | Nullable; set on revocation |
### secrets
| Column | Type (PG / libSQL) | Notes |
|--------|--------------------|-------|
| `id` | `UUID` / `TEXT` | Primary key |
| `user_id` | `TEXT NOT NULL` | Scoped to user |
| `name` | `TEXT NOT NULL` | Unique per user (lowercase normalized) |
| `encrypted_value` | `BYTEA` / `BLOB` | AES-256-GCM (nonce + ciphertext + tag) |
| `key_salt` | `BYTEA` / `BLOB` | Per-secret HKDF salt |
| `provider` | `TEXT` | Optional grouping tag |
| `expires_at` | `TIMESTAMPTZ` / `TEXT` | Nullable |
| `last_used_at` | `TIMESTAMPTZ` / `TEXT` | Audit: last injection time |
| `usage_count` | `BIGINT` / `INTEGER` | Audit: total injections |
| `created_at` | `TIMESTAMPTZ` / `TEXT` | |
| `updated_at` | `TIMESTAMPTZ` / `TEXT` | |
-31
View File
@@ -1,31 +0,0 @@
-- User management tables for multi-tenant deployments.
--
-- Replaces the static GATEWAY_USER_TOKENS env var with DB-backed
-- user registration, API token management, and invitation flow.
CREATE TABLE users (
id TEXT PRIMARY KEY, -- matches existing user_id pattern (string, not UUID)
email TEXT UNIQUE, -- nullable for token-only users
display_name TEXT NOT NULL,
status TEXT NOT NULL DEFAULT 'active', -- active | suspended | deactivated
role TEXT NOT NULL DEFAULT 'member', -- admin | member
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
last_login_at TIMESTAMPTZ,
created_by TEXT REFERENCES users(id), -- who invited this user (nullable for bootstrap)
metadata JSONB NOT NULL DEFAULT '{}' -- extensible profile data
);
CREATE TABLE api_tokens (
id UUID PRIMARY KEY,
user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE,
token_hash BYTEA NOT NULL, -- SHA-256 hash (never store plaintext)
token_prefix TEXT NOT NULL, -- first 8 hex chars for display
name TEXT NOT NULL, -- human label ("my-laptop", "ci-bot")
expires_at TIMESTAMPTZ, -- nullable = never expires
last_used_at TIMESTAMPTZ,
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
revoked_at TIMESTAMPTZ -- soft-revoke: set this instead of deleting
);
CREATE INDEX idx_api_tokens_user ON api_tokens(user_id);
CREATE INDEX idx_api_tokens_hash ON api_tokens(token_hash);
+3 -3
View File
@@ -2,7 +2,7 @@
"name": "feishu", "name": "feishu",
"display_name": "Feishu / Lark Channel", "display_name": "Feishu / Lark Channel",
"kind": "channel", "kind": "channel",
"version": "0.1.3", "version": "0.1.1",
"wit_version": "0.3.0", "wit_version": "0.3.0",
"description": "Talk to your agent through a Feishu or Lark bot", "description": "Talk to your agent through a Feishu or Lark bot",
"keywords": [ "keywords": [
@@ -19,8 +19,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"sha256": "a66ff0dafb67d2216d8161bb7e96e724a94acb0ab993b85d2782d30412f8fe94", "sha256": "5fca74022264d1c8e78a0853766276f7ffa3cf0d8065b2f51ca10985acad4714",
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/channel-feishu-0.1.3-wasm32-wasip2.tar.gz" "url": "https://github.com/nearai/ironclaw/releases/download/v0.19.0/channel-feishu-0.1.1-wasm32-wasip2.tar.gz"
} }
}, },
"auth_summary": { "auth_summary": {
+2 -2
View File
@@ -18,8 +18,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.20.0/channel-telegram-0.2.5-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/download/v0.19.0/channel-telegram-0.2.4-wasm32-wasip2.tar.gz",
"sha256": "1ef20a538f55b379e049356e4d6758006251846bc3365ceaa1c87eba8379a329" "sha256": "a7cb300ec1c946831cfceaa95c1dc8f30d0f42a3924f3cb5de8098821573f4b8"
} }
}, },
"auth_summary": { "auth_summary": {
+3 -3
View File
@@ -2,7 +2,7 @@
"name": "github", "name": "github",
"display_name": "GitHub", "display_name": "GitHub",
"kind": "tool", "kind": "tool",
"version": "0.2.2", "version": "0.2.1",
"wit_version": "0.3.0", "wit_version": "0.3.0",
"description": "GitHub integration for issues, PRs, repos, and code search", "description": "GitHub integration for issues, PRs, repos, and code search",
"keywords": [ "keywords": [
@@ -19,8 +19,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-github-0.2.2-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/download/v0.19.0/tool-github-0.2.1-wasm32-wasip2.tar.gz",
"sha256": "70b55af593193d8fa495c0f702ea23284d83a624124f8a5f7564916ec5032c3f" "sha256": "92c530b3ad172e2372d819744b5233f1d8f65768e26eb5a6c213eba3ce1de758"
} }
}, },
"auth_summary": { "auth_summary": {
+3 -3
View File
@@ -2,7 +2,7 @@
"name": "gmail", "name": "gmail",
"display_name": "Gmail", "display_name": "Gmail",
"kind": "tool", "kind": "tool",
"version": "0.2.1", "version": "0.2.0",
"wit_version": "0.3.0", "wit_version": "0.3.0",
"description": "Read, send, and manage Gmail messages and threads", "description": "Read, send, and manage Gmail messages and threads",
"keywords": [ "keywords": [
@@ -18,8 +18,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-gmail-0.2.1-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/gmail-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "79025b40ee70ce1120acc4320bae50da095d7afb0ef67bd56d99b064b72ea779" "sha256": "ee9574e02e92bc1d481f1310eb88afd99ee52bf6971074ab33bd76bf99b34b1d"
} }
}, },
"auth_summary": { "auth_summary": {
+3 -3
View File
@@ -2,7 +2,7 @@
"name": "google-calendar", "name": "google-calendar",
"display_name": "Google Calendar", "display_name": "Google Calendar",
"kind": "tool", "kind": "tool",
"version": "0.2.1", "version": "0.2.0",
"wit_version": "0.3.0", "wit_version": "0.3.0",
"description": "Create, read, update, and delete Google Calendar events", "description": "Create, read, update, and delete Google Calendar events",
"keywords": [ "keywords": [
@@ -18,8 +18,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-google-calendar-0.2.1-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-calendar-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "86bcc075010b08f5ab2f98f504cec1c6c9e0ca144857d185cbecf72a11f504bf" "sha256": "2fa47150ea222e787c122182ad6f4dfa2ffaf5fe490d05e8de887a76445f8d2d"
} }
}, },
"auth_summary": { "auth_summary": {
+3 -3
View File
@@ -2,7 +2,7 @@
"name": "google-docs", "name": "google-docs",
"display_name": "Google Docs", "display_name": "Google Docs",
"kind": "tool", "kind": "tool",
"version": "0.2.1", "version": "0.2.0",
"wit_version": "0.3.0", "wit_version": "0.3.0",
"description": "Create and edit Google Docs documents", "description": "Create and edit Google Docs documents",
"keywords": [ "keywords": [
@@ -18,8 +18,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-google-docs-0.2.1-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-docs-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "39d476029764949498a53a6a223f9952b5f4df151be7b8b19bf3fe4d401a57cd" "sha256": "40e134a1c1564f832ca861c3396895d4e33ec67b99313fc1f97baf8d971423a9"
} }
}, },
"auth_summary": { "auth_summary": {
+3 -3
View File
@@ -2,7 +2,7 @@
"name": "google-drive", "name": "google-drive",
"display_name": "Google Drive", "display_name": "Google Drive",
"kind": "tool", "kind": "tool",
"version": "0.2.1", "version": "0.2.0",
"wit_version": "0.3.0", "wit_version": "0.3.0",
"description": "Upload, download, search, and manage Google Drive files and folders", "description": "Upload, download, search, and manage Google Drive files and folders",
"keywords": [ "keywords": [
@@ -18,8 +18,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-google-drive-0.2.1-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-drive-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "6e9a700fab93865c852af718666af64c5b534ad6a419fb4b736e07740188f494" "sha256": "002a341a1d58125563a7c69561b26fbc2629b04ea723cade744102bdc0fbb71f"
} }
}, },
"auth_summary": { "auth_summary": {
+3 -3
View File
@@ -2,7 +2,7 @@
"name": "google-sheets", "name": "google-sheets",
"display_name": "Google Sheets", "display_name": "Google Sheets",
"kind": "tool", "kind": "tool",
"version": "0.2.1", "version": "0.2.0",
"wit_version": "0.3.0", "wit_version": "0.3.0",
"description": "Read and write Google Sheets spreadsheet data", "description": "Read and write Google Sheets spreadsheet data",
"keywords": [ "keywords": [
@@ -18,8 +18,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-google-sheets-0.2.1-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-sheets-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "1f8c381799a916be83263cac9d497d52946e21b1b588592a3a42ca94a73b7051" "sha256": "8aa2c9d52f033edea3a6c2311b0ec694ccb6d0a54ef07e94d72bf8be1ce8009a"
} }
}, },
"auth_summary": { "auth_summary": {
+3 -3
View File
@@ -2,7 +2,7 @@
"name": "google-slides", "name": "google-slides",
"display_name": "Google Slides", "display_name": "Google Slides",
"kind": "tool", "kind": "tool",
"version": "0.2.1", "version": "0.2.0",
"wit_version": "0.3.0", "wit_version": "0.3.0",
"description": "Create and edit Google Slides presentations", "description": "Create and edit Google Slides presentations",
"keywords": [ "keywords": [
@@ -17,8 +17,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-google-slides-0.2.1-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-slides-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "e2528be5da02f1b8cfc8ee9b0cdd849516c53d412e2f75c6175b3bded7f512cb" "sha256": "e931a97d4fd0b0b938e464dc7c7f2be6ea6b4d1508f5ea3cd931d44db23f05f5"
} }
}, },
"auth_summary": { "auth_summary": {
+3 -3
View File
@@ -2,7 +2,7 @@
"name": "llm-context", "name": "llm-context",
"display_name": "LLM Context", "display_name": "LLM Context",
"kind": "tool", "kind": "tool",
"version": "0.1.1", "version": "0.1.0",
"wit_version": "0.3.0", "wit_version": "0.3.0",
"description": "Fetch pre-extracted web content from Brave Search for grounding LLM answers (RAG, fact-checking)", "description": "Fetch pre-extracted web content from Brave Search for grounding LLM answers (RAG, fact-checking)",
"keywords": [ "keywords": [
@@ -21,8 +21,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-llm-context-0.1.1-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/download/v0.19.0/tool-llm-context-0.1.0-wasm32-wasip2.tar.gz",
"sha256": "9b19e2fd05dbbbe3c8bd55309a91db09124e8415eb0f767828b6e10b55771e63" "sha256": "d9ced2b1226b879135891e0ee40e072c7c95412e1b2462925a23853e1f92497e"
} }
}, },
"auth_summary": { "auth_summary": {
+3 -3
View File
@@ -2,7 +2,7 @@
"name": "slack-tool", "name": "slack-tool",
"display_name": "Slack Tool", "display_name": "Slack Tool",
"kind": "tool", "kind": "tool",
"version": "0.2.1", "version": "0.2.0",
"wit_version": "0.3.0", "wit_version": "0.3.0",
"description": "Your agent uses Slack to post and read messages in your workspace", "description": "Your agent uses Slack to post and read messages in your workspace",
"keywords": [ "keywords": [
@@ -17,8 +17,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-slack-0.2.1-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/download/v0.19.0/tool-slack-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "927519e5b7734beeb022d3b8bbd152e0e6b9f67c9452a8ad47809d3c4221a137" "sha256": "ccfb0415d7a04f9497726c712d15216de36e86f498b849101283c017f5ab4efb"
} }
}, },
"auth_summary": { "auth_summary": {
+3 -3
View File
@@ -2,7 +2,7 @@
"name": "telegram-mtproto", "name": "telegram-mtproto",
"display_name": "Telegram Tool", "display_name": "Telegram Tool",
"kind": "tool", "kind": "tool",
"version": "0.2.1", "version": "0.2.0",
"wit_version": "0.3.0", "wit_version": "0.3.0",
"description": "Your agent uses your Telegram account to read and send messages", "description": "Your agent uses your Telegram account to read and send messages",
"keywords": [ "keywords": [
@@ -18,8 +18,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-telegram-0.2.1-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/download/v0.19.0/tool-telegram-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "1e57d0755fc9c7b3ec013d079f30168898b484a6919f9edd105f0cd80131c1cd" "sha256": "c17065ca41fae5f2a7c43b36144686718cd310a2f22442313bb1aa82bbad0ae4"
} }
}, },
"auth_summary": { "auth_summary": {
+3 -3
View File
@@ -2,7 +2,7 @@
"name": "web-search", "name": "web-search",
"display_name": "Web Search", "display_name": "Web Search",
"kind": "tool", "kind": "tool",
"version": "0.2.2", "version": "0.2.1",
"wit_version": "0.3.0", "wit_version": "0.3.0",
"description": "Search the web using Brave Search API", "description": "Search the web using Brave Search API",
"keywords": [ "keywords": [
@@ -18,8 +18,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-web-search-0.2.2-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/download/v0.19.0/tool-web-search-0.2.1-wasm32-wasip2.tar.gz",
"sha256": "47382b50c1ea7525b20d59dc02fab04e336d018665826c2f24710bdf460779ae" "sha256": "bad275ca4ec314adea5241d6b92c44ccf9cebcbca8e30ba2493cc0bcb4b57218"
} }
}, },
"auth_summary": { "auth_summary": {
+5
View File
@@ -1,2 +1,7 @@
[workspace] [workspace]
git_release_enable = false git_release_enable = false
[[package]]
name = "ironclaw_safety"
publish = false
release = false
-94
View File
@@ -1,94 +0,0 @@
---
name: abound-remittance
version: 0.1.0
description: Smart remittance assistant for Abound — helps users send money to India with intelligent forex timing and transfer management.
activation:
keywords:
- send money
- transfer
- remittance
- exchange rate
- forex
- INR
- India
- wire
- schedule trade
- trade tomorrow
- convert currency
- send dollars
- rupees
- beneficiary
- funding source
- payment
- how much
- rate today
- best time
- family maintenance
patterns:
- "send \\$?\\d+"
- "schedule.*(trade|transfer|send|wire)"
- "how much.*(INR|rupees|India)"
- "best time to (send|transfer|convert)"
- "(rate|forex).*(good|bad|high|low|today|now)"
- "transfer.*tomorrow|tomorrow.*transfer"
tags:
- fintech
- remittance
- forex
max_context_tokens: 2500
---
# Abound Remittance Assistant
You are a smart remittance assistant for Abound, helping users send money from USD to INR (India) with intelligent timing advice.
## Available Tools
You have these Abound-specific tools:
- **abound_get_account_info** — Get the user's account: limits, recipients, funding sources, payment reasons
- **abound_get_exchange_rate** — Get current USD/INR exchange rate (current + effective after fees)
- **abound_get_forex_score** — Get a 0-100 forex timing score with a signal (convert_now / split_transfer / wait)
- **abound_send_wire** — Execute a wire transfer (requires: funding_source_id, beneficiary_ref_id, amount, payment_reason_key)
- **abound_create_notification** — Send a notification to the user's Abound app
You also have **routine_create** for scheduling future/recurring transfers.
## Workflow: "Send $X" or "Transfer money"
1. **Always check the rate first.** Call `abound_get_exchange_rate` to get the current rate.
2. **Check the forex score.** Call `abound_get_forex_score` to assess timing.
3. **Get account info.** Call `abound_get_account_info` to know the user's limits, recipients, and funding sources.
4. **Advise based on the score:**
- **Score >= 60 (convert_now):** Tell the user it's a good time. Show the rate, the INR equivalent of their amount, and recommend proceeding.
- **Score 40-59 (split_transfer):** Suggest splitting — send half now at the current rate, schedule the rest for later when the rate may improve.
- **Score < 40 (wait):** Unless the transfer is urgent, recommend waiting. Explain why (rate below average, unfavorable season).
5. **Execute if user confirms.** Use `abound_send_wire` with the correct funding_source_id, beneficiary_ref_id, amount, and payment_reason_key from the account info.
6. **Notify.** After a successful wire, call `abound_create_notification` with relevant metadata.
## Workflow: "Schedule a trade" or "Send tomorrow morning"
1. Gather the same info (rate, score, account).
2. Use **routine_create** to schedule the transfer:
- For "tomorrow morning": use cron `"0 9 * * *"` with the user's timezone, set to fire once
- For "every week": use cron `"0 9 * * MON"` (or the user's preferred day)
- The routine prompt should instruct the agent to check the rate and execute the wire
3. Confirm the schedule with the user, showing when it will fire.
## Presentation Rules
- Always show amounts in **both USD and INR**: "$1,000 (~INR 85,420 at today's rate of 85.42)"
- Show the **effective rate** (after fees), not just the market rate
- When showing the forex score, explain it simply: "The forex timing score is 72/100 — this is a good time to send."
- If the user's amount exceeds their limit ($5,000), tell them and suggest splitting into multiple transfers
- Always mention the **estimated delivery time** (1-3 business days) after a wire
## Payment Reasons
When asking about the purpose, offer these options:
- Family Maintenance
- Gift
- Education Support
- Medical Support
If the user doesn't specify, ask which applies.
+48 -226
View File
@@ -13,10 +13,9 @@ use futures::StreamExt;
use uuid::Uuid; use uuid::Uuid;
use crate::agent::context_monitor::ContextMonitor; use crate::agent::context_monitor::ContextMonitor;
use crate::agent::heartbeat::{spawn_heartbeat, spawn_multi_user_heartbeat}; use crate::agent::heartbeat::spawn_heartbeat;
use crate::agent::routine_engine::{RoutineEngine, spawn_cron_ticker}; use crate::agent::routine_engine::{RoutineEngine, spawn_cron_ticker};
use crate::agent::self_repair::{DefaultSelfRepair, RepairResult, SelfRepair}; use crate::agent::self_repair::{DefaultSelfRepair, RepairResult, SelfRepair};
use crate::agent::session::ThreadState;
use crate::agent::session_manager::SessionManager; use crate::agent::session_manager::SessionManager;
use crate::agent::submission::{Submission, SubmissionParser, SubmissionResult}; use crate::agent::submission::{Submission, SubmissionParser, SubmissionResult};
use crate::agent::{HeartbeatConfig as AgentHeartbeatConfig, Router, Scheduler, SchedulerDeps}; use crate::agent::{HeartbeatConfig as AgentHeartbeatConfig, Router, Scheduler, SchedulerDeps};
@@ -85,15 +84,6 @@ fn resolve_owner_scope_notification_user(
trimmed_option(explicit_user).or_else(|| trimmed_option(owner_fallback)) 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( async fn resolve_channel_notification_user(
extension_manager: Option<&Arc<ExtensionManager>>, extension_manager: Option<&Arc<ExtensionManager>>,
channel: Option<&str>, channel: Option<&str>,
@@ -167,8 +157,8 @@ pub struct AgentDeps {
pub hooks: Arc<HookRegistry>, pub hooks: Arc<HookRegistry>,
/// Cost enforcement guardrails (daily budget, hourly rate limits). /// Cost enforcement guardrails (daily budget, hourly rate limits).
pub cost_guard: Arc<crate::agent::cost_guard::CostGuard>, pub cost_guard: Arc<crate::agent::cost_guard::CostGuard>,
/// SSE manager for live job event streaming to the web gateway. /// SSE broadcast sender for live job event streaming to the web gateway.
pub sse_tx: Option<Arc<crate::channels::web::sse::SseManager>>, pub sse_tx: Option<tokio::sync::broadcast::Sender<crate::channels::web::types::SseEvent>>,
/// HTTP interceptor for trace recording/replay. /// HTTP interceptor for trace recording/replay.
pub http_interceptor: Option<Arc<dyn crate::llm::recording::HttpInterceptor>>, pub http_interceptor: Option<Arc<dyn crate::llm::recording::HttpInterceptor>>,
/// Audio transcription middleware for voice messages. /// Audio transcription middleware for voice messages.
@@ -179,11 +169,6 @@ pub struct AgentDeps {
pub sandbox_readiness: crate::agent::routine_engine::SandboxReadiness, pub sandbox_readiness: crate::agent::routine_engine::SandboxReadiness,
/// Software builder for self-repair tool rebuilding. /// Software builder for self-repair tool rebuilding.
pub builder: Option<Arc<dyn crate::tools::SoftwareBuilder>>, pub builder: Option<Arc<dyn crate::tools::SoftwareBuilder>>,
/// 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<crate::tenant::TenantRateRegistry>,
} }
/// The main agent that coordinates all components. /// The main agent that coordinates all components.
@@ -246,15 +231,12 @@ impl Agent {
SchedulerDeps { SchedulerDeps {
tools: deps.tools.clone(), tools: deps.tools.clone(),
extension_manager: deps.extension_manager.clone(), extension_manager: deps.extension_manager.clone(),
store: deps store: deps.store.clone(),
.store
.as_ref()
.map(|db| crate::tenant::AdminScope::new(Arc::clone(db))),
hooks: deps.hooks.clone(), hooks: deps.hooks.clone(),
}, },
); );
if let Some(ref sse) = deps.sse_tx { if let Some(ref tx) = deps.sse_tx {
scheduler.set_sse_sender(Arc::clone(sse)); scheduler.set_sse_sender(tx.clone());
} }
if let Some(ref interceptor) = deps.http_interceptor { if let Some(ref interceptor) = deps.http_interceptor {
scheduler.set_http_interceptor(Arc::clone(interceptor)); scheduler.set_http_interceptor(Arc::clone(interceptor));
@@ -330,50 +312,6 @@ impl Agent {
&self.deps.cost_guard &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<crate::tenant::AdminScope> {
self.deps
.store
.as_ref()
.map(|db| crate::tenant::AdminScope::new(Arc::clone(db)))
}
pub(super) fn skill_registry(&self) -> Option<&Arc<std::sync::RwLock<SkillRegistry>>> { pub(super) fn skill_registry(&self) -> Option<&Arc<std::sync::RwLock<SkillRegistry>>> {
self.deps.skill_registry.as_ref() self.deps.skill_registry.as_ref()
} }
@@ -459,8 +397,8 @@ impl Agent {
self.config.stuck_threshold, self.config.stuck_threshold,
self.config.max_repair_attempts, self.config.max_repair_attempts,
); );
if let Some(admin) = self.admin_store() { if let Some(ref store) = self.deps.store {
self_repair = self_repair.with_store(admin); self_repair = self_repair.with_store(Arc::clone(store));
} }
if let Some(ref builder) = self.deps.builder { if let Some(ref builder) = self.deps.builder {
self_repair = self_repair.with_builder(Arc::clone(builder), Arc::clone(self.tools())); self_repair = self_repair.with_builder(Arc::clone(builder), Arc::clone(self.tools()));
@@ -567,7 +505,6 @@ impl Agent {
.with_interval(std::time::Duration::from_secs(hb_config.interval_secs)); .with_interval(std::time::Duration::from_secs(hb_config.interval_secs));
config.quiet_hours_start = hb_config.quiet_hours_start; config.quiet_hours_start = hb_config.quiet_hours_start;
config.quiet_hours_end = hb_config.quiet_hours_end; config.quiet_hours_end = hb_config.quiet_hours_end;
config.multi_tenant = hb_config.multi_tenant;
config.timezone = hb_config config.timezone = hb_config
.timezone .timezone
.clone() .clone()
@@ -597,52 +534,30 @@ impl Agent {
.await; .await;
let notify_user = heartbeat_notify_user; let notify_user = heartbeat_notify_user;
let channels = self.channels.clone(); let channels = self.channels.clone();
let is_multi_tenant = hb_config.multi_tenant;
tokio::spawn(async move { tokio::spawn(async move {
while let Some(response) = notify_rx.recv().await { 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 // Try the configured channel first, fall back to
// broadcasting on all channels. // broadcasting on all channels.
let targeted_ok = if let Some(ref channel) = notify_channel { let targeted_ok = if let Some(ref channel) = notify_channel
let target = effective_user.as_deref().or(notify_target.as_deref()); && let Some(ref user) = notify_target
if let Some(user) = target { {
channels channels
.broadcast(channel, user, response.clone()) .broadcast(channel, user, response.clone())
.await .await
.is_ok() .is_ok()
} else {
false
}
} else { } else {
false false
}; };
if !targeted_ok { if !targeted_ok && let Some(ref user) = notify_user {
let fallback = effective_user.as_deref().or(notify_user.as_deref()); let results = channels.broadcast_all(user, response).await;
if let Some(user) = fallback { for (ch, result) in results {
let results = channels.broadcast_all(user, response).await; if let Err(e) = result {
for (ch, result) in results { tracing::warn!(
if let Err(e) = result { "Failed to broadcast heartbeat to {}: {}",
tracing::warn!( ch,
"Failed to broadcast heartbeat to {}: {}", e
ch, );
e
);
}
} }
} }
} }
@@ -655,29 +570,14 @@ impl Agent {
.map(|h| h.to_workspace_config()) .map(|h| h.to_workspace_config())
.unwrap_or_default(); .unwrap_or_default();
if config.multi_tenant { Some(spawn_heartbeat(
if let Some(admin) = self.admin_store() { config,
Some(spawn_multi_user_heartbeat( hygiene,
config, workspace.clone(),
hygiene, self.cheap_llm().clone(),
self.cheap_llm().clone(), Some(notify_tx),
Some(notify_tx), self.store().map(Arc::clone),
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 { } else {
tracing::warn!("Heartbeat enabled but no workspace available"); tracing::warn!("Heartbeat enabled but no workspace available");
None None
@@ -699,7 +599,7 @@ impl Agent {
let engine = Arc::new(RoutineEngine::new( let engine = Arc::new(RoutineEngine::new(
rt_config.clone(), rt_config.clone(),
crate::tenant::AdminScope::new(Arc::clone(store)), Arc::clone(store),
self.llm().clone(), self.llm().clone(),
Arc::clone(workspace), Arc::clone(workspace),
notify_tx, notify_tx,
@@ -1152,11 +1052,10 @@ impl Agent {
} else { } else {
drop(sess); drop(sess);
self.session_manager self.session_manager
.resolve_thread_with_parsed_uuid( .resolve_thread(
&message.user_id, &message.user_id,
&message.channel, &message.channel,
message.conversation_scope(), message.conversation_scope(),
approval_thread_uuid,
) )
.await .await
} }
@@ -1237,14 +1136,9 @@ impl Agent {
&& let Submission::UserInput { ref content } = submission && let Submission::UserInput { ref content } = submission
&& let Some(engine) = self.routine_engine().await && let Some(engine) = self.routine_engine().await
{ {
let single_message_repl = is_single_message_repl(message); let fired = engine
// Use post-hook content so that BeforeInbound hooks that rewrite .check_event_triggers(&message.user_id, &message.channel, content)
// input are respected by event trigger matching. .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 { if fired > 0 {
tracing::debug!( tracing::debug!(
channel = %message.channel, channel = %message.channel,
@@ -1252,30 +1146,15 @@ impl Agent {
fired, fired,
"Consumed inbound user message with matching event-triggered routine(s)" "Consumed inbound user message with matching event-triggered routine(s)"
); );
return if single_message_repl { return Ok(Some(String::new()));
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 // Process based on submission type
let result = match submission { let result = match submission {
Submission::UserInput { content } => { Submission::UserInput { content } => {
let mut result = self let mut result = self
.process_user_input( .process_user_input(message, session.clone(), thread_id, &content)
message,
tenant.clone(),
session.clone(),
thread_id,
&content,
)
.await; .await;
// Drain any messages queued during processing. // Drain any messages queued during processing.
@@ -1342,13 +1221,7 @@ impl Agent {
let mut queued_msg = message.clone(); let mut queued_msg = message.clone();
queued_msg.attachments.clear(); queued_msg.attachments.clear();
result = self result = self
.process_user_input( .process_user_input(&queued_msg, session.clone(), thread_id, &next_content)
&queued_msg,
tenant.clone(),
session.clone(),
thread_id,
&next_content,
)
.await; .await;
// If processing failed, re-queue the drained content so it // If processing failed, re-queue the drained content so it
@@ -1373,30 +1246,8 @@ impl Agent {
command, command,
message.channel 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 // Authorization checks (including restart channel check) are enforced in handle_system_command
self.handle_system_command(&command, &args, &message.channel, &tenant) self.handle_system_command(&command, &args, &message.channel)
.await .await
} }
Submission::Undo => self.process_undo(session, thread_id).await, Submission::Undo => self.process_undo(session, thread_id).await,
@@ -1409,9 +1260,12 @@ impl Agent {
Submission::Summarize => self.process_summarize(session, thread_id).await, Submission::Summarize => self.process_summarize(session, thread_id).await,
Submission::Suggest => self.process_suggest(session, thread_id).await, Submission::Suggest => self.process_suggest(session, thread_id).await,
Submission::JobStatus { job_id } => { Submission::JobStatus { job_id } => {
self.process_job_status(&tenant, job_id.as_deref()).await 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
} }
Submission::JobCancel { job_id } => self.process_job_cancel(&tenant, &job_id).await,
Submission::Quit => return Ok(None), Submission::Quit => return Ok(None),
Submission::SwitchThread { thread_id: target } => { Submission::SwitchThread { thread_id: target } => {
self.process_switch_thread(message, target).await self.process_switch_thread(message, target).await
@@ -1451,26 +1305,7 @@ impl Agent {
Ok(Some(content)) Ok(Some(content))
} }
} }
SubmissionResult::Ok { SubmissionResult::Ok { message } => Ok(message),
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::Error { message } => Ok(Some(format!("Error: {}", message))),
SubmissionResult::Interrupted => Ok(Some("Interrupted.".into())), SubmissionResult::Interrupted => Ok(Some("Interrupted.".into())),
SubmissionResult::NeedApproval { .. } => { SubmissionResult::NeedApproval { .. } => {
@@ -1486,7 +1321,7 @@ impl Agent {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::{ use super::{
chat_tool_execution_metadata, is_single_message_repl, resolve_routine_notification_user, chat_tool_execution_metadata, resolve_routine_notification_user,
should_fallback_routine_notification, truncate_for_preview, should_fallback_routine_notification, truncate_for_preview,
}; };
use crate::channels::IncomingMessage; use crate::channels::IncomingMessage;
@@ -1648,17 +1483,4 @@ mod tests {
assert!(should_fallback_routine_notification(&error)); // safety: test-only assertion 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
}
} }
+1 -126
View File
@@ -10,7 +10,7 @@ use std::borrow::Cow;
use crate::agent::session::PendingApproval; use crate::agent::session::PendingApproval;
use crate::error::Error; use crate::error::Error;
use crate::llm::{ChatMessage, FinishReason, Reasoning, ReasoningContext, RespondResult}; use crate::llm::{ChatMessage, Reasoning, ReasoningContext, RespondResult};
/// Signal from the delegate indicating how the loop should proceed. /// Signal from the delegate indicating how the loop should proceed.
pub enum LoopSignal { pub enum LoopSignal {
@@ -134,9 +134,6 @@ pub async fn run_agentic_loop(
config: &AgenticLoopConfig, config: &AgenticLoopConfig,
) -> Result<LoopOutcome, Error> { ) -> Result<LoopOutcome, Error> {
let mut consecutive_tool_intent_nudges: u32 = 0; 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 { for iteration in 1..=config.max_iterations {
// Check for external signals (stop, cancellation, user messages) // Check for external signals (stop, cancellation, user messages)
@@ -218,35 +215,7 @@ pub async fn run_agentic_loop(
tool_calls, tool_calls,
content, 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; consecutive_tool_intent_nudges = 0;
truncation_count = 0;
if let Some(outcome) = delegate if let Some(outcome) = delegate
.execute_tool_calls(tool_calls, content, reason_ctx) .execute_tool_calls(tool_calls, content, reason_ctx)
@@ -302,7 +271,6 @@ mod tests {
RespondOutput { RespondOutput {
result: RespondResult::Text(text.to_string()), result: RespondResult::Text(text.to_string()),
usage: zero_usage(), usage: zero_usage(),
finish_reason: FinishReason::Stop,
} }
} }
@@ -313,7 +281,6 @@ mod tests {
content: None, content: None,
}, },
usage: zero_usage(), usage: zero_usage(),
finish_reason: FinishReason::ToolUse,
} }
} }
@@ -447,7 +414,6 @@ mod tests {
id: "call_1".to_string(), id: "call_1".to_string(),
name: "echo".to_string(), name: "echo".to_string(),
arguments: serde_json::json!({}), arguments: serde_json::json!({}),
reasoning: None,
}; };
let delegate = MockDelegate::new(vec![ let delegate = MockDelegate::new(vec![
tool_calls_output(vec![tool_call]), tool_calls_output(vec![tool_call]),
@@ -655,95 +621,4 @@ mod tests {
let result = truncate_for_preview("café", 4); let result = truncate_for_preview("café", 4);
assert_eq!(result, "caf..."); 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"
);
}
} }
+54 -224
View File
@@ -33,7 +33,6 @@ impl Agent {
&self, &self,
intent: MessageIntent, intent: MessageIntent,
message: &IncomingMessage, message: &IncomingMessage,
tenant: &crate::tenant::TenantCtx,
) -> Result<SubmissionResult, Error> { ) -> Result<SubmissionResult, Error> {
// Send thinking status for non-trivial operations // Send thinking status for non-trivial operations
if let MessageIntent::CreateJob { .. } = &intent { if let MessageIntent::CreateJob { .. } = &intent {
@@ -53,18 +52,24 @@ impl Agent {
description, description,
category, category,
} => { } => {
self.handle_create_job(tenant, title, description, category) self.handle_create_job(&message.user_id, title, description, category)
.await? .await?
} }
MessageIntent::CheckJobStatus { job_id } => { MessageIntent::CheckJobStatus { job_id } => {
self.handle_check_status(tenant, job_id).await? 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?
} }
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 } => { MessageIntent::Command { command, args } => {
match self match self
.handle_command(&command, &args, &message.channel, tenant) .handle_command(&command, &args, &message.channel)
.await? .await?
{ {
Some(s) => s, Some(s) => s,
@@ -78,14 +83,14 @@ impl Agent {
async fn handle_create_job( async fn handle_create_job(
&self, &self,
tenant: &crate::tenant::TenantCtx, user_id: &str,
title: String, title: String,
description: String, description: String,
category: Option<String>, category: Option<String>,
) -> Result<String, Error> { ) -> Result<String, Error> {
let job_id = self let job_id = self
.scheduler .scheduler
.dispatch_job(tenant.user_id(), &title, &description, None) .dispatch_job(user_id, &title, &description, None)
.await?; .await?;
// Set the dedicated category field (not stored in metadata) // Set the dedicated category field (not stored in metadata)
@@ -108,7 +113,7 @@ impl Agent {
async fn handle_check_status( async fn handle_check_status(
&self, &self,
tenant: &crate::tenant::TenantCtx, user_id: &str,
job_id: Option<String>, job_id: Option<String>,
) -> Result<String, Error> { ) -> Result<String, Error> {
match job_id { match job_id {
@@ -117,8 +122,7 @@ impl Agent {
.map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?; .map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?;
// Try DB first for persistent state, fall back to ContextManager. // Try DB first for persistent state, fall back to ContextManager.
// TenantScope.get_job() auto-filters by ownership — no manual check needed. if let Some(store) = self.store()
if let Some(store) = tenant.store()
&& let Ok(Some(ctx)) = store.get_job(uuid).await && let Ok(Some(ctx)) = store.get_job(uuid).await
{ {
return Ok(format!( return Ok(format!(
@@ -134,7 +138,7 @@ impl Agent {
} }
let ctx = self.context_manager.get_context(uuid).await?; let ctx = self.context_manager.get_context(uuid).await?;
if ctx.user_id != tenant.user_id() { if ctx.user_id != user_id {
return Err(crate::error::JobError::NotFound { id: uuid }.into()); return Err(crate::error::JobError::NotFound { id: uuid }.into());
} }
@@ -151,8 +155,7 @@ impl Agent {
} }
None => { None => {
// Show summary from DB for consistency with Jobs tab. // Show summary from DB for consistency with Jobs tab.
// TenantScope methods auto-scope to user — no user_id parameter needed. if let Some(store) = self.store() {
if let Some(store) = tenant.store() {
let mut total = 0; let mut total = 0;
let mut in_progress = 0; let mut in_progress = 0;
let mut completed = 0; let mut completed = 0;
@@ -180,7 +183,7 @@ impl Agent {
} }
// Fallback to ContextManager if no DB. // Fallback to ContextManager if no DB.
let summary = self.context_manager.summary_for(tenant.user_id()).await; let summary = self.context_manager.summary_for(user_id).await;
Ok(format!( Ok(format!(
"Jobs summary: Total: {} In Progress: {} Completed: {} Failed: {} Stuck: {}", "Jobs summary: Total: {} In Progress: {} Completed: {} Failed: {} Stuck: {}",
summary.total, summary.total,
@@ -193,24 +196,19 @@ impl Agent {
} }
} }
async fn handle_cancel_job( async fn handle_cancel_job(&self, user_id: &str, job_id: &str) -> Result<String, Error> {
&self,
tenant: &crate::tenant::TenantCtx,
job_id: &str,
) -> Result<String, Error> {
let uuid = Uuid::parse_str(job_id) let uuid = Uuid::parse_str(job_id)
.map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?; .map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?;
let ctx = self.context_manager.get_context(uuid).await?; let ctx = self.context_manager.get_context(uuid).await?;
if ctx.user_id != tenant.user_id() { if ctx.user_id != user_id {
return Err(crate::error::JobError::NotFound { id: uuid }.into()); return Err(crate::error::JobError::NotFound { id: uuid }.into());
} }
self.scheduler.stop(uuid).await?; self.scheduler.stop(uuid).await?;
// Also update DB so the Jobs tab reflects cancellation immediately. // Also update DB so the Jobs tab reflects cancellation immediately.
// Use TenantScope — ownership already verified above. if let Some(store) = self.store()
if let Some(store) = tenant.store()
&& let Err(e) = store && let Err(e) = store
.update_job_status(uuid, JobState::Cancelled, Some("Cancelled by user")) .update_job_status(uuid, JobState::Cancelled, Some("Cancelled by user"))
.await .await
@@ -223,12 +221,11 @@ impl Agent {
async fn handle_list_jobs( async fn handle_list_jobs(
&self, &self,
tenant: &crate::tenant::TenantCtx, user_id: &str,
_filter: Option<String>, _filter: Option<String>,
) -> Result<String, Error> { ) -> Result<String, Error> {
// List from DB for consistency with Jobs tab. // List from DB for consistency with Jobs tab.
// TenantScope methods auto-scope to user. if let Some(store) = self.store() {
if let Some(store) = tenant.store() {
let agent_jobs = match store.list_agent_jobs().await { let agent_jobs = match store.list_agent_jobs().await {
Ok(jobs) => jobs, Ok(jobs) => jobs,
Err(e) => { Err(e) => {
@@ -259,7 +256,7 @@ impl Agent {
} }
// Fallback to ContextManager if no DB. // Fallback to ContextManager if no DB.
let jobs = self.context_manager.all_jobs_for(tenant.user_id()).await; let jobs = self.context_manager.all_jobs_for(user_id).await;
if jobs.is_empty() { if jobs.is_empty() {
return Ok("No jobs found.".to_string()); return Ok("No jobs found.".to_string());
} }
@@ -273,16 +270,12 @@ impl Agent {
Ok(output) Ok(output)
} }
async fn handle_help_job( async fn handle_help_job(&self, user_id: &str, job_id: &str) -> Result<String, Error> {
&self,
tenant: &crate::tenant::TenantCtx,
job_id: &str,
) -> Result<String, Error> {
let uuid = Uuid::parse_str(job_id) let uuid = Uuid::parse_str(job_id)
.map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?; .map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?;
let ctx = self.context_manager.get_context(uuid).await?; let ctx = self.context_manager.get_context(uuid).await?;
if ctx.user_id != tenant.user_id() { if ctx.user_id != user_id {
return Err(crate::error::JobError::NotFound { id: uuid }.into()); return Err(crate::error::JobError::NotFound { id: uuid }.into());
} }
@@ -315,11 +308,11 @@ impl Agent {
/// Show job status inline — either all jobs (no id) or a specific job. /// Show job status inline — either all jobs (no id) or a specific job.
pub(super) async fn process_job_status( pub(super) async fn process_job_status(
&self, &self,
tenant: &crate::tenant::TenantCtx, user_id: &str,
job_id: Option<&str>, job_id: Option<&str>,
) -> Result<SubmissionResult, Error> { ) -> Result<SubmissionResult, Error> {
match self match self
.handle_check_status(tenant, job_id.map(|s| s.to_string())) .handle_check_status(user_id, job_id.map(|s| s.to_string()))
.await .await
{ {
Ok(text) => Ok(SubmissionResult::response(text)), Ok(text) => Ok(SubmissionResult::response(text)),
@@ -330,10 +323,10 @@ impl Agent {
/// Cancel a job by ID. /// Cancel a job by ID.
pub(super) async fn process_job_cancel( pub(super) async fn process_job_cancel(
&self, &self,
tenant: &crate::tenant::TenantCtx, user_id: &str,
job_id: &str, job_id: &str,
) -> Result<SubmissionResult, Error> { ) -> Result<SubmissionResult, Error> {
match self.handle_cancel_job(tenant, job_id).await { match self.handle_cancel_job(user_id, job_id).await {
Ok(text) => Ok(SubmissionResult::response(text)), Ok(text) => Ok(SubmissionResult::response(text)),
Err(e) => Ok(SubmissionResult::error(format!("Cancel error: {}", e))), Err(e) => Ok(SubmissionResult::error(format!("Cancel error: {}", e))),
} }
@@ -472,101 +465,12 @@ 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<Mutex<Session>>,
thread_id: Uuid,
) -> SubmissionResult {
// Clone the turn data we need, then drop the session lock.
let turns_snapshot: Vec<(
usize,
Option<String>,
Vec<crate::agent::session::TurnToolCall>,
)>;
{
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::<usize>() {
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. /// Handle system commands that bypass thread-state checks entirely.
pub(super) async fn handle_system_command( pub(super) async fn handle_system_command(
&self, &self,
command: &str, command: &str,
args: &[String], args: &[String],
channel: &str, channel: &str,
tenant: &crate::tenant::TenantCtx,
) -> Result<SubmissionResult, Error> { ) -> Result<SubmissionResult, Error> {
match command { match command {
"help" => Ok(SubmissionResult::response(concat!( "help" => Ok(SubmissionResult::response(concat!(
@@ -576,7 +480,6 @@ impl Agent {
" /version Show version info\n", " /version Show version info\n",
" /tools List available tools\n", " /tools List available tools\n",
" /debug Toggle debug mode\n", " /debug Toggle debug mode\n",
" /reasoning [N|all] Show agent reasoning for turns\n",
" /ping Connectivity check\n", " /ping Connectivity check\n",
"\n", "\n",
"Jobs:\n", "Jobs:\n",
@@ -760,32 +663,19 @@ impl Agent {
} }
} }
if self.config.multi_tenant { match self.llm().set_model(requested) {
// Multi-tenant: only persist to per-user DB settings. Ok(()) => {
// Do NOT call set_model() on the shared provider — that // Persist the model choice so it survives restarts.
// would change the default for all users. The per-request self.persist_selected_model(requested).await;
// model_override in the dispatcher reads from the same Ok(SubmissionResult::response(format!(
// "selected_model" setting and applies it per-user. "Switched model to: {}",
self.persist_selected_model(tenant, requested).await; requested
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
))),
} }
} }
} }
@@ -927,14 +817,10 @@ impl Agent {
command: &str, command: &str,
args: &[String], args: &[String],
channel: &str, channel: &str,
tenant: &crate::tenant::TenantCtx,
) -> Result<Option<String>, Error> { ) -> Result<Option<String>, Error> {
// System commands are now handled directly via Submission::SystemCommand, // System commands are now handled directly via Submission::SystemCommand,
// but the router may still send us unknown /commands. // but the router may still send us unknown /commands.
match self match self.handle_system_command(command, args, channel).await? {
.handle_system_command(command, args, channel, tenant)
.await?
{
SubmissionResult::Response { content } => Ok(Some(content)), SubmissionResult::Response { content } => Ok(Some(content)),
SubmissionResult::Ok { message } => Ok(message), SubmissionResult::Ok { message } => Ok(message),
SubmissionResult::Error { message } => Ok(Some(format!("Error: {}", message))), SubmissionResult::Error { message } => Ok(Some(format!("Error: {}", message))),
@@ -946,69 +832,21 @@ impl Agent {
/// ///
/// Best-effort: logs warnings on failure but does not propagate errors, /// Best-effort: logs warnings on failure but does not propagate errors,
/// since the in-memory model switch already succeeded. /// since the in-memory model switch already succeeded.
/// async fn persist_selected_model(&self, model: &str) {
/// In multi-tenant mode, only the per-user DB setting is written — global // 1. Persist to DB if available.
/// .env and TOML files are shared across users and must not be mutated. if let Some(store) = self.store() {
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()); let value = serde_json::Value::String(model.to_string());
if let Err(e) = store.set_setting("selected_model", &value).await { if let Err(e) = store
.set_setting(self.owner_id(), "selected_model", &value)
.await
{
tracing::warn!("Failed to persist model to DB: {}", e); tracing::warn!("Failed to persist model to DB: {}", e);
} else {
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. In multi-tenant mode, skip .env/TOML writes — these are global // 2. Update TOML config file if it exists (sync I/O in spawn_blocking).
// 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 model_owned = model.to_string();
let backend = self.deps.llm_backend.clone();
if let Err(e) = tokio::task::spawn_blocking(move || { if let Err(e) = tokio::task::spawn_blocking(move || {
// 2a. Update the backend-specific model env var in ~/.ironclaw/.env.
//
// Env vars have the HIGHEST priority in LlmConfig::resolve_model()
// (env var > TOML > DB > default). If the .env file has e.g.
// NEARAI_MODEL=old-model, it shadows everything else. We must
// update this var or the /model change is invisible on restart.
let registry = crate::llm::ProviderRegistry::load();
let model_env = registry.model_env_var(&backend);
let env_var_prefix = format!("{}=", model_env);
// Only update the .env file if the var is actually set there
// (avoid injecting new vars the user never configured).
let env_path = crate::bootstrap::ironclaw_env_path();
let env_has_var = std::fs::read_to_string(&env_path)
.ok()
.is_some_and(|content| {
content.lines().any(|line| {
let trimmed = line.trim_start();
!trimmed.starts_with('#') && trimmed.starts_with(&env_var_prefix)
})
});
if env_has_var {
if let Err(e) = crate::bootstrap::upsert_bootstrap_var(model_env, &model_owned) {
tracing::warn!("Failed to update {} in .env: {}", model_env, e);
} else {
tracing::debug!("Updated {} in .env to {}", model_env, model_owned);
}
}
// 2b. Update (or create) the TOML config file.
//
// The TOML overlay has higher priority than DB settings on
// startup, so it MUST stay in sync with the DB.
let toml_path = crate::settings::Settings::default_toml_path(); let toml_path = crate::settings::Settings::default_toml_path();
match crate::settings::Settings::load_toml(&toml_path) { match crate::settings::Settings::load_toml(&toml_path) {
Ok(Some(mut settings)) => { Ok(Some(mut settings)) => {
@@ -1018,15 +856,7 @@ impl Agent {
} }
} }
Ok(None) => { Ok(None) => {
// No config file yet — create one so the model choice // No config file on disk; nothing to update.
// survives restarts even when the DB is unavailable.
let settings = crate::settings::Settings {
selected_model: Some(model_owned),
..Default::default()
};
if let Err(e) = settings.save_toml(&toml_path) {
tracing::warn!("Failed to create config.toml for model persistence: {}", e);
}
} }
Err(e) => { Err(e) => {
tracing::warn!("Failed to load config.toml for model persistence: {}", e); tracing::warn!("Failed to load config.toml for model persistence: {}", e);
@@ -1035,7 +865,7 @@ impl Agent {
}) })
.await .await
{ {
tracing::warn!("Model persistence task failed: {}", e); tracing::warn!("Model TOML persistence task failed: {}", e);
} }
} }
} }
+3 -236
View File
@@ -21,9 +21,6 @@ pub struct CostGuardConfig {
pub max_cost_per_day_cents: Option<u64>, pub max_cost_per_day_cents: Option<u64>,
/// Maximum LLM calls per hour. None = unlimited. /// Maximum LLM calls per hour. None = unlimited.
pub max_actions_per_hour: Option<u64>, pub max_actions_per_hour: Option<u64>,
/// 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<u64>,
} }
/// Error returned when a cost limit is exceeded. /// Error returned when a cost limit is exceeded.
@@ -33,12 +30,6 @@ pub enum CostLimitExceeded {
DailyBudget { spent_cents: u64, limit_cents: u64 }, DailyBudget { spent_cents: u64, limit_cents: u64 },
/// Hourly action rate limit reached. /// Hourly action rate limit reached.
HourlyRate { actions: u64, limit: u64 }, 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 { impl std::fmt::Display for CostLimitExceeded {
@@ -58,17 +49,6 @@ impl std::fmt::Display for CostLimitExceeded {
"Hourly action limit exceeded: {} actions of {} allowed per hour", "Hourly action limit exceeded: {} actions of {} allowed per hour",
actions, limit 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
),
} }
} }
} }
@@ -98,9 +78,6 @@ pub struct CostGuard {
/// Per-model token usage since startup. /// Per-model token usage since startup.
model_tokens: Mutex<HashMap<String, ModelTokens>>, model_tokens: Mutex<HashMap<String, ModelTokens>>,
/// Per-user daily cost tracking. Each entry resets independently at midnight UTC.
per_user_daily_cost: Mutex<HashMap<String, DailyCost>>,
} }
struct DailyCost { struct DailyCost {
@@ -120,7 +97,6 @@ impl CostGuard {
action_window: Mutex::new(VecDeque::new()), action_window: Mutex::new(VecDeque::new()),
budget_exceeded: AtomicBool::new(false), budget_exceeded: AtomicBool::new(false),
model_tokens: Mutex::new(HashMap::new()), model_tokens: Mutex::new(HashMap::new()),
per_user_daily_cost: Mutex::new(HashMap::new()),
} }
} }
@@ -227,11 +203,6 @@ impl CostGuard {
daily.reset_date = today; daily.reset_date = today;
self.budget_exceeded.store(false, Ordering::Relaxed); self.budget_exceeded.store(false, Ordering::Relaxed);
tracing::info!("Cost guard: daily counter reset for {}", today); 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; daily.total += cost;
@@ -277,85 +248,6 @@ impl CostGuard {
cost 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). /// Current daily spend in USD (as Decimal).
pub async fn daily_spend(&self) -> Decimal { pub async fn daily_spend(&self) -> Decimal {
let daily = self.daily_cost.lock().await; let daily = self.daily_cost.lock().await;
@@ -367,16 +259,6 @@ 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. /// Number of actions in the current hourly window.
pub async fn actions_this_hour(&self) -> u64 { pub async fn actions_this_hour(&self) -> u64 {
let mut window = self.action_window.lock().await; let mut window = self.action_window.lock().await;
@@ -432,7 +314,7 @@ mod tests {
async fn test_daily_budget_enforcement() { async fn test_daily_budget_enforcement() {
let guard = CostGuard::new(CostGuardConfig { let guard = CostGuard::new(CostGuardConfig {
max_cost_per_day_cents: Some(1), // $0.01 limit max_cost_per_day_cents: Some(1), // $0.01 limit
..CostGuardConfig::default() max_actions_per_hour: None,
}); });
// First call allowed // First call allowed
@@ -468,8 +350,8 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_hourly_rate_enforcement() { async fn test_hourly_rate_enforcement() {
let guard = CostGuard::new(CostGuardConfig { let guard = CostGuard::new(CostGuardConfig {
max_cost_per_day_cents: None,
max_actions_per_hour: Some(3), max_actions_per_hour: Some(3),
..CostGuardConfig::default()
}); });
// First 3 actions allowed // First 3 actions allowed
@@ -751,8 +633,8 @@ mod tests {
// A fresh CostGuard with rate limits should not panic even if // A fresh CostGuard with rate limits should not panic even if
// checked_sub returns None (simulating short uptime). // checked_sub returns None (simulating short uptime).
let guard = CostGuard::new(CostGuardConfig { let guard = CostGuard::new(CostGuardConfig {
max_cost_per_day_cents: None,
max_actions_per_hour: Some(100), max_actions_per_hour: Some(100),
..CostGuardConfig::default()
}); });
// These must not panic regardless of system uptime // These must not panic regardless of system uptime
@@ -774,119 +656,4 @@ mod tests {
let result = Instant::now().checked_sub(std::time::Duration::MAX); let result = Instant::now().checked_sub(std::time::Duration::MAX);
assert!(result.is_none()); 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"));
}
} }
+19 -173
View File
@@ -42,7 +42,6 @@ impl Agent {
pub(super) async fn run_agentic_loop( pub(super) async fn run_agentic_loop(
&self, &self,
message: &IncomingMessage, message: &IncomingMessage,
tenant: crate::tenant::TenantCtx,
session: Arc<Mutex<Session>>, session: Arc<Mutex<Session>>,
thread_id: Uuid, thread_id: Uuid,
initial_messages: Vec<ChatMessage>, initial_messages: Vec<ChatMessage>,
@@ -64,12 +63,7 @@ impl Agent {
); );
let system_prompt = if let Some(ws) = self.workspace() { let system_prompt = if let Some(ws) = self.workspace() {
let scoped_workspace = if ws.user_id() == message.user_id { match ws
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) .system_prompt_for_context_tz(is_group_chat, user_tz)
.await .await
{ {
@@ -169,7 +163,6 @@ impl Agent {
let delegate = ChatDelegate { let delegate = ChatDelegate {
agent: self, agent: self,
tenant,
session: session.clone(), session: session.clone(),
thread_id, thread_id,
message, message,
@@ -242,7 +235,6 @@ impl Agent {
/// auth intercept, and cost tracking. /// auth intercept, and cost tracking.
struct ChatDelegate<'a> { struct ChatDelegate<'a> {
agent: &'a Agent, agent: &'a Agent,
tenant: crate::tenant::TenantCtx,
session: Arc<Mutex<Session>>, session: Arc<Mutex<Session>>,
thread_id: Uuid, thread_id: Uuid,
message: &'a IncomingMessage, message: &'a IncomingMessage,
@@ -306,8 +298,6 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
// Update context for this iteration // Update context for this iteration
reason_ctx.available_tools = tool_defs; 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 { reason_ctx.system_prompt = Some(if force_text {
self.cached_prompt_no_tools.clone() self.cached_prompt_no_tools.clone()
} else { } else {
@@ -341,8 +331,8 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
reason_ctx: &mut ReasoningContext, reason_ctx: &mut ReasoningContext,
iteration: usize, iteration: usize,
) -> Result<crate::llm::RespondOutput, Error> { ) -> Result<crate::llm::RespondOutput, Error> {
// Enforce cost guardrails before the LLM call (global + per-user) // Enforce cost guardrails before the LLM call
if let Err(limit) = self.tenant.check_cost_allowed().await { if let Err(limit) = self.agent.cost_guard().check_allowed().await {
return Err(crate::error::LlmError::InvalidResponse { return Err(crate::error::LlmError::InvalidResponse {
provider: "agent".to_string(), provider: "agent".to_string(),
reason: limit.to_string(), reason: limit.to_string(),
@@ -350,21 +340,6 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
.into()); .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 { let output = match reasoning.respond_with_tools(reason_ctx).await {
Ok(output) => output, Ok(output) => output,
Err(crate::error::LlmError::ContextLengthExceeded { used, limit }) => { Err(crate::error::LlmError::ContextLengthExceeded { used, limit }) => {
@@ -399,27 +374,13 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
Err(e) => return Err(e.into()), Err(e) => return Err(e.into()),
}; };
// Record cost and track token usage (global + per-user). // Record cost and track token usage
// Use the provider's effective_model_name so cost attribution matches let model_name = self.agent.llm().active_model_name();
// the model that actually served the request. When the override is
// honoured (e.g. NearAI), this returns the override name; when the
// provider ignores overrides (e.g. Rig-based), it returns the active
// model, keeping attribution accurate in both cases.
let model_name = self
.agent
.llm()
.effective_model_name(reason_ctx.model_override.as_deref());
let cost_per_token = if reason_ctx.model_override.is_some() {
// Override may use different pricing; let CostGuard fall back to
// costs::model_cost() for the effective model.
None
} else {
Some(self.agent.llm().cost_per_token())
};
let read_discount = self.agent.llm().cache_read_discount(); let read_discount = self.agent.llm().cache_read_discount();
let write_multiplier = self.agent.llm().cache_write_multiplier(); let write_multiplier = self.agent.llm().cache_write_multiplier();
let call_cost = self let call_cost = self
.tenant .agent
.cost_guard()
.record_llm_call( .record_llm_call(
&model_name, &model_name,
output.usage.input_tokens, output.usage.input_tokens,
@@ -428,7 +389,7 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
output.usage.cache_creation_input_tokens, output.usage.cache_creation_input_tokens,
read_discount, read_discount,
write_multiplier, write_multiplier,
cost_per_token, Some(self.agent.llm().cost_per_token()),
) )
.await; .await;
tracing::debug!( tracing::debug!(
@@ -459,19 +420,6 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
content: Option<String>, content: Option<String>,
reason_ctx: &mut ReasoningContext, reason_ctx: &mut ReasoningContext,
) -> Result<Option<LoopOutcome>, Error> { ) -> Result<Option<LoopOutcome>, 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. // Add the assistant message with tool_calls to context.
// OpenAI protocol requires this before tool-result messages. // OpenAI protocol requires this before tool-result messages.
reason_ctx reason_ctx
@@ -492,41 +440,6 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
) )
.await; .await;
// Build per-tool decisions for the reasoning update.
// Sanitize each rationale through SafetyLayer (parity with JobDelegate).
let decisions: Vec<crate::channels::ToolDecision> = 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. // Record tool calls in the thread with sensitive params redacted.
{ {
let mut redacted_args: Vec<serde_json::Value> = Vec::with_capacity(tool_calls.len()); let mut redacted_args: Vec<serde_json::Value> = Vec::with_capacity(tool_calls.len());
@@ -542,23 +455,8 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
if let Some(thread) = sess.threads.get_mut(&self.thread_id) if let Some(thread) = sess.threads.get_mut(&self.thread_id)
&& let Some(turn) = thread.last_turn_mut() && 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) { for (tc, safe_args) in tool_calls.iter().zip(redacted_args) {
let sanitized_rationale = tc.reasoning.as_ref().map(|r| { turn.record_tool_call(&tc.name, safe_args);
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()),
);
} }
} }
} }
@@ -828,7 +726,7 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
if let Some(thread) = sess.threads.get_mut(&self.thread_id) if let Some(thread) = sess.threads.get_mut(&self.thread_id)
&& let Some(turn) = thread.last_turn_mut() && let Some(turn) = thread.last_turn_mut()
{ {
turn.record_tool_error_for(&tc.id, error_msg.clone()); turn.record_tool_error(error_msg.clone());
} }
} }
reason_ctx reason_ctx
@@ -954,19 +852,16 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
Err(e) => format!("Tool '{}' failed: {}", tc.name, e), Err(e) => format!("Tool '{}' failed: {}", tc.name, e),
}; };
// Record sanitized result in thread (identity-based matching). // Record sanitized result in thread
{ {
let mut sess = self.session.lock().await; let mut sess = self.session.lock().await;
if let Some(thread) = sess.threads.get_mut(&self.thread_id) if let Some(thread) = sess.threads.get_mut(&self.thread_id)
&& let Some(turn) = thread.last_turn_mut() && let Some(turn) = thread.last_turn_mut()
{ {
if is_tool_error { if is_tool_error {
turn.record_tool_error_for(&tc.id, result_content.clone()); turn.record_tool_error(result_content.clone());
} else { } else {
turn.record_tool_result_for( turn.record_tool_result(serde_json::json!(result_content));
&tc.id,
serde_json::json!(result_content),
);
} }
} }
} }
@@ -1203,23 +1098,15 @@ pub(crate) fn extract_suggestions(text: &str) -> (String, Vec<String>) {
Regex::new(r"(?s)<suggestions>\s*(.*?)\s*</suggestions>").expect("valid regex") // safety: constant pattern Regex::new(r"(?s)<suggestions>\s*(.*?)\s*</suggestions>").expect("valid regex") // safety: constant pattern
}); });
// Build a sorted list of code fence positions to determine open/close pairing. // Find the position of the last closing code fence to avoid matching inside code blocks
// A position is "inside" a fenced block when it falls between an odd-numbered let last_code_fence = text.rfind("```").unwrap_or(0);
// fence (opening) and the next even-numbered fence (closing).
let fence_positions: Vec<usize> = text.match_indices("```").map(|(pos, _)| pos).collect();
let is_inside_fence = |pos: usize| -> bool { // Find all matches, take the last one that's after the last code fence
// Count how many fences appear before `pos`. If odd, we're inside a fence.
let count = fence_positions.iter().take_while(|&&fp| fp <= pos).count();
count % 2 == 1
};
// Find all matches, take the last one that's outside any code fence
let mut best_match: Option<regex::Match<'_>> = None; let mut best_match: Option<regex::Match<'_>> = None;
let mut best_capture: Option<String> = None; let mut best_capture: Option<String> = None;
for caps in RE.captures_iter(text) { for caps in RE.captures_iter(text) {
if let (Some(full), Some(inner)) = (caps.get(0), caps.get(1)) if let (Some(full), Some(inner)) = (caps.get(0), caps.get(1))
&& !is_inside_fence(full.start()) && full.start() >= last_code_fence
{ {
best_match = Some(full); best_match = Some(full);
best_capture = Some(inner.as_str().to_string()); best_capture = Some(inner.as_str().to_string());
@@ -1338,8 +1225,6 @@ mod tests {
document_extraction: None, document_extraction: None,
sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig,
builder: None, builder: None,
llm_backend: "nearai".to_string(),
tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)),
}; };
Agent::new( Agent::new(
@@ -1355,15 +1240,10 @@ mod tests {
allow_local_tools: false, allow_local_tools: false,
max_cost_per_day_cents: None, max_cost_per_day_cents: None,
max_actions_per_hour: None, max_actions_per_hour: None,
max_cost_per_user_per_day_cents: None,
max_tool_iterations: 50, max_tool_iterations: 50,
auto_approve_tools: false, auto_approve_tools: false,
default_timezone: "UTC".to_string(), default_timezone: "UTC".to_string(),
max_jobs_per_user: None,
max_tokens_per_job: 0, max_tokens_per_job: 0,
multi_tenant: false,
max_llm_concurrent_per_user: None,
max_jobs_concurrent_per_user: None,
}, },
deps, deps,
Arc::new(ChannelManager::new()), Arc::new(ChannelManager::new()),
@@ -1573,13 +1453,11 @@ mod tests {
id: "call_2".to_string(), id: "call_2".to_string(),
name: "http".to_string(), name: "http".to_string(),
arguments: serde_json::json!({"url": "https://example.com"}), arguments: serde_json::json!({"url": "https://example.com"}),
reasoning: None,
}, },
ToolCall { ToolCall {
id: "call_3".to_string(), id: "call_3".to_string(),
name: "echo".to_string(), name: "echo".to_string(),
arguments: serde_json::json!({"message": "done"}), arguments: serde_json::json!({"message": "done"}),
reasoning: None,
}, },
], ],
user_timezone: None, user_timezone: None,
@@ -1765,7 +1643,6 @@ mod tests {
id: "call_1".to_string(), id: "call_1".to_string(),
name: "echo".to_string(), name: "echo".to_string(),
arguments: serde_json::json!({"message": "hi"}), arguments: serde_json::json!({"message": "hi"}),
reasoning: None,
}], }],
), ),
ChatMessage::tool_result("call_1", "echo", "hi"), ChatMessage::tool_result("call_1", "echo", "hi"),
@@ -1858,13 +1735,11 @@ mod tests {
id: "c1".to_string(), id: "c1".to_string(),
name: "http".to_string(), name: "http".to_string(),
arguments: serde_json::json!({}), arguments: serde_json::json!({}),
reasoning: None,
}, },
ToolCall { ToolCall {
id: "c2".to_string(), id: "c2".to_string(),
name: "echo".to_string(), name: "echo".to_string(),
arguments: serde_json::json!({}), arguments: serde_json::json!({}),
reasoning: None,
}, },
], ],
), ),
@@ -1898,7 +1773,6 @@ mod tests {
id: "c1".to_string(), id: "c1".to_string(),
name: "echo".to_string(), name: "echo".to_string(),
arguments: serde_json::json!({}), arguments: serde_json::json!({}),
reasoning: None,
}], }],
), ),
ChatMessage::tool_result("c1", "echo", "done"), ChatMessage::tool_result("c1", "echo", "done"),
@@ -2029,7 +1903,6 @@ mod tests {
id: crate::llm::generate_tool_call_id(0, 0), id: crate::llm::generate_tool_call_id(0, 0),
name: "echo".to_string(), name: "echo".to_string(),
arguments: serde_json::json!({"message": "looping"}), arguments: serde_json::json!({"message": "looping"}),
reasoning: None,
}], }],
input_tokens: 0, input_tokens: 0,
output_tokens: 5, output_tokens: 5,
@@ -2183,7 +2056,6 @@ mod tests {
id: crate::llm::generate_tool_call_id(0, 0), id: crate::llm::generate_tool_call_id(0, 0),
name: "nonexistent_tool".to_string(), name: "nonexistent_tool".to_string(),
arguments: serde_json::json!({}), arguments: serde_json::json!({}),
reasoning: None,
}], }],
input_tokens: 0, input_tokens: 0,
output_tokens: 5, output_tokens: 5,
@@ -2220,8 +2092,6 @@ mod tests {
document_extraction: None, document_extraction: None,
sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig,
builder: None, builder: None,
llm_backend: "nearai".to_string(),
tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)),
}; };
Agent::new( Agent::new(
@@ -2237,15 +2107,10 @@ mod tests {
allow_local_tools: false, allow_local_tools: false,
max_cost_per_day_cents: None, max_cost_per_day_cents: None,
max_actions_per_hour: None, max_actions_per_hour: None,
max_cost_per_user_per_day_cents: None,
max_tool_iterations, max_tool_iterations,
auto_approve_tools: true, auto_approve_tools: true,
default_timezone: "UTC".to_string(), default_timezone: "UTC".to_string(),
max_jobs_per_user: None,
max_tokens_per_job: 0, max_tokens_per_job: 0,
multi_tenant: false,
max_llm_concurrent_per_user: None,
max_jobs_concurrent_per_user: None,
}, },
deps, deps,
Arc::new(ChannelManager::new()), Arc::new(ChannelManager::new()),
@@ -2280,14 +2145,13 @@ mod tests {
let message = IncomingMessage::new("test", "test-user", "do something"); let message = IncomingMessage::new("test", "test-user", "do something");
let initial_messages = vec![ChatMessage::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 // The dispatcher must terminate within 5 seconds. If there is an
// infinite loop bug (e.g., index not advancing on tool failure), the // infinite loop bug (e.g., index not advancing on tool failure), the
// timeout will fire and the test will fail. // timeout will fire and the test will fail.
let result = tokio::time::timeout( let result = tokio::time::timeout(
Duration::from_secs(5), Duration::from_secs(5),
agent.run_agentic_loop(&message, tenant, session, thread_id, initial_messages), agent.run_agentic_loop(&message, session, thread_id, initial_messages),
) )
.await; .await;
@@ -2348,8 +2212,6 @@ mod tests {
document_extraction: None, document_extraction: None,
sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig,
builder: None, builder: None,
llm_backend: "nearai".to_string(),
tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)),
}; };
Agent::new( Agent::new(
@@ -2365,15 +2227,10 @@ mod tests {
allow_local_tools: false, allow_local_tools: false,
max_cost_per_day_cents: None, max_cost_per_day_cents: None,
max_actions_per_hour: None, max_actions_per_hour: None,
max_cost_per_user_per_day_cents: None,
max_tool_iterations: max_iter, max_tool_iterations: max_iter,
auto_approve_tools: true, auto_approve_tools: true,
default_timezone: "UTC".to_string(), default_timezone: "UTC".to_string(),
max_jobs_per_user: None,
max_tokens_per_job: 0, max_tokens_per_job: 0,
multi_tenant: false,
max_llm_concurrent_per_user: None,
max_jobs_concurrent_per_user: None,
}, },
deps, deps,
Arc::new(ChannelManager::new()), Arc::new(ChannelManager::new()),
@@ -2393,14 +2250,13 @@ mod tests {
let message = IncomingMessage::new("test", "test-user", "keep calling tools"); let message = IncomingMessage::new("test", "test-user", "keep calling tools");
let initial_messages = vec![ChatMessage::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 // Even with an LLM that always wants to call tools, the dispatcher
// must terminate within the timeout thanks to force_text at // must terminate within the timeout thanks to force_text at
// max_tool_iterations. // max_tool_iterations.
let result = tokio::time::timeout( let result = tokio::time::timeout(
Duration::from_secs(5), Duration::from_secs(5),
agent.run_agentic_loop(&message, tenant, session, thread_id, initial_messages), agent.run_agentic_loop(&message, session, thread_id, initial_messages),
) )
.await; .await;
@@ -2489,16 +2345,6 @@ mod tests {
assert!(suggestions.is_empty()); // safety: test assert!(suggestions.is_empty()); // safety: test
} }
#[test]
fn test_extract_suggestions_inside_unclosed_code_fence() {
// Regression: odd number of fences (unclosed fence) must still be
// treated as "inside a code block".
let input = "```\ncode\n<suggestions>[\"bar\"]</suggestions>";
let (text, suggestions) = super::extract_suggestions(input);
assert_eq!(text, input); // safety: test
assert!(suggestions.is_empty()); // safety: test
}
#[test] #[test]
fn test_extract_suggestions_after_code_fence() { fn test_extract_suggestions_after_code_fence() {
let input = "```\ncode\n```\nAnswer.\n<suggestions>[\"foo\"]</suggestions>"; let input = "```\ncode\n```\nAnswer.\n<suggestions>[\"foo\"]</suggestions>";
+6 -185
View File
@@ -31,8 +31,8 @@ use chrono_tz::Tz;
use tokio::sync::mpsc; use tokio::sync::mpsc;
use crate::channels::OutgoingResponse; use crate::channels::OutgoingResponse;
use crate::db::Database;
use crate::llm::{ChatMessage, CompletionRequest, LlmProvider, Reasoning}; use crate::llm::{ChatMessage, CompletionRequest, LlmProvider, Reasoning};
use crate::tenant::AdminScope;
use crate::workspace::Workspace; use crate::workspace::Workspace;
use crate::workspace::hygiene::HygieneConfig; use crate::workspace::hygiene::HygieneConfig;
@@ -57,9 +57,6 @@ pub struct HeartbeatConfig {
pub quiet_hours_end: Option<u32>, pub quiet_hours_end: Option<u32>,
/// Timezone for fire_at and quiet hours evaluation (IANA name). /// Timezone for fire_at and quiet hours evaluation (IANA name).
pub timezone: Option<String>, pub timezone: Option<String>,
/// 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 { impl Default for HeartbeatConfig {
@@ -74,7 +71,6 @@ impl Default for HeartbeatConfig {
quiet_hours_start: None, quiet_hours_start: None,
quiet_hours_end: None, quiet_hours_end: None,
timezone: None, timezone: None,
multi_tenant: false,
} }
} }
} }
@@ -182,7 +178,7 @@ pub struct HeartbeatRunner {
workspace: Arc<Workspace>, workspace: Arc<Workspace>,
llm: Arc<dyn LlmProvider>, llm: Arc<dyn LlmProvider>,
response_tx: Option<mpsc::Sender<OutgoingResponse>>, response_tx: Option<mpsc::Sender<OutgoingResponse>>,
store: Option<AdminScope>, store: Option<Arc<dyn Database>>,
consecutive_failures: u32, consecutive_failures: u32,
} }
@@ -211,8 +207,8 @@ impl HeartbeatRunner {
self self
} }
/// Set the admin-scoped database store for persistent heartbeat conversations. /// Set the database store for persistent heartbeat conversations.
pub fn with_store(mut self, store: AdminScope) -> Self { pub fn with_store(mut self, store: Arc<dyn Database>) -> Self {
self.store = Some(store); self.store = Some(store);
self self
} }
@@ -497,7 +493,7 @@ pub fn spawn_heartbeat(
workspace: Arc<Workspace>, workspace: Arc<Workspace>,
llm: Arc<dyn LlmProvider>, llm: Arc<dyn LlmProvider>,
response_tx: Option<mpsc::Sender<OutgoingResponse>>, response_tx: Option<mpsc::Sender<OutgoingResponse>>,
store: Option<AdminScope>, store: Option<Arc<dyn Database>>,
) -> tokio::task::JoinHandle<()> { ) -> tokio::task::JoinHandle<()> {
let mut runner = HeartbeatRunner::new(config, hygiene_config, workspace, llm); let mut runner = HeartbeatRunner::new(config, hygiene_config, workspace, llm);
if let Some(tx) = response_tx { if let Some(tx) = response_tx {
@@ -512,181 +508,6 @@ pub fn spawn_heartbeat(
}) })
} }
/// Spawn a multi-user heartbeat runner that cycles through all users who
/// have 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<dyn LlmProvider>,
response_tx: Option<mpsc::Sender<OutgoingResponse>>,
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<String, u32> =
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<String> = routines
.iter()
.map(|r| r.user_id.clone())
.collect::<std::collections::HashSet<_>>()
.into_iter()
.collect();
ids.sort();
ids
}
Err(e) => {
tracing::error!("Multi-user heartbeat: failed to list routines: {}", e);
continue;
}
};
// Run user heartbeats (and hygiene) concurrently so one slow LLM
// call doesn't block others. Cap concurrency to avoid flooding the
// LLM provider. Hygiene runs inside the same JoinSet so it is
// tracked and bounded by the same concurrency cap.
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())));
// 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();
// In multi-tenant mode, clear notify_user_id so that
// HeartbeatRunner::send_notification falls back to
// workspace.user_id() — each user's heartbeat should persist
// and notify that user, not the shared config target.
let mut cfg = config.clone();
cfg.notify_user_id = None;
let hyg = hygiene_config.clone();
let llm_clone = llm.clone();
let tx = response_tx.clone();
let admin = store.clone();
join_set.spawn(async move {
// Run memory hygiene per user (same as single-user heartbeat)
// inside the tracked task so concurrency is bounded.
let report = crate::workspace::hygiene::run_if_due(&workspace, &hyg).await;
if report.had_work() {
tracing::info!(
user_id = uid,
daily_logs_deleted = report.daily_logs_deleted,
conversation_docs_deleted = report.conversation_docs_deleted,
"multi-user heartbeat: memory hygiene deleted stale documents"
);
}
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<String, u32>,
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)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
@@ -905,7 +726,7 @@ mod tests {
Arc<crate::workspace::Workspace>, Arc<crate::workspace::Workspace>,
Arc<dyn crate::llm::LlmProvider>, Arc<dyn crate::llm::LlmProvider>,
Option<tokio::sync::mpsc::Sender<crate::channels::OutgoingResponse>>, Option<tokio::sync::mpsc::Sender<crate::channels::OutgoingResponse>>,
Option<AdminScope>, Option<Arc<dyn crate::db::Database>>,
) -> tokio::task::JoinHandle<()> = spawn_heartbeat; ) -> tokio::task::JoinHandle<()> = spawn_heartbeat;
let _ = _fn_ptr; let _ = _fn_ptr;
} }
+25 -35
View File
@@ -21,8 +21,8 @@ use tokio::task::JoinHandle;
use uuid::Uuid; use uuid::Uuid;
use crate::channels::IncomingMessage; use crate::channels::IncomingMessage;
use crate::channels::web::types::SseEvent;
use crate::context::{ContextManager, JobState}; use crate::context::{ContextManager, JobState};
use ironclaw_common::AppEvent;
/// Route context for forwarding job monitor events back to the user's channel. /// Route context for forwarding job monitor events back to the user's channel.
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
@@ -36,15 +36,15 @@ pub struct JobMonitorRoute {
/// injects assistant messages into the agent loop. /// injects assistant messages into the agent loop.
/// ///
/// The monitor forwards: /// The monitor forwards:
/// - `AppEvent::JobMessage` (assistant role): injected as incoming messages so /// - `SseEvent::JobMessage` (assistant role): injected as incoming messages so
/// the main agent can read and relay to the user. /// the main agent can read and relay to the user.
/// - `AppEvent::JobResult`: injected as a completion notice, then the task exits. /// - `SseEvent::JobResult`: injected as a completion notice, then the task exits.
/// ///
/// Tool use/result and status events are intentionally skipped (too noisy for /// Tool use/result and status events are intentionally skipped (too noisy for
/// the main agent's context window). /// the main agent's context window).
pub fn spawn_job_monitor( pub fn spawn_job_monitor(
job_id: Uuid, job_id: Uuid,
event_rx: broadcast::Receiver<(Uuid, String, AppEvent)>, event_rx: broadcast::Receiver<(Uuid, SseEvent)>,
inject_tx: mpsc::Sender<IncomingMessage>, inject_tx: mpsc::Sender<IncomingMessage>,
route: JobMonitorRoute, route: JobMonitorRoute,
) -> JoinHandle<()> { ) -> JoinHandle<()> {
@@ -56,7 +56,7 @@ pub fn spawn_job_monitor(
/// jobs don't stay `InProgress` forever in the `ContextManager`. /// jobs don't stay `InProgress` forever in the `ContextManager`.
pub fn spawn_job_monitor_with_context( pub fn spawn_job_monitor_with_context(
job_id: Uuid, job_id: Uuid,
mut event_rx: broadcast::Receiver<(Uuid, String, AppEvent)>, mut event_rx: broadcast::Receiver<(Uuid, SseEvent)>,
inject_tx: mpsc::Sender<IncomingMessage>, inject_tx: mpsc::Sender<IncomingMessage>,
route: JobMonitorRoute, route: JobMonitorRoute,
context_manager: Option<Arc<ContextManager>>, context_manager: Option<Arc<ContextManager>>,
@@ -68,13 +68,13 @@ pub fn spawn_job_monitor_with_context(
loop { loop {
match event_rx.recv().await { match event_rx.recv().await {
Ok((ev_job_id, _user_id, event)) => { Ok((ev_job_id, event)) => {
if ev_job_id != job_id { if ev_job_id != job_id {
continue; continue;
} }
match event { match event {
AppEvent::JobMessage { role, content, .. } if role == "assistant" => { SseEvent::JobMessage { role, content, .. } if role == "assistant" => {
let mut msg = IncomingMessage::new( let mut msg = IncomingMessage::new(
route.channel.clone(), route.channel.clone(),
route.user_id.clone(), route.user_id.clone(),
@@ -92,7 +92,7 @@ pub fn spawn_job_monitor_with_context(
break; break;
} }
} }
AppEvent::JobResult { status, .. } => { SseEvent::JobResult { status, .. } => {
// Transition in-memory state so the job frees its // Transition in-memory state so the job frees its
// max_jobs slot and query tools show the final state. // max_jobs slot and query tools show the final state.
if let Some(ref cm) = context_manager { 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. /// inject messages into) but we still need to free the `max_jobs` slot.
pub fn spawn_completion_watcher( pub fn spawn_completion_watcher(
job_id: Uuid, job_id: Uuid,
mut event_rx: broadcast::Receiver<(Uuid, String, AppEvent)>, mut event_rx: broadcast::Receiver<(Uuid, SseEvent)>,
context_manager: Arc<ContextManager>, context_manager: Arc<ContextManager>,
) -> JoinHandle<()> { ) -> JoinHandle<()> {
let short_id = job_id.to_string()[..8].to_string(); let short_id = job_id.to_string()[..8].to_string();
@@ -170,9 +170,7 @@ pub fn spawn_completion_watcher(
tokio::spawn(async move { tokio::spawn(async move {
loop { loop {
match event_rx.recv().await { match event_rx.recv().await {
Ok((ev_job_id, _user_id, AppEvent::JobResult { status, .. })) Ok((ev_job_id, SseEvent::JobResult { status, .. })) if ev_job_id == job_id => {
if ev_job_id == job_id =>
{
let target = if status == "completed" { let target = if status == "completed" {
JobState::Completed JobState::Completed
} else { } else {
@@ -229,7 +227,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_monitor_forwards_assistant_messages() { async fn test_monitor_forwards_assistant_messages() {
let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16);
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16); let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let job_id = Uuid::new_v4(); let job_id = Uuid::new_v4();
@@ -239,8 +237,7 @@ mod tests {
event_tx event_tx
.send(( .send((
job_id, job_id,
"test-user".to_string(), SseEvent::JobMessage {
AppEvent::JobMessage {
job_id: job_id.to_string(), job_id: job_id.to_string(),
role: "assistant".to_string(), role: "assistant".to_string(),
content: "I found a bug".to_string(), content: "I found a bug".to_string(),
@@ -262,7 +259,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_monitor_ignores_other_jobs() { async fn test_monitor_ignores_other_jobs() {
let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16);
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16); let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let job_id = Uuid::new_v4(); let job_id = Uuid::new_v4();
@@ -273,8 +270,7 @@ mod tests {
event_tx event_tx
.send(( .send((
other_job_id, other_job_id,
"test-user".to_string(), SseEvent::JobMessage {
AppEvent::JobMessage {
job_id: other_job_id.to_string(), job_id: other_job_id.to_string(),
role: "assistant".to_string(), role: "assistant".to_string(),
content: "wrong job".to_string(), content: "wrong job".to_string(),
@@ -293,7 +289,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_monitor_exits_on_job_result() { async fn test_monitor_exits_on_job_result() {
let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16);
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16); let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let job_id = Uuid::new_v4(); let job_id = Uuid::new_v4();
@@ -303,8 +299,7 @@ mod tests {
event_tx event_tx
.send(( .send((
job_id, job_id,
"test-user".to_string(), SseEvent::JobResult {
AppEvent::JobResult {
job_id: job_id.to_string(), job_id: job_id.to_string(),
status: "completed".to_string(), status: "completed".to_string(),
session_id: None, session_id: None,
@@ -329,7 +324,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_monitor_skips_tool_events() { async fn test_monitor_skips_tool_events() {
let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16);
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16); let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let job_id = Uuid::new_v4(); let job_id = Uuid::new_v4();
@@ -339,8 +334,7 @@ mod tests {
event_tx event_tx
.send(( .send((
job_id, job_id,
"test-user".to_string(), SseEvent::JobToolUse {
AppEvent::JobToolUse {
job_id: job_id.to_string(), job_id: job_id.to_string(),
tool_name: "shell".to_string(), tool_name: "shell".to_string(),
input: serde_json::json!({"command": "ls"}), input: serde_json::json!({"command": "ls"}),
@@ -352,8 +346,7 @@ mod tests {
event_tx event_tx
.send(( .send((
job_id, job_id,
"test-user".to_string(), SseEvent::JobMessage {
AppEvent::JobMessage {
job_id: job_id.to_string(), job_id: job_id.to_string(),
role: "user".to_string(), role: "user".to_string(),
content: "user prompt".to_string(), content: "user prompt".to_string(),
@@ -409,7 +402,7 @@ mod tests {
.await .await
.unwrap(); .unwrap();
let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16);
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16); let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let handle = spawn_job_monitor_with_context( let handle = spawn_job_monitor_with_context(
@@ -424,8 +417,7 @@ mod tests {
event_tx event_tx
.send(( .send((
job_id, job_id,
"test-user".to_string(), SseEvent::JobResult {
AppEvent::JobResult {
job_id: job_id.to_string(), job_id: job_id.to_string(),
status: "completed".to_string(), status: "completed".to_string(),
session_id: None, session_id: None,
@@ -458,7 +450,7 @@ mod tests {
.await .await
.unwrap(); .unwrap();
let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16);
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16); let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let handle = spawn_job_monitor_with_context( let handle = spawn_job_monitor_with_context(
@@ -473,8 +465,7 @@ mod tests {
event_tx event_tx
.send(( .send((
job_id, job_id,
"test-user".to_string(), SseEvent::JobResult {
AppEvent::JobResult {
job_id: job_id.to_string(), job_id: job_id.to_string(),
status: "failed".to_string(), status: "failed".to_string(),
session_id: None, session_id: None,
@@ -507,14 +498,13 @@ mod tests {
.await .await
.unwrap(); .unwrap();
let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16);
let handle = spawn_completion_watcher(job_id, event_tx.subscribe(), Arc::clone(&cm)); let handle = spawn_completion_watcher(job_id, event_tx.subscribe(), Arc::clone(&cm));
event_tx event_tx
.send(( .send((
job_id, job_id,
"test-user".to_string(), SseEvent::JobResult {
AppEvent::JobResult {
job_id: job_id.to_string(), job_id: job_id.to_string(),
status: "completed".to_string(), status: "completed".to_string(),
session_id: None, session_id: None,
+1 -3
View File
@@ -36,9 +36,7 @@ pub(crate) use agent_loop::truncate_for_preview;
pub use agent_loop::{Agent, AgentDeps}; pub use agent_loop::{Agent, AgentDeps};
pub use compaction::{CompactionResult, ContextCompactor}; pub use compaction::{CompactionResult, ContextCompactor};
pub use context_monitor::{CompactionStrategy, ContextBreakdown, ContextMonitor}; pub use context_monitor::{CompactionStrategy, ContextBreakdown, ContextMonitor};
pub use heartbeat::{ pub use heartbeat::{HeartbeatConfig, HeartbeatResult, HeartbeatRunner, spawn_heartbeat};
HeartbeatConfig, HeartbeatResult, HeartbeatRunner, spawn_heartbeat, spawn_multi_user_heartbeat,
};
pub use router::{MessageIntent, Router}; pub use router::{MessageIntent, Router};
pub use routine::{Routine, RoutineAction, RoutineRun, Trigger}; pub use routine::{Routine, RoutineAction, RoutineRun, Trigger};
pub use routine_engine::{RoutineEngine, SandboxReadiness}; pub use routine_engine::{RoutineEngine, SandboxReadiness};
+39 -297
View File
@@ -18,22 +18,21 @@ use std::time::Duration;
use chrono::Utc; use chrono::Utc;
use regex::Regex; use regex::Regex;
use tokio::sync::{RwLock, mpsc}; use tokio::sync::{RwLock, mpsc};
use tokio::task::JoinHandle;
use uuid::Uuid; use uuid::Uuid;
use crate::agent::Scheduler; use crate::agent::Scheduler;
use crate::agent::routine::{ use crate::agent::routine::{
NotifyConfig, Routine, RoutineAction, RoutineRun, RunStatus, Trigger, next_cron_fire, NotifyConfig, Routine, RoutineAction, RoutineRun, RunStatus, Trigger, next_cron_fire,
}; };
use crate::channels::{IncomingMessage, OutgoingResponse}; use crate::channels::OutgoingResponse;
use crate::config::RoutineConfig; use crate::config::RoutineConfig;
use crate::context::{JobContext, JobState}; use crate::context::{JobContext, JobState};
use crate::db::Database;
use crate::error::RoutineError; use crate::error::RoutineError;
use crate::extensions::ExtensionManager; use crate::extensions::ExtensionManager;
use crate::llm::{ use crate::llm::{
ChatMessage, CompletionRequest, FinishReason, LlmProvider, ToolCall, ToolCompletionRequest, ChatMessage, CompletionRequest, FinishReason, LlmProvider, ToolCall, ToolCompletionRequest,
}; };
use crate::tenant::AdminScope;
use crate::tools::{ use crate::tools::{
ToolError, ToolRegistry, autonomous_allowed_tool_names, autonomous_unavailable_message, ToolError, ToolRegistry, autonomous_allowed_tool_names, autonomous_unavailable_message,
prepare_tool_params, prepare_tool_params,
@@ -46,11 +45,6 @@ enum EventMatcher {
System { routine: Routine }, System { routine: Routine },
} }
struct TriggeredRoutine {
routine: Routine,
detail: String,
}
/// Distinguishes why sandbox is unavailable so error messages are accurate. /// Distinguishes why sandbox is unavailable so error messages are accurate.
#[derive(Debug, Clone, Copy, PartialEq, Eq)] #[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SandboxReadiness { pub enum SandboxReadiness {
@@ -62,44 +56,10 @@ pub enum SandboxReadiness {
DockerUnavailable, DockerUnavailable,
} }
/// Check whether an event-triggered routine's user/channel filters match an
/// incoming message.
///
/// Returns `true` if:
/// - The routine has an `Event` trigger (non-Event routines always return `false`)
/// - The routine's `user_id` matches the message's user scope
/// - The routine's channel filter (if any) matches the message channel
/// case-insensitively
///
/// This is a pure function extracted from `check_event_triggers` so the
/// filter logic can be unit-tested without async infrastructure.
pub(crate) fn routine_matches_message(routine: &Routine, message: &IncomingMessage) -> bool {
// Only Event-triggered routines can match incoming messages.
if !matches!(routine.trigger, Trigger::Event { .. }) {
return false;
}
// User ownership filter — only fire routines scoped to this user.
if routine.user_id != message.user_id {
return false;
}
// Channel filter (case-insensitive, matching emit_system_event behavior)
if let Trigger::Event {
channel: Some(ch), ..
} = &routine.trigger
&& !ch.eq_ignore_ascii_case(&message.channel)
{
return false;
}
true
}
/// The routine execution engine. /// The routine execution engine.
pub struct RoutineEngine { pub struct RoutineEngine {
config: RoutineConfig, config: RoutineConfig,
store: AdminScope, store: Arc<dyn Database>,
llm: Arc<dyn LlmProvider>, llm: Arc<dyn LlmProvider>,
workspace: Arc<Workspace>, workspace: Arc<Workspace>,
/// Sender for notifications (routed to channel manager). /// Sender for notifications (routed to channel manager).
@@ -128,7 +88,7 @@ impl RoutineEngine {
#[allow(clippy::too_many_arguments)] #[allow(clippy::too_many_arguments)]
pub fn new( pub fn new(
config: RoutineConfig, config: RoutineConfig,
store: AdminScope, store: Arc<dyn Database>,
llm: Arc<dyn LlmProvider>, llm: Arc<dyn LlmProvider>,
workspace: Arc<Workspace>, workspace: Arc<Workspace>,
notify_tx: mpsc::Sender<OutgoingResponse>, notify_tx: mpsc::Sender<OutgoingResponse>,
@@ -207,45 +167,10 @@ impl RoutineEngine {
} }
/// Check incoming message against event triggers. Returns number of routines fired. /// 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 /// Accepts only the three fields needed for matching (user scope, channel,
/// background event-triggered routines finish. /// message content) so callers never need to clone a full `IncomingMessage`.
pub async fn check_event_triggers_and_wait( pub async fn check_event_triggers(&self, user_id: &str, channel: &str, content: &str) -> usize {
&self,
message: &IncomingMessage,
content: &str,
) -> usize {
let triggered = self.matching_event_triggers(message, content).await;
let fired = triggered.len();
let handles: Vec<JoinHandle<()>> = 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<TriggeredRoutine> {
let cache = self.event_cache.read().await; let cache = self.event_cache.read().await;
// Early return if there are no message matchers at all. // Early return if there are no message matchers at all.
@@ -253,9 +178,10 @@ impl RoutineEngine {
.iter() .iter()
.any(|m| matches!(m, EventMatcher::Message { .. })) .any(|m| matches!(m, EventMatcher::Message { .. }))
{ {
return Vec::new(); return 0;
} }
let mut triggered = Vec::new();
let mut fired = 0;
// Collect routine IDs for batch query // Collect routine IDs for batch query
let routine_ids: Vec<Uuid> = cache let routine_ids: Vec<Uuid> = cache
@@ -267,13 +193,13 @@ impl RoutineEngine {
.collect(); .collect();
if routine_ids.is_empty() { if routine_ids.is_empty() {
return Vec::new(); return 0;
} }
// Single batch query instead of N queries // Single batch query instead of N queries
let concurrent_counts = match self.batch_concurrent_counts(&routine_ids).await { let concurrent_counts = match self.batch_concurrent_counts(&routine_ids).await {
Some(counts) => counts, Some(counts) => counts,
None => return Vec::new(), None => return 0,
}; };
for matcher in cache.iter() { for matcher in cache.iter() {
@@ -282,24 +208,16 @@ impl RoutineEngine {
EventMatcher::System { .. } => continue, EventMatcher::System { .. } => continue,
}; };
// User ownership + channel filter (extracted for testability). if routine.user_id != user_id {
if !routine_matches_message(routine, message) { continue;
// User mismatch is expected for multi-user setups — keep at }
// trace to avoid one log per routine per inbound message.
if routine.user_id != message.user_id { // Channel filter
tracing::trace!( if let Trigger::Event {
routine = %routine.name, channel: Some(ch), ..
routine_user = %routine.user_id, } = &routine.trigger
message_user = %message.user_id, && ch != channel
"Skipped: user scope mismatch" {
);
} else {
tracing::debug!(
routine = %routine.name,
channel = %message.channel,
"Skipped: channel mismatch"
);
}
continue; continue;
} }
@@ -310,14 +228,14 @@ impl RoutineEngine {
// Cooldown check // Cooldown check
if !self.check_cooldown(routine) { if !self.check_cooldown(routine) {
tracing::debug!(routine = %routine.name, "Skipped: cooldown active"); tracing::trace!(routine = %routine.name, "Skipped: cooldown active");
continue; continue;
} }
// Concurrent run check (using batch-loaded counts) // Concurrent run check (using batch-loaded counts)
let running_count = concurrent_counts.get(&routine.id).copied().unwrap_or(0); let running_count = concurrent_counts.get(&routine.id).copied().unwrap_or(0);
if running_count >= routine.guardrails.max_concurrent as i64 { if running_count >= routine.guardrails.max_concurrent as i64 {
tracing::debug!(routine = %routine.name, "Skipped: max concurrent reached"); tracing::trace!(routine = %routine.name, "Skipped: max concurrent reached");
continue; continue;
} }
@@ -328,13 +246,11 @@ impl RoutineEngine {
} }
let detail = truncate(content, 200); let detail = truncate(content, 200);
triggered.push(TriggeredRoutine { self.spawn_fire(routine.clone(), "event", Some(detail));
routine: routine.clone(), fired += 1;
detail,
});
} }
triggered fired
} }
/// Emit a structured event to system-event routines. /// Emit a structured event to system-event routines.
@@ -782,22 +698,12 @@ 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) // Execute inline for manual triggers (caller wants to wait)
let engine = EngineContext { let engine = EngineContext {
config: self.config.clone(), config: self.config.clone(),
store: self.store.clone(), store: self.store.clone(),
llm: self.llm.clone(), llm: self.llm.clone(),
workspace: routine_workspace, workspace: self.workspace.clone(),
notify_tx: self.notify_tx.clone(), notify_tx: self.notify_tx.clone(),
running_count: self.running_count.clone(), running_count: self.running_count.clone(),
scheduler: self.scheduler.clone(), scheduler: self.scheduler.clone(),
@@ -900,12 +806,7 @@ impl RoutineEngine {
} }
/// Spawn a fire in a background task. /// Spawn a fire in a background task.
fn spawn_fire( fn spawn_fire(&self, routine: Routine, trigger_type: &str, trigger_detail: Option<String>) {
&self,
routine: Routine,
trigger_type: &str,
trigger_detail: Option<String>,
) -> JoinHandle<()> {
let run = RoutineRun { let run = RoutineRun {
id: Uuid::new_v4(), id: Uuid::new_v4(),
routine_id: routine.id, routine_id: routine.id,
@@ -920,23 +821,11 @@ impl RoutineEngine {
created_at: Utc::now(), 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 { let engine = EngineContext {
config: self.config.clone(), config: self.config.clone(),
store: self.store.clone(), store: self.store.clone(),
llm: self.llm.clone(), llm: self.llm.clone(),
workspace: routine_workspace, workspace: self.workspace.clone(),
notify_tx: self.notify_tx.clone(), notify_tx: self.notify_tx.clone(),
running_count: self.running_count.clone(), running_count: self.running_count.clone(),
scheduler: self.scheduler.clone(), scheduler: self.scheduler.clone(),
@@ -954,7 +843,7 @@ impl RoutineEngine {
return; return;
} }
execute_routine(engine, routine, run).await; execute_routine(engine, routine, run).await;
}) });
} }
fn check_cooldown(&self, routine: &Routine) -> bool { fn check_cooldown(&self, routine: &Routine) -> bool {
@@ -989,7 +878,7 @@ impl RoutineEngine {
/// an active state (Pending/InProgress/Stuck). Maps the final `JobState` to /// an active state (Pending/InProgress/Stuck). Maps the final `JobState` to
/// a `RunStatus` for the routine run. /// a `RunStatus` for the routine run.
struct FullJobWatcher { struct FullJobWatcher {
store: AdminScope, store: Arc<dyn Database>,
job_id: Uuid, job_id: Uuid,
routine_name: String, routine_name: String,
} }
@@ -1000,7 +889,7 @@ impl FullJobWatcher {
/// Safety ceiling: 24 hours, derived from POLL_INTERVAL. /// Safety ceiling: 24 hours, derived from POLL_INTERVAL.
const MAX_POLLS: u32 = (24 * 60 * 60) / Self::POLL_INTERVAL.as_secs() as u32; const MAX_POLLS: u32 = (24 * 60 * 60) / Self::POLL_INTERVAL.as_secs() as u32;
fn new(store: AdminScope, job_id: Uuid, routine_name: String) -> Self { fn new(store: Arc<dyn Database>, job_id: Uuid, routine_name: String) -> Self {
Self { Self {
store, store,
job_id, job_id,
@@ -1072,7 +961,7 @@ impl FullJobWatcher {
/// Shared context passed to the execution function. /// Shared context passed to the execution function.
struct EngineContext { struct EngineContext {
config: RoutineConfig, config: RoutineConfig,
store: AdminScope, store: Arc<dyn Database>,
llm: Arc<dyn LlmProvider>, llm: Arc<dyn LlmProvider>,
workspace: Arc<Workspace>, workspace: Arc<Workspace>,
notify_tx: mpsc::Sender<OutgoingResponse>, notify_tx: mpsc::Sender<OutgoingResponse>,
@@ -1416,19 +1305,6 @@ async fn execute_lightweight(
} }
} }
/// Sanitize a user-controlled string before interpolation into an LLM prompt.
/// Strips newlines (which could break prompt structure) and truncates to a
/// reasonable length to limit abuse surface.
fn sanitize_prompt_field(value: &str) -> String {
const MAX_LEN: usize = 128;
value
.chars()
.filter(|&c| c != '\n' && c != '\r')
.take(MAX_LEN)
.map(|c| if c == '`' { '\'' } else { c })
.collect()
}
fn build_lightweight_prompt( fn build_lightweight_prompt(
prompt: &str, prompt: &str,
context_parts: &[String], context_parts: &[String],
@@ -1447,16 +1323,14 @@ fn build_lightweight_prompt(
); );
if let Some(channel) = notify.channel.as_deref() { if let Some(channel) = notify.channel.as_deref() {
let sanitized = sanitize_prompt_field(channel);
full_prompt.push_str(&format!( full_prompt.push_str(&format!(
"The configured delivery channel for this routine is `{sanitized}`.\n" "The configured delivery channel for this routine is `{channel}`.\n"
)); ));
} }
if let Some(user) = notify.user.as_deref() { if let Some(user) = notify.user.as_deref() {
let sanitized = sanitize_prompt_field(user);
full_prompt.push_str(&format!( full_prompt.push_str(&format!(
"The configured delivery target for this routine is `{sanitized}`.\n" "The configured delivery target for this routine is `{user}`.\n"
)); ));
} }
@@ -1566,7 +1440,6 @@ fn handle_text_response(
/// This is a simplified version of the full dispatcher loop: /// This is a simplified version of the full dispatcher loop:
/// - Max 3-5 iterations (configurable) /// - Max 3-5 iterations (configurable)
/// - Sequential tool execution (not parallel) /// - Sequential tool execution (not parallel)
/// - Uses the owner's live autonomous tool scope when lightweight tools are enabled
/// - Auto-approval of non-Always tools /// - Auto-approval of non-Always tools
/// - No hooks or approval dialogs /// - No hooks or approval dialogs
async fn execute_lightweight_with_tools( async fn execute_lightweight_with_tools(
@@ -1613,10 +1486,7 @@ async fn execute_lightweight_with_tools(
let force_text = iteration >= max_iterations; let force_text = iteration >= max_iterations;
if force_text { 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) let request = CompletionRequest::new(messages)
.with_max_tokens(effective_max_tokens) .with_max_tokens(effective_max_tokens)
.with_temperature(0.3); .with_temperature(0.3);
@@ -1895,13 +1765,6 @@ pub fn spawn_cron_ticker(
engine.check_cron_triggers().await; engine.check_cron_triggers().await;
let mut ticker = tokio::time::interval(interval); let mut ticker = tokio::time::interval(interval);
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
// Periodic event cache refresh so web/CLI mutations are picked up
// without requiring tool-path code to call refresh_event_cache().
// Uses wall-clock elapsed time so the refresh cadence is stable
// regardless of the cron tick interval configuration.
let refresh_interval = Duration::from_secs(60);
let mut last_refresh = tokio::time::Instant::now();
loop { loop {
ticker.tick().await; ticker.tick().await;
@@ -1909,11 +1772,7 @@ pub fn spawn_cron_ticker(
// never races with FullJobWatcher instances from this process. // never races with FullJobWatcher instances from this process.
engine.sync_dispatched_runs().await; engine.sync_dispatched_runs().await;
engine.check_cron_triggers().await; engine.check_cron_triggers().await;
engine.sync_dispatched_runs().await;
if last_refresh.elapsed() >= refresh_interval {
engine.refresh_event_cache().await;
last_refresh = tokio::time::Instant::now();
}
} }
}) })
} }
@@ -1979,13 +1838,7 @@ fn strip_html_tags(s: &str) -> String {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use chrono::Utc; use crate::agent::routine::{NotifyConfig, RunStatus};
use uuid::Uuid;
use crate::agent::routine::{
NotifyConfig, Routine, RoutineAction, RoutineGuardrails, RunStatus, Trigger,
};
use crate::channels::IncomingMessage;
use crate::config::RoutineConfig; use crate::config::RoutineConfig;
#[test] #[test]
@@ -2183,117 +2036,6 @@ mod tests {
} }
} }
/// Helper to build a test routine with the given user_id and trigger.
fn make_routine(user_id: &str, trigger: Trigger) -> Routine {
Routine {
id: Uuid::new_v4(),
name: "test".to_string(),
description: String::new(),
user_id: user_id.to_string(),
enabled: true,
trigger,
action: RoutineAction::Lightweight {
prompt: String::new(),
context_paths: vec![],
max_tokens: 1000,
use_tools: false,
max_tool_rounds: 0,
},
guardrails: RoutineGuardrails::default(),
notify: Default::default(),
last_run_at: None,
next_fire_at: None,
run_count: 0,
consecutive_failures: 0,
state: serde_json::Value::Null,
created_at: Utc::now(),
updated_at: Utc::now(),
}
}
/// Helper to build a test IncomingMessage.
fn make_message(user_id: &str, channel: &str, content: &str) -> IncomingMessage {
IncomingMessage {
id: Uuid::new_v4(),
channel: channel.to_string(),
user_id: user_id.to_string(),
owner_id: user_id.to_string(),
sender_id: user_id.to_string(),
user_name: None,
content: content.to_string(),
thread_id: None,
conversation_scope_id: None,
received_at: Utc::now(),
metadata: serde_json::Value::Null,
timezone: None,
attachments: vec![],
is_internal: false,
}
}
/// Regression test for issue #1051: event triggers used case-sensitive
/// channel comparison, so "Telegram" != "telegram" caused silent mismatch.
/// Tests the actual `routine_matches_message` function used in `check_event_triggers`.
#[test]
fn test_channel_filter_is_case_insensitive() {
let routine = make_routine(
"user1",
Trigger::Event {
pattern: ".*".to_string(),
channel: Some("Telegram".to_string()),
},
);
let msg = make_message("user1", "telegram", "hello");
// Case-insensitive channel match must succeed
assert!(super::routine_matches_message(&routine, &msg));
// Exact case must also work
let msg_exact = make_message("user1", "Telegram", "hello");
assert!(super::routine_matches_message(&routine, &msg_exact));
// Different channel must not match
let msg_wrong = make_message("user1", "discord", "hello");
assert!(!super::routine_matches_message(&routine, &msg_wrong));
}
/// Regression test for issue #1051: event triggers did not filter by
/// user_id, so routines from user A could fire on messages from user B.
/// Tests the actual `routine_matches_message` function used in `check_event_triggers`.
#[test]
fn test_event_trigger_requires_user_match() {
let routine = make_routine(
"alice",
Trigger::Event {
pattern: ".*".to_string(),
channel: None,
},
);
// Different user must not match
let msg_bob = make_message("bob", "telegram", "hello");
assert!(!super::routine_matches_message(&routine, &msg_bob));
// Same user must match
let msg_alice = make_message("alice", "telegram", "hello");
assert!(super::routine_matches_message(&routine, &msg_alice));
}
/// When no channel filter is set, any channel should match (given user matches).
#[test]
fn test_no_channel_filter_matches_any_channel() {
let routine = make_routine(
"user1",
Trigger::Event {
pattern: ".*".to_string(),
channel: None,
},
);
let msg = make_message("user1", "whatever_channel", "hello");
assert!(super::routine_matches_message(&routine, &msg));
}
#[test] #[test]
fn test_routine_tool_denylist_blocks_self_management_tools() { fn test_routine_tool_denylist_blocks_self_management_tools() {
let denylisted = vec![ let denylisted = vec![
+9 -27
View File
@@ -9,14 +9,15 @@ use tokio::task::JoinHandle;
use uuid::Uuid; use uuid::Uuid;
use crate::agent::task::{Task, TaskContext, TaskOutput}; use crate::agent::task::{Task, TaskContext, TaskOutput};
use crate::channels::web::types::SseEvent;
use crate::config::AgentConfig; use crate::config::AgentConfig;
use crate::context::{ContextManager, JobContext, JobState}; use crate::context::{ContextManager, JobContext, JobState};
use crate::db::Database;
use crate::error::{Error, JobError}; use crate::error::{Error, JobError};
use crate::extensions::ExtensionManager; use crate::extensions::ExtensionManager;
use crate::hooks::HookRegistry; use crate::hooks::HookRegistry;
use crate::llm::LlmProvider; use crate::llm::LlmProvider;
use crate::safety::SafetyLayer; use crate::safety::SafetyLayer;
use crate::tenant::AdminScope;
use crate::tools::{ use crate::tools::{
ApprovalContext, ToolRegistry, autonomous_allowed_tool_names, autonomous_unavailable_error, ApprovalContext, ToolRegistry, autonomous_allowed_tool_names, autonomous_unavailable_error,
prepare_tool_params, prepare_tool_params,
@@ -52,7 +53,7 @@ struct ScheduledSubtask {
pub struct SchedulerDeps { pub struct SchedulerDeps {
pub tools: Arc<ToolRegistry>, pub tools: Arc<ToolRegistry>,
pub extension_manager: Option<Arc<ExtensionManager>>, pub extension_manager: Option<Arc<ExtensionManager>>,
pub store: Option<AdminScope>, pub store: Option<Arc<dyn Database>>,
pub hooks: Arc<HookRegistry>, pub hooks: Arc<HookRegistry>,
} }
@@ -64,10 +65,10 @@ pub struct Scheduler {
safety: Arc<SafetyLayer>, safety: Arc<SafetyLayer>,
tools: Arc<ToolRegistry>, tools: Arc<ToolRegistry>,
extension_manager: Option<Arc<ExtensionManager>>, extension_manager: Option<Arc<ExtensionManager>>,
store: Option<AdminScope>, store: Option<Arc<dyn Database>>,
hooks: Arc<HookRegistry>, hooks: Arc<HookRegistry>,
/// SSE manager for live job event streaming. /// SSE broadcast sender for live job event streaming.
sse_tx: Option<Arc<crate::channels::web::sse::SseManager>>, sse_tx: Option<tokio::sync::broadcast::Sender<SseEvent>>,
/// HTTP interceptor for trace recording/replay (propagated to workers). /// HTTP interceptor for trace recording/replay (propagated to workers).
http_interceptor: Option<Arc<dyn crate::llm::recording::HttpInterceptor>>, http_interceptor: Option<Arc<dyn crate::llm::recording::HttpInterceptor>>,
/// Running jobs (main LLM-driven jobs). /// Running jobs (main LLM-driven jobs).
@@ -101,9 +102,9 @@ impl Scheduler {
} }
} }
/// Set the SSE manager for live job event streaming. /// Set the SSE broadcast sender for live job event streaming.
pub fn set_sse_sender(&mut self, sse: Arc<crate::channels::web::sse::SseManager>) { pub fn set_sse_sender(&mut self, tx: tokio::sync::broadcast::Sender<SseEvent>) {
self.sse_tx = Some(sse); self.sse_tx = Some(tx);
} }
/// Set the HTTP interceptor for trace recording/replay. /// Set the HTTP interceptor for trace recording/replay.
@@ -267,20 +268,6 @@ impl Scheduler {
}); });
} }
// Per-user concurrency check
if let Some(max_per_user) = self.config.max_jobs_per_user
&& let Ok(ctx) = self.context_manager.get_context(job_id).await
{
let user_active = self
.context_manager
.active_jobs_for(&ctx.user_id)
.await
.len();
if user_active >= max_per_user {
return Err(JobError::MaxJobsExceeded { max: max_per_user });
}
}
// Transition job to in_progress // Transition job to in_progress
self.context_manager self.context_manager
.update_context(job_id, |ctx| { .update_context(job_id, |ctx| {
@@ -794,15 +781,10 @@ mod tests {
allow_local_tools: true, allow_local_tools: true,
max_cost_per_day_cents: None, max_cost_per_day_cents: None,
max_actions_per_hour: None, max_actions_per_hour: None,
max_cost_per_user_per_day_cents: None,
max_tool_iterations: 10, max_tool_iterations: 10,
auto_approve_tools: true, auto_approve_tools: true,
default_timezone: "UTC".to_string(), default_timezone: "UTC".to_string(),
max_jobs_per_user: None,
max_tokens_per_job, 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 cm = Arc::new(ContextManager::new(5));
let llm: Arc<dyn LlmProvider> = Arc::new(StubLlm); let llm: Arc<dyn LlmProvider> = Arc::new(StubLlm);
+5 -5
View File
@@ -8,8 +8,8 @@ use chrono::{DateTime, Utc};
use uuid::Uuid; use uuid::Uuid;
use crate::context::{ContextManager, JobState}; use crate::context::{ContextManager, JobState};
use crate::db::Database;
use crate::error::RepairError; use crate::error::RepairError;
use crate::tenant::AdminScope;
use crate::tools::{BuildRequirement, Language, SoftwareBuilder, SoftwareType, ToolRegistry}; use crate::tools::{BuildRequirement, Language, SoftwareBuilder, SoftwareType, ToolRegistry};
/// A job that has been detected as stuck. /// 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. /// Jobs in `InProgress` longer than this are treated as stuck.
stuck_threshold: Duration, stuck_threshold: Duration,
max_repair_attempts: u32, max_repair_attempts: u32,
store: Option<AdminScope>, store: Option<Arc<dyn Database>>,
builder: Option<Arc<dyn SoftwareBuilder>>, builder: Option<Arc<dyn SoftwareBuilder>>,
tools: Option<Arc<ToolRegistry>>, tools: Option<Arc<ToolRegistry>>,
} }
@@ -91,8 +91,8 @@ impl DefaultSelfRepair {
} }
} }
/// Add an admin-scoped store for tool failure tracking. /// Add a Store for tool failure tracking.
pub fn with_store(mut self, store: AdminScope) -> Self { pub fn with_store(mut self, store: Arc<dyn Database>) -> Self {
self.store = Some(store); self.store = Some(store);
self self
} }
@@ -806,7 +806,7 @@ mod tests {
// Create self-repair with zero threshold (detect immediately), // Create self-repair with zero threshold (detect immediately),
// wired with store, builder, and tools. // wired with store, builder, and tools.
let repair = DefaultSelfRepair::new(Arc::clone(&cm), Duration::from_secs(0), 3) let repair = DefaultSelfRepair::new(Arc::clone(&cm), Duration::from_secs(0), 3)
.with_store(crate::tenant::AdminScope::new(Arc::clone(&db))) .with_store(Arc::clone(&db))
.with_builder( .with_builder(
Arc::clone(&builder) as Arc<dyn crate::tools::SoftwareBuilder>, Arc::clone(&builder) as Arc<dyn crate::tools::SoftwareBuilder>,
tools, tools,
+2 -193
View File
@@ -16,8 +16,8 @@ use chrono::{DateTime, TimeDelta, Utc};
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use uuid::Uuid; use uuid::Uuid;
use crate::channels::web::util::truncate_preview;
use crate::llm::{ChatMessage, ToolCall, generate_tool_call_id}; use crate::llm::{ChatMessage, ToolCall, generate_tool_call_id};
use ironclaw_common::truncate_preview;
/// A session containing one or more threads. /// A session containing one or more threads.
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
@@ -449,7 +449,6 @@ impl Thread {
id: call_id.clone(), id: call_id.clone(),
name: tc.name.clone(), name: tc.name.clone(),
arguments: tc.parameters.clone(), arguments: tc.parameters.clone(),
reasoning: None,
}) })
.collect(); .collect();
@@ -523,12 +522,7 @@ impl Thread {
&& let Some(ref tcs) = assistant_msg.tool_calls && let Some(ref tcs) = assistant_msg.tool_calls
{ {
for tc in tcs { for tc in tcs {
turn.record_tool_call_with_reasoning( turn.record_tool_call(&tc.name, tc.arguments.clone());
&tc.name,
tc.arguments.clone(),
tc.reasoning.clone(),
Some(tc.id.clone()),
);
} }
} }
@@ -608,10 +602,6 @@ pub struct Turn {
pub completed_at: Option<DateTime<Utc>>, pub completed_at: Option<DateTime<Utc>>,
/// Error message (if failed). /// Error message (if failed).
pub error: Option<String>, pub error: Option<String>,
/// 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<String>,
/// Transient image content parts for multimodal LLM input. /// Transient image content parts for multimodal LLM input.
/// Not serialized — images are only needed for the current LLM call. /// Not serialized — images are only needed for the current LLM call.
/// The text description in `user_input` persists for compaction/context. /// The text description in `user_input` persists for compaction/context.
@@ -631,7 +621,6 @@ impl Turn {
started_at: Utc::now(), started_at: Utc::now(),
completed_at: None, completed_at: None,
error: None, error: None,
narrative: None,
image_content_parts: Vec::new(), image_content_parts: Vec::new(),
} }
} }
@@ -667,26 +656,6 @@ impl Turn {
parameters: params, parameters: params,
result: None, result: None,
error: 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<String>,
params: serde_json::Value,
rationale: Option<String>,
tool_call_id: Option<String>,
) {
self.tool_calls.push(TurnToolCall {
name: name.into(),
parameters: params,
result: None,
error: None,
rationale,
tool_call_id,
}); });
} }
@@ -703,60 +672,6 @@ impl Turn {
call.error = Some(error.into()); 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<String>) {
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. /// Record of a tool call made during a turn.
@@ -770,12 +685,6 @@ pub struct TurnToolCall {
pub result: Option<serde_json::Value>, pub result: Option<serde_json::Value>,
/// Error from the tool (if failed). /// Error from the tool (if failed).
pub error: Option<String>, pub error: Option<String>,
/// Agent's reasoning for choosing this tool.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub rationale: Option<String>,
/// 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<String>,
} }
#[cfg(test)] #[cfg(test)]
@@ -1400,7 +1309,6 @@ mod tests {
id: "call_0".to_string(), id: "call_0".to_string(),
name: "search".to_string(), name: "search".to_string(),
arguments: serde_json::json!({"q": "test"}), arguments: serde_json::json!({"q": "test"}),
reasoning: None,
}; };
let messages = vec![ let messages = vec![
ChatMessage::user("Find test"), ChatMessage::user("Find test"),
@@ -1431,7 +1339,6 @@ mod tests {
id: "call_0".to_string(), id: "call_0".to_string(),
name: "http".to_string(), name: "http".to_string(),
arguments: serde_json::json!({}), arguments: serde_json::json!({}),
reasoning: None,
}; };
let messages = vec![ let messages = vec![
ChatMessage::user("Fetch URL"), ChatMessage::user("Fetch URL"),
@@ -1497,13 +1404,11 @@ mod tests {
id: "call_a".to_string(), id: "call_a".to_string(),
name: "search".to_string(), name: "search".to_string(),
arguments: serde_json::json!({"q": "data"}), arguments: serde_json::json!({"q": "data"}),
reasoning: None,
}; };
let tc2 = ToolCall { let tc2 = ToolCall {
id: "call_b".to_string(), id: "call_b".to_string(),
name: "write".to_string(), name: "write".to_string(),
arguments: serde_json::json!({"path": "out.txt"}), arguments: serde_json::json!({"path": "out.txt"}),
reasoning: None,
}; };
let messages = vec![ let messages = vec![
ChatMessage::user("Find and save"), ChatMessage::user("Find and save"),
@@ -1715,100 +1620,4 @@ mod tests {
let merged = thread.drain_pending_messages().unwrap(); let merged = thread.drain_pending_messages().unwrap();
assert_eq!(merged, "failed batch\nnew msg"); 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")
);
}
} }
+34 -189
View File
@@ -102,30 +102,11 @@ impl SessionManager {
/// Resolve an external thread ID to an internal thread. /// Resolve an external thread ID to an internal thread.
/// ///
/// Returns the session and thread ID. Creates both if they don't exist. /// 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( pub async fn resolve_thread(
&self, &self,
user_id: &str, user_id: &str,
channel: &str, channel: &str,
external_thread_id: Option<&str>, external_thread_id: Option<&str>,
) -> (Arc<Mutex<Session>>, 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<Uuid>,
) -> (Arc<Mutex<Session>>, Uuid) { ) -> (Arc<Mutex<Session>>, Uuid) {
let session = self.get_or_create_session(user_id).await; let session = self.get_or_create_session(user_id).await;
@@ -135,65 +116,51 @@ impl SessionManager {
external_thread_id: external_thread_id.map(String::from), external_thread_id: external_thread_id.map(String::from),
}; };
// Use pre-parsed UUID if available, otherwise parse from string. // Check if we have a mapping
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; let thread_map = self.thread_map.read().await;
// Fast path: exact key match
if let Some(&thread_id) = thread_map.get(&key) { if let Some(&thread_id) = thread_map.get(&key) {
// Verify thread still exists in session
let sess = session.lock().await; let sess = session.lock().await;
if sess.threads.contains_key(&thread_id) { if sess.threads.contains_key(&thread_id) {
return (Arc::clone(&session), thread_id); return (Arc::clone(&session), thread_id);
} }
} }
}
// UUID adoption check (still under the same read lock). // Check if external_thread_id is itself a known thread UUID that
// If external_thread_id is a valid UUID not mapped elsewhere, // exists in the session but was never registered in the thread_map
// it may be a thread created by chat_new_thread_handler or // (e.g. created by chat_new_thread_handler or hydrated from DB).
// hydrated from DB that we can adopt. // We only adopt it if no thread_map entry maps to this UUID —
// Only attempt adoption when external_thread_id is Some, preserving // otherwise it belongs to a different channel scope.
// the invariant that None external_thread_id never triggers adoption. if let Some(ext_tid) = external_thread_id
if external_thread_id.is_some() { && let Ok(ext_uuid) = Uuid::parse_str(ext_tid)
ext_uuid.filter(|&uuid| !thread_map.values().any(|&v| v == uuid)) {
} else { let thread_map = self.thread_map.read().await;
None let mapped_elsewhere = thread_map.values().any(|&v| v == ext_uuid);
} drop(thread_map);
}; // Single read lock dropped here
// If we found an adoptable UUID, verify it exists in session and acquire write lock if !mapped_elsewhere {
if let Some(ext_uuid) = adoptable_uuid { let sess = session.lock().await;
let sess = session.lock().await; if sess.threads.contains_key(&ext_uuid) {
if sess.threads.contains_key(&ext_uuid) { drop(sess);
drop(sess);
let mut thread_map = self.thread_map.write().await; let mut thread_map = self.thread_map.write().await;
// Re-check after acquiring write lock to prevent race condition // Re-check after acquiring write lock to prevent race condition
// where another task mapped this UUID between our read and write. // where another task mapped this UUID between our read and write.
if !thread_map.values().any(|&v| v == ext_uuid) { if !thread_map.values().any(|&v| v == ext_uuid) {
thread_map.insert(key, ext_uuid); thread_map.insert(key, ext_uuid);
drop(thread_map); drop(thread_map);
// Ensure undo manager exists // Ensure undo manager exists
let mut undo_managers = self.undo_managers.write().await; let mut undo_managers = self.undo_managers.write().await;
undo_managers undo_managers
.entry(ext_uuid) .entry(ext_uuid)
.or_insert_with(|| Arc::new(Mutex::new(UndoManager::new()))); .or_insert_with(|| Arc::new(Mutex::new(UndoManager::new())));
return (session, ext_uuid); return (session, ext_uuid);
}
// If it was mapped elsewhere while we were unlocked, fall through
// to create a new thread, preserving channel isolation.
} }
// If mapped elsewhere while unlocked, fall through to create new thread
} }
} }
@@ -942,44 +909,6 @@ 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] #[tokio::test]
async fn test_resolve_thread_finds_existing_session_thread_by_uuid() { async fn test_resolve_thread_finds_existing_session_thread_by_uuid() {
use crate::agent::session::{Session, Thread}; use crate::agent::session::{Session, Thread};
@@ -1018,88 +947,4 @@ mod tests {
"should have exactly 1 thread, not a duplicate" "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"
);
}
} }
-11
View File
@@ -92,17 +92,6 @@ impl SubmissionParser {
args: vec![], args: vec![],
}; };
} }
if lower == "/reasoning" || lower.starts_with("/reasoning ") {
let args: Vec<String> = trimmed
.split_whitespace()
.skip(1)
.map(|s| s.to_string())
.collect();
return Submission::SystemCommand {
command: "reasoning".to_string(),
args,
};
}
if lower == "/restart" { if lower == "/restart" {
tracing::debug!("[SubmissionParser::parse] Recognized /restart command"); tracing::debug!("[SubmissionParser::parse] Recognized /restart command");
return Submission::SystemCommand { return Submission::SystemCommand {
+163 -95
View File
@@ -16,12 +16,12 @@ use crate::agent::dispatcher::{
}; };
use crate::agent::session::{MAX_PENDING_MESSAGES, PendingApproval, Session, ThreadState}; use crate::agent::session::{MAX_PENDING_MESSAGES, PendingApproval, Session, ThreadState};
use crate::agent::submission::SubmissionResult; use crate::agent::submission::SubmissionResult;
use crate::channels::web::util::truncate_preview;
use crate::channels::{IncomingMessage, StatusUpdate}; use crate::channels::{IncomingMessage, StatusUpdate};
use crate::context::JobContext; use crate::context::JobContext;
use crate::error::Error; use crate::error::Error;
use crate::llm::{ChatMessage, ToolCall}; use crate::llm::{ChatMessage, ToolCall};
use crate::tools::redact_params; use crate::tools::redact_params;
use ironclaw_common::truncate_preview;
const FORGED_THREAD_ID_ERROR: &str = "Invalid or unauthorized thread ID."; const FORGED_THREAD_ID_ERROR: &str = "Invalid or unauthorized thread ID.";
@@ -175,7 +175,6 @@ impl Agent {
pub(super) async fn process_user_input( pub(super) async fn process_user_input(
&self, &self,
message: &IncomingMessage, message: &IncomingMessage,
tenant: crate::tenant::TenantCtx,
session: Arc<Mutex<Session>>, session: Arc<Mutex<Session>>,
thread_id: Uuid, thread_id: Uuid,
content: &str, content: &str,
@@ -352,7 +351,7 @@ impl Agent {
if let Some(intent) = self.router.route_command(&temp_message) { if let Some(intent) = self.router.route_command(&temp_message) {
// Explicit command like /status, /job, /list - handle directly // Explicit command like /status, /job, /list - handle directly
return self.handle_job_or_command(intent, message, &tenant).await; return self.handle_job_or_command(intent, message).await;
} }
// Natural language goes through the agentic loop // Natural language goes through the agentic loop
@@ -463,7 +462,7 @@ impl Agent {
// Run the agentic tool execution loop // Run the agentic tool execution loop
let result = self let result = self
.run_agentic_loop(message, tenant, session.clone(), thread_id, turn_messages) .run_agentic_loop(message, session.clone(), thread_id, turn_messages)
.await; .await;
// Re-acquire lock and check if interrupted // Re-acquire lock and check if interrupted
@@ -514,10 +513,10 @@ impl Agent {
}; };
thread.complete_turn(&response); thread.complete_turn(&response);
let (turn_number, tool_calls, narrative) = thread let (turn_number, tool_calls) = thread
.turns .turns
.last() .last()
.map(|t| (t.turn_number, t.tool_calls.clone(), t.narrative.clone())) .map(|t| (t.turn_number, t.tool_calls.clone()))
.unwrap_or_default(); .unwrap_or_default();
let _ = self let _ = self
.channels .channels
@@ -535,7 +534,6 @@ impl Agent {
&message.user_id, &message.user_id,
turn_number, turn_number,
&tool_calls, &tool_calls,
narrative.as_deref(),
) )
.await; .await;
self.persist_assistant_response( self.persist_assistant_response(
@@ -727,9 +725,7 @@ impl Agent {
/// ///
/// Stored between the user and assistant messages so that /// Stored between the user and assistant messages so that
/// `build_turns_from_db_messages` can reconstruct the tool call history. /// `build_turns_from_db_messages` can reconstruct the tool call history.
/// Content is a JSON object: `{ "calls": [...], "narrative": "..." }`. /// Content is a JSON array of tool call summaries.
/// 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( pub(super) async fn persist_tool_calls(
&self, &self,
thread_id: Uuid, thread_id: Uuid,
@@ -737,7 +733,6 @@ impl Agent {
user_id: &str, user_id: &str,
turn_number: usize, turn_number: usize,
tool_calls: &[crate::agent::session::TurnToolCall], tool_calls: &[crate::agent::session::TurnToolCall],
narrative: Option<&str>,
) { ) {
if tool_calls.is_empty() { if tool_calls.is_empty() {
return; return;
@@ -772,30 +767,11 @@ impl Agent {
if let Some(ref error) = tc.error { if let Some(ref error) = tc.error {
obj["error"] = serde_json::Value::String(truncate_preview(error, 200)); 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 obj
}) })
.collect(); .collect();
// Wrap in an object with optional narrative so it can be reconstructed. let content = match serde_json::to_string(&summaries) {
// 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, Ok(c) => c,
Err(e) => { Err(e) => {
tracing::warn!("Failed to serialize tool calls: {}", e); tracing::warn!("Failed to serialize tool calls: {}", e);
@@ -1016,8 +992,16 @@ impl Agent {
{ {
// Put it back and return error // Put it back and return error
let mut sess = session.lock().await; let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&thread_id) { match sess.threads.get_mut(&thread_id) {
thread.await_approval(pending); Some(thread) => {
thread.await_approval(pending);
}
None => {
tracing::warn!(
%thread_id,
"Thread disappeared while restoring pending approval after request ID mismatch"
);
}
} }
return Ok(SubmissionResult::error( return Ok(SubmissionResult::error(
"Request ID mismatch. Use the correct request ID.", "Request ID mismatch. Use the correct request ID.",
@@ -1039,8 +1023,19 @@ impl Agent {
// Reset thread state to processing // Reset thread state to processing
{ {
let mut sess = session.lock().await; let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&thread_id) { match sess.threads.get_mut(&thread_id) {
thread.state = ThreadState::Processing; Some(thread) => {
thread.state = ThreadState::Processing;
}
None => {
tracing::error!(
%thread_id,
"Thread disappeared while setting state to Processing during approval"
);
return Ok(SubmissionResult::error(
"Internal error: thread no longer exists",
));
}
} }
} }
@@ -1124,15 +1119,20 @@ impl Agent {
// Record sanitized result in thread // Record sanitized result in thread
{ {
let mut sess = session.lock().await; let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&thread_id) match sess.threads.get_mut(&thread_id) {
&& let Some(turn) = thread.last_turn_mut() Some(thread) => {
{ if let Some(turn) = thread.last_turn_mut() {
if is_tool_error { if is_tool_error {
turn.record_tool_error_for(&pending.tool_call_id, result_content.clone()); turn.record_tool_error(result_content.clone());
} else { } else {
turn.record_tool_result_for( turn.record_tool_result(serde_json::json!(result_content));
&pending.tool_call_id, }
serde_json::json!(result_content), }
}
None => {
tracing::error!(
%thread_id,
"Thread disappeared while recording tool result during approval"
); );
} }
} }
@@ -1381,15 +1381,21 @@ impl Agent {
// Record sanitized result in thread // Record sanitized result in thread
{ {
let mut sess = session.lock().await; let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&thread_id) match sess.threads.get_mut(&thread_id) {
&& let Some(turn) = thread.last_turn_mut() Some(thread) => {
{ if let Some(turn) = thread.last_turn_mut() {
if is_deferred_error { if is_deferred_error {
turn.record_tool_error_for(&tc.id, deferred_content.clone()); turn.record_tool_error(deferred_content.clone());
} else { } else {
turn.record_tool_result_for( turn.record_tool_result(serde_json::json!(deferred_content));
&tc.id, }
serde_json::json!(deferred_content), }
}
None => {
tracing::error!(
%thread_id,
tool_name = %tc.name,
"Thread disappeared while recording deferred tool result during approval"
); );
} }
} }
@@ -1443,8 +1449,19 @@ impl Agent {
{ {
let mut sess = session.lock().await; let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&thread_id) { match sess.threads.get_mut(&thread_id) {
thread.await_approval(new_pending); Some(thread) => {
thread.await_approval(new_pending);
}
None => {
tracing::error!(
%thread_id,
"Thread disappeared while setting up deferred tool approval"
);
return Ok(SubmissionResult::error(
"Internal error: thread no longer exists",
));
}
} }
} }
@@ -1474,13 +1491,7 @@ impl Agent {
// Continue the agentic loop (a tool was already executed this turn) // Continue the agentic loop (a tool was already executed this turn)
let result = self let result = self
.run_agentic_loop( .run_agentic_loop(message, session.clone(), thread_id, context_messages)
message,
self.tenant_ctx(&message.user_id).await,
session.clone(),
thread_id,
context_messages,
)
.await; .await;
// Handle the result // Handle the result
@@ -1495,10 +1506,10 @@ impl Agent {
let (response, suggestions) = let (response, suggestions) =
crate::agent::dispatcher::extract_suggestions(&response); crate::agent::dispatcher::extract_suggestions(&response);
thread.complete_turn(&response); thread.complete_turn(&response);
let (turn_number, tool_calls, narrative) = thread let (turn_number, tool_calls) = thread
.turns .turns
.last() .last()
.map(|t| (t.turn_number, t.tool_calls.clone(), t.narrative.clone())) .map(|t| (t.turn_number, t.tool_calls.clone()))
.unwrap_or_default(); .unwrap_or_default();
// User message already persisted at turn start; save tool calls then assistant response // User message already persisted at turn start; save tool calls then assistant response
self.persist_tool_calls( self.persist_tool_calls(
@@ -1507,7 +1518,6 @@ impl Agent {
&message.user_id, &message.user_id,
turn_number, turn_number,
&tool_calls, &tool_calls,
narrative.as_deref(),
) )
.await; .await;
self.persist_assistant_response( self.persist_assistant_response(
@@ -1583,17 +1593,28 @@ impl Agent {
); );
{ {
let mut sess = session.lock().await; let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&thread_id) { match sess.threads.get_mut(&thread_id) {
thread.clear_pending_approval(); Some(thread) => {
thread.complete_turn(&rejection); thread.clear_pending_approval();
// User message already persisted at turn start; save rejection response thread.complete_turn(&rejection);
self.persist_assistant_response( // User message already persisted at turn start; save rejection response
thread_id, self.persist_assistant_response(
&message.channel, thread_id,
&message.user_id, &message.channel,
&rejection, &message.user_id,
) &rejection,
.await; )
.await;
}
None => {
tracing::error!(
%thread_id,
"Thread disappeared during approval rejection"
);
return Ok(SubmissionResult::error(
"Internal error: thread no longer exists",
));
}
} }
} }
@@ -1683,7 +1704,7 @@ impl Agent {
}; };
match ext_mgr match ext_mgr
.configure_token(&pending.extension_name, token, &message.user_id) .configure_token(&pending.extension_name, token)
.await .await
{ {
Ok(result) if result.activated => { Ok(result) if result.activated => {
@@ -1853,20 +1874,7 @@ fn rebuild_chat_messages_from_db(
"assistant" => result.push(ChatMessage::assistant(&msg.content)), "assistant" => result.push(ChatMessage::assistant(&msg.content)),
"tool_calls" => { "tool_calls" => {
// Try to parse the enriched JSON and rebuild tool messages. // Try to parse the enriched JSON and rebuild tool messages.
// Supports two formats: if let Ok(calls) = serde_json::from_str::<Vec<serde_json::Value>>(&msg.content) {
// - Old: plain JSON array of tool call summaries
// - New: wrapped object { "calls": [...], "narrative": "..." }
let calls: Vec<serde_json::Value> =
match serde_json::from_str::<serde_json::Value>(&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() { if calls.is_empty() {
continue; continue;
} }
@@ -1889,10 +1897,6 @@ fn rebuild_chat_messages_from_db(
.get("parameters") .get("parameters")
.cloned() .cloned()
.unwrap_or(serde_json::json!({})), .unwrap_or(serde_json::json!({})),
reasoning: c
.get("rationale")
.and_then(|v| v.as_str())
.map(String::from),
}) })
.collect(); .collect();
@@ -2152,6 +2156,70 @@ mod tests {
} }
} }
#[tokio::test]
async fn test_approval_on_missing_thread_should_error() {
// Regression for #1487: when a thread disappears from the session
// during approval processing, the code must return a visible error
// rather than silently succeeding.
//
// We can't call process_approval() directly (requires full Agent),
// so we simulate the exact code pattern used in the rejection and
// state-setting paths: lock session, match on get_mut, verify the
// None arm produces an error.
use crate::agent::session::{Session, Thread, ThreadState};
use std::sync::Arc;
use tokio::sync::Mutex;
use uuid::Uuid;
let thread_id = Uuid::new_v4();
let session_id = Uuid::new_v4();
let session = Arc::new(Mutex::new(Session::new("test-user")));
// Scenario 1: Thread never existed
{
let sess = session.lock().await;
let result = match sess.threads.get(&thread_id) {
Some(_) => Ok("processed"),
None => Err("Internal error: thread no longer exists"),
};
assert!(result.is_err());
assert_eq!(
result.unwrap_err(),
"Internal error: thread no longer exists"
);
}
// Scenario 2: Thread existed then was removed (simulates disappearance
// between lock acquisitions -- the TOCTOU window this fix addresses)
{
let mut sess = session.lock().await;
let mut thread = Thread::with_id(thread_id, session_id);
thread.start_turn("pending approval");
thread.state = ThreadState::AwaitingApproval;
sess.threads.insert(thread_id, thread);
}
{
let mut sess = session.lock().await;
// Simulate thread disappearing (e.g., pruned by another task)
sess.threads.remove(&thread_id);
// The rejection path must detect this and return an error
let result = match sess.threads.get_mut(&thread_id) {
Some(thread) => {
thread.clear_pending_approval();
thread.complete_turn("rejected");
Ok("rejection persisted")
}
None => Err("Internal error: thread no longer exists"),
};
assert!(result.is_err());
assert_eq!(
result.unwrap_err(),
"Internal error: thread no longer exists"
);
}
}
#[test] #[test]
fn test_queue_cap_rejects_at_capacity() { fn test_queue_cap_rejects_at_capacity() {
use crate::agent::session::{MAX_PENDING_MESSAGES, Thread, ThreadState}; use crate::agent::session::{MAX_PENDING_MESSAGES, Thread, ThreadState};
+9 -27
View File
@@ -312,7 +312,13 @@ impl AppBuilder {
.create_provider(&self.config.llm.nearai.base_url, self.session.clone()); .create_provider(&self.config.llm.nearai.base_url, self.session.clone());
// Register memory tools if database is available // Register memory tools if database is available
let workspace_user_id = self.config.owner_id.as_str(); let workspace_user_id = self
.config
.channels
.gateway
.as_ref()
.map(|gw| gw.user_id.as_str())
.unwrap_or("default");
let workspace = if let Some(ref db) = self.db { let workspace = if let Some(ref db) = self.db {
let emb_cache_config = EmbeddingCacheConfig { let emb_cache_config = EmbeddingCacheConfig {
max_entries: self.config.embeddings.cache_size, max_entries: self.config.embeddings.cache_size,
@@ -321,7 +327,7 @@ impl AppBuilder {
.with_search_config(&self.config.search); .with_search_config(&self.config.search);
if let Some(ref emb) = embeddings { if let Some(ref emb) = embeddings {
ws = ws.with_embeddings_cached(emb.clone(), emb_cache_config.clone()); ws = ws.with_embeddings_cached(emb.clone(), emb_cache_config);
} }
// Wire workspace-level settings (read scopes, memory layers) // Wire workspace-level settings (read scopes, memory layers)
@@ -335,30 +341,7 @@ impl AppBuilder {
} }
ws = ws.with_memory_layers(self.config.workspace.memory_layers.clone()); ws = ws.with_memory_layers(self.config.workspace.memory_layers.clone());
let ws = Arc::new(ws); let ws = Arc::new(ws);
tools.register_memory_tools(Arc::clone(&ws));
// Detect multi-tenant mode: when the database has registered users,
// each authenticated user needs their own workspace scope. Use
// WorkspacePool (which implements WorkspaceResolver) to create
// per-user workspaces on demand instead of sharing the startup
// workspace across all users.
let is_multi_tenant = db.has_any_users().await.unwrap_or(false);
if is_multi_tenant {
let pool = Arc::new(crate::channels::web::server::WorkspacePool::new(
Arc::clone(db),
embeddings.clone(),
emb_cache_config,
self.config.search.clone(),
self.config.workspace.clone(),
));
tools.register_memory_tools_with_resolver(pool);
tracing::info!(
"Memory tools configured with per-user workspace resolver (multi-tenant mode)"
);
} else {
tools.register_memory_tools(Arc::clone(&ws));
}
Some(ws) Some(ws)
} else { } else {
None None
@@ -875,7 +858,6 @@ impl AppBuilder {
crate::agent::cost_guard::CostGuardConfig { crate::agent::cost_guard::CostGuardConfig {
max_cost_per_day_cents: self.config.agent.max_cost_per_day_cents, 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_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,
}, },
)); ));
-16
View File
@@ -265,15 +265,6 @@ 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. /// Status update types for showing agent activity.
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub enum StatusUpdate { pub enum StatusUpdate {
@@ -342,13 +333,6 @@ pub enum StatusUpdate {
}, },
/// Suggested follow-up messages for the user. /// Suggested follow-up messages for the user.
Suggestions { suggestions: Vec<String> }, Suggestions { suggestions: Vec<String> },
/// Agent reasoning update (why it chose specific tools).
ReasoningUpdate {
/// Human-readable summary of the agent's decision.
narrative: String,
/// Per-tool decisions.
decisions: Vec<ToolDecision>,
},
/// Per-turn token usage and cost summary (shown as subtle metadata). /// Per-turn token usage and cost summary (shown as subtle metadata).
TurnCost { TurnCost {
input_tokens: u64, input_tokens: u64,
+1 -1
View File
@@ -39,7 +39,7 @@ mod webhook_server;
pub use channel::{ pub use channel::{
AttachmentKind, Channel, ChannelSecretUpdater, IncomingAttachment, IncomingMessage, AttachmentKind, Channel, ChannelSecretUpdater, IncomingAttachment, IncomingMessage,
MessageStream, OutgoingResponse, StatusUpdate, ToolDecision, routing_target_from_metadata, MessageStream, OutgoingResponse, StatusUpdate, routing_target_from_metadata,
}; };
pub use http::{HttpChannel, HttpChannelState}; pub use http::{HttpChannel, HttpChannelState};
pub use manager::ChannelManager; pub use manager::ChannelManager;
+9 -65
View File
@@ -75,7 +75,6 @@ const SLASH_COMMANDS: &[&str] = &[
"/suggest", "/suggest",
"/thread", "/thread",
"/resume", "/resume",
"/reasoning",
]; ];
/// Rustyline helper for slash-command tab completion. /// Rustyline helper for slash-command tab completion.
@@ -431,18 +430,6 @@ impl ReplChannel {
let _ = execute!(stderr, terminal::Clear(terminal::ClearType::FromCursorDown)); 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 { impl Default for ReplChannel {
@@ -492,9 +479,7 @@ impl Channel for ReplChannel {
async fn start(&self) -> Result<MessageStream, ChannelError> { async fn start(&self) -> Result<MessageStream, ChannelError> {
let (tx, rx) = mpsc::channel(32); let (tx, rx) = mpsc::channel(32);
// Approval prompts inject responses back through this sender. // Store tx so send_status can inject approval responses directly
// 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() { if let Ok(mut guard) = self.msg_tx.lock() {
*guard = Some(tx.clone()); *guard = Some(tx.clone());
} }
@@ -510,10 +495,11 @@ impl Channel for ReplChannel {
// Single message mode: send it and return // Single message mode: send it and return
if let Some(msg) = single_message { if let Some(msg) = single_message {
let incoming = IncomingMessage::new("repl", &user_id, &msg) let incoming = IncomingMessage::new("repl", &user_id, &msg).with_timezone(&sys_tz);
.with_metadata(serde_json::json!({ "single_message_mode": true }))
.with_timezone(&sys_tz);
let _ = tx.blocking_send(incoming); 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; return;
} }
@@ -676,7 +662,6 @@ impl Channel for ReplChannel {
println!(); println!();
println!(); println!();
self.stdin_locked.store(false, Ordering::Relaxed); self.stdin_locked.store(false, Ordering::Relaxed);
self.finish_single_message_turn().await;
return Ok(()); return Ok(());
} }
@@ -695,7 +680,6 @@ impl Channel for ReplChannel {
println!(); println!();
// Unlock stdin so readline can resume // Unlock stdin so readline can resume
self.stdin_locked.store(false, Ordering::Relaxed); self.stdin_locked.store(false, Ordering::Relaxed);
self.finish_single_message_turn().await;
Ok(()) Ok(())
} }
@@ -795,7 +779,6 @@ impl Channel for ReplChannel {
let msg_tx = Arc::clone(&self.msg_tx); let msg_tx = Arc::clone(&self.msg_tx);
let user_id = self.user_id.clone(); let user_id = self.user_id.clone();
let lock_flag = Arc::clone(&self.stdin_locked); let lock_flag = Arc::clone(&self.stdin_locked);
let single_message_mode = self.single_message.is_some();
tokio::task::spawn_blocking(move || { tokio::task::spawn_blocking(move || {
let action = run_approval_selector(allow_always).unwrap_or("n"); let action = run_approval_selector(allow_always).unwrap_or("n");
// Unlock stdin so readline can resume after approval // Unlock stdin so readline can resume after approval
@@ -804,12 +787,7 @@ impl Channel for ReplChannel {
return; return;
}; };
if let Some(tx) = guard.as_ref() { if let Some(tx) = guard.as_ref() {
let msg = if single_message_mode { let msg = IncomingMessage::new("repl", &user_id, action);
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); let _ = tx.blocking_send(msg);
} }
}); });
@@ -863,19 +841,6 @@ impl Channel for ReplChannel {
StatusUpdate::Suggestions { .. } => { StatusUpdate::Suggestions { .. } => {
// Suggestions are only rendered by the web gateway // 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 { .. } => { StatusUpdate::TurnCost { .. } => {
// Cost display is handled by the TUI channel // Cost display is handled by the TUI channel
} }
@@ -910,7 +875,6 @@ impl Channel for ReplChannel {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use futures::StreamExt; use futures::StreamExt;
use tokio::time::{Duration, timeout};
use super::*; use super::*;
@@ -919,36 +883,16 @@ mod tests {
let repl = ReplChannel::with_message("hi".to_string()); let repl = ReplChannel::with_message("hi".to_string());
let mut stream = repl.start().await.expect("repl start should succeed"); let mut stream = repl.start().await.expect("repl start should succeed");
let first = timeout(Duration::from_secs(1), stream.next()) let first = stream.next().await.expect("first message missing");
.await
.expect("timed out waiting for first message")
.expect("first message missing");
assert_eq!(first.channel, "repl"); assert_eq!(first.channel, "repl");
assert_eq!(first.content, "hi"); assert_eq!(first.content, "hi");
assert!( let second = stream.next().await.expect("quit message missing");
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.channel, "repl");
assert_eq!(second.content, "/quit"); assert_eq!(second.content, "/quit");
assert!( assert!(
timeout(Duration::from_secs(1), stream.next()) stream.next().await.is_none(),
.await
.expect("timed out waiting for stream to close")
.is_none(),
"stream should end after /quit" "stream should end after /quit"
); );
} }
+2 -9
View File
@@ -333,9 +333,6 @@ async fn webhook_handler(
let channel_name = channel.channel_name(); let channel_name = channel.channel_name();
// Track whether any authentication was performed and passed.
let mut did_authenticate = false;
// Check if secret is required // Check if secret is required
if state.router.requires_secret(channel_name).await { if state.router.requires_secret(channel_name).await {
// Get the secret header name for this channel (from capabilities or default) // Get the secret header name for this channel (from capabilities or default)
@@ -385,7 +382,6 @@ async fn webhook_handler(
); );
} }
tracing::debug!(channel = %channel_name, "Webhook secret validated"); tracing::debug!(channel = %channel_name, "Webhook secret validated");
did_authenticate = true;
} }
None => { None => {
tracing::warn!( tracing::warn!(
@@ -437,7 +433,6 @@ async fn webhook_handler(
); );
} }
tracing::debug!(channel = %channel_name, "Ed25519 signature verified"); tracing::debug!(channel = %channel_name, "Ed25519 signature verified");
did_authenticate = true;
} }
_ => { _ => {
tracing::warn!( tracing::warn!(
@@ -489,7 +484,6 @@ async fn webhook_handler(
); );
} }
tracing::debug!(channel = %channel_name, "HMAC-SHA256 signature verified"); tracing::debug!(channel = %channel_name, "HMAC-SHA256 signature verified");
did_authenticate = true;
} }
_ => { _ => {
tracing::warn!( tracing::warn!(
@@ -516,9 +510,8 @@ async fn webhook_handler(
}) })
.collect(); .collect();
// Call the WASM channel. `did_authenticate` was set above by whichever // Call the WASM channel
// auth guard (secret / Ed25519 / HMAC) successfully validated the request. let secret_validated = state.router.requires_secret(channel_name).await;
let secret_validated = did_authenticate;
tracing::info!( tracing::info!(
channel = %channel_name, channel = %channel_name,
-14
View File
@@ -3061,20 +3061,6 @@ fn status_to_wit(
}, },
// Suggestions and turn cost are web-gateway-only; skip for WASM channels // Suggestions and turn cost are web-gateway-only; skip for WASM channels
StatusUpdate::Suggestions { .. } | StatusUpdate::TurnCost { .. } => return None, 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,
}
}
}) })
} }
-28
View File
@@ -91,34 +91,6 @@ Browser-facing HTTP API and SSE/WebSocket real-time streaming. Axum-based, singl
| DELETE | `/api/routines/{id}` | Delete a routine | | DELETE | `/api/routines/{id}` | Delete a routine |
| GET | `/api/routines/{id}/runs` | List runs for a specific routine | | GET | `/api/routines/{id}/runs` | List runs for a specific routine |
### User Management (admin — requires `admin` role, see `docs/USER_MANAGEMENT_API.md`)
| Method | Path | Description |
|--------|------|-------------|
| POST | `/api/admin/users` | Create a new user (returns one-time token) |
| GET | `/api/admin/users` | List all users |
| GET | `/api/admin/users/{id}` | Get a single user |
| PATCH | `/api/admin/users/{id}` | Update user profile/metadata |
| DELETE | `/api/admin/users/{id}` | Delete user and all data |
| POST | `/api/admin/users/{id}/suspend` | Suspend a user |
| POST | `/api/admin/users/{id}/activate` | Re-activate a user |
| GET | `/api/admin/usage` | Per-user LLM usage stats |
| GET | `/api/admin/users/{id}/secrets` | List a user's secrets (names only) |
| PUT | `/api/admin/users/{id}/secrets/{name}` | Create or update a user's secret |
| DELETE | `/api/admin/users/{id}/secrets/{name}` | Delete a user's secret |
### Profile (self-service)
| Method | Path | Description |
|--------|------|-------------|
| GET | `/api/profile` | Get own profile |
| PATCH | `/api/profile` | Update own display name/metadata |
### Tokens (self-service)
| Method | Path | Description |
|--------|------|-------------|
| POST | `/api/tokens` | Create API token (returns plaintext once) |
| GET | `/api/tokens` | List own tokens |
| DELETE | `/api/tokens/{id}` | Revoke a token |
### Settings ### Settings
| Method | Path | Description | | Method | Path | Description |
|--------|------|-------------| |--------|------|-------------|
+35 -577
View File
@@ -1,278 +1,17 @@
//! Bearer token authentication middleware for the web gateway. //! Bearer token authentication middleware for the web gateway.
//!
//! Supports multi-user mode: each token maps to a `UserIdentity` that carries
//! the user_id. The identity is inserted into request extensions so downstream
//! handlers can extract it via `AuthenticatedUser`.
use std::collections::HashMap;
use std::num::NonZeroUsize;
use axum::{ use axum::{
extract::{FromRequestParts, Request, State}, extract::{Request, State},
http::{HeaderMap, Method, StatusCode, request::Parts}, http::{HeaderMap, Method, StatusCode},
middleware::Next, middleware::Next,
response::{IntoResponse, Response}, response::{IntoResponse, Response},
}; };
use sha2::{Digest, Sha256};
use std::sync::Arc;
use std::time::Instant;
use subtle::ConstantTimeEq; use subtle::ConstantTimeEq;
use tokio::sync::RwLock;
use crate::db::Database; /// Shared auth state injected via axum middleware state.
/// Identity resolved from a bearer token.
#[derive(Debug, Clone)]
pub struct UserIdentity {
pub user_id: String,
/// `admin` or `member`.
pub role: String,
/// Additional user scopes this identity can read from.
pub workspace_read_scopes: Vec<String>,
}
/// Hash a token with SHA-256 for constant-size, timing-safe storage.
pub fn hash_token(token: &str) -> [u8; 32] {
let mut hasher = Sha256::new();
hasher.update(token.as_bytes());
hasher.finalize().into()
}
/// Multi-user auth state: maps token hashes to user identities.
///
/// Tokens are SHA-256 hashed on construction so they are never stored in
/// plaintext. Authentication compares fixed-size (32-byte) digests using
/// constant-time comparison, eliminating both length-oracle timing leaks
/// and accidental token exposure in memory dumps.
///
/// In single-user mode (the default), contains exactly one entry.
#[derive(Clone)] #[derive(Clone)]
pub struct MultiAuthState { pub struct AuthState {
/// Maps SHA-256(token) → identity. Tokens are never stored in cleartext. pub token: String,
hashed_tokens: Vec<([u8; 32], UserIdentity)>,
/// Original first token kept only for single-user startup printing.
/// Not used for authentication.
display_token: Option<String>,
}
impl MultiAuthState {
/// Create a single-user auth state (backwards compatible).
pub fn single(token: String, user_id: String) -> Self {
let hash = hash_token(&token);
Self {
hashed_tokens: vec![(
hash,
UserIdentity {
user_id,
role: "admin".to_string(),
workspace_read_scopes: Vec::new(),
},
)],
display_token: Some(token),
}
}
/// Create a multi-user auth state from a map of tokens to identities.
///
/// **Test-only** — production multi-user auth is DB-backed via
/// `DbAuthenticator`. This constructor is kept public (not `#[cfg(test)]`)
/// because integration tests in `tests/` compile the crate as a library
/// where `cfg(test)` is not set.
pub fn multi(tokens: HashMap<String, UserIdentity>) -> Self {
let hashed_tokens: Vec<([u8; 32], UserIdentity)> = tokens
.into_iter()
.map(|(tok, identity)| (hash_token(&tok), identity))
.collect();
Self {
hashed_tokens,
display_token: None,
}
}
/// Authenticate a token, returning the associated identity if valid.
///
/// Uses SHA-256 hashing + constant-time comparison (`subtle::ConstantTimeEq`)
/// to prevent timing side-channels. Both the candidate and stored tokens are
/// hashed to 32-byte digests, eliminating length-oracle leaks. Iterates all
/// entries regardless of match to avoid early-exit timing differences.
/// O(n) in the number of configured users — negligible for typical
/// deployments (< 10 users).
pub fn authenticate(&self, candidate: &str) -> Option<&UserIdentity> {
let candidate_hash = hash_token(candidate);
let mut matched: Option<&UserIdentity> = None;
for (stored_hash, identity) in &self.hashed_tokens {
if bool::from(candidate_hash.ct_eq(stored_hash)) {
matched = Some(identity);
}
}
matched
}
/// Get the first token for backwards-compatible printing at startup.
///
/// Only available in single-user mode; returns `None` in multi-user mode
/// to avoid exposing tokens.
pub fn first_token(&self) -> Option<&str> {
self.display_token.as_deref()
}
/// Get the first user identity (for single-user fallback).
pub fn first_identity(&self) -> Option<&UserIdentity> {
self.hashed_tokens.first().map(|(_, id)| id)
}
}
/// DB-backed token authenticator with a bounded LRU cache.
///
/// Checks an LRU cache first (TTL 60s), then falls back to a DB query.
/// The cache is bounded to `MAX_CACHE_ENTRIES` — when full, the least
/// recently used entry is evicted regardless of TTL.
///
/// Revoking a token or suspending a user has at most 60s of stale
/// authentication before the cache entry expires.
#[derive(Clone)]
#[allow(clippy::type_complexity)]
pub struct DbAuthenticator {
store: Arc<dyn Database>,
/// Bounded LRU cache: token_hash → (identity, inserted_at).
cache: Arc<RwLock<lru::LruCache<[u8; 32], (UserIdentity, Instant)>>>,
}
impl DbAuthenticator {
/// Cache TTL — how long a successful auth is cached before re-querying the DB.
const CACHE_TTL_SECS: u64 = 60;
/// Maximum cache entries to prevent unbounded growth.
// SAFETY: 1024 is non-zero, so the unwrap in `new()` is infallible.
const MAX_CACHE_ENTRIES: NonZeroUsize = match NonZeroUsize::new(1024) {
Some(v) => v,
None => unreachable!(),
};
pub fn new(store: Arc<dyn Database>) -> Self {
Self {
store,
cache: Arc::new(RwLock::new(lru::LruCache::new(Self::MAX_CACHE_ENTRIES))),
}
}
/// Authenticate a token against the database, using cache when possible.
///
/// Returns `Ok(Some(identity))` on success, `Ok(None)` if the token is
/// not found, or `Err(())` if the database is unreachable (so the caller
/// can return 503 instead of 401).
pub async fn authenticate(&self, candidate: &str) -> Result<Option<UserIdentity>, ()> {
let hash = hash_token(candidate);
// Check cache first (promotes to most-recent on hit)
{
let mut cache = self.cache.write().await;
if let Some((identity, inserted_at)) = cache.get(&hash) {
if inserted_at.elapsed().as_secs() < Self::CACHE_TTL_SECS {
return Ok(Some(identity.clone()));
}
// Expired — remove stale entry
cache.pop(&hash);
}
}
// Cache miss or expired — query DB
let (token_record, user_record) = match self.store.authenticate_token(&hash).await {
Ok(Some(pair)) => pair,
Ok(None) => return Ok(None),
Err(e) => {
tracing::error!(error = %e, "DB auth lookup failed, returning 503");
return Err(());
}
};
let identity = UserIdentity {
user_id: user_record.id.clone(),
role: user_record.role.clone(),
workspace_read_scopes: Vec::new(),
};
// Record token usage (best-effort, don't block auth)
let store = self.store.clone();
let token_id = token_record.id;
let user_id = user_record.id;
tokio::spawn(async move {
let _ = store.record_token_usage(token_id).await;
let _ = store.record_login(&user_id).await;
});
// Insert into bounded LRU — if full, least-recently-used entry is evicted
{
let mut cache = self.cache.write().await;
cache.put(hash, (identity.clone(), Instant::now()));
}
Ok(Some(identity))
}
}
/// Combined auth state: tries env-var tokens first, then DB-backed tokens.
#[derive(Clone)]
pub struct CombinedAuthState {
/// In-memory tokens from GATEWAY_AUTH_TOKEN.
pub env_auth: MultiAuthState,
/// DB-backed token authenticator (optional — only when a database is available).
pub db_auth: Option<DbAuthenticator>,
}
impl From<MultiAuthState> for CombinedAuthState {
fn from(env_auth: MultiAuthState) -> Self {
Self {
env_auth,
db_auth: None,
}
}
}
/// Axum extractor that provides the authenticated user identity.
///
/// Only available on routes behind `auth_middleware`. Extracts the
/// `UserIdentity` that the middleware inserted into request extensions.
pub struct AuthenticatedUser(pub UserIdentity);
impl<S> FromRequestParts<S> for AuthenticatedUser
where
S: Send + Sync,
{
type Rejection = (StatusCode, &'static str);
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
parts
.extensions
.get::<UserIdentity>()
.cloned()
.map(AuthenticatedUser)
.ok_or((StatusCode::UNAUTHORIZED, "Not authenticated"))
}
}
/// Axum extractor that requires the authenticated user to have the `admin` role.
///
/// Use instead of `AuthenticatedUser` on endpoints that modify system-wide
/// state (user management, model selection, extension/skill installation).
pub struct AdminUser(pub UserIdentity);
impl<S> FromRequestParts<S> for AdminUser
where
S: Send + Sync,
{
type Rejection = (StatusCode, &'static str);
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
let identity = parts
.extensions
.get::<UserIdentity>()
.cloned()
.ok_or((StatusCode::UNAUTHORIZED, "Not authenticated"))?;
if identity.role != "admin" {
return Err((StatusCode::FORBIDDEN, "Admin role required"));
}
Ok(AdminUser(identity))
}
} }
/// Whether query-string token auth is allowed for this request. /// Whether query-string token auth is allowed for this request.
@@ -311,127 +50,48 @@ fn query_token(request: &Request) -> Option<String> {
/// Auth middleware that validates bearer token from header or query param. /// Auth middleware that validates bearer token from header or query param.
/// ///
/// Tries env-var tokens first (constant-time, in-memory), then falls back /// SSE connections can't set headers from `EventSource`, so we also accept
/// to DB-backed token lookup if configured. SSE connections can't set /// `?token=xxx` as a query parameter, but only on SSE endpoints.
/// headers from `EventSource`, so we also accept `?token=xxx` as a query
/// parameter, but only on SSE/WS endpoints.
///
/// On successful authentication, inserts the matching `UserIdentity` into
/// request extensions for downstream extraction via `AuthenticatedUser`.
pub async fn auth_middleware( pub async fn auth_middleware(
State(auth): State<CombinedAuthState>, State(auth): State<AuthState>,
headers: HeaderMap, headers: HeaderMap,
mut request: Request, request: Request,
next: Next, next: Next,
) -> Response { ) -> Response {
// Extract the candidate token from header or query param. // Try Authorization header first (constant-time comparison).
let token = extract_token(&headers, &request); // RFC 6750 Section 2.1: auth-scheme comparison is case-insensitive.
if let Some(auth_header) = headers.get("authorization")
&& let Ok(value) = auth_header.to_str()
&& value.len() > 7
&& value[..7].eq_ignore_ascii_case("Bearer ")
&& bool::from(value.as_bytes()[7..].ct_eq(auth.token.as_bytes()))
{
return next.run(request).await;
}
if let Some(ref tok) = token { // Fall back to query parameter, but only for SSE endpoints (constant-time comparison).
// 1. Try env-var tokens first (fast, constant-time, in-memory). if allows_query_token_auth(&request)
if let Some(identity) = auth.env_auth.authenticate(tok) { && let Some(token) = query_token(&request)
request.extensions_mut().insert(identity.clone()); && bool::from(token.as_bytes().ct_eq(auth.token.as_bytes()))
return next.run(request).await; {
} return next.run(request).await;
// 2. Fall back to DB-backed token lookup.
if let Some(ref db_auth) = auth.db_auth {
match db_auth.authenticate(tok).await {
Ok(Some(identity)) => {
request.extensions_mut().insert(identity);
return next.run(request).await;
}
Err(()) => {
return (StatusCode::SERVICE_UNAVAILABLE, "Database unavailable")
.into_response();
}
Ok(None) => {}
}
}
} }
(StatusCode::UNAUTHORIZED, "Invalid or missing auth token").into_response() (StatusCode::UNAUTHORIZED, "Invalid or missing auth token").into_response()
} }
/// Extract a bearer token from the Authorization header or query parameter.
fn extract_token(headers: &HeaderMap, request: &Request) -> Option<String> {
// Try Authorization header first (RFC 6750).
if let Some(auth_header) = headers.get("authorization")
&& let Ok(value) = auth_header.to_str()
&& value.len() > 7
&& value[..7].eq_ignore_ascii_case("Bearer ")
{
return Some(value[7..].to_string());
}
// Fall back to query parameter for SSE/WS endpoints.
if allows_query_token_auth(request) {
return query_token(request);
}
None
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use crate::testing::credentials::TEST_AUTH_SECRET_TOKEN; use crate::testing::credentials::{TEST_AUTH_SECRET_TOKEN, TEST_BEARER_TOKEN};
#[test] #[test]
fn test_multi_auth_state_single() { fn test_auth_state_clone() {
let state = MultiAuthState::single("tok-123".to_string(), "alice".to_string()); let state = AuthState {
let identity = state.authenticate("tok-123"); token: TEST_BEARER_TOKEN.to_string(),
assert!(identity.is_some()); };
assert_eq!(identity.unwrap().user_id, "alice"); let cloned = state.clone();
} assert_eq!(cloned.token, TEST_BEARER_TOKEN);
#[test]
fn test_multi_auth_state_reject_wrong_token() {
let state = MultiAuthState::single("tok-123".to_string(), "alice".to_string());
assert!(state.authenticate("wrong-token").is_none());
}
#[test]
fn test_multi_auth_state_multi_users() {
let mut tokens = HashMap::new();
tokens.insert(
"tok-alice".to_string(),
UserIdentity {
user_id: "alice".to_string(),
role: "admin".to_string(),
workspace_read_scopes: Vec::new(),
},
);
tokens.insert(
"tok-bob".to_string(),
UserIdentity {
user_id: "bob".to_string(),
role: "admin".to_string(),
workspace_read_scopes: Vec::new(),
},
);
let state = MultiAuthState::multi(tokens);
let alice = state.authenticate("tok-alice").unwrap();
assert_eq!(alice.user_id, "alice");
let bob = state.authenticate("tok-bob").unwrap();
assert_eq!(bob.user_id, "bob");
assert!(state.authenticate("tok-charlie").is_none());
}
#[test]
fn test_multi_auth_state_first_token() {
let state = MultiAuthState::single("my-token".to_string(), "user1".to_string());
assert_eq!(state.first_token(), Some("my-token"));
}
#[test]
fn test_multi_auth_state_first_identity() {
let state = MultiAuthState::single("my-token".to_string(), "user1".to_string());
let identity = state.first_identity().unwrap();
assert_eq!(identity.user_id, "user1");
} }
use axum::Router; use axum::Router;
@@ -447,10 +107,9 @@ mod tests {
/// Router with streaming endpoints (query auth allowed) and regular /// Router with streaming endpoints (query auth allowed) and regular
/// endpoints (query auth rejected). /// endpoints (query auth rejected).
fn test_app(token: &str) -> Router { fn test_app(token: &str) -> Router {
let state = CombinedAuthState::from(MultiAuthState::single( let state = AuthState {
token.to_string(), token: token.to_string(),
"test-user".to_string(), };
));
Router::new() Router::new()
.route("/api/chat/events", get(dummy_handler)) .route("/api/chat/events", get(dummy_handler))
.route("/api/logs/events", get(dummy_handler)) .route("/api/logs/events", get(dummy_handler))
@@ -647,205 +306,4 @@ mod tests {
let resp = app.oneshot(req).await.unwrap(); let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
} }
// --- Multi-tenant auth integration tests ---
/// Handler that extracts `AuthenticatedUser` and returns the resolved user_id.
async fn identity_handler(AuthenticatedUser(identity): AuthenticatedUser) -> String {
identity.user_id
}
/// Handler that extracts `AuthenticatedUser` and returns workspace_read_scopes as JSON.
async fn scopes_handler(AuthenticatedUser(identity): AuthenticatedUser) -> String {
serde_json::to_string(&identity.workspace_read_scopes).unwrap()
}
/// Build a multi-user router where each token maps to a distinct identity.
fn multi_user_app(tokens: HashMap<String, UserIdentity>) -> Router {
let state = CombinedAuthState::from(MultiAuthState::multi(tokens));
Router::new()
.route("/api/chat/events", get(identity_handler))
.route("/api/chat/send", post(identity_handler))
.route("/api/scopes", get(scopes_handler))
.layer(middleware::from_fn_with_state(state, auth_middleware))
}
fn two_user_tokens() -> HashMap<String, UserIdentity> {
let mut tokens = HashMap::new();
tokens.insert(
"tok-alice".to_string(),
UserIdentity {
user_id: "alice".to_string(),
role: "admin".to_string(),
workspace_read_scopes: vec!["shared".to_string()],
},
);
tokens.insert(
"tok-bob".to_string(),
UserIdentity {
user_id: "bob".to_string(),
role: "admin".to_string(),
workspace_read_scopes: vec!["shared".to_string(), "alice".to_string()],
},
);
tokens
}
#[tokio::test]
async fn test_multi_user_alice_token_resolves_to_alice() {
let app = multi_user_app(two_user_tokens());
let req = Request::builder()
.uri("/api/chat/events")
.header("Authorization", "Bearer tok-alice")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
assert_eq!(body, "alice");
}
#[tokio::test]
async fn test_multi_user_bob_token_resolves_to_bob() {
let app = multi_user_app(two_user_tokens());
let req = Request::builder()
.uri("/api/chat/events")
.header("Authorization", "Bearer tok-bob")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
assert_eq!(body, "bob");
}
#[tokio::test]
async fn test_multi_user_sequential_tokens_resolve_independently() {
// Send both alice and bob tokens sequentially and verify each gets
// the correct identity — guards against token map corruption.
let tokens = two_user_tokens();
let app1 = multi_user_app(tokens.clone());
let req = Request::builder()
.uri("/api/chat/events")
.header("Authorization", "Bearer tok-alice")
.body(Body::empty())
.unwrap();
let resp = app1.oneshot(req).await.unwrap();
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
assert_eq!(body, "alice");
let app2 = multi_user_app(tokens);
let req = Request::builder()
.uri("/api/chat/events")
.header("Authorization", "Bearer tok-bob")
.body(Body::empty())
.unwrap();
let resp = app2.oneshot(req).await.unwrap();
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
assert_eq!(body, "bob");
}
#[tokio::test]
async fn test_multi_user_unknown_token_rejected() {
let app = multi_user_app(two_user_tokens());
let req = Request::builder()
.uri("/api/chat/events")
.header("Authorization", "Bearer tok-charlie")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test]
async fn test_multi_user_workspace_read_scopes_propagated() {
let app = multi_user_app(two_user_tokens());
// Alice has ["shared"]
let req = Request::builder()
.uri("/api/scopes")
.header("Authorization", "Bearer tok-alice")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
let scopes: Vec<String> = serde_json::from_slice(&body).unwrap();
assert_eq!(scopes, vec!["shared"]);
}
#[tokio::test]
async fn test_multi_user_bob_has_two_scopes() {
let app = multi_user_app(two_user_tokens());
// Bob has ["shared", "alice"]
let req = Request::builder()
.uri("/api/scopes")
.header("Authorization", "Bearer tok-bob")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
let scopes: Vec<String> = serde_json::from_slice(&body).unwrap();
assert_eq!(scopes, vec!["shared", "alice"]);
}
#[tokio::test]
async fn test_multi_user_query_param_resolves_correct_identity() {
let app = multi_user_app(two_user_tokens());
let req = Request::builder()
.uri("/api/chat/events?token=tok-bob")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
assert_eq!(body, "bob");
}
#[tokio::test]
async fn test_multi_user_post_with_bearer_resolves_identity() {
let app = multi_user_app(two_user_tokens());
let req = Request::builder()
.method(Method::POST)
.uri("/api/chat/send")
.header("Authorization", "Bearer tok-alice")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
assert_eq!(body, "alice");
}
#[tokio::test]
async fn test_multi_user_empty_scopes_for_single_user() {
// Single-user mode creates identity with empty workspace_read_scopes.
let state = CombinedAuthState::from(MultiAuthState::single(
"tok-only".to_string(),
"solo".to_string(),
));
let app = Router::new()
.route("/api/scopes", get(scopes_handler))
.layer(middleware::from_fn_with_state(state, auth_middleware));
let req = Request::builder()
.uri("/api/scopes")
.header("Authorization", "Bearer tok-only")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
let scopes: Vec<String> = serde_json::from_slice(&body).unwrap();
assert!(scopes.is_empty());
}
#[tokio::test]
async fn test_prefix_and_extension_tokens_rejected() {
// Verifies that prefix/suffix variants of valid tokens are rejected.
// Note: the constant-time property is enforced structurally by use of
// subtle::ConstantTimeEq and cannot be verified via outcome testing.
let state = MultiAuthState::single("long-secret-token".to_string(), "user".to_string());
assert!(state.authenticate("long-secret").is_none());
assert!(state.authenticate("long-secret-token-extra").is_none());
}
} }
+35 -64
View File
@@ -12,24 +12,22 @@ use serde::Deserialize;
use uuid::Uuid; use uuid::Uuid;
use crate::channels::IncomingMessage; use crate::channels::IncomingMessage;
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::server::GatewayState; use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*; use crate::channels::web::types::*;
use crate::channels::web::util::{build_turns_from_db_messages, truncate_preview}; use crate::channels::web::util::{build_turns_from_db_messages, truncate_preview};
pub async fn chat_send_handler( pub async fn chat_send_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(identity): AuthenticatedUser,
Json(req): Json<SendMessageRequest>, Json(req): Json<SendMessageRequest>,
) -> Result<(StatusCode, Json<SendMessageResponse>), (StatusCode, String)> { ) -> Result<(StatusCode, Json<SendMessageResponse>), (StatusCode, String)> {
if !state.chat_rate_limiter.check(&identity.user_id) { if !state.chat_rate_limiter.check() {
return Err(( return Err((
StatusCode::TOO_MANY_REQUESTS, StatusCode::TOO_MANY_REQUESTS,
"Rate limit exceeded. Try again shortly.".to_string(), "Rate limit exceeded. Try again shortly.".to_string(),
)); ));
} }
let mut msg = IncomingMessage::new("gateway", &identity.user_id, &req.content); let mut msg = IncomingMessage::new("gateway", &state.user_id, &req.content);
if let Some(ref thread_id) = req.thread_id { if let Some(ref thread_id) = req.thread_id {
msg = msg.with_thread(thread_id); msg = msg.with_thread(thread_id);
@@ -76,7 +74,6 @@ pub async fn chat_send_handler(
pub async fn chat_approval_handler( pub async fn chat_approval_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(identity): AuthenticatedUser,
Json(req): Json<ApprovalRequest>, Json(req): Json<ApprovalRequest>,
) -> Result<(StatusCode, Json<SendMessageResponse>), (StatusCode, String)> { ) -> Result<(StatusCode, Json<SendMessageResponse>), (StatusCode, String)> {
let (approved, always) = match req.action.as_str() { let (approved, always) = match req.action.as_str() {
@@ -112,7 +109,7 @@ pub async fn chat_approval_handler(
) )
})?; })?;
let mut msg = IncomingMessage::new("gateway", &identity.user_id, content); let mut msg = IncomingMessage::new("gateway", &state.user_id, content);
if let Some(ref thread_id) = req.thread_id { if let Some(ref thread_id) = req.thread_id {
msg = msg.with_thread(thread_id); msg = msg.with_thread(thread_id);
@@ -153,7 +150,6 @@ pub async fn chat_approval_handler(
/// The token never touches the LLM, chat history, or SSE stream. /// The token never touches the LLM, chat history, or SSE stream.
pub async fn chat_auth_token_handler( pub async fn chat_auth_token_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Json(req): Json<AuthTokenRequest>, Json(req): Json<AuthTokenRequest>,
) -> Result<Json<ActionResponse>, (StatusCode, String)> { ) -> Result<Json<ActionResponse>, (StatusCode, String)> {
let ext_mgr = state.extension_manager.as_ref().ok_or(( let ext_mgr = state.extension_manager.as_ref().ok_or((
@@ -162,7 +158,7 @@ pub async fn chat_auth_token_handler(
))?; ))?;
match ext_mgr match ext_mgr
.configure_token(&req.extension_name, &req.token, &user.user_id) .configure_token(&req.extension_name, &req.token)
.await .await
{ {
Ok(result) => { Ok(result) => {
@@ -173,26 +169,20 @@ pub async fn chat_auth_token_handler(
resp.instructions = result.verification.as_ref().map(|v| v.instructions.clone()); resp.instructions = result.verification.as_ref().map(|v| v.instructions.clone());
if result.verification.is_some() { if result.verification.is_some() {
state.sse.broadcast_for_user( state.sse.broadcast(SseEvent::AuthRequired {
&user.user_id, extension_name: req.extension_name.clone(),
AppEvent::AuthRequired { instructions: Some(result.message),
extension_name: req.extension_name.clone(), auth_url: None,
instructions: Some(result.message), setup_url: None,
auth_url: None, });
setup_url: None,
},
);
} else { } else {
clear_auth_mode(&state, &user.user_id).await; clear_auth_mode(&state).await;
state.sse.broadcast_for_user( state.sse.broadcast(SseEvent::AuthCompleted {
&user.user_id, extension_name: req.extension_name.clone(),
AppEvent::AuthCompleted { success: true,
extension_name: req.extension_name.clone(), message: result.message,
success: true, });
message: result.message,
},
);
} }
Ok(Json(resp)) Ok(Json(resp))
@@ -200,15 +190,12 @@ pub async fn chat_auth_token_handler(
Err(e) => { Err(e) => {
let msg = e.to_string(); let msg = e.to_string();
if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) { if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) {
state.sse.broadcast_for_user( state.sse.broadcast(SseEvent::AuthRequired {
&user.user_id, extension_name: req.extension_name.clone(),
AppEvent::AuthRequired { instructions: Some(msg.clone()),
extension_name: req.extension_name.clone(), auth_url: None,
instructions: Some(msg.clone()), setup_url: None,
auth_url: None, });
setup_url: None,
},
);
} }
Ok(Json(ActionResponse::fail(msg))) Ok(Json(ActionResponse::fail(msg)))
} }
@@ -218,17 +205,16 @@ pub async fn chat_auth_token_handler(
/// Cancel an in-progress auth flow. /// Cancel an in-progress auth flow.
pub async fn chat_auth_cancel_handler( pub async fn chat_auth_cancel_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(identity): AuthenticatedUser,
Json(_req): Json<AuthCancelRequest>, Json(_req): Json<AuthCancelRequest>,
) -> Result<Json<ActionResponse>, (StatusCode, String)> { ) -> Result<Json<ActionResponse>, (StatusCode, String)> {
clear_auth_mode(&state, &identity.user_id).await; clear_auth_mode(&state).await;
Ok(Json(ActionResponse::ok("Auth cancelled"))) Ok(Json(ActionResponse::ok("Auth cancelled")))
} }
/// Clear pending auth mode on the active thread. /// Clear pending auth mode on the active thread.
pub async fn clear_auth_mode(state: &GatewayState, user_id: &str) { pub async fn clear_auth_mode(state: &GatewayState) {
if let Some(ref sm) = state.session_manager { if let Some(ref sm) = state.session_manager {
let session = sm.get_or_create_session(user_id).await; let session = sm.get_or_create_session(&state.user_id).await;
let mut sess = session.lock().await; let mut sess = session.lock().await;
if let Some(thread_id) = sess.active_thread if let Some(thread_id) = sess.active_thread
&& let Some(thread) = sess.threads.get_mut(&thread_id) && let Some(thread) = sess.threads.get_mut(&thread_id)
@@ -240,9 +226,8 @@ pub async fn clear_auth_mode(state: &GatewayState, user_id: &str) {
pub async fn chat_events_handler( pub async fn chat_events_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<impl IntoResponse, (StatusCode, String)> { ) -> Result<impl IntoResponse, (StatusCode, String)> {
state.sse.subscribe(Some(user.user_id)).ok_or(( state.sse.subscribe().ok_or((
StatusCode::SERVICE_UNAVAILABLE, StatusCode::SERVICE_UNAVAILABLE,
"Too many connections".to_string(), "Too many connections".to_string(),
)) ))
@@ -252,7 +237,6 @@ pub async fn chat_ws_handler(
headers: axum::http::HeaderMap, headers: axum::http::HeaderMap,
ws: WebSocketUpgrade, ws: WebSocketUpgrade,
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(identity): AuthenticatedUser,
) -> Result<impl IntoResponse, (StatusCode, String)> { ) -> Result<impl IntoResponse, (StatusCode, String)> {
// Validate Origin header to prevent cross-site WebSocket hijacking. // Validate Origin header to prevent cross-site WebSocket hijacking.
let origin = headers let origin = headers
@@ -278,9 +262,7 @@ pub async fn chat_ws_handler(
"WebSocket origin not allowed".to_string(), "WebSocket origin not allowed".to_string(),
)); ));
} }
Ok(ws.on_upgrade(move |socket| { Ok(ws.on_upgrade(move |socket| crate::channels::web::ws::handle_ws_connection(socket, state)))
crate::channels::web::ws::handle_ws_connection(socket, state, identity)
}))
} }
#[derive(Deserialize)] #[derive(Deserialize)]
@@ -292,7 +274,6 @@ pub struct HistoryQuery {
pub async fn chat_history_handler( pub async fn chat_history_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(identity): AuthenticatedUser,
Query(query): Query<HistoryQuery>, Query(query): Query<HistoryQuery>,
) -> Result<Json<HistoryResponse>, (StatusCode, String)> { ) -> Result<Json<HistoryResponse>, (StatusCode, String)> {
let session_manager = state.session_manager.as_ref().ok_or(( let session_manager = state.session_manager.as_ref().ok_or((
@@ -300,9 +281,7 @@ pub async fn chat_history_handler(
"Session manager not available".to_string(), "Session manager not available".to_string(),
))?; ))?;
let session = session_manager let session = session_manager.get_or_create_session(&state.user_id).await;
.get_or_create_session(&identity.user_id)
.await;
let limit = query.limit.unwrap_or(50); let limit = query.limit.unwrap_or(50);
let before_cursor = query let before_cursor = query
@@ -335,7 +314,7 @@ pub async fn chat_history_handler(
&& let Some(ref store) = state.store && let Some(ref store) = state.store
{ {
let owned = store let owned = store
.conversation_belongs_to_user(thread_id, &identity.user_id) .conversation_belongs_to_user(thread_id, &state.user_id)
.await .await
.unwrap_or(false); .unwrap_or(false);
if !owned { if !owned {
@@ -398,10 +377,8 @@ pub async fn chat_history_handler(
truncate_preview(&s, 500) truncate_preview(&s, 500)
}), }),
error: tc.error.clone(), error: tc.error.clone(),
rationale: tc.rationale.clone(),
}) })
.collect(), .collect(),
narrative: t.narrative.clone(),
}) })
.collect(); .collect();
@@ -457,27 +434,24 @@ pub async fn chat_history_handler(
pub async fn chat_threads_handler( pub async fn chat_threads_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(identity): AuthenticatedUser,
) -> Result<Json<ThreadListResponse>, (StatusCode, String)> { ) -> Result<Json<ThreadListResponse>, (StatusCode, String)> {
let session_manager = state.session_manager.as_ref().ok_or(( let session_manager = state.session_manager.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE, StatusCode::SERVICE_UNAVAILABLE,
"Session manager not available".to_string(), "Session manager not available".to_string(),
))?; ))?;
let session = session_manager let session = session_manager.get_or_create_session(&state.user_id).await;
.get_or_create_session(&identity.user_id)
.await;
// Try DB first for persistent thread list // Try DB first for persistent thread list
if let Some(ref store) = state.store { if let Some(ref store) = state.store {
// Auto-create assistant thread if it doesn't exist // Auto-create assistant thread if it doesn't exist
let assistant_id = store let assistant_id = store
.get_or_create_assistant_conversation(&identity.user_id, "gateway") .get_or_create_assistant_conversation(&state.user_id, "gateway")
.await .await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
if let Ok(summaries) = store if let Ok(summaries) = store
.list_conversations_all_channels(&identity.user_id, 50) .list_conversations_all_channels(&state.user_id, 50)
.await .await
{ {
let mut assistant_thread = None; let mut assistant_thread = None;
@@ -560,16 +534,13 @@ pub async fn chat_threads_handler(
pub async fn chat_new_thread_handler( pub async fn chat_new_thread_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(identity): AuthenticatedUser,
) -> Result<Json<ThreadInfo>, (StatusCode, String)> { ) -> Result<Json<ThreadInfo>, (StatusCode, String)> {
let session_manager = state.session_manager.as_ref().ok_or(( let session_manager = state.session_manager.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE, StatusCode::SERVICE_UNAVAILABLE,
"Session manager not available".to_string(), "Session manager not available".to_string(),
))?; ))?;
let session = session_manager let session = session_manager.get_or_create_session(&state.user_id).await;
.get_or_create_session(&identity.user_id)
.await;
let (thread_id, info) = { let (thread_id, info) = {
let mut sess = session.lock().await; let mut sess = session.lock().await;
let thread = sess.create_thread(); let thread = sess.create_thread();
@@ -591,12 +562,12 @@ pub async fn chat_new_thread_handler(
// so that the subsequent loadThreads() call from the frontend sees it. // so that the subsequent loadThreads() call from the frontend sees it.
if let Some(ref store) = state.store { if let Some(ref store) = state.store {
match store match store
.ensure_conversation(thread_id, "gateway", &identity.user_id, None) .ensure_conversation(thread_id, "gateway", &state.user_id, None)
.await .await
{ {
Ok(true) => {} Ok(true) => {}
Ok(false) => tracing::warn!( Ok(false) => tracing::warn!(
user = %identity.user_id, user = %state.user_id,
thread_id = %thread_id, thread_id = %thread_id,
"Skipped persisting new thread due to ownership/channel conflict" "Skipped persisting new thread due to ownership/channel conflict"
), ),
+3 -8
View File
@@ -8,13 +8,11 @@ use axum::{
http::StatusCode, http::StatusCode,
}; };
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::server::GatewayState; use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*; use crate::channels::web::types::*;
pub async fn extensions_list_handler( pub async fn extensions_list_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<ExtensionListResponse>, (StatusCode, String)> { ) -> Result<Json<ExtensionListResponse>, (StatusCode, String)> {
let ext_mgr = state.extension_manager.as_ref().ok_or(( let ext_mgr = state.extension_manager.as_ref().ok_or((
StatusCode::NOT_IMPLEMENTED, StatusCode::NOT_IMPLEMENTED,
@@ -22,7 +20,7 @@ pub async fn extensions_list_handler(
))?; ))?;
let installed = ext_mgr let installed = ext_mgr
.list(None, false, &user.user_id) .list(None, false)
.await .await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
@@ -82,7 +80,6 @@ pub async fn extensions_list_handler(
pub async fn extensions_tools_handler( pub async fn extensions_tools_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(_user): AuthenticatedUser,
) -> Result<Json<ToolListResponse>, (StatusCode, String)> { ) -> Result<Json<ToolListResponse>, (StatusCode, String)> {
let registry = state.tool_registry.as_ref().ok_or(( let registry = state.tool_registry.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE, StatusCode::SERVICE_UNAVAILABLE,
@@ -103,7 +100,6 @@ pub async fn extensions_tools_handler(
pub async fn extensions_install_handler( pub async fn extensions_install_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Json(req): Json<InstallExtensionRequest>, Json(req): Json<InstallExtensionRequest>,
) -> Result<Json<ActionResponse>, (StatusCode, String)> { ) -> Result<Json<ActionResponse>, (StatusCode, String)> {
let ext_mgr = state.extension_manager.as_ref().ok_or(( let ext_mgr = state.extension_manager.as_ref().ok_or((
@@ -120,7 +116,7 @@ pub async fn extensions_install_handler(
}); });
match ext_mgr match ext_mgr
.install(&req.name, req.url.as_deref(), kind_hint, &user.user_id) .install(&req.name, req.url.as_deref(), kind_hint)
.await .await
{ {
Ok(result) => Ok(Json(ActionResponse::ok(result.message))), Ok(result) => Ok(Json(ActionResponse::ok(result.message))),
@@ -130,7 +126,6 @@ pub async fn extensions_install_handler(
pub async fn extensions_remove_handler( pub async fn extensions_remove_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(name): Path<String>, Path(name): Path<String>,
) -> Result<Json<ActionResponse>, (StatusCode, String)> { ) -> Result<Json<ActionResponse>, (StatusCode, String)> {
let ext_mgr = state.extension_manager.as_ref().ok_or(( let ext_mgr = state.extension_manager.as_ref().ok_or((
@@ -138,7 +133,7 @@ pub async fn extensions_remove_handler(
"Extension manager not available (secrets store required)".to_string(), "Extension manager not available (secrets store required)".to_string(),
))?; ))?;
match ext_mgr.remove(&name, &user.user_id).await { match ext_mgr.remove(&name).await {
Ok(message) => Ok(Json(ActionResponse::ok(message))), Ok(message) => Ok(Json(ActionResponse::ok(message))),
Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))), Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))),
} }
+277 -400
View File
@@ -11,13 +11,11 @@ use axum::{
use serde::Deserialize; use serde::Deserialize;
use uuid::Uuid; use uuid::Uuid;
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::server::GatewayState; use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*; use crate::channels::web::types::*;
pub async fn jobs_list_handler( pub async fn jobs_list_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<JobListResponse>, (StatusCode, String)> { ) -> Result<Json<JobListResponse>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or(( let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE, StatusCode::SERVICE_UNAVAILABLE,
@@ -27,8 +25,8 @@ pub async fn jobs_list_handler(
let mut jobs: Vec<JobInfo> = Vec::new(); let mut jobs: Vec<JobInfo> = Vec::new();
let mut seen_ids: HashSet<Uuid> = HashSet::new(); let mut seen_ids: HashSet<Uuid> = HashSet::new();
// Fetch sandbox jobs scoped to this user. // Fetch sandbox jobs from database.
match store.list_sandbox_jobs_for_user(&user.user_id).await { match store.list_sandbox_jobs().await {
Ok(sandbox_jobs) => { Ok(sandbox_jobs) => {
for j in &sandbox_jobs { for j in &sandbox_jobs {
let ui_state = match j.status.as_str() { let ui_state = match j.status.as_str() {
@@ -52,8 +50,8 @@ pub async fn jobs_list_handler(
} }
} }
// Fetch agent (non-sandbox) jobs scoped to this user, deduplicating by ID. // Fetch agent (non-sandbox) jobs from database, deduplicating by ID.
match store.list_agent_jobs_for_user(&user.user_id).await { match store.list_agent_jobs().await {
Ok(agent_jobs) => { Ok(agent_jobs) => {
for j in &agent_jobs { for j in &agent_jobs {
if seen_ids.contains(&j.id) { if seen_ids.contains(&j.id) {
@@ -82,7 +80,6 @@ pub async fn jobs_list_handler(
pub async fn jobs_summary_handler( pub async fn jobs_summary_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<JobSummaryResponse>, (StatusCode, String)> { ) -> Result<Json<JobSummaryResponse>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or(( let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE, StatusCode::SERVICE_UNAVAILABLE,
@@ -96,8 +93,8 @@ pub async fn jobs_summary_handler(
let mut failed = 0; let mut failed = 0;
let mut stuck = 0; let mut stuck = 0;
// Sandbox job counts scoped to this user. // Sandbox job counts.
match store.sandbox_job_summary_for_user(&user.user_id).await { match store.sandbox_job_summary().await {
Ok(s) => { Ok(s) => {
total += s.total; total += s.total;
pending += s.creating; pending += s.creating;
@@ -110,8 +107,8 @@ pub async fn jobs_summary_handler(
} }
} }
// Agent job counts scoped to this user. // Agent job counts.
match store.agent_job_summary_for_user(&user.user_id).await { match store.agent_job_summary().await {
Ok(s) => { Ok(s) => {
total += s.total; total += s.total;
pending += s.pending; pending += s.pending;
@@ -137,7 +134,6 @@ pub async fn jobs_summary_handler(
pub async fn jobs_detail_handler( pub async fn jobs_detail_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
) -> Result<Json<JobDetailResponse>, (StatusCode, String)> { ) -> Result<Json<JobDetailResponse>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or(( let store = state.store.as_ref().ok_or((
@@ -149,213 +145,169 @@ pub async fn jobs_detail_handler(
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?; .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
// Try sandbox job from DB first. // Try sandbox job from DB first.
match store.get_sandbox_job(job_id).await { if let Ok(Some(job)) = store.get_sandbox_job(job_id).await {
Ok(Some(job)) => { let browse_id = std::path::Path::new(&job.project_dir)
if job.user_id != user.user_id { .file_name()
return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); .map(|n| n.to_string_lossy().to_string())
} .unwrap_or_else(|| job.id.to_string());
let browse_id = std::path::Path::new(&job.project_dir)
.file_name()
.map(|n| n.to_string_lossy().to_string())
.unwrap_or_else(|| job.id.to_string());
let ui_state = match job.status.as_str() { let ui_state = match job.status.as_str() {
"creating" => "pending", "creating" => "pending",
"running" => "in_progress", "running" => "in_progress",
s => s, s => s,
}; };
let elapsed_secs = job.started_at.map(|start| { let elapsed_secs = job.started_at.map(|start| {
let end = job.completed_at.unwrap_or_else(chrono::Utc::now); let end = job.completed_at.unwrap_or_else(chrono::Utc::now);
(end - start).num_seconds().max(0) as u64 (end - start).num_seconds().max(0) as u64
});
// Synthesize transitions from timestamps.
let mut transitions = Vec::new();
if let Some(started) = job.started_at {
transitions.push(TransitionInfo {
from: "creating".to_string(),
to: "running".to_string(),
timestamp: started.to_rfc3339(),
reason: None,
}); });
// Synthesize transitions from timestamps.
let mut transitions = Vec::new();
if let Some(started) = job.started_at {
transitions.push(TransitionInfo {
from: "creating".to_string(),
to: "running".to_string(),
timestamp: started.to_rfc3339(),
reason: None,
});
}
if let Some(completed) = job.completed_at {
transitions.push(TransitionInfo {
from: "running".to_string(),
to: job.status.clone(),
timestamp: completed.to_rfc3339(),
reason: job.failure_reason.clone(),
});
}
let mode = store.get_sandbox_job_mode(job.id).await.ok().flatten();
let is_claude_code = mode.as_deref() == Some("claude_code");
return Ok(Json(JobDetailResponse {
id: job.id,
title: job.task.clone(),
description: String::new(),
state: ui_state.to_string(),
user_id: job.user_id.clone(),
created_at: job.created_at.to_rfc3339(),
started_at: job.started_at.map(|dt| dt.to_rfc3339()),
completed_at: job.completed_at.map(|dt| dt.to_rfc3339()),
elapsed_secs,
project_dir: Some(job.project_dir.clone()),
browse_url: Some(format!("/projects/{}/", browse_id)),
job_mode: mode.filter(|m| m != "worker"),
transitions,
can_restart: state.job_manager.is_some(),
can_prompt: is_claude_code && state.prompt_queue.is_some(),
job_kind: Some("sandbox".to_string()),
}));
} }
Ok(None) => {} if let Some(completed) = job.completed_at {
Err(e) => { transitions.push(TransitionInfo {
return Err(( from: "running".to_string(),
StatusCode::INTERNAL_SERVER_ERROR, to: job.status.clone(),
format!("Database error: {}", e), timestamp: completed.to_rfc3339(),
)); reason: job.failure_reason.clone(),
});
} }
let mode = store.get_sandbox_job_mode(job.id).await.ok().flatten();
let is_claude_code = mode.as_deref() == Some("claude_code");
return Ok(Json(JobDetailResponse {
id: job.id,
title: job.task.clone(),
description: String::new(),
state: ui_state.to_string(),
user_id: job.user_id.clone(),
created_at: job.created_at.to_rfc3339(),
started_at: job.started_at.map(|dt| dt.to_rfc3339()),
completed_at: job.completed_at.map(|dt| dt.to_rfc3339()),
elapsed_secs,
project_dir: Some(job.project_dir.clone()),
browse_url: Some(format!("/projects/{}/", browse_id)),
job_mode: mode.filter(|m| m != "worker"),
transitions,
can_restart: state.job_manager.is_some(),
can_prompt: is_claude_code && state.prompt_queue.is_some(),
job_kind: Some("sandbox".to_string()),
}));
} }
// Fall back to agent job from DB. // Fall back to agent job from DB.
match store.get_job(job_id).await { if let Ok(Some(ctx)) = store.get_job(job_id).await {
Ok(Some(ctx)) => { let elapsed_secs = ctx.started_at.map(|start| {
if ctx.user_id != user.user_id { let end = ctx.completed_at.unwrap_or_else(chrono::Utc::now);
return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); (end - start).num_seconds().max(0) as u64
} });
let elapsed_secs = ctx.started_at.map(|start| {
let end = ctx.completed_at.unwrap_or_else(chrono::Utc::now);
(end - start).num_seconds().max(0) as u64
});
// Only show prompt bar for jobs that have a running worker (Pending/InProgress). // Only show prompt bar for jobs that have a running worker (Pending/InProgress).
// Stuck jobs have no active worker loop, so messages would be silently dropped. // Stuck jobs have no active worker loop, so messages would be silently dropped.
let is_promptable = matches!( let is_promptable = matches!(
ctx.state, ctx.state,
crate::context::JobState::Pending | crate::context::JobState::InProgress crate::context::JobState::Pending | crate::context::JobState::InProgress
); );
Ok(Json(JobDetailResponse { return Ok(Json(JobDetailResponse {
id: ctx.job_id, id: ctx.job_id,
title: ctx.title.clone(), title: ctx.title.clone(),
description: ctx.description.clone(), description: ctx.description.clone(),
state: ctx.state.to_string(), state: ctx.state.to_string(),
user_id: ctx.user_id.clone(), user_id: ctx.user_id.clone(),
created_at: ctx.created_at.to_rfc3339(), created_at: ctx.created_at.to_rfc3339(),
started_at: ctx.started_at.map(|dt| dt.to_rfc3339()), started_at: ctx.started_at.map(|dt| dt.to_rfc3339()),
completed_at: ctx.completed_at.map(|dt| dt.to_rfc3339()), completed_at: ctx.completed_at.map(|dt| dt.to_rfc3339()),
elapsed_secs, elapsed_secs,
project_dir: None, project_dir: None,
browse_url: None, browse_url: None,
job_mode: None, job_mode: None,
transitions: Vec::new(), transitions: Vec::new(),
can_restart: state.scheduler.is_some(), can_restart: state.scheduler.is_some(),
can_prompt: is_promptable && state.scheduler.is_some(), can_prompt: is_promptable && state.scheduler.is_some(),
job_kind: Some("agent".to_string()), job_kind: Some("agent".to_string()),
})) }));
}
Ok(None) => Err((StatusCode::NOT_FOUND, "Job not found".to_string())),
Err(e) => Err((
StatusCode::INTERNAL_SERVER_ERROR,
format!("Database error: {}", e),
)),
} }
Err((StatusCode::NOT_FOUND, "Job not found".to_string()))
} }
pub async fn jobs_cancel_handler( pub async fn jobs_cancel_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> { ) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let job_id = Uuid::parse_str(&id) let job_id = Uuid::parse_str(&id)
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?; .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
// Try sandbox job cancellation. // Try sandbox job cancellation.
if let Some(ref store) = state.store { if let Some(ref store) = state.store
match store.get_sandbox_job(job_id).await { && let Ok(Some(job)) = store.get_sandbox_job(job_id).await
Ok(Some(job)) => { {
if job.user_id != user.user_id { if job.status == "running" || job.status == "creating" {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); // Stop the container if we have a job manager.
} if let Some(ref jm) = state.job_manager
if job.status == "running" || job.status == "creating" { && let Err(e) = jm.stop_job(job_id).await
if let Some(ref jm) = state.job_manager {
&& let Err(e) = jm.stop_job(job_id).await tracing::warn!(job_id = %job_id, error = %e, "Failed to stop container during cancellation");
{
tracing::warn!(job_id = %job_id, error = %e, "Failed to stop container during cancellation");
}
store
.update_sandbox_job_status(
job_id,
"failed",
Some(false),
Some("Cancelled by user"),
None,
Some(chrono::Utc::now()),
)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
}
return Ok(Json(serde_json::json!({
"status": "cancelled",
"job_id": job_id,
})));
}
Ok(None) => {}
Err(e) => {
return Err((
StatusCode::INTERNAL_SERVER_ERROR,
format!("Database error: {}", e),
));
} }
store
.update_sandbox_job_status(
job_id,
"failed",
Some(false),
Some("Cancelled by user"),
None,
Some(chrono::Utc::now()),
)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
} }
return Ok(Json(serde_json::json!({
"status": "cancelled",
"job_id": job_id,
})));
} }
// Fall back to agent job cancellation: stop the worker via the scheduler // Fall back to agent job cancellation: stop the worker via the scheduler
// (which updates the in-memory ContextManager AND aborts the task handle), // (which updates the in-memory ContextManager AND aborts the task handle),
// then persist the status to the DB as a fallback. // then persist the status to the DB as a fallback.
if let Some(ref store) = state.store { if let Some(ref store) = state.store
match store.get_job(job_id).await { && let Ok(Some(job)) = store.get_job(job_id).await
Ok(Some(job)) => { {
if job.user_id != user.user_id { if job.state.is_active() {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); // Try to stop via scheduler (aborts the worker task + updates
} // in-memory ContextManager). This is best-effort — the job may
if job.state.is_active() { // not be in the scheduler map if it already finished.
// Try to stop via scheduler (aborts the worker task + updates if let Some(ref slot) = state.scheduler
// in-memory ContextManager). This is best-effort — the job may && let Some(ref scheduler) = *slot.read().await
// not be in the scheduler map if it already finished. {
if let Some(ref slot) = state.scheduler let _ = scheduler.stop(job_id).await;
&& let Some(ref scheduler) = *slot.read().await }
{
let _ = scheduler.stop(job_id).await;
}
// Always persist cancellation to the DB so the state is // Always persist cancellation to the DB so the state is
// consistent even if the scheduler wasn't available or the // consistent even if the scheduler wasn't available or the
// job wasn't in its in-memory map. // job wasn't in its in-memory map.
store store
.update_job_status( .update_job_status(
job_id, job_id,
crate::context::JobState::Cancelled, crate::context::JobState::Cancelled,
Some("Cancelled by user"), Some("Cancelled by user"),
) )
.await .await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
}
return Ok(Json(serde_json::json!({
"status": "cancelled",
"job_id": job_id,
})));
}
Ok(None) => {}
Err(e) => {
return Err((
StatusCode::INTERNAL_SERVER_ERROR,
format!("Database error: {}", e),
));
}
} }
return Ok(Json(serde_json::json!({
"status": "cancelled",
"job_id": job_id,
})));
} }
Err((StatusCode::NOT_FOUND, "Job not found".to_string())) Err((StatusCode::NOT_FOUND, "Job not found".to_string()))
@@ -363,7 +315,6 @@ pub async fn jobs_cancel_handler(
pub async fn jobs_restart_handler( pub async fn jobs_restart_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> { ) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or(( let store = state.store.as_ref().ok_or((
@@ -375,166 +326,146 @@ pub async fn jobs_restart_handler(
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?; .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
// Try sandbox job restart first. // Try sandbox job restart first.
match store.get_sandbox_job(old_job_id).await { if let Ok(Some(old_job)) = store.get_sandbox_job(old_job_id).await {
Ok(Some(old_job)) => { if old_job.status != "interrupted" && old_job.status != "failed" {
if old_job.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
if old_job.status != "interrupted" && old_job.status != "failed" {
return Err((
StatusCode::CONFLICT,
format!("Cannot restart job in state '{}'", old_job.status),
));
}
let jm = state.job_manager.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Sandbox not enabled".to_string(),
))?;
// Enrich the task with failure context.
let task = if let Some(ref reason) = old_job.failure_reason {
format!(
"Previous attempt failed: {}. Retry: {}",
reason, old_job.task
)
} else {
old_job.task.clone()
};
let new_job_id = Uuid::new_v4();
let now = chrono::Utc::now();
let record = crate::history::SandboxJobRecord {
id: new_job_id,
task: task.clone(),
status: "creating".to_string(),
user_id: old_job.user_id.clone(),
project_dir: old_job.project_dir.clone(),
success: None,
failure_reason: None,
created_at: now,
started_at: None,
completed_at: None,
credential_grants_json: old_job.credential_grants_json.clone(),
};
store
.save_sandbox_job(&record)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
let mode = match store.get_sandbox_job_mode(old_job_id).await {
Ok(Some(m)) if m == "claude_code" => {
crate::orchestrator::job_manager::JobMode::ClaudeCode
}
_ => crate::orchestrator::job_manager::JobMode::Worker,
};
let credential_grants: Vec<crate::orchestrator::auth::CredentialGrant> =
serde_json::from_str(&old_job.credential_grants_json).unwrap_or_else(|e| {
tracing::warn!(
job_id = %old_job.id,
"Failed to deserialize credential grants from stored job: {}. \
Restarted job will have no credentials.",
e
);
vec![]
});
let project_dir = std::path::PathBuf::from(&old_job.project_dir);
let _token = jm
.create_job(
new_job_id,
&task,
Some(project_dir),
mode,
credential_grants,
)
.await
.map_err(|e| {
(
StatusCode::INTERNAL_SERVER_ERROR,
format!("Failed to create container: {}", e),
)
})?;
store
.update_sandbox_job_status(new_job_id, "running", None, None, Some(now), None)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
return Ok(Json(serde_json::json!({
"status": "restarted",
"old_job_id": old_job_id,
"new_job_id": new_job_id,
})));
}
Ok(None) => {}
Err(e) => {
return Err(( return Err((
StatusCode::INTERNAL_SERVER_ERROR, StatusCode::CONFLICT,
format!("Database error: {}", e), format!("Cannot restart job in state '{}'", old_job.status),
)); ));
} }
let jm = state.job_manager.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Sandbox not enabled".to_string(),
))?;
// Enrich the task with failure context.
let task = if let Some(ref reason) = old_job.failure_reason {
format!(
"Previous attempt failed: {}. Retry: {}",
reason, old_job.task
)
} else {
old_job.task.clone()
};
let new_job_id = Uuid::new_v4();
let now = chrono::Utc::now();
let record = crate::history::SandboxJobRecord {
id: new_job_id,
task: task.clone(),
status: "creating".to_string(),
user_id: old_job.user_id.clone(),
project_dir: old_job.project_dir.clone(),
success: None,
failure_reason: None,
created_at: now,
started_at: None,
completed_at: None,
credential_grants_json: old_job.credential_grants_json.clone(),
};
store
.save_sandbox_job(&record)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
let mode = match store.get_sandbox_job_mode(old_job_id).await {
Ok(Some(m)) if m == "claude_code" => {
crate::orchestrator::job_manager::JobMode::ClaudeCode
}
_ => crate::orchestrator::job_manager::JobMode::Worker,
};
let credential_grants: Vec<crate::orchestrator::auth::CredentialGrant> =
serde_json::from_str(&old_job.credential_grants_json).unwrap_or_else(|e| {
tracing::warn!(
job_id = %old_job.id,
"Failed to deserialize credential grants from stored job: {}. \
Restarted job will have no credentials.",
e
);
vec![]
});
let project_dir = std::path::PathBuf::from(&old_job.project_dir);
let _token = jm
.create_job(
new_job_id,
&task,
Some(project_dir),
mode,
credential_grants,
)
.await
.map_err(|e| {
(
StatusCode::INTERNAL_SERVER_ERROR,
format!("Failed to create container: {}", e),
)
})?;
store
.update_sandbox_job_status(new_job_id, "running", None, None, Some(now), None)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
return Ok(Json(serde_json::json!({
"status": "restarted",
"old_job_id": old_job_id,
"new_job_id": new_job_id,
})));
} }
// Try agent job restart: dispatch a new job via the scheduler. // Try agent job restart: dispatch a new job via the scheduler.
match store.get_job(old_job_id).await { if let Ok(Some(old_job)) = store.get_job(old_job_id).await {
Ok(Some(old_job)) => { if old_job.state.is_active() {
if old_job.user_id != user.user_id { return Err((
return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); StatusCode::CONFLICT,
} format!("Cannot restart job in state '{}'", old_job.state),
if old_job.state.is_active() { ));
return Err((
StatusCode::CONFLICT,
format!("Cannot restart job in state '{}'", old_job.state),
));
}
let slot = state.scheduler.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Scheduler not available".to_string(),
))?;
let scheduler_guard = slot.read().await;
let scheduler = scheduler_guard.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Agent not started yet".to_string(),
))?;
// Look up failure reason (O(1) point lookup).
let failure_reason = store
.get_agent_job_failure_reason(old_job_id)
.await
.ok()
.flatten()
.unwrap_or_default();
let title = if !failure_reason.is_empty() {
format!(
"Previous attempt failed: {}. Retry: {}",
failure_reason, old_job.title
)
} else {
old_job.title.clone()
};
let new_job_id = scheduler
.dispatch_job(&old_job.user_id, &title, &old_job.description, None)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
Ok(Json(serde_json::json!({
"status": "restarted",
"old_job_id": old_job_id,
"new_job_id": new_job_id,
})))
} }
Ok(None) => Err((StatusCode::NOT_FOUND, "Job not found".to_string())),
Err(e) => Err(( let slot = state.scheduler.as_ref().ok_or((
StatusCode::INTERNAL_SERVER_ERROR, StatusCode::SERVICE_UNAVAILABLE,
format!("Database error: {}", e), "Scheduler not available".to_string(),
)), ))?;
let scheduler_guard = slot.read().await;
let scheduler = scheduler_guard.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Agent not started yet".to_string(),
))?;
// Look up failure reason (O(1) point lookup).
let failure_reason = store
.get_agent_job_failure_reason(old_job_id)
.await
.ok()
.flatten()
.unwrap_or_default();
let title = if !failure_reason.is_empty() {
format!(
"Previous attempt failed: {}. Retry: {}",
failure_reason, old_job.title
)
} else {
old_job.title.clone()
};
let new_job_id = scheduler
.dispatch_job(&old_job.user_id, &title, &old_job.description, None)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
return Ok(Json(serde_json::json!({
"status": "restarted",
"old_job_id": old_job_id,
"new_job_id": new_job_id,
})));
} }
Err((StatusCode::NOT_FOUND, "Job not found".to_string()))
} }
/// Submit a follow-up prompt to a running job. /// Submit a follow-up prompt to a running job.
@@ -545,7 +476,6 @@ pub async fn jobs_restart_handler(
/// - Worker-mode sandbox jobs → not supported (no mechanism to inject) /// - Worker-mode sandbox jobs → not supported (no mechanism to inject)
pub async fn jobs_prompt_handler( pub async fn jobs_prompt_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
Json(body): Json<serde_json::Value>, Json(body): Json<serde_json::Value>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> { ) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
@@ -564,15 +494,10 @@ pub async fn jobs_prompt_handler(
let done = body.get("done").and_then(|v| v.as_bool()).unwrap_or(false); let done = body.get("done").and_then(|v| v.as_bool()).unwrap_or(false);
// Try sandbox job path first: verify ownership, then route to Claude Code or reject. // Try sandbox job path: check if we have a sandbox record for this ID.
if let Some(ref s) = state.store if let Some(ref s) = state.store
&& let Ok(Some(sandbox_job)) = s.get_sandbox_job(job_id).await && let Ok(Some(_)) = s.get_sandbox_job(job_id).await
{ {
// Verify ownership.
if sandbox_job.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
// It's a sandbox job. Check if Claude Code mode. // It's a sandbox job. Check if Claude Code mode.
let mode = s.get_sandbox_job_mode(job_id).await.ok().flatten(); let mode = s.get_sandbox_job_mode(job_id).await.ok().flatten();
if mode.as_deref() == Some("claude_code") { if mode.as_deref() == Some("claude_code") {
@@ -597,26 +522,7 @@ pub async fn jobs_prompt_handler(
} }
} }
// Try agent job path: verify ownership, then send via scheduler. // Try agent job path: send via scheduler.
if let Some(ref store) = state.store {
match store.get_job(job_id).await {
Ok(Some(agent_job)) => {
if agent_job.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
}
Ok(None) => {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
Err(e) => {
return Err((
StatusCode::INTERNAL_SERVER_ERROR,
format!("Database error: {}", e),
));
}
}
}
let slot = state.scheduler.as_ref().ok_or(( let slot = state.scheduler.as_ref().ok_or((
StatusCode::NOT_IMPLEMENTED, StatusCode::NOT_IMPLEMENTED,
"Agent job prompts require the scheduler to be configured".to_string(), "Agent job prompts require the scheduler to be configured".to_string(),
@@ -644,7 +550,6 @@ pub async fn jobs_prompt_handler(
/// Load persisted job events for a job (for history replay on page open). /// Load persisted job events for a job (for history replay on page open).
pub async fn jobs_events_handler( pub async fn jobs_events_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> { ) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or(( let store = state.store.as_ref().ok_or((
@@ -656,24 +561,6 @@ pub async fn jobs_events_handler(
.parse() .parse()
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?; .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
// Verify ownership before returning events.
match store.get_sandbox_job(job_id).await {
Ok(Some(job)) => {
if job.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
}
Ok(None) => {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
Err(e) => {
return Err((
StatusCode::INTERNAL_SERVER_ERROR,
format!("Database error: {}", e),
));
}
}
let events = store let events = store
.list_job_events(job_id, None) .list_job_events(job_id, None)
.await .await
@@ -706,7 +593,6 @@ pub struct FilePathQuery {
pub async fn job_files_list_handler( pub async fn job_files_list_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
Query(query): Query<FilePathQuery>, Query(query): Query<FilePathQuery>,
) -> Result<Json<ProjectFilesResponse>, (StatusCode, String)> { ) -> Result<Json<ProjectFilesResponse>, (StatusCode, String)> {
@@ -724,10 +610,6 @@ pub async fn job_files_list_handler(
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "Job not found".to_string()))?; .ok_or((StatusCode::NOT_FOUND, "Job not found".to_string()))?;
if job.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
let base = std::path::PathBuf::from(&job.project_dir); let base = std::path::PathBuf::from(&job.project_dir);
let rel_path = query.path.as_deref().unwrap_or(""); let rel_path = query.path.as_deref().unwrap_or("");
let target = base.join(rel_path); let target = base.join(rel_path);
@@ -774,7 +656,6 @@ pub async fn job_files_list_handler(
pub async fn job_files_read_handler( pub async fn job_files_read_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
Query(query): Query<FilePathQuery>, Query(query): Query<FilePathQuery>,
) -> Result<Json<ProjectFileReadResponse>, (StatusCode, String)> { ) -> Result<Json<ProjectFileReadResponse>, (StatusCode, String)> {
@@ -792,10 +673,6 @@ pub async fn job_files_read_handler(
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "Job not found".to_string()))?; .ok_or((StatusCode::NOT_FOUND, "Job not found".to_string()))?;
if job.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
let path = query.path.as_deref().ok_or(( let path = query.path.as_deref().ok_or((
StatusCode::BAD_REQUEST, StatusCode::BAD_REQUEST,
"path parameter required".to_string(), "path parameter required".to_string(),
+21 -92
View File
@@ -9,27 +9,8 @@ use axum::{
}; };
use serde::Deserialize; use serde::Deserialize;
use crate::channels::web::auth::{AuthenticatedUser, UserIdentity};
use crate::channels::web::server::GatewayState; use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*; use crate::channels::web::types::*;
use crate::workspace::Workspace;
/// Resolve the workspace for the authenticated user.
///
/// Prefers `workspace_pool` (multi-user mode) when available, falling back
/// to the single-user `state.workspace`.
pub(crate) async fn resolve_workspace(
state: &GatewayState,
user: &UserIdentity,
) -> Result<Arc<Workspace>, (StatusCode, String)> {
if let Some(ref pool) = state.workspace_pool {
return Ok(pool.get_or_create(user).await);
}
state.workspace.as_ref().cloned().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Workspace not available".to_string(),
))
}
#[derive(Deserialize)] #[derive(Deserialize)]
pub struct TreeQuery { pub struct TreeQuery {
@@ -39,10 +20,12 @@ pub struct TreeQuery {
pub async fn memory_tree_handler( pub async fn memory_tree_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Query(_query): Query<TreeQuery>, Query(_query): Query<TreeQuery>,
) -> Result<Json<MemoryTreeResponse>, (StatusCode, String)> { ) -> Result<Json<MemoryTreeResponse>, (StatusCode, String)> {
let workspace = resolve_workspace(&state, &user).await?; let workspace = state.workspace.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Workspace not available".to_string(),
))?;
// Build tree from list_all (flat list of all paths) // Build tree from list_all (flat list of all paths)
let all_paths = workspace let all_paths = workspace
@@ -85,10 +68,12 @@ pub struct ListQuery {
pub async fn memory_list_handler( pub async fn memory_list_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Query(query): Query<ListQuery>, Query(query): Query<ListQuery>,
) -> Result<Json<MemoryListResponse>, (StatusCode, String)> { ) -> Result<Json<MemoryListResponse>, (StatusCode, String)> {
let workspace = resolve_workspace(&state, &user).await?; let workspace = state.workspace.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Workspace not available".to_string(),
))?;
let path = query.path.as_deref().unwrap_or(""); let path = query.path.as_deref().unwrap_or("");
let entries = workspace let entries = workspace
@@ -119,10 +104,12 @@ pub struct ReadQuery {
pub async fn memory_read_handler( pub async fn memory_read_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Query(query): Query<ReadQuery>, Query(query): Query<ReadQuery>,
) -> Result<Json<MemoryReadResponse>, (StatusCode, String)> { ) -> Result<Json<MemoryReadResponse>, (StatusCode, String)> {
let workspace = resolve_workspace(&state, &user).await?; let workspace = state.workspace.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Workspace not available".to_string(),
))?;
let doc = workspace let doc = workspace
.read(&query.path) .read(&query.path)
@@ -136,75 +123,17 @@ pub async fn memory_read_handler(
})) }))
} }
pub async fn memory_write_handler( // memory_write_handler lives in server.rs (layer-aware version with append,
State(state): State<Arc<GatewayState>>, // privacy redirect, and proper error status codes).
AuthenticatedUser(user): AuthenticatedUser,
Json(req): Json<MemoryWriteRequest>,
) -> Result<Json<MemoryWriteResponse>, (StatusCode, String)> {
let workspace = resolve_workspace(&state, &user).await?;
// Route through layer-aware methods when a layer is specified.
//
// Note: unlike MemoryWriteTool, this endpoint does NOT block writes to
// identity files (IDENTITY.md, SOUL.md, etc.). The HTTP API is an
// authenticated admin interface; the supervisor uses it to seed identity
// files at startup. Identity-file protection is enforced at the tool
// layer (LLM-facing) where the write originates from an untrusted agent.
if let Some(ref layer_name) = req.layer {
let result = if req.append {
workspace
.append_to_layer(layer_name, &req.path, &req.content, req.force)
.await
} else {
workspace
.write_to_layer(layer_name, &req.path, &req.content, req.force)
.await
}
.map_err(|e| {
use crate::error::WorkspaceError;
let status = match &e {
WorkspaceError::LayerNotFound { .. } => StatusCode::BAD_REQUEST,
WorkspaceError::LayerReadOnly { .. } => StatusCode::FORBIDDEN,
WorkspaceError::PrivacyRedirectFailed => StatusCode::UNPROCESSABLE_ENTITY,
_ => StatusCode::INTERNAL_SERVER_ERROR,
};
(status, e.to_string())
})?;
return Ok(Json(MemoryWriteResponse {
path: req.path,
status: "written",
redirected: Some(result.redirected),
actual_layer: Some(result.actual_layer),
}));
}
// Non-layer path: honor the append field
if req.append {
workspace
.append(&req.path, &req.content)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
} else {
workspace
.write(&req.path, &req.content)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
}
Ok(Json(MemoryWriteResponse {
path: req.path,
status: "written",
redirected: None,
actual_layer: None,
}))
}
pub async fn memory_search_handler( pub async fn memory_search_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Json(req): Json<MemorySearchRequest>, Json(req): Json<MemorySearchRequest>,
) -> Result<Json<MemorySearchResponse>, (StatusCode, String)> { ) -> Result<Json<MemorySearchResponse>, (StatusCode, String)> {
let workspace = resolve_workspace(&state, &user).await?; let workspace = state.workspace.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Workspace not available".to_string(),
))?;
let limit = req.limit.unwrap_or(10); let limit = req.limit.unwrap_or(10);
let results = workspace let results = workspace
@@ -213,10 +142,10 @@ pub async fn memory_search_handler(
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
let hits: Vec<SearchHit> = results let hits: Vec<SearchHit> = results
.iter() .into_iter()
.map(|r| SearchHit { .map(|r| SearchHit {
path: r.document_id.to_string(), path: r.document_path,
content: r.content.clone(), content: r.content,
score: r.score as f64, score: r.score as f64,
}) })
.collect(); .collect();
+12 -6
View File
@@ -1,14 +1,14 @@
//! Handler modules for the web gateway API. //! Handler modules for the web gateway API.
//! //!
//! Each module groups related endpoint handlers by domain. //! Each module groups related endpoint handlers by domain.
//!
//! # Migration status
//!
//! `skills` is the canonical implementation used by `server.rs`.
//! The remaining modules are in-progress migrations from inline server.rs
//! handlers; their functions are not yet wired up, hence the `dead_code` allow.
pub mod jobs;
pub mod memory;
pub mod routines;
pub mod secrets;
pub mod skills; pub mod skills;
pub mod tokens;
pub mod users;
// Modules not yet wired into server.rs router -- suppress dead_code until // Modules not yet wired into server.rs router -- suppress dead_code until
// they replace their inline counterparts. // they replace their inline counterparts.
@@ -17,6 +17,12 @@ pub mod chat;
#[allow(dead_code)] #[allow(dead_code)]
pub mod extensions; pub mod extensions;
#[allow(dead_code)] #[allow(dead_code)]
pub mod jobs;
#[allow(dead_code)]
pub mod memory;
#[allow(dead_code)]
pub mod routines;
#[allow(dead_code)]
pub mod settings; pub mod settings;
#[allow(dead_code)] #[allow(dead_code)]
pub mod static_files; pub mod static_files;
+5 -44
View File
@@ -11,14 +11,12 @@ use serde::Deserialize;
use uuid::Uuid; use uuid::Uuid;
use crate::agent::routine::{Trigger, next_cron_fire}; use crate::agent::routine::{Trigger, next_cron_fire};
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::server::GatewayState; use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*; use crate::channels::web::types::*;
use crate::error::RoutineError; use crate::error::RoutineError;
pub async fn routines_list_handler( pub async fn routines_list_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<RoutineListResponse>, (StatusCode, String)> { ) -> Result<Json<RoutineListResponse>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or(( let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE, StatusCode::SERVICE_UNAVAILABLE,
@@ -26,7 +24,7 @@ pub async fn routines_list_handler(
))?; ))?;
let routines = store let routines = store
.list_routines(&user.user_id) .list_all_routines()
.await .await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
@@ -37,7 +35,6 @@ pub async fn routines_list_handler(
pub async fn routines_summary_handler( pub async fn routines_summary_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<RoutineSummaryResponse>, (StatusCode, String)> { ) -> Result<Json<RoutineSummaryResponse>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or(( let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE, StatusCode::SERVICE_UNAVAILABLE,
@@ -45,7 +42,7 @@ pub async fn routines_summary_handler(
))?; ))?;
let routines = store let routines = store
.list_routines(&user.user_id) .list_all_routines()
.await .await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
@@ -81,7 +78,6 @@ pub async fn routines_summary_handler(
pub async fn routines_detail_handler( pub async fn routines_detail_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
) -> Result<Json<RoutineDetailResponse>, (StatusCode, String)> { ) -> Result<Json<RoutineDetailResponse>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or(( let store = state.store.as_ref().ok_or((
@@ -98,10 +94,6 @@ pub async fn routines_detail_handler(
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?; .ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
if routine.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Routine not found".to_string()));
}
let runs = store let runs = store
.list_routine_runs(routine_id, 20) .list_routine_runs(routine_id, 20)
.await .await
@@ -114,7 +106,7 @@ pub async fn routines_detail_handler(
trigger_type: run.trigger_type.clone(), trigger_type: run.trigger_type.clone(),
started_at: run.started_at.to_rfc3339(), started_at: run.started_at.to_rfc3339(),
completed_at: run.completed_at.map(|dt| dt.to_rfc3339()), completed_at: run.completed_at.map(|dt| dt.to_rfc3339()),
status: run.status.to_string(), status: format!("{:?}", run.status),
result_summary: run.result_summary.clone(), result_summary: run.result_summary.clone(),
tokens_used: run.tokens_used, tokens_used: run.tokens_used,
job_id: run.job_id, job_id: run.job_id,
@@ -145,7 +137,6 @@ pub async fn routines_detail_handler(
pub async fn routines_trigger_handler( pub async fn routines_trigger_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> { ) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
// Clone the Arc out of the lock to avoid holding the RwLock across .await. // Clone the Arc out of the lock to avoid holding the RwLock across .await.
@@ -161,7 +152,7 @@ pub async fn routines_trigger_handler(
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?; .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?;
let run_id = engine let run_id = engine
.fire_manual(routine_id, Some(&user.user_id)) .fire_manual(routine_id, Some(&state.user_id))
.await .await
.map_err(|e| (routine_error_status(&e), e.to_string()))?; .map_err(|e| (routine_error_status(&e), e.to_string()))?;
@@ -179,7 +170,6 @@ pub struct ToggleRequest {
pub async fn routines_toggle_handler( pub async fn routines_toggle_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
body: Option<Json<ToggleRequest>>, body: Option<Json<ToggleRequest>>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> { ) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
@@ -197,10 +187,6 @@ pub async fn routines_toggle_handler(
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?; .ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
if routine.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Routine not found".to_string()));
}
let was_enabled = routine.enabled; let was_enabled = routine.enabled;
// If a specific value was provided, use it; otherwise toggle. // If a specific value was provided, use it; otherwise toggle.
routine.enabled = match body { routine.enabled = match body {
@@ -244,7 +230,6 @@ pub async fn routines_toggle_handler(
pub async fn routines_delete_handler( pub async fn routines_delete_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> { ) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or(( let store = state.store.as_ref().ok_or((
@@ -255,17 +240,6 @@ pub async fn routines_delete_handler(
let routine_id = Uuid::parse_str(&id) let routine_id = Uuid::parse_str(&id)
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?; .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?;
// Verify ownership before deleting.
let routine = store
.get_routine(routine_id)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
if routine.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Routine not found".to_string()));
}
let deleted = store let deleted = store
.delete_routine(routine_id) .delete_routine(routine_id)
.await .await
@@ -287,10 +261,8 @@ pub async fn routines_delete_handler(
} }
} }
#[allow(dead_code)] // Used by server.rs inline version; kept in sync here for future migration.
pub async fn routines_runs_handler( pub async fn routines_runs_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> { ) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or(( let store = state.store.as_ref().ok_or((
@@ -301,17 +273,6 @@ pub async fn routines_runs_handler(
let routine_id = Uuid::parse_str(&id) let routine_id = Uuid::parse_str(&id)
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?; .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?;
// Verify ownership before listing runs.
let routine = store
.get_routine(routine_id)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
if routine.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Routine not found".to_string()));
}
let runs = store let runs = store
.list_routine_runs(routine_id, 50) .list_routine_runs(routine_id, 50)
.await .await
@@ -324,7 +285,7 @@ pub async fn routines_runs_handler(
trigger_type: run.trigger_type.clone(), trigger_type: run.trigger_type.clone(),
started_at: run.started_at.to_rfc3339(), started_at: run.started_at.to_rfc3339(),
completed_at: run.completed_at.map(|dt| dt.to_rfc3339()), completed_at: run.completed_at.map(|dt| dt.to_rfc3339()),
status: run.status.to_string(), status: format!("{:?}", run.status),
result_summary: run.result_summary.clone(), result_summary: run.result_summary.clone(),
tokens_used: run.tokens_used, tokens_used: run.tokens_used,
job_id: run.job_id, job_id: run.job_id,
-134
View File
@@ -1,134 +0,0 @@
//! Admin secrets provisioning handlers.
//!
//! Allows an admin (typically an application backend) to create, list, and
//! delete secrets on behalf of individual users so their IronClaw agent can
//! call back to external services with per-user credentials.
use std::sync::Arc;
use axum::{
Json,
extract::{Path, State},
http::StatusCode,
};
use crate::channels::web::auth::AdminUser;
use crate::channels::web::server::GatewayState;
use crate::secrets::CreateSecretParams;
/// PUT /api/admin/users/{user_id}/secrets/{name} — create or update a secret.
///
/// Upserts: if a secret with the same (user_id, name) already exists it is
/// overwritten. The plaintext value is encrypted at rest (AES-256-GCM) and
/// never returned by any endpoint.
pub async fn secrets_put_handler(
State(state): State<Arc<GatewayState>>,
AdminUser(_admin): AdminUser,
Path((user_id, name)): Path<(String, String)>,
Json(body): Json<serde_json::Value>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let secrets = state.secrets_store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Secrets store not available".to_string(),
))?;
let value = body
.get("value")
.and_then(|v| v.as_str())
.ok_or((
StatusCode::BAD_REQUEST,
"Missing required field 'value'".to_string(),
))?
.to_string();
let provider = body
.get("provider")
.and_then(|v| v.as_str())
.map(String::from);
let expires_at = body
.get("expires_in_days")
.and_then(|v| v.as_u64())
.map(|d| d.min(36500))
.map(|days| chrono::Utc::now() + chrono::Duration::days(days as i64));
let mut params = CreateSecretParams::new(name.clone(), value);
if let Some(p) = provider {
params = params.with_provider(p);
}
if let Some(exp) = expires_at {
params = params.with_expiry(exp);
}
secrets
.create(&user_id, params)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
Ok(Json(serde_json::json!({
"user_id": user_id,
"name": name.to_lowercase(),
"status": "created",
})))
}
/// GET /api/admin/users/{user_id}/secrets — list a user's secrets (names only).
///
/// Never returns secret values or hashes.
pub async fn secrets_list_handler(
State(state): State<Arc<GatewayState>>,
AdminUser(_admin): AdminUser,
Path(user_id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let secrets = state.secrets_store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Secrets store not available".to_string(),
))?;
let refs = secrets
.list(&user_id)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
let secrets_json: Vec<serde_json::Value> = refs
.into_iter()
.map(|r| {
serde_json::json!({
"name": r.name,
"provider": r.provider,
})
})
.collect();
Ok(Json(serde_json::json!({
"user_id": user_id,
"secrets": secrets_json,
})))
}
/// DELETE /api/admin/users/{user_id}/secrets/{name} — delete a user's secret.
pub async fn secrets_delete_handler(
State(state): State<Arc<GatewayState>>,
AdminUser(_admin): AdminUser,
Path((user_id, name)): Path<(String, String)>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let secrets = state.secrets_store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Secrets store not available".to_string(),
))?;
let deleted = secrets
.delete(&user_id, &name)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
if !deleted {
return Err((StatusCode::NOT_FOUND, "Secret not found".to_string()));
}
Ok(Json(serde_json::json!({
"user_id": user_id,
"name": name,
"deleted": true,
})))
}
+6 -13
View File
@@ -8,19 +8,17 @@ use axum::{
http::StatusCode, http::StatusCode,
}; };
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::server::GatewayState; use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*; use crate::channels::web::types::*;
pub async fn settings_list_handler( pub async fn settings_list_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<SettingsListResponse>, StatusCode> { ) -> Result<Json<SettingsListResponse>, StatusCode> {
let store = state let store = state
.store .store
.as_ref() .as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?; .ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
let rows = store.list_settings(&user.user_id).await.map_err(|e| { let rows = store.list_settings(&state.user_id).await.map_err(|e| {
tracing::error!("Failed to list settings: {}", e); tracing::error!("Failed to list settings: {}", e);
StatusCode::INTERNAL_SERVER_ERROR StatusCode::INTERNAL_SERVER_ERROR
})?; })?;
@@ -39,7 +37,6 @@ pub async fn settings_list_handler(
pub async fn settings_get_handler( pub async fn settings_get_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(key): Path<String>, Path(key): Path<String>,
) -> Result<Json<SettingResponse>, StatusCode> { ) -> Result<Json<SettingResponse>, StatusCode> {
let store = state let store = state
@@ -47,7 +44,7 @@ pub async fn settings_get_handler(
.as_ref() .as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?; .ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
let row = store let row = store
.get_setting_full(&user.user_id, &key) .get_setting_full(&state.user_id, &key)
.await .await
.map_err(|e| { .map_err(|e| {
tracing::error!("Failed to get setting '{}': {}", key, e); tracing::error!("Failed to get setting '{}': {}", key, e);
@@ -64,7 +61,6 @@ pub async fn settings_get_handler(
pub async fn settings_set_handler( pub async fn settings_set_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(key): Path<String>, Path(key): Path<String>,
Json(body): Json<SettingWriteRequest>, Json(body): Json<SettingWriteRequest>,
) -> Result<StatusCode, StatusCode> { ) -> Result<StatusCode, StatusCode> {
@@ -73,7 +69,7 @@ pub async fn settings_set_handler(
.as_ref() .as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?; .ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
store store
.set_setting(&user.user_id, &key, &body.value) .set_setting(&state.user_id, &key, &body.value)
.await .await
.map_err(|e| { .map_err(|e| {
tracing::error!("Failed to set setting '{}': {}", key, e); tracing::error!("Failed to set setting '{}': {}", key, e);
@@ -85,7 +81,6 @@ pub async fn settings_set_handler(
pub async fn settings_delete_handler( pub async fn settings_delete_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(key): Path<String>, Path(key): Path<String>,
) -> Result<StatusCode, StatusCode> { ) -> Result<StatusCode, StatusCode> {
let store = state let store = state
@@ -93,7 +88,7 @@ pub async fn settings_delete_handler(
.as_ref() .as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?; .ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
store store
.delete_setting(&user.user_id, &key) .delete_setting(&state.user_id, &key)
.await .await
.map_err(|e| { .map_err(|e| {
tracing::error!("Failed to delete setting '{}': {}", key, e); tracing::error!("Failed to delete setting '{}': {}", key, e);
@@ -105,13 +100,12 @@ pub async fn settings_delete_handler(
pub async fn settings_export_handler( pub async fn settings_export_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<SettingsExportResponse>, StatusCode> { ) -> Result<Json<SettingsExportResponse>, StatusCode> {
let store = state let store = state
.store .store
.as_ref() .as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?; .ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
let settings = store.get_all_settings(&user.user_id).await.map_err(|e| { let settings = store.get_all_settings(&state.user_id).await.map_err(|e| {
tracing::error!("Failed to export settings: {}", e); tracing::error!("Failed to export settings: {}", e);
StatusCode::INTERNAL_SERVER_ERROR StatusCode::INTERNAL_SERVER_ERROR
})?; })?;
@@ -121,7 +115,6 @@ pub async fn settings_export_handler(
pub async fn settings_import_handler( pub async fn settings_import_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Json(body): Json<SettingsImportRequest>, Json(body): Json<SettingsImportRequest>,
) -> Result<StatusCode, StatusCode> { ) -> Result<StatusCode, StatusCode> {
let store = state let store = state
@@ -129,7 +122,7 @@ pub async fn settings_import_handler(
.as_ref() .as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?; .ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
store store
.set_all_settings(&user.user_id, &body.settings) .set_all_settings(&state.user_id, &body.settings)
.await .await
.map_err(|e| { .map_err(|e| {
tracing::error!("Failed to import settings: {}", e); tracing::error!("Failed to import settings: {}", e);
-9
View File
@@ -8,13 +8,11 @@ use axum::{
http::StatusCode, http::StatusCode,
}; };
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::server::GatewayState; use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*; use crate::channels::web::types::*;
pub async fn skills_list_handler( pub async fn skills_list_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(_user): AuthenticatedUser,
) -> Result<Json<SkillListResponse>, (StatusCode, String)> { ) -> Result<Json<SkillListResponse>, (StatusCode, String)> {
let registry = state.skill_registry.as_ref().ok_or(( let registry = state.skill_registry.as_ref().ok_or((
StatusCode::NOT_IMPLEMENTED, StatusCode::NOT_IMPLEMENTED,
@@ -47,7 +45,6 @@ pub async fn skills_list_handler(
pub async fn skills_search_handler( pub async fn skills_search_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(_user): AuthenticatedUser,
Json(req): Json<SkillSearchRequest>, Json(req): Json<SkillSearchRequest>,
) -> Result<Json<SkillSearchResponse>, (StatusCode, String)> { ) -> Result<Json<SkillSearchResponse>, (StatusCode, String)> {
let registry = state.skill_registry.as_ref().ok_or(( let registry = state.skill_registry.as_ref().ok_or((
@@ -122,7 +119,6 @@ pub async fn skills_search_handler(
pub async fn skills_install_handler( pub async fn skills_install_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
headers: axum::http::HeaderMap, headers: axum::http::HeaderMap,
Json(req): Json<SkillInstallRequest>, Json(req): Json<SkillInstallRequest>,
) -> Result<Json<ActionResponse>, (StatusCode, String)> { ) -> Result<Json<ActionResponse>, (StatusCode, String)> {
@@ -139,8 +135,6 @@ pub async fn skills_install_handler(
)); ));
} }
tracing::info!(user_id = %user.user_id, skill = %req.name, "skill install requested");
let registry = state.skill_registry.as_ref().ok_or(( let registry = state.skill_registry.as_ref().ok_or((
StatusCode::NOT_IMPLEMENTED, StatusCode::NOT_IMPLEMENTED,
"Skills system not enabled".to_string(), "Skills system not enabled".to_string(),
@@ -225,7 +219,6 @@ pub async fn skills_install_handler(
pub async fn skills_remove_handler( pub async fn skills_remove_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
headers: axum::http::HeaderMap, headers: axum::http::HeaderMap,
Path(name): Path<String>, Path(name): Path<String>,
) -> Result<Json<ActionResponse>, (StatusCode, String)> { ) -> Result<Json<ActionResponse>, (StatusCode, String)> {
@@ -241,8 +234,6 @@ pub async fn skills_remove_handler(
)); ));
} }
tracing::info!(user_id = %user.user_id, skill = %name, "skill remove requested");
let registry = state.skill_registry.as_ref().ok_or(( let registry = state.skill_registry.as_ref().ok_or((
StatusCode::NOT_IMPLEMENTED, StatusCode::NOT_IMPLEMENTED,
"Skills system not enabled".to_string(), "Skills system not enabled".to_string(),
@@ -7,7 +7,6 @@ use axum::{
}; };
use crate::bootstrap::ironclaw_base_dir; use crate::bootstrap::ironclaw_base_dir;
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::types::*; use crate::channels::web::types::*;
// --- Static file handlers --- // --- Static file handlers ---
@@ -114,7 +113,6 @@ use crate::channels::web::server::GatewayState;
pub async fn logs_events_handler( pub async fn logs_events_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(_user): AuthenticatedUser,
) -> Result< ) -> Result<
Sse<impl futures::Stream<Item = Result<Event, Infallible>> + Send + 'static>, Sse<impl futures::Stream<Item = Result<Event, Infallible>> + Send + 'static>,
(StatusCode, String), (StatusCode, String),
@@ -154,7 +152,6 @@ pub async fn logs_events_handler(
pub async fn gateway_status_handler( pub async fn gateway_status_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(_user): AuthenticatedUser,
) -> Json<GatewayStatusResponse> { ) -> Json<GatewayStatusResponse> {
let sse_connections = state.sse.connection_count(); let sse_connections = state.sse.connection_count();
let ws_connections = state let ws_connections = state
-150
View File
@@ -1,150 +0,0 @@
//! API token management handlers.
use std::sync::Arc;
use axum::{
Json,
extract::{Path, State},
http::StatusCode,
};
use rand::RngCore;
use rand::rngs::OsRng;
use uuid::Uuid;
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::server::GatewayState;
/// POST /api/tokens — create a new API token (returns plaintext ONCE).
pub async fn tokens_create_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Json(body): Json<serde_json::Value>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Database not available".to_string(),
))?;
let name = body
.get("name")
.and_then(|v| v.as_str())
.ok_or((
StatusCode::BAD_REQUEST,
"Missing required field 'name'".to_string(),
))?
.to_string();
let expires_in_days: Option<i64> = body
.get("expires_in_days")
.and_then(|v| v.as_u64())
.map(|d| d.min(36500) as i64);
let expires_at = expires_in_days.map(|days| chrono::Utc::now() + chrono::Duration::days(days));
// Generate 32 random bytes for the token.
// Hash the hex-encoded plaintext (what the user sends as Bearer token),
// NOT the raw bytes — must match hash_token() in auth.rs.
let mut token_bytes = [0u8; 32];
OsRng.fill_bytes(&mut token_bytes);
let plaintext_token = hex::encode(token_bytes);
let hash = crate::channels::web::auth::hash_token(&plaintext_token);
// First 8 chars of the hex token as a prefix for identification.
let token_prefix = &plaintext_token[..8];
// Admin users can create tokens for other users via optional "user_id" field.
let target_user = body
.get("user_id")
.and_then(|v| v.as_str())
.filter(|_| user.role == "admin")
.unwrap_or(&user.user_id);
// Verify the target user exists to prevent orphan tokens.
if target_user != user.user_id {
store
.get_user(target_user)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((
StatusCode::BAD_REQUEST,
format!("Target user '{target_user}' not found"),
))?;
}
let record = store
.create_api_token(target_user, &name, &hash, token_prefix, expires_at)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
// Return the plaintext token — this is the ONLY time it is shown.
Ok(Json(serde_json::json!({
"token": plaintext_token,
"id": record.id.to_string(),
"name": record.name,
"token_prefix": record.token_prefix,
"expires_at": record.expires_at.map(|dt| dt.to_rfc3339()),
"created_at": record.created_at.to_rfc3339(),
})))
}
/// GET /api/tokens — list the current user's tokens (no hashes).
pub async fn tokens_list_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Database not available".to_string(),
))?;
let tokens = store
.list_api_tokens(&user.user_id)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
let tokens_json: Vec<serde_json::Value> = tokens
.into_iter()
.map(|t| {
serde_json::json!({
"id": t.id.to_string(),
"name": t.name,
"token_prefix": t.token_prefix,
"expires_at": t.expires_at.map(|dt| dt.to_rfc3339()),
"last_used_at": t.last_used_at.map(|dt| dt.to_rfc3339()),
"created_at": t.created_at.to_rfc3339(),
"revoked_at": t.revoked_at.map(|dt| dt.to_rfc3339()),
})
})
.collect();
Ok(Json(serde_json::json!({ "tokens": tokens_json })))
}
/// DELETE /api/tokens/{id} — revoke a token.
pub async fn tokens_revoke_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Database not available".to_string(),
))?;
let token_id = Uuid::parse_str(&id)
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid token ID".to_string()))?;
let revoked = store
.revoke_api_token(token_id, &user.user_id)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
if !revoked {
return Err((StatusCode::NOT_FOUND, "Token not found".to_string()));
}
Ok(Json(serde_json::json!({
"status": "revoked",
"id": token_id.to_string(),
})))
}
-406
View File
@@ -1,406 +0,0 @@
//! User management API handlers (admin).
use std::sync::Arc;
use axum::{
Json,
extract::{Path, State},
http::StatusCode,
};
use rand::RngCore;
use rand::rngs::OsRng;
use uuid::Uuid;
use crate::channels::web::auth::{AdminUser, AuthenticatedUser};
use crate::channels::web::server::GatewayState;
use crate::db::UserRecord;
/// POST /api/admin/users — create a new user.
pub async fn users_create_handler(
State(state): State<Arc<GatewayState>>,
AdminUser(user): AdminUser,
Json(body): Json<serde_json::Value>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Database not available".to_string(),
))?;
let display_name = body
.get("display_name")
.and_then(|v| v.as_str())
.ok_or((
StatusCode::BAD_REQUEST,
"Missing required field 'display_name'".to_string(),
))?
.to_string();
let email = body.get("email").and_then(|v| v.as_str()).map(String::from);
let role = body
.get("role")
.and_then(|v| v.as_str())
.unwrap_or("member")
.to_string();
if role != "admin" && role != "member" {
return Err((
StatusCode::BAD_REQUEST,
"role must be 'admin' or 'member'".to_string(),
));
}
let user_id = Uuid::new_v4().to_string();
let now = chrono::Utc::now();
let user_record = UserRecord {
id: user_id.clone(),
email,
display_name: display_name.clone(),
status: "active".to_string(),
role,
created_at: now,
updated_at: now,
last_login_at: None,
created_by: match store.get_user(&user.user_id).await {
Ok(Some(_)) => Some(user.user_id.clone()),
_ => None,
},
metadata: serde_json::json!({}),
};
store
.create_user(&user_record)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
// Generate a first API token so the new user can authenticate immediately.
// Hash the hex-encoded plaintext (what the user sends as Bearer token),
// NOT the raw bytes — must match hash_token() in auth.rs.
let mut token_bytes = [0u8; 32];
OsRng.fill_bytes(&mut token_bytes);
let plaintext_token = hex::encode(token_bytes);
let token_hash = crate::channels::web::auth::hash_token(&plaintext_token);
let token_prefix = &plaintext_token[..8];
let _token_record = store
.create_api_token(&user_id, "initial", &token_hash, token_prefix, None)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
Ok(Json(serde_json::json!({
"id": user_record.id,
"email": user_record.email,
"display_name": user_record.display_name,
"status": user_record.status,
"role": user_record.role,
"token": plaintext_token,
"created_at": user_record.created_at.to_rfc3339(),
"created_by": user_record.created_by,
})))
}
/// GET /api/admin/users — list all users.
pub async fn users_list_handler(
State(state): State<Arc<GatewayState>>,
AdminUser(_user): AdminUser,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Database not available".to_string(),
))?;
let users = store
.list_users(None)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
let users_json: Vec<serde_json::Value> = users
.into_iter()
.map(|u| {
serde_json::json!({
"id": u.id,
"email": u.email,
"display_name": u.display_name,
"status": u.status,
"role": u.role,
"created_at": u.created_at.to_rfc3339(),
"updated_at": u.updated_at.to_rfc3339(),
"last_login_at": u.last_login_at.map(|dt| dt.to_rfc3339()),
"created_by": u.created_by,
})
})
.collect();
Ok(Json(serde_json::json!({ "users": users_json })))
}
/// GET /api/admin/users/{id} — get a single user.
pub async fn users_detail_handler(
State(state): State<Arc<GatewayState>>,
AdminUser(_user): AdminUser,
Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Database not available".to_string(),
))?;
let user_record = store
.get_user(&id)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "User not found".to_string()))?;
Ok(Json(serde_json::json!({
"id": user_record.id,
"email": user_record.email,
"display_name": user_record.display_name,
"status": user_record.status,
"role": user_record.role,
"created_at": user_record.created_at.to_rfc3339(),
"updated_at": user_record.updated_at.to_rfc3339(),
"last_login_at": user_record.last_login_at.map(|dt| dt.to_rfc3339()),
"created_by": user_record.created_by,
"metadata": user_record.metadata,
})))
}
/// PATCH /api/admin/users/{id} — update a user's profile.
pub async fn users_update_handler(
State(state): State<Arc<GatewayState>>,
AdminUser(_user): AdminUser,
Path(id): Path<String>,
Json(body): Json<serde_json::Value>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Database not available".to_string(),
))?;
// Verify the user exists.
let existing = store
.get_user(&id)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "User not found".to_string()))?;
let display_name = body
.get("display_name")
.and_then(|v| v.as_str())
.unwrap_or(&existing.display_name);
let metadata = body.get("metadata").unwrap_or(&existing.metadata);
store
.update_user_profile(&id, display_name, metadata)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
// Re-fetch the updated record to return consistent data.
let updated = store
.get_user(&id)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "User not found".to_string()))?;
Ok(Json(serde_json::json!({
"id": updated.id,
"email": updated.email,
"display_name": updated.display_name,
"status": updated.status,
"role": updated.role,
"created_at": updated.created_at.to_rfc3339(),
"updated_at": updated.updated_at.to_rfc3339(),
"metadata": updated.metadata,
})))
}
/// POST /api/admin/users/{id}/suspend — suspend a user.
pub async fn users_suspend_handler(
State(state): State<Arc<GatewayState>>,
AdminUser(_user): AdminUser,
Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Database not available".to_string(),
))?;
// Verify the user exists.
store
.get_user(&id)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "User not found".to_string()))?;
store
.update_user_status(&id, "suspended")
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
Ok(Json(serde_json::json!({
"id": id,
"status": "suspended",
})))
}
/// POST /api/admin/users/{id}/activate — activate a user.
pub async fn users_activate_handler(
State(state): State<Arc<GatewayState>>,
AdminUser(_user): AdminUser,
Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Database not available".to_string(),
))?;
// Verify the user exists.
store
.get_user(&id)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "User not found".to_string()))?;
store
.update_user_status(&id, "active")
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
Ok(Json(serde_json::json!({
"id": id,
"status": "active",
})))
}
/// DELETE /api/admin/users/{id} — delete a user and all their data.
pub async fn users_delete_handler(
State(state): State<Arc<GatewayState>>,
AdminUser(_user): AdminUser,
Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Database not available".to_string(),
))?;
let deleted = store
.delete_user(&id)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
if !deleted {
return Err((StatusCode::NOT_FOUND, "User not found".to_string()));
}
Ok(Json(serde_json::json!({
"id": id,
"deleted": true,
})))
}
/// GET /api/profile — get the authenticated user's own profile.
pub async fn profile_get_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Database not available".to_string(),
))?;
let record = store
.get_user(&user.user_id)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "User not found".to_string()))?;
Ok(Json(serde_json::json!({
"id": record.id,
"email": record.email,
"display_name": record.display_name,
"status": record.status,
"role": record.role,
"created_at": record.created_at.to_rfc3339(),
"last_login_at": record.last_login_at.map(|dt| dt.to_rfc3339()),
})))
}
/// PATCH /api/profile — update the authenticated user's own profile.
pub async fn profile_update_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Json(body): Json<serde_json::Value>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Database not available".to_string(),
))?;
let current = store
.get_user(&user.user_id)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "User not found".to_string()))?;
let display_name = body
.get("display_name")
.and_then(|v| v.as_str())
.unwrap_or(&current.display_name);
let metadata = body.get("metadata").unwrap_or(&current.metadata);
store
.update_user_profile(&user.user_id, display_name, metadata)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
Ok(Json(serde_json::json!({
"id": user.user_id,
"display_name": display_name,
"updated": true,
})))
}
/// GET /api/admin/usage — per-user LLM usage stats.
pub async fn usage_stats_handler(
State(state): State<Arc<GatewayState>>,
AdminUser(_user): AdminUser,
axum::extract::Query(params): axum::extract::Query<std::collections::HashMap<String, String>>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Database not available".to_string(),
))?;
let user_id = params.get("user_id").map(|s| s.as_str());
let period = params.get("period").map(|s| s.as_str()).unwrap_or("day");
let since = match period {
"week" => chrono::Utc::now() - chrono::Duration::days(7),
"month" => chrono::Utc::now() - chrono::Duration::days(30),
_ => chrono::Utc::now() - chrono::Duration::days(1),
};
let stats = store
.user_usage_stats(user_id, since)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
let entries: Vec<serde_json::Value> = stats
.iter()
.map(|s| {
serde_json::json!({
"user_id": s.user_id,
"model": s.model,
"call_count": s.call_count,
"input_tokens": s.input_tokens,
"output_tokens": s.output_tokens,
"total_cost": s.total_cost.to_string(),
})
})
.collect();
Ok(Json(serde_json::json!({
"period": period,
"since": since.to_rfc3339(),
"usage": entries,
})))
}
+3 -30
View File
@@ -54,37 +54,10 @@ fn validate_webhook_secret(
/// ///
/// This endpoint is **public** (no gateway auth token required) but protected /// This endpoint is **public** (no gateway auth token required) but protected
/// by the per-routine webhook secret sent via the `X-Webhook-Secret` header. /// 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( pub async fn webhook_trigger_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
Path(path): Path<String>, Path(path): Path<String>,
headers: HeaderMap, headers: HeaderMap,
) -> Result<Json<serde_json::Value>, (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<Arc<GatewayState>>,
Path((user_id, path)): Path<(String, String)>,
headers: HeaderMap,
) -> Result<Json<serde_json::Value>, (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<GatewayState>,
path: &str,
user_id: Option<&str>,
headers: &HeaderMap,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> { ) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
// Rate limit check // Rate limit check
if !state.webhook_rate_limiter.check() { if !state.webhook_rate_limiter.check() {
@@ -99,9 +72,9 @@ async fn fire_webhook_inner(
"Database not available".to_string(), "Database not available".to_string(),
))?; ))?;
// Targeted query — when user_id is provided, restrict to that user's routines // Targeted query instead of loading all routines
let routine = store let routine = store
.get_webhook_routine_by_path(path, user_id) .get_webhook_routine_by_path(&path)
.await .await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or(( .ok_or((
@@ -126,7 +99,7 @@ async fn fire_webhook_inner(
))? ))?
}; };
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 { let status = match &e {
crate::error::RoutineError::NotFound { .. } => StatusCode::NOT_FOUND, crate::error::RoutineError::NotFound { .. } => StatusCode::NOT_FOUND,
crate::error::RoutineError::Disabled { .. } crate::error::RoutineError::Disabled { .. }
+37 -101
View File
@@ -18,7 +18,6 @@ pub mod auth;
pub(crate) mod handlers; pub(crate) mod handlers;
pub mod log_layer; pub mod log_layer;
pub mod openai_compat; pub mod openai_compat;
pub mod responses_api;
pub mod server; pub mod server;
pub mod sse; pub mod sse;
pub mod types; pub mod types;
@@ -32,9 +31,6 @@ pub mod ws;
/// [`TestGatewayBuilder`](test_helpers::TestGatewayBuilder). /// [`TestGatewayBuilder`](test_helpers::TestGatewayBuilder).
pub mod test_helpers; pub mod test_helpers;
#[cfg(test)]
mod tests;
use std::net::SocketAddr; use std::net::SocketAddr;
use std::sync::Arc; use std::sync::Arc;
@@ -56,25 +52,23 @@ use crate::workspace::Workspace;
use self::log_layer::{LogBroadcaster, LogLevelHandle}; use self::log_layer::{LogBroadcaster, LogLevelHandle};
use self::auth::{CombinedAuthState, DbAuthenticator, MultiAuthState};
use self::server::GatewayState; use self::server::GatewayState;
use self::sse::SseManager; use self::sse::SseManager;
use self::types::AppEvent; use self::types::SseEvent;
/// Web gateway channel implementing the Channel trait. /// Web gateway channel implementing the Channel trait.
pub struct GatewayChannel { pub struct GatewayChannel {
config: GatewayConfig, config: GatewayConfig,
state: Arc<GatewayState>, state: Arc<GatewayState>,
/// Combined auth state: env-var tokens + optional DB-backed tokens. /// The actual auth token in use (generated or from config).
auth: CombinedAuthState, auth_token: String,
} }
impl GatewayChannel { impl GatewayChannel {
/// Create a new gateway channel. /// Create a new gateway channel.
/// ///
/// If no auth token is configured, generates a random one and prints it. /// If no auth token is configured, generates a random one and prints it.
/// Builds a single-user `MultiAuthState` from the config. pub fn new(config: GatewayConfig) -> Self {
pub fn new(config: GatewayConfig, owner_id: String) -> Self {
let auth_token = config.auth_token.clone().unwrap_or_else(|| { let auth_token = config.auth_token.clone().unwrap_or_else(|| {
use rand::RngCore; use rand::RngCore;
use rand::rngs::OsRng; use rand::rngs::OsRng;
@@ -83,16 +77,10 @@ impl GatewayChannel {
bytes.iter().map(|b| format!("{b:02x}")).collect() bytes.iter().map(|b| format!("{b:02x}")).collect()
}); });
let auth = CombinedAuthState {
env_auth: MultiAuthState::single(auth_token, owner_id.clone()),
db_auth: None,
};
let state = Arc::new(GatewayState { let state = Arc::new(GatewayState {
msg_tx: tokio::sync::RwLock::new(None), msg_tx: tokio::sync::RwLock::new(None),
sse: Arc::new(SseManager::new()), sse: SseManager::new(),
workspace: None, workspace: None,
workspace_pool: None,
session_manager: None, session_manager: None,
log_broadcaster: None, log_broadcaster: None,
log_level_handle: None, log_level_handle: None,
@@ -102,13 +90,13 @@ impl GatewayChannel {
job_manager: None, job_manager: None,
prompt_queue: None, prompt_queue: None,
scheduler: None, scheduler: None,
owner_id, user_id: config.user_id.clone(),
shutdown_tx: tokio::sync::RwLock::new(None), shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())), ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())),
llm_provider: None, llm_provider: None,
skill_registry: None, skill_registry: None,
skill_catalog: None, skill_catalog: None,
chat_rate_limiter: server::PerUserRateLimiter::new(30, 60), chat_rate_limiter: server::RateLimiter::new(30, 60),
oauth_rate_limiter: server::RateLimiter::new(10, 60), oauth_rate_limiter: server::RateLimiter::new(10, 60),
webhook_rate_limiter: server::RateLimiter::new(10, 60), webhook_rate_limiter: server::RateLimiter::new(10, 60),
registry_entries: Vec::new(), registry_entries: Vec::new(),
@@ -116,13 +104,12 @@ impl GatewayChannel {
routine_engine: Arc::new(tokio::sync::RwLock::new(None)), routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
startup_time: std::time::Instant::now(), startup_time: std::time::Instant::now(),
active_config: server::ActiveConfigSnapshot::default(), active_config: server::ActiveConfigSnapshot::default(),
secrets_store: None,
}); });
Self { Self {
config, config,
state, state,
auth, auth_token,
} }
} }
@@ -131,9 +118,8 @@ impl GatewayChannel {
let mut new_state = GatewayState { let mut new_state = GatewayState {
msg_tx: tokio::sync::RwLock::new(None), msg_tx: tokio::sync::RwLock::new(None),
// Preserve the existing broadcast channel so sender handles remain valid. // Preserve the existing broadcast channel so sender handles remain valid.
sse: Arc::new(SseManager::from_sender(self.state.sse.sender())), sse: SseManager::from_sender(self.state.sse.sender()),
workspace: self.state.workspace.clone(), workspace: self.state.workspace.clone(),
workspace_pool: self.state.workspace_pool.clone(),
session_manager: self.state.session_manager.clone(), session_manager: self.state.session_manager.clone(),
log_broadcaster: self.state.log_broadcaster.clone(), log_broadcaster: self.state.log_broadcaster.clone(),
log_level_handle: self.state.log_level_handle.clone(), log_level_handle: self.state.log_level_handle.clone(),
@@ -143,13 +129,13 @@ impl GatewayChannel {
job_manager: self.state.job_manager.clone(), job_manager: self.state.job_manager.clone(),
prompt_queue: self.state.prompt_queue.clone(), prompt_queue: self.state.prompt_queue.clone(),
scheduler: self.state.scheduler.clone(), scheduler: self.state.scheduler.clone(),
owner_id: self.state.owner_id.clone(), user_id: self.state.user_id.clone(),
shutdown_tx: tokio::sync::RwLock::new(None), shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: self.state.ws_tracker.clone(), ws_tracker: self.state.ws_tracker.clone(),
llm_provider: self.state.llm_provider.clone(), llm_provider: self.state.llm_provider.clone(),
skill_registry: self.state.skill_registry.clone(), skill_registry: self.state.skill_registry.clone(),
skill_catalog: self.state.skill_catalog.clone(), skill_catalog: self.state.skill_catalog.clone(),
chat_rate_limiter: server::PerUserRateLimiter::new(30, 60), chat_rate_limiter: server::RateLimiter::new(30, 60),
oauth_rate_limiter: server::RateLimiter::new(10, 60), oauth_rate_limiter: server::RateLimiter::new(10, 60),
webhook_rate_limiter: server::RateLimiter::new(10, 60), webhook_rate_limiter: server::RateLimiter::new(10, 60),
registry_entries: self.state.registry_entries.clone(), registry_entries: self.state.registry_entries.clone(),
@@ -157,7 +143,6 @@ impl GatewayChannel {
routine_engine: Arc::clone(&self.state.routine_engine), routine_engine: Arc::clone(&self.state.routine_engine),
startup_time: self.state.startup_time, startup_time: self.state.startup_time,
active_config: self.state.active_config.clone(), active_config: self.state.active_config.clone(),
secrets_store: self.state.secrets_store.clone(),
}; };
mutate(&mut new_state); mutate(&mut new_state);
self.state = Arc::new(new_state); self.state = Arc::new(new_state);
@@ -205,12 +190,6 @@ impl GatewayChannel {
self self
} }
/// Enable DB-backed token authentication alongside env-var tokens.
pub fn with_db_auth(mut self, store: Arc<dyn Database>) -> Self {
self.auth.db_auth = Some(DbAuthenticator::new(store));
self
}
/// Inject the container job manager for sandbox operations. /// Inject the container job manager for sandbox operations.
pub fn with_job_manager(mut self, jm: Arc<ContainerJobManager>) -> Self { pub fn with_job_manager(mut self, jm: Arc<ContainerJobManager>) -> Self {
self.rebuild_state(|s| s.job_manager = Some(jm)); self.rebuild_state(|s| s.job_manager = Some(jm));
@@ -281,24 +260,9 @@ impl GatewayChannel {
self self
} }
/// Inject the secrets store for admin secret provisioning. /// Get the auth token (for printing to console on startup).
pub fn with_secrets_store(
mut self,
store: Arc<dyn crate::secrets::SecretsStore + Send + Sync>,
) -> Self {
self.rebuild_state(|s| s.secrets_store = Some(store));
self
}
/// Inject the per-user workspace pool for multi-user mode.
pub fn with_workspace_pool(mut self, pool: Arc<server::WorkspacePool>) -> Self {
self.rebuild_state(|s| s.workspace_pool = Some(pool));
self
}
/// Get the first auth token (for printing to console on startup).
pub fn auth_token(&self) -> &str { pub fn auth_token(&self) -> &str {
self.auth.env_auth.first_token().unwrap_or("") &self.auth_token
} }
/// Get a reference to the shared gateway state (for the agent to push SSE events). /// Get a reference to the shared gateway state (for the agent to push SSE events).
@@ -327,7 +291,7 @@ impl Channel for GatewayChannel {
), ),
})?; })?;
server::start_server(addr, self.state.clone(), self.auth.clone()).await?; server::start_server(addr, self.state.clone(), self.auth_token.clone()).await?;
Ok(Box::pin(ReceiverStream::new(rx))) Ok(Box::pin(ReceiverStream::new(rx)))
} }
@@ -347,13 +311,10 @@ impl Channel for GatewayChannel {
} }
}; };
self.state.sse.broadcast_for_user( self.state.sse.broadcast(SseEvent::Response {
&msg.user_id, content: response.content,
AppEvent::Response { thread_id,
content: response.content, });
thread_id,
},
);
Ok(()) Ok(())
} }
@@ -368,11 +329,11 @@ impl Channel for GatewayChannel {
.and_then(|v| v.as_str()) .and_then(|v| v.as_str())
.map(String::from); .map(String::from);
let event = match status { let event = match status {
StatusUpdate::Thinking(msg) => AppEvent::Thinking { StatusUpdate::Thinking(msg) => SseEvent::Thinking {
message: msg, message: msg,
thread_id: thread_id.clone(), thread_id: thread_id.clone(),
}, },
StatusUpdate::ToolStarted { name } => AppEvent::ToolStarted { StatusUpdate::ToolStarted { name } => SseEvent::ToolStarted {
name, name,
thread_id: thread_id.clone(), thread_id: thread_id.clone(),
}, },
@@ -381,23 +342,23 @@ impl Channel for GatewayChannel {
success, success,
error, error,
parameters, parameters,
} => AppEvent::ToolCompleted { } => SseEvent::ToolCompleted {
name, name,
success, success,
error, error,
parameters, parameters,
thread_id: thread_id.clone(), thread_id: thread_id.clone(),
}, },
StatusUpdate::ToolResult { name, preview } => AppEvent::ToolResult { StatusUpdate::ToolResult { name, preview } => SseEvent::ToolResult {
name, name,
preview, preview,
thread_id: thread_id.clone(), thread_id: thread_id.clone(),
}, },
StatusUpdate::StreamChunk(content) => AppEvent::StreamChunk { StatusUpdate::StreamChunk(content) => SseEvent::StreamChunk {
content, content,
thread_id: thread_id.clone(), thread_id: thread_id.clone(),
}, },
StatusUpdate::Status(msg) => AppEvent::Status { StatusUpdate::Status(msg) => SseEvent::Status {
message: msg, message: msg,
thread_id: thread_id.clone(), thread_id: thread_id.clone(),
}, },
@@ -405,7 +366,7 @@ impl Channel for GatewayChannel {
job_id, job_id,
title, title,
browse_url, browse_url,
} => AppEvent::JobStarted { } => SseEvent::JobStarted {
job_id, job_id,
title, title,
browse_url, browse_url,
@@ -416,7 +377,7 @@ impl Channel for GatewayChannel {
description, description,
parameters, parameters,
allow_always, allow_always,
} => AppEvent::ApprovalNeeded { } => SseEvent::ApprovalNeeded {
request_id, request_id,
tool_name, tool_name,
description, description,
@@ -430,7 +391,7 @@ impl Channel for GatewayChannel {
instructions, instructions,
auth_url, auth_url,
setup_url, setup_url,
} => AppEvent::AuthRequired { } => SseEvent::AuthRequired {
extension_name, extension_name,
instructions, instructions,
auth_url, auth_url,
@@ -440,39 +401,25 @@ impl Channel for GatewayChannel {
extension_name, extension_name,
success, success,
message, message,
} => AppEvent::AuthCompleted { } => SseEvent::AuthCompleted {
extension_name, extension_name,
success, success,
message, message,
}, },
StatusUpdate::ImageGenerated { data_url, path } => AppEvent::ImageGenerated { StatusUpdate::ImageGenerated { data_url, path } => SseEvent::ImageGenerated {
data_url, data_url,
path, path,
thread_id: thread_id.clone(), thread_id: thread_id.clone(),
}, },
StatusUpdate::Suggestions { suggestions } => AppEvent::Suggestions { StatusUpdate::Suggestions { suggestions } => SseEvent::Suggestions {
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, thread_id,
}, },
StatusUpdate::TurnCost { StatusUpdate::TurnCost {
input_tokens, input_tokens,
output_tokens, output_tokens,
cost_usd, cost_usd,
} => AppEvent::TurnCost { } => SseEvent::TurnCost {
input_tokens, input_tokens,
output_tokens, output_tokens,
cost_usd, cost_usd,
@@ -480,21 +427,13 @@ impl Channel for GatewayChannel {
}, },
}; };
// Scope events to the user when user_id is available in metadata. self.state.sse.broadcast(event);
// When user_id is missing (heartbeat, routines), events go to all
// subscribers. In multi-tenant mode this leaks status across users.
if let Some(uid) = metadata.get("user_id").and_then(|v| v.as_str()) {
self.state.sse.broadcast_for_user(uid, event);
} else {
tracing::debug!("Status event missing user_id in metadata; broadcasting globally");
self.state.sse.broadcast(event);
}
Ok(()) Ok(())
} }
async fn broadcast( async fn broadcast(
&self, &self,
user_id: &str, _user_id: &str,
response: OutgoingResponse, response: OutgoingResponse,
) -> Result<(), ChannelError> { ) -> Result<(), ChannelError> {
let thread_id = match response.thread_id { let thread_id = match response.thread_id {
@@ -506,13 +445,10 @@ impl Channel for GatewayChannel {
return Ok(()); return Ok(());
} }
}; };
self.state.sse.broadcast_for_user( self.state.sse.broadcast(SseEvent::Response {
user_id, content: response.content,
AppEvent::Response { thread_id,
content: response.content, });
thread_id,
},
);
Ok(()) Ok(())
} }
+1 -4
View File
@@ -231,7 +231,6 @@ pub fn convert_messages(messages: &[OpenAiMessage]) -> Result<Vec<ChatMessage>,
name: tc.function.name.clone(), name: tc.function.name.clone(),
arguments: serde_json::from_str(&tc.function.arguments) arguments: serde_json::from_str(&tc.function.arguments)
.unwrap_or(serde_json::Value::Object(Default::default())), .unwrap_or(serde_json::Value::Object(Default::default())),
reasoning: None,
}) })
.collect(); .collect();
Ok(ChatMessage::assistant_with_tool_calls( Ok(ChatMessage::assistant_with_tool_calls(
@@ -464,10 +463,9 @@ fn build_tool_request(
pub async fn chat_completions_handler( pub async fn chat_completions_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
super::auth::AuthenticatedUser(user): super::auth::AuthenticatedUser,
Json(req): Json<OpenAiChatRequest>, Json(req): Json<OpenAiChatRequest>,
) -> Result<impl IntoResponse, (StatusCode, Json<OpenAiErrorResponse>)> { ) -> Result<impl IntoResponse, (StatusCode, Json<OpenAiErrorResponse>)> {
if !state.chat_rate_limiter.check(&user.user_id) { if !state.chat_rate_limiter.check() {
return Err(openai_error( return Err(openai_error(
StatusCode::TOO_MANY_REQUESTS, StatusCode::TOO_MANY_REQUESTS,
"Rate limit exceeded. Please try again later.", "Rate limit exceeded. Please try again later.",
@@ -955,7 +953,6 @@ mod tests {
id: "call_abc".to_string(), id: "call_abc".to_string(),
name: "search".to_string(), name: "search".to_string(),
arguments: serde_json::json!({"query": "rust"}), arguments: serde_json::json!({"query": "rust"}),
reasoning: None,
}]; }];
let converted = convert_tool_calls_to_openai(&calls); let converted = convert_tool_calls_to_openai(&calls);
File diff suppressed because it is too large Load Diff
+378 -613
View File
File diff suppressed because it is too large Load Diff
+61 -138
View File
@@ -11,31 +11,15 @@ use tokio::sync::broadcast;
use tokio_stream::StreamExt; use tokio_stream::StreamExt;
use tokio_stream::wrappers::BroadcastStream; use tokio_stream::wrappers::BroadcastStream;
use crate::channels::web::types::AppEvent; use crate::channels::web::types::SseEvent;
/// Maximum number of concurrent SSE/WebSocket connections. /// Maximum number of concurrent SSE/WebSocket connections.
/// Prevents resource exhaustion from connection flooding. /// Prevents resource exhaustion from connection flooding.
const MAX_CONNECTIONS: u64 = 100; const MAX_CONNECTIONS: u64 = 100;
/// Envelope for broadcast events: carries an optional user scope.
///
/// `user_id = None` means the event is global (e.g. Heartbeat) and delivered
/// to all subscribers. `user_id = Some(id)` means the event is only delivered
/// to subscribers that match that user_id.
#[derive(Debug, Clone)]
pub(crate) struct ScopedEvent {
pub(crate) user_id: Option<String>,
pub(crate) event: AppEvent,
}
/// Manages SSE broadcast to all connected browser tabs. /// Manages SSE broadcast to all connected browser tabs.
///
/// In multi-user mode, events are scoped by user_id so that each subscriber
/// only receives events intended for their user (plus global events like
/// Heartbeat). In single-user mode, all events are delivered to all subscribers
/// (backwards compatible).
pub struct SseManager { pub struct SseManager {
tx: broadcast::Sender<ScopedEvent>, tx: broadcast::Sender<SseEvent>,
connection_count: Arc<AtomicU64>, connection_count: Arc<AtomicU64>,
max_connections: u64, max_connections: u64,
} }
@@ -61,7 +45,7 @@ impl SseManager {
/// only be called before the server starts accepting connections (i.e., /// only be called before the server starts accepting connections (i.e.,
/// during startup wiring). Calling it after connections are established /// during startup wiring). Calling it after connections are established
/// will break connection tracking and allow exceeding `MAX_CONNECTIONS`. /// will break connection tracking and allow exceeding `MAX_CONNECTIONS`.
pub(crate) fn from_sender(tx: broadcast::Sender<ScopedEvent>) -> Self { pub fn from_sender(tx: broadcast::Sender<SseEvent>) -> Self {
Self { Self {
tx, tx,
connection_count: Arc::new(AtomicU64::new(0)), connection_count: Arc::new(AtomicU64::new(0)),
@@ -69,30 +53,17 @@ impl SseManager {
} }
} }
/// Broadcast an event to all connected clients.
pub fn broadcast(&self, event: SseEvent) {
// Ignore send errors (no receivers is fine)
let _ = self.tx.send(event);
}
/// Get a clone of the broadcast sender for use by other components. /// Get a clone of the broadcast sender for use by other components.
pub(crate) fn sender(&self) -> broadcast::Sender<ScopedEvent> { pub fn sender(&self) -> broadcast::Sender<SseEvent> {
self.tx.clone() self.tx.clone()
} }
/// Broadcast an event to all connected clients (global/unscoped).
pub fn broadcast(&self, event: AppEvent) {
let _ = self.tx.send(ScopedEvent {
user_id: None,
event,
});
}
/// Broadcast an event scoped to a specific user.
///
/// Only subscribers for this user_id (or unscoped subscribers) will
/// receive the event.
pub fn broadcast_for_user(&self, user_id: &str, event: AppEvent) {
let _ = self.tx.send(ScopedEvent {
user_id: Some(user_id.to_string()),
event,
});
}
/// Get current number of active connections. /// Get current number of active connections.
pub fn connection_count(&self) -> u64 { pub fn connection_count(&self) -> u64 {
self.connection_count.load(Ordering::Relaxed) self.connection_count.load(Ordering::Relaxed)
@@ -100,15 +71,11 @@ impl SseManager {
/// Create a raw broadcast subscription for non-SSE consumers (e.g. WebSocket). /// Create a raw broadcast subscription for non-SSE consumers (e.g. WebSocket).
/// ///
/// When `user_id` is `Some`, only events scoped to that user (or global /// Returns a stream of `SseEvent` values and increments/decrements the
/// events) are delivered. When `None`, all events are delivered (single-user /// connection counter on creation/drop, just like `subscribe()` does for SSE.
/// backwards compatibility).
/// ///
/// Returns `None` if the maximum connection limit has been reached. /// Returns `None` if the maximum connection limit has been reached.
pub fn subscribe_raw( pub fn subscribe_raw(&self) -> Option<impl Stream<Item = SseEvent> + Send + 'static + use<>> {
&self,
user_id: Option<String>,
) -> Option<impl Stream<Item = AppEvent> + Send + 'static + use<>> {
// Atomically increment only if below the limit. This prevents // Atomically increment only if below the limit. This prevents
// concurrent callers from overshooting max_connections. // concurrent callers from overshooting max_connections.
let counter = Arc::clone(&self.connection_count); let counter = Arc::clone(&self.connection_count);
@@ -124,19 +91,7 @@ impl SseManager {
.ok()?; .ok()?;
let rx = self.tx.subscribe(); let rx = self.tx.subscribe();
let stream = BroadcastStream::new(rx).filter_map(move |result| match result { let stream = BroadcastStream::new(rx).filter_map(|result| result.ok());
Ok(scoped) => {
// Global events (user_id=None) always pass through.
// Scoped events only pass if the subscriber matches (or subscriber is unscoped).
match (&user_id, &scoped.user_id) {
(_, None) => Some(scoped.event), // global -> all
(None, _) => Some(scoped.event), // unscoped subscriber -> all
(Some(sub), Some(ev)) if sub == ev => Some(scoped.event), // match
_ => None, // different user -> skip
}
}
Err(_) => None,
});
Some(CountedStream { Some(CountedStream {
inner: stream, inner: stream,
@@ -146,13 +101,9 @@ impl SseManager {
/// Create a new SSE stream for a client connection. /// Create a new SSE stream for a client connection.
/// ///
/// When `user_id` is `Some`, only events for that user (or global events)
/// are delivered. When `None`, all events are delivered.
///
/// Returns `None` if the maximum connection limit has been reached. /// Returns `None` if the maximum connection limit has been reached.
pub fn subscribe( pub fn subscribe(
&self, &self,
user_id: Option<String>,
) -> Option<Sse<impl Stream<Item = Result<Event, Infallible>> + Send + 'static + use<>>> { ) -> Option<Sse<impl Stream<Item = Result<Event, Infallible>> + Send + 'static + use<>>> {
// Atomically increment only if below the limit. // Atomically increment only if below the limit.
let counter = Arc::clone(&self.connection_count); let counter = Arc::clone(&self.connection_count);
@@ -169,25 +120,34 @@ impl SseManager {
let rx = self.tx.subscribe(); let rx = self.tx.subscribe();
let stream = BroadcastStream::new(rx) let stream = BroadcastStream::new(rx)
.filter_map(move |result| match result { .filter_map(|result| result.ok())
Ok(scoped) => match (&user_id, &scoped.user_id) { .map(|event| {
(_, None) => Some(scoped.event), let data = serde_json::to_string(&event).unwrap_or_default();
(None, _) => Some(scoped.event), let event_type = match &event {
(Some(sub), Some(ev)) if sub == ev => Some(scoped.event), SseEvent::Response { .. } => "response",
_ => None, SseEvent::Thinking { .. } => "thinking",
}, SseEvent::ToolStarted { .. } => "tool_started",
Err(_) => None, SseEvent::ToolCompleted { .. } => "tool_completed",
}) SseEvent::ToolResult { .. } => "tool_result",
.filter_map(|event| { SseEvent::StreamChunk { .. } => "stream_chunk",
let data = match serde_json::to_string(&event) { SseEvent::Status { .. } => "status",
Ok(s) => s, SseEvent::ApprovalNeeded { .. } => "approval_needed",
Err(e) => { SseEvent::AuthRequired { .. } => "auth_required",
tracing::warn!("Failed to serialize SSE event: {}", e); SseEvent::AuthCompleted { .. } => "auth_completed",
return None; 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(); Ok(Event::default().event(event_type).data(data))
Some(Ok(Event::default().event(event_type).data(data)))
}); });
// Wrap in a stream that decrements on drop // Wrap in a stream that decrements on drop
@@ -249,22 +209,24 @@ mod tests {
fn test_broadcast_without_receivers() { fn test_broadcast_without_receivers() {
let manager = SseManager::new(); let manager = SseManager::new();
// Should not panic even with no receivers // Should not panic even with no receivers
manager.broadcast(AppEvent::Heartbeat); manager.broadcast(SseEvent::Heartbeat);
} }
#[tokio::test] #[tokio::test]
async fn test_broadcast_to_receiver() { async fn test_broadcast_to_receiver() {
let manager = SseManager::new(); let manager = SseManager::new();
let mut stream = Box::pin(manager.subscribe_raw(None).expect("should subscribe")); let mut rx = BroadcastStream::new(manager.tx.subscribe());
manager.broadcast(AppEvent::Status { manager.broadcast(SseEvent::Status {
message: "test".to_string(), message: "test".to_string(),
thread_id: None, thread_id: None,
}); });
let event = stream.next().await.unwrap(); let event = rx.next().await;
assert!(event.is_some());
let event = event.unwrap().unwrap();
match event { match event {
AppEvent::Status { message, .. } => assert_eq!(message, "test"), SseEvent::Status { message, .. } => assert_eq!(message, "test"),
_ => panic!("unexpected event type"), _ => panic!("unexpected event type"),
} }
} }
@@ -272,18 +234,18 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_subscribe_raw_receives_events() { async fn test_subscribe_raw_receives_events() {
let manager = SseManager::new(); let manager = SseManager::new();
let mut stream = Box::pin(manager.subscribe_raw(None).expect("should subscribe")); let mut stream = Box::pin(manager.subscribe_raw().expect("should subscribe"));
assert_eq!(manager.connection_count(), 1); assert_eq!(manager.connection_count(), 1);
manager.broadcast(AppEvent::Thinking { manager.broadcast(SseEvent::Thinking {
message: "working".to_string(), message: "working".to_string(),
thread_id: None, thread_id: None,
}); });
let event = stream.next().await.unwrap(); let event = stream.next().await.unwrap();
match event { match event {
AppEvent::Thinking { message, .. } => assert_eq!(message, "working"), SseEvent::Thinking { message, .. } => assert_eq!(message, "working"),
_ => panic!("Expected Thinking event"), _ => panic!("Expected Thinking event"),
} }
} }
@@ -292,7 +254,7 @@ mod tests {
async fn test_subscribe_raw_decrements_on_drop() { async fn test_subscribe_raw_decrements_on_drop() {
let manager = SseManager::new(); let manager = SseManager::new();
{ {
let _stream = Box::pin(manager.subscribe_raw(None).expect("should subscribe")); let _stream = Box::pin(manager.subscribe_raw().expect("should subscribe"));
assert_eq!(manager.connection_count(), 1); assert_eq!(manager.connection_count(), 1);
} }
// Stream dropped, counter should decrement // Stream dropped, counter should decrement
@@ -302,16 +264,16 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_subscribe_raw_multiple_subscribers() { async fn test_subscribe_raw_multiple_subscribers() {
let manager = SseManager::new(); let manager = SseManager::new();
let mut s1 = Box::pin(manager.subscribe_raw(None).expect("should subscribe")); let mut s1 = Box::pin(manager.subscribe_raw().expect("should subscribe"));
let mut s2 = Box::pin(manager.subscribe_raw(None).expect("should subscribe")); let mut s2 = Box::pin(manager.subscribe_raw().expect("should subscribe"));
assert_eq!(manager.connection_count(), 2); assert_eq!(manager.connection_count(), 2);
manager.broadcast(AppEvent::Heartbeat); manager.broadcast(SseEvent::Heartbeat);
let e1 = s1.next().await.unwrap(); let e1 = s1.next().await.unwrap();
let e2 = s2.next().await.unwrap(); let e2 = s2.next().await.unwrap();
assert!(matches!(e1, AppEvent::Heartbeat)); assert!(matches!(e1, SseEvent::Heartbeat));
assert!(matches!(e2, AppEvent::Heartbeat)); assert!(matches!(e2, SseEvent::Heartbeat));
drop(s1); drop(s1);
assert_eq!(manager.connection_count(), 1); assert_eq!(manager.connection_count(), 1);
@@ -324,51 +286,12 @@ mod tests {
let mut manager = SseManager::new(); let mut manager = SseManager::new();
manager.max_connections = 2; // Low limit for testing manager.max_connections = 2; // Low limit for testing
let _s1 = Box::pin(manager.subscribe_raw(None).expect("first should succeed")); let _s1 = Box::pin(manager.subscribe_raw().expect("first should succeed"));
let _s2 = Box::pin(manager.subscribe_raw(None).expect("second should succeed")); let _s2 = Box::pin(manager.subscribe_raw().expect("second should succeed"));
assert_eq!(manager.connection_count(), 2); assert_eq!(manager.connection_count(), 2);
// Third should be rejected // Third should be rejected
assert!(manager.subscribe_raw(None).is_none()); assert!(manager.subscribe_raw().is_none());
assert!(manager.subscribe(None).is_none()); assert!(manager.subscribe().is_none());
}
#[tokio::test]
async fn test_scoped_events_filtered_by_user() {
let manager = SseManager::new();
let mut alice = Box::pin(
manager
.subscribe_raw(Some("alice".to_string()))
.expect("subscribe"),
);
let mut bob = Box::pin(
manager
.subscribe_raw(Some("bob".to_string()))
.expect("subscribe"),
);
// Send event scoped to alice
manager.broadcast_for_user(
"alice",
AppEvent::Status {
message: "alice only".to_string(),
thread_id: None,
},
);
// Send global event
manager.broadcast(AppEvent::Heartbeat);
// Alice gets her scoped event
let e = alice.next().await.unwrap();
assert!(matches!(e, AppEvent::Status { .. }));
// Alice also gets the global heartbeat
let e = alice.next().await.unwrap();
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, AppEvent::Heartbeat)); // safety: test assertion
} }
} }
+3 -147
View File
@@ -186,13 +186,6 @@ function authenticate() {
connectSSE(); connectSSE();
connectLogSSE(); connectLogSSE();
startGatewayStatusPolling(); startGatewayStatusPolling();
// Hide the Users settings tab for non-admin users.
apiFetch('/api/profile').then(function(profile) {
if (profile && profile.role !== 'admin') {
var usersTab = document.querySelector('[data-settings-subtab="users"]');
if (usersTab) usersTab.style.display = 'none';
}
}).catch(function() {});
checkTeeStatus(); checkTeeStatus();
loadThreads(); loadThreads();
loadMemoryTree(); loadMemoryTree();
@@ -4272,9 +4265,9 @@ function renderRoutineDetail(routine) {
+ '<th>Trigger</th><th>Started</th><th>Completed</th><th>Status</th><th>Summary</th><th>Tokens</th>' + '<th>Trigger</th><th>Started</th><th>Completed</th><th>Status</th><th>Summary</th><th>Tokens</th>'
+ '</tr></thead><tbody>'; + '</tr></thead><tbody>';
for (const run of routine.recent_runs) { for (const run of routine.recent_runs) {
const runStatusClass = run.status === 'ok' ? 'completed' const runStatusClass = run.status === 'Ok' ? 'completed'
: run.status === 'failed' ? 'failed' : run.status === 'Failed' ? 'failed'
: run.status === 'attention' ? 'stuck' : run.status === 'Attention' ? 'stuck'
: 'in_progress'; : 'in_progress';
html += '<tr>' html += '<tr>'
+ '<td>' + escapeHtml(run.trigger_type) + '</td>' + '<td>' + escapeHtml(run.trigger_type) + '</td>'
@@ -4346,142 +4339,6 @@ function formatRelativeTime(isoString) {
return future ? I18n.t('time.daysFromNow', { n: days }) : I18n.t('time.daysAgo', { n: days }); return future ? I18n.t('time.daysFromNow', { n: days }) : I18n.t('time.daysAgo', { n: days });
} }
// --- Users (admin) ---
function loadUsers() {
apiFetch('/api/admin/users').then(function(data) {
renderUsersList(data.users || []);
}).catch(function(err) {
// Non-admin users get 403 — show a message instead of an error
var tbody = document.getElementById('users-tbody');
var empty = document.getElementById('users-empty');
if (tbody) tbody.innerHTML = '';
if (empty) {
empty.style.display = 'block';
empty.textContent = 'Admin access required to manage users.';
}
});
}
function renderUsersList(users) {
var tbody = document.getElementById('users-tbody');
var empty = document.getElementById('users-empty');
if (!users || users.length === 0) {
tbody.innerHTML = '';
empty.style.display = 'block';
empty.textContent = 'No users found. Create the first user to get started.';
return;
}
empty.style.display = 'none';
tbody.innerHTML = users.map(function(u) {
var statusClass = u.status === 'active' ? 'active' : 'failed';
var roleLabel = u.role === 'admin' ? '<span class="badge badge-admin">admin</span>' : '<span class="badge">member</span>';
var actions = '';
if (u.status === 'active') {
actions += '<button class="btn-small btn-danger" data-action="suspend-user" data-user-id="' + escapeHtml(u.id) + '">Suspend</button> ';
} else {
actions += '<button class="btn-small btn-primary" data-action="activate-user" data-user-id="' + escapeHtml(u.id) + '">Activate</button> ';
}
actions += '<button class="btn-small" data-action="create-token" data-user-id="' + escapeHtml(u.id) + '" data-user-name="' + escapeHtml(u.display_name) + '">+ Token</button>';
return '<tr>'
+ '<td class="user-id" title="' + escapeHtml(u.id) + '">' + escapeHtml(u.id.substring(0, 8)) + '…</td>'
+ '<td>' + escapeHtml(u.display_name) + '</td>'
+ '<td>' + escapeHtml(u.email || '—') + '</td>'
+ '<td>' + roleLabel + '</td>'
+ '<td><span class="status-badge ' + statusClass + '">' + escapeHtml(u.status) + '</span></td>'
+ '<td>' + formatRelativeTime(u.created_at) + '</td>'
+ '<td>' + actions + '</td>'
+ '</tr>';
}).join('');
}
function suspendUser(userId) {
apiFetch('/api/admin/users/' + userId + '/suspend', { method: 'POST' })
.then(function() { loadUsers(); })
.catch(function(e) { alert('Failed to suspend user: ' + e.message); });
}
function activateUser(userId) {
apiFetch('/api/admin/users/' + userId + '/activate', { method: 'POST' })
.then(function() { loadUsers(); })
.catch(function(e) { alert('Failed to activate user: ' + e.message); });
}
function createTokenForUser(userId, displayName) {
var tokenName = prompt('Token name for ' + displayName + ':', 'api-token');
if (!tokenName) return;
apiFetch('/api/tokens', {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify({ name: tokenName, user_id: userId }),
}).then(function(data) {
showTokenBanner(data.token);
}).catch(function(e) { alert('Failed to create token: ' + e.message); });
}
function showTokenBanner(tokenValue) {
var banner = document.getElementById('users-token-result');
if (!banner) return;
var loginUrl = window.location.origin + '/?token=' + encodeURIComponent(tokenValue);
banner.style.display = 'block';
banner.innerHTML = '<strong>User created!</strong> Share this login link — it won\'t be shown again:<br>'
+ '<code class="token-display" id="token-copy-value">' + escapeHtml(loginUrl) + '</code>'
+ '<button class="btn-small" id="token-copy-link">Copy Link</button>'
+ '<br><span style="font-size:0.8em;color:var(--text-muted)">Raw token: ' + escapeHtml(tokenValue) + '</span>';
document.getElementById('token-copy-link').addEventListener('click', function() {
navigator.clipboard.writeText(loginUrl);
this.textContent = 'Copied!';
});
}
// Delegated click handler for user action buttons (CSP-safe, no inline onclick)
document.getElementById('users-table')?.addEventListener('click', function(e) {
var btn = e.target.closest('[data-action]');
if (!btn) return;
var action = btn.getAttribute('data-action');
var userId = btn.getAttribute('data-user-id');
var userName = btn.getAttribute('data-user-name');
if (action === 'suspend-user') suspendUser(userId);
else if (action === 'activate-user') activateUser(userId);
else if (action === 'create-token') createTokenForUser(userId, userName || '');
});
// Wire up Users tab create form
document.getElementById('users-create-btn')?.addEventListener('click', function() {
document.getElementById('users-create-form').style.display = 'flex';
document.getElementById('users-token-result').style.display = 'none';
document.getElementById('user-display-name').focus();
});
document.getElementById('users-create-cancel')?.addEventListener('click', function() {
document.getElementById('users-create-form').style.display = 'none';
});
document.getElementById('users-create-submit')?.addEventListener('click', function() {
var displayName = document.getElementById('user-display-name').value.trim();
var email = document.getElementById('user-email').value.trim();
var role = document.getElementById('user-role').value;
if (!displayName) { alert('Display name is required'); return; }
apiFetch('/api/admin/users', {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify({
display_name: displayName,
email: email || undefined,
role: role,
}),
}).then(function(data) {
document.getElementById('users-create-form').style.display = 'none';
document.getElementById('user-display-name').value = '';
document.getElementById('user-email').value = '';
if (data.token) {
showTokenBanner(data.token);
}
loadUsers();
}).catch(function(e) { alert('Failed to create user: ' + e.message); });
});
// --- Gateway status widget --- // --- Gateway status widget ---
let gatewayStatusInterval = null; let gatewayStatusInterval = null;
@@ -5171,7 +5028,6 @@ function loadSettingsSubtab(subtab) {
else if (subtab === 'extensions') { loadExtensions(); startPairingPoll(); } else if (subtab === 'extensions') { loadExtensions(); startPairingPoll(); }
else if (subtab === 'mcp') loadMcpServers(); else if (subtab === 'mcp') loadMcpServers();
else if (subtab === 'skills') loadSkills(); else if (subtab === 'skills') loadSkills();
else if (subtab === 'users') loadUsers();
if (subtab !== 'extensions' && subtab !== 'channels') stopPairingPoll(); if (subtab !== 'extensions' && subtab !== 'channels') stopPairingPoll();
} }
+3 -2
View File
@@ -9,6 +9,8 @@ I18n.register('en', {
'auth.connect': 'Connect', 'auth.connect': 'Connect',
'auth.errorRequired': 'Token required', 'auth.errorRequired': 'Token required',
'auth.errorInvalid': 'Invalid token', 'auth.errorInvalid': 'Invalid token',
'auth.hint': 'Enter the GATEWAY_AUTH_TOKEN from your .env file',
// Chat // Chat
'chat.inputPlaceholder': 'Message or / for commands...', 'chat.inputPlaceholder': 'Message or / for commands...',
@@ -42,8 +44,7 @@ I18n.register('en', {
'settings.channels': 'Channels', 'settings.channels': 'Channels',
'settings.networking': 'Networking', 'settings.networking': 'Networking',
'settings.mcp': 'MCP', 'settings.mcp': 'MCP',
'settings.users': 'Users',
// Status // Status
'status.connected': 'Connected', 'status.connected': 'Connected',
'status.disconnected': 'Disconnected', 'status.disconnected': 'Disconnected',
+3 -2
View File
@@ -9,6 +9,8 @@ I18n.register('zh-CN', {
'auth.connect': '连接', 'auth.connect': '连接',
'auth.errorRequired': '请输入令牌', 'auth.errorRequired': '请输入令牌',
'auth.errorInvalid': '令牌无效', 'auth.errorInvalid': '令牌无效',
'auth.hint': '输入 .env 配置文件中的 GATEWAY_AUTH_TOKEN',
// 聊天 // 聊天
'chat.inputPlaceholder': '输入消息或 / 以使用命令...', 'chat.inputPlaceholder': '输入消息或 / 以使用命令...',
@@ -42,8 +44,7 @@ I18n.register('zh-CN', {
'settings.channels': '频道', 'settings.channels': '频道',
'settings.networking': '网络', 'settings.networking': '网络',
'settings.mcp': 'MCP', 'settings.mcp': 'MCP',
'settings.users': '用户管理',
// 状态 // 状态
'status.connected': '已连接', 'status.connected': '已连接',
'status.disconnected': '已断开', 'status.disconnected': '已断开',
+1 -24
View File
@@ -41,6 +41,7 @@
<button id="auth-connect-btn" data-i18n="auth.connect">Connect</button> <button id="auth-connect-btn" data-i18n="auth.connect">Connect</button>
</div> </div>
<div id="auth-error"></div> <div id="auth-error"></div>
<p class="auth-hint" data-i18n="auth.hint">Enter the GATEWAY_AUTH_TOKEN from your .env configuration.</p>
</div> </div>
</div> </div>
@@ -292,7 +293,6 @@
<button class="settings-subtab" data-settings-subtab="extensions" data-i18n="tab.extensions">Extensions</button> <button class="settings-subtab" data-settings-subtab="extensions" data-i18n="tab.extensions">Extensions</button>
<button class="settings-subtab" data-settings-subtab="mcp" data-i18n="settings.mcp">MCP</button> <button class="settings-subtab" data-settings-subtab="mcp" data-i18n="settings.mcp">MCP</button>
<button class="settings-subtab" data-settings-subtab="skills" data-i18n="tab.skills">Skills</button> <button class="settings-subtab" data-settings-subtab="skills" data-i18n="tab.skills">Skills</button>
<button class="settings-subtab" data-settings-subtab="users" data-i18n="settings.users">Users</button>
<button class="settings-theme-toggle" id="settings-theme-toggle" data-i18n="theme.tooltipSystem" title="Toggle theme">Theme</button> <button class="settings-theme-toggle" id="settings-theme-toggle" data-i18n="theme.tooltipSystem" title="Toggle theme">Theme</button>
</div> </div>
<div class="settings-content"> <div class="settings-content">
@@ -390,29 +390,6 @@
</div> </div>
</div> </div>
</div> </div>
<div class="settings-subpanel" id="settings-users">
<div class="users-container">
<div class="users-header">
<h3>User Management</h3>
<button id="users-create-btn" class="btn-primary">+ New User</button>
</div>
<div id="users-create-form" style="display:none" class="users-form" autocomplete="off">
<input type="text" id="user-display-name" placeholder="Display name" autocomplete="off" />
<input type="text" id="user-email" placeholder="Email (optional)" autocomplete="off" />
<select id="user-role"><option value="member">Member</option><option value="admin">Admin</option></select>
<button id="users-create-submit" class="btn-primary">Create</button>
<button id="users-create-cancel" class="btn-secondary">Cancel</button>
</div>
<div id="users-token-result" style="display:none" class="users-token-banner"></div>
<table class="routines-table" id="users-table">
<thead><tr>
<th>ID</th><th>Display Name</th><th>Email</th><th>Role</th><th>Status</th><th>Created</th><th>Actions</th>
</tr></thead>
<tbody id="users-tbody"></tbody>
</table>
<div id="users-empty" class="empty-state" style="display:none">No users found. Create the first user to get started.</div>
</div>
</div>
</div> </div>
</div> </div>
</div> </div>
-19
View File
@@ -5429,22 +5429,3 @@ body.theme-transition *:not(svg):not(path):not(line):not(circle):not(rect) {
--text-muted: #a1a1aa; --text-muted: #a1a1aa;
} }
} }
/* --- Users Tab --- */
.users-container { padding: 1rem; }
.users-header { display: flex; align-items: center; justify-content: space-between; margin-bottom: 1rem; }
.users-header h3 { margin: 0; font-size: 1.1rem; }
.users-form { display: flex; gap: 0.5rem; align-items: center; margin-bottom: 1rem; flex-wrap: wrap; }
.users-form input, .users-form select { padding: 0.4rem 0.6rem; border-radius: 6px; border: 1px solid var(--border); background: var(--bg-secondary); color: var(--text-primary); font-size: 0.85rem; }
.users-token-banner { background: var(--bg-tertiary); border: 1px solid var(--accent); border-radius: 8px; padding: 0.75rem 1rem; margin-bottom: 1rem; font-size: 0.85rem; }
.token-display { display: inline-block; padding: 0.3rem 0.6rem; background: var(--bg-primary); border-radius: 4px; font-family: var(--font-mono); word-break: break-all; margin: 0.4rem 0; user-select: all; }
.user-id { font-family: var(--font-mono); font-size: 0.8rem; color: var(--text-muted); }
.badge { display: inline-block; padding: 0.15rem 0.5rem; border-radius: 10px; font-size: 0.75rem; background: var(--bg-tertiary); color: var(--text-secondary); }
.badge-admin { background: var(--accent); color: #fff; }
.btn-small { padding: 0.25rem 0.5rem; font-size: 0.75rem; border-radius: 4px; border: 1px solid var(--border); background: var(--bg-secondary); color: var(--text-primary); cursor: pointer; }
.btn-small:hover { background: var(--bg-tertiary); }
.btn-danger { border-color: #ef4444; color: #ef4444; }
.btn-danger:hover { background: #ef4444; color: #fff; }
.btn-primary { background: var(--accent); color: #fff; border: none; padding: 0.4rem 0.8rem; border-radius: 6px; cursor: pointer; font-size: 0.85rem; }
.btn-primary:hover { opacity: 0.9; }
.btn-secondary { background: var(--bg-tertiary); color: var(--text-primary); border: 1px solid var(--border); padding: 0.4rem 0.8rem; border-radius: 6px; cursor: pointer; font-size: 0.85rem; }
+6 -24
View File
@@ -10,8 +10,7 @@ use std::sync::Arc;
use tokio::sync::mpsc; use tokio::sync::mpsc;
use crate::channels::IncomingMessage; use crate::channels::IncomingMessage;
use crate::channels::web::auth::MultiAuthState; use crate::channels::web::server::{GatewayState, RateLimiter, start_server};
use crate::channels::web::server::{GatewayState, PerUserRateLimiter, RateLimiter, start_server};
use crate::channels::web::sse::SseManager; use crate::channels::web::sse::SseManager;
use crate::channels::web::ws::WsConnectionTracker; use crate::channels::web::ws::WsConnectionTracker;
@@ -65,9 +64,8 @@ impl TestGatewayBuilder {
pub fn build(self) -> Arc<GatewayState> { pub fn build(self) -> Arc<GatewayState> {
Arc::new(GatewayState { Arc::new(GatewayState {
msg_tx: tokio::sync::RwLock::new(self.msg_tx), msg_tx: tokio::sync::RwLock::new(self.msg_tx),
sse: Arc::new(SseManager::new()), sse: SseManager::new(),
workspace: None, workspace: None,
workspace_pool: None,
session_manager: None, session_manager: None,
log_broadcaster: None, log_broadcaster: None,
log_level_handle: None, log_level_handle: None,
@@ -76,14 +74,14 @@ impl TestGatewayBuilder {
store: None, store: None,
job_manager: None, job_manager: None,
prompt_queue: None, prompt_queue: None,
owner_id: self.user_id.clone(), user_id: self.user_id,
shutdown_tx: tokio::sync::RwLock::new(None), shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())), ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: self.llm_provider, llm_provider: self.llm_provider,
skill_registry: None, skill_registry: None,
skill_catalog: None, skill_catalog: None,
scheduler: None, scheduler: None,
chat_rate_limiter: PerUserRateLimiter::new(30, 60), chat_rate_limiter: RateLimiter::new(30, 60),
oauth_rate_limiter: RateLimiter::new(10, 60), oauth_rate_limiter: RateLimiter::new(10, 60),
webhook_rate_limiter: RateLimiter::new(10, 60), webhook_rate_limiter: RateLimiter::new(10, 60),
registry_entries: Vec::new(), registry_entries: Vec::new(),
@@ -91,7 +89,6 @@ impl TestGatewayBuilder {
routine_engine: Arc::new(tokio::sync::RwLock::new(None)), routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
startup_time: std::time::Instant::now(), startup_time: std::time::Instant::now(),
active_config: crate::channels::web::server::ActiveConfigSnapshot::default(), active_config: crate::channels::web::server::ActiveConfigSnapshot::default(),
secrets_store: None,
}) })
} }
@@ -101,26 +98,11 @@ impl TestGatewayBuilder {
self, self,
auth_token: &str, auth_token: &str,
) -> Result<(SocketAddr, Arc<GatewayState>), crate::error::ChannelError> { ) -> Result<(SocketAddr, Arc<GatewayState>), crate::error::ChannelError> {
let auth = MultiAuthState::single(auth_token.to_string(), "test-user".to_string());
let state = self.build(); let state = self.build();
let addr: SocketAddr = "127.0.0.1:0" let addr: SocketAddr = "127.0.0.1:0"
.parse() .parse()
.expect("hard-coded address must parse"); // safety: constant literal .expect("hard-coded address must parse");
let bound = start_server(addr, state.clone(), auth.into()).await?; let bound = start_server(addr, state.clone(), auth_token.to_string()).await?;
Ok((bound, state))
}
/// Build the state and start a gateway server with multi-user auth.
/// Returns the bound address and the shared state.
pub async fn start_multi(
self,
auth: MultiAuthState,
) -> Result<(SocketAddr, Arc<GatewayState>), crate::error::ChannelError> {
let state = self.build();
let addr: SocketAddr = "127.0.0.1:0"
.parse()
.expect("hard-coded address must parse"); // safety: constant literal
let bound = start_server(addr, state.clone(), auth.into()).await?;
Ok((bound, state)) Ok((bound, state))
} }
} }
-3
View File
@@ -1,3 +0,0 @@
//! Integration tests for the web gateway module.
mod multi_tenant;
-951
View File
@@ -1,951 +0,0 @@
//! Multi-tenant isolation tests for the web gateway.
//!
//! Tests cover workspace pool scoping, job handler isolation, and auth
//! enforcement on protected endpoints. Uses `LibSqlBackend::new_local()`
//! with a temporary directory for a real (but ephemeral) database.
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use axum::Router;
use axum::body::Body;
use axum::http::{Method, Request, StatusCode};
use axum::middleware;
use axum::routing::{delete, get, post};
use tower::ServiceExt;
use uuid::Uuid;
use crate::channels::web::auth::{
AuthenticatedUser, MultiAuthState, UserIdentity, auth_middleware,
};
use crate::channels::web::server::{
ActiveConfigSnapshot, GatewayState, PerUserRateLimiter, PromptQueue, RateLimiter, WorkspacePool,
};
use crate::channels::web::sse::SseManager;
// ── Helpers ────────────────────────────────────────────────────────────
/// Create a two-user `MultiAuthState` for alice and bob.
fn two_user_auth() -> MultiAuthState {
let mut tokens = HashMap::new();
tokens.insert(
"tok-alice".to_string(),
UserIdentity {
user_id: "alice".to_string(),
role: "admin".to_string(),
workspace_read_scopes: vec!["shared".to_string()],
},
);
tokens.insert(
"tok-bob".to_string(),
UserIdentity {
user_id: "bob".to_string(),
role: "admin".to_string(),
workspace_read_scopes: vec!["shared".to_string(), "alice".to_string()],
},
);
MultiAuthState::multi(tokens)
}
/// Build a `GatewayState` with configurable store and prompt queue.
fn build_state(
store: Option<Arc<dyn crate::db::Database>>,
prompt_queue: Option<PromptQueue>,
) -> Arc<GatewayState> {
Arc::new(GatewayState {
msg_tx: tokio::sync::RwLock::new(None),
sse: Arc::new(SseManager::new()),
workspace: None,
workspace_pool: None,
session_manager: None,
log_broadcaster: None,
log_level_handle: None,
extension_manager: None,
tool_registry: None,
store,
job_manager: None,
prompt_queue,
owner_id: "test".to_string(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: None,
llm_provider: None,
skill_registry: None,
skill_catalog: None,
scheduler: None,
chat_rate_limiter: PerUserRateLimiter::new(30, 60),
oauth_rate_limiter: RateLimiter::new(10, 60),
webhook_rate_limiter: RateLimiter::new(10, 60),
registry_entries: Vec::new(),
cost_guard: None,
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
startup_time: std::time::Instant::now(),
active_config: ActiveConfigSnapshot::default(),
secrets_store: None,
})
}
/// Create a libSQL-backed test database in a temporary directory.
///
/// Returns the database and a `TempDir` guard — the database file is
/// deleted when the guard is dropped.
#[cfg(feature = "libsql")]
async fn test_db() -> (Arc<dyn crate::db::Database>, tempfile::TempDir) {
use crate::db::Database;
let dir = tempfile::tempdir().expect("failed to create temp dir"); // safety: test-only
let path = dir.path().join("test.db");
let backend = crate::db::libsql::LibSqlBackend::new_local(&path)
.await
.expect("failed to create test LibSqlBackend"); // safety: test-only
backend
.run_migrations()
.await
.expect("failed to run migrations"); // safety: test-only
(Arc::new(backend) as Arc<dyn crate::db::Database>, dir)
}
/// Build a minimal Routine for testing.
fn make_routine(user_id: &str, name: &str) -> crate::agent::routine::Routine {
let now = chrono::Utc::now();
crate::agent::routine::Routine {
id: Uuid::new_v4(),
name: name.to_string(),
description: format!("Test routine: {name}"),
user_id: user_id.to_string(),
enabled: true,
trigger: crate::agent::routine::Trigger::Cron {
schedule: "0 9 * * *".to_string(),
timezone: None,
},
action: crate::agent::routine::RoutineAction::Lightweight {
prompt: "hello".to_string(),
context_paths: vec![],
max_tokens: 1024,
use_tools: false,
max_tool_rounds: 3,
},
guardrails: crate::agent::routine::RoutineGuardrails {
cooldown: Duration::from_secs(60),
max_concurrent: 1,
dedup_window: None,
},
notify: crate::agent::routine::NotifyConfig {
channel: None,
user: None,
on_success: false,
on_failure: true,
on_attention: true,
},
last_run_at: None,
next_fire_at: None,
run_count: 0,
consecutive_failures: 0,
state: serde_json::json!({}),
created_at: now,
updated_at: now,
}
}
/// Build a minimal SandboxJobRecord for testing.
fn make_sandbox_job(user_id: &str, task: &str) -> crate::history::SandboxJobRecord {
let now = chrono::Utc::now();
crate::history::SandboxJobRecord {
id: Uuid::new_v4(),
task: task.to_string(),
status: "completed".to_string(),
user_id: user_id.to_string(),
project_dir: format!("/tmp/test-{}", Uuid::new_v4()),
success: Some(true),
failure_reason: None,
created_at: now,
started_at: Some(now),
completed_at: Some(now),
credential_grants_json: "[]".to_string(),
}
}
// ═══════════════════════════════════════════════════════════════════════
// WorkspacePool Tests
// ═══════════════════════════════════════════════════════════════════════
#[cfg(feature = "libsql")]
mod workspace_pool {
use super::*;
use crate::config::{WorkspaceConfig, WorkspaceSearchConfig};
use crate::workspace::EmbeddingCacheConfig;
use crate::workspace::layer::MemoryLayer;
#[tokio::test]
async fn test_workspace_pool_applies_search_config() {
let (db, _dir) = test_db().await;
let search_config = WorkspaceSearchConfig {
rrf_k: 42,
..Default::default()
};
let pool = WorkspacePool::new(
db,
None,
EmbeddingCacheConfig::default(),
search_config,
WorkspaceConfig::default(),
);
let identity = UserIdentity {
user_id: "alice".to_string(),
role: "admin".to_string(),
workspace_read_scopes: vec![],
};
let ws = pool.get_or_create(&identity).await;
assert_eq!(ws.user_id(), "alice");
}
#[tokio::test]
async fn test_workspace_pool_applies_memory_layers() {
let (db, _dir) = test_db().await;
let layers = vec![MemoryLayer {
name: "shared-layer".to_string(),
scope: "shared".to_string(),
writable: false,
sensitivity: Default::default(),
}];
let ws_config = WorkspaceConfig {
memory_layers: layers,
read_scopes: vec![],
};
let pool = WorkspacePool::new(
db,
None,
EmbeddingCacheConfig::default(),
WorkspaceSearchConfig::default(),
ws_config,
);
let identity = UserIdentity {
user_id: "alice".to_string(),
role: "admin".to_string(),
workspace_read_scopes: vec![],
};
let ws = pool.get_or_create(&identity).await;
// Memory layer scope "shared" should appear in read_user_ids.
assert!(
ws.read_user_ids().contains(&"shared".to_string()),
"expected 'shared' in read_user_ids, got {:?}",
ws.read_user_ids()
);
}
#[tokio::test]
async fn test_workspace_pool_applies_identity_read_scopes() {
let (db, _dir) = test_db().await;
let pool = WorkspacePool::new(
db,
None,
EmbeddingCacheConfig::default(),
WorkspaceSearchConfig::default(),
WorkspaceConfig::default(),
);
let identity = UserIdentity {
user_id: "bob".to_string(),
role: "admin".to_string(),
workspace_read_scopes: vec!["alice".to_string(), "shared".to_string()],
};
let ws = pool.get_or_create(&identity).await;
assert_eq!(ws.user_id(), "bob");
assert!(
ws.read_user_ids().contains(&"alice".to_string()),
"expected 'alice' in read_user_ids from identity scopes"
);
assert!(
ws.read_user_ids().contains(&"shared".to_string()),
"expected 'shared' in read_user_ids from identity scopes"
);
}
#[tokio::test]
async fn test_workspace_pool_caches_per_user() {
let (db, _dir) = test_db().await;
let pool = WorkspacePool::new(
db,
None,
EmbeddingCacheConfig::default(),
WorkspaceSearchConfig::default(),
WorkspaceConfig::default(),
);
let alice_id = UserIdentity {
user_id: "alice".to_string(),
role: "admin".to_string(),
workspace_read_scopes: vec![],
};
let bob_id = UserIdentity {
user_id: "bob".to_string(),
role: "admin".to_string(),
workspace_read_scopes: vec![],
};
let alice_ws1 = pool.get_or_create(&alice_id).await;
let alice_ws2 = pool.get_or_create(&alice_id).await;
let bob_ws = pool.get_or_create(&bob_id).await;
// Same user gets the same Arc.
assert!(Arc::ptr_eq(&alice_ws1, &alice_ws2));
// Different users get different instances.
assert!(!Arc::ptr_eq(&alice_ws1, &bob_ws));
assert_eq!(alice_ws1.user_id(), "alice");
assert_eq!(bob_ws.user_id(), "bob");
}
#[tokio::test]
async fn test_workspace_pool_combines_global_and_identity_scopes() {
let (db, _dir) = test_db().await;
let ws_config = WorkspaceConfig {
memory_layers: vec![],
read_scopes: vec!["global-shared".to_string()],
};
let pool = WorkspacePool::new(
db,
None,
EmbeddingCacheConfig::default(),
WorkspaceSearchConfig::default(),
ws_config,
);
let identity = UserIdentity {
user_id: "alice".to_string(),
role: "admin".to_string(),
workspace_read_scopes: vec!["token-scope".to_string()],
};
let ws = pool.get_or_create(&identity).await;
let scopes = ws.read_user_ids();
// Primary scope
assert!(scopes.contains(&"alice".to_string()));
// Global config scope
assert!(
scopes.contains(&"global-shared".to_string()),
"expected global scope 'global-shared', got {:?}",
scopes
);
// Token identity scope
assert!(
scopes.contains(&"token-scope".to_string()),
"expected token scope 'token-scope', got {:?}",
scopes
);
}
}
// ═══════════════════════════════════════════════════════════════════════
// Jobs Handler Isolation Tests
// ═══════════════════════════════════════════════════════════════════════
#[cfg(feature = "libsql")]
mod jobs_isolation {
use super::*;
use crate::channels::web::handlers::jobs::{
jobs_cancel_handler, jobs_prompt_handler, jobs_restart_handler, jobs_summary_handler,
};
// SandboxStore methods are accessed through the Database supertrait.
/// Build a router with job endpoints behind multi-user auth.
fn jobs_router(state: Arc<GatewayState>, auth: MultiAuthState) -> Router {
Router::new()
.route("/api/jobs/summary", get(jobs_summary_handler))
.route("/api/jobs/{id}/cancel", post(jobs_cancel_handler))
.route("/api/jobs/{id}/restart", post(jobs_restart_handler))
.route("/api/jobs/{id}/prompt", post(jobs_prompt_handler))
.layer(middleware::from_fn_with_state(
crate::channels::web::auth::CombinedAuthState::from(auth),
auth_middleware,
))
.with_state(state)
}
#[tokio::test]
async fn test_jobs_summary_scoped_to_user() {
let (db, _dir) = test_db().await;
// Insert sandbox jobs for alice and bob.
let alice_job = make_sandbox_job("alice", "alice task");
let bob_job = make_sandbox_job("bob", "bob task");
db.save_sandbox_job(&alice_job).await.unwrap();
db.save_sandbox_job(&bob_job).await.unwrap();
let state = build_state(Some(db), None);
let auth = two_user_auth();
let app = jobs_router(state, auth);
// Alice should see 1 job.
let req = Request::builder()
.uri("/api/jobs/summary")
.header("Authorization", "Bearer tok-alice")
.body(Body::empty())
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body: serde_json::Value =
serde_json::from_slice(&axum::body::to_bytes(resp.into_body(), 4096).await.unwrap())
.unwrap();
assert_eq!(body["total"], 1, "alice should see only her own jobs");
// Bob should see 1 job.
let req = Request::builder()
.uri("/api/jobs/summary")
.header("Authorization", "Bearer tok-bob")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body: serde_json::Value =
serde_json::from_slice(&axum::body::to_bytes(resp.into_body(), 4096).await.unwrap())
.unwrap();
assert_eq!(body["total"], 1, "bob should see only his own jobs");
}
#[tokio::test]
async fn test_jobs_restart_rejects_other_user() {
let (db, _dir) = test_db().await;
// Insert a failed sandbox job owned by alice.
let mut alice_job = make_sandbox_job("alice", "alice task");
alice_job.status = "failed".to_string();
alice_job.success = Some(false);
db.save_sandbox_job(&alice_job).await.unwrap();
let state = build_state(Some(db), None);
let auth = two_user_auth();
let app = jobs_router(state, auth);
// Bob tries to restart alice's job.
let req = Request::builder()
.method(Method::POST)
.uri(format!("/api/jobs/{}/restart", alice_job.id))
.header("Authorization", "Bearer tok-bob")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::NOT_FOUND,
"bob should not be able to restart alice's job"
);
}
#[tokio::test]
async fn test_jobs_prompt_works_for_agent_jobs() {
let (db, _dir) = test_db().await;
// Insert a running sandbox job owned by alice in claude_code mode.
let mut alice_job = make_sandbox_job("alice", "prompt test");
alice_job.status = "running".to_string();
alice_job.success = None;
alice_job.completed_at = None;
db.save_sandbox_job(&alice_job).await.unwrap();
db.update_sandbox_job_mode(alice_job.id, "claude_code")
.await
.unwrap();
let prompt_queue: PromptQueue =
Arc::new(tokio::sync::Mutex::new(std::collections::HashMap::new()));
let state = build_state(Some(db), Some(prompt_queue.clone()));
let auth = two_user_auth();
let app = jobs_router(state, auth);
// Alice prompts her own job.
let req = Request::builder()
.method(Method::POST)
.uri(format!("/api/jobs/{}/prompt", alice_job.id))
.header("Authorization", "Bearer tok-alice")
.header("Content-Type", "application/json")
.body(Body::from(
serde_json::to_string(&serde_json::json!({"content": "hello"})).unwrap(),
))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::OK,
"alice should be able to prompt her own job"
);
// Verify prompt was enqueued.
let queue = prompt_queue.lock().await;
assert!(
queue.contains_key(&alice_job.id),
"prompt queue should contain alice's job"
);
}
#[tokio::test]
async fn test_jobs_prompt_rejects_other_user() {
let (db, _dir) = test_db().await;
let mut alice_job = make_sandbox_job("alice", "alice task");
alice_job.status = "running".to_string();
alice_job.success = None;
alice_job.completed_at = None;
db.save_sandbox_job(&alice_job).await.unwrap();
db.update_sandbox_job_mode(alice_job.id, "claude_code")
.await
.unwrap();
let prompt_queue: PromptQueue =
Arc::new(tokio::sync::Mutex::new(std::collections::HashMap::new()));
let state = build_state(Some(db), Some(prompt_queue));
let auth = two_user_auth();
let app = jobs_router(state, auth);
// Bob tries to prompt alice's job.
let req = Request::builder()
.method(Method::POST)
.uri(format!("/api/jobs/{}/prompt", alice_job.id))
.header("Authorization", "Bearer tok-bob")
.header("Content-Type", "application/json")
.body(Body::from(
serde_json::to_string(&serde_json::json!({"content": "sneaky"})).unwrap(),
))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::NOT_FOUND,
"bob should not be able to prompt alice's job"
);
}
#[tokio::test]
async fn test_jobs_cancel_rejects_other_user() {
let (db, _dir) = test_db().await;
let mut alice_job = make_sandbox_job("alice", "alice running");
alice_job.status = "running".to_string();
alice_job.success = None;
alice_job.completed_at = None;
db.save_sandbox_job(&alice_job).await.unwrap();
let state = build_state(Some(db), None);
let auth = two_user_auth();
let app = jobs_router(state, auth);
// Bob tries to cancel alice's job.
let req = Request::builder()
.method(Method::POST)
.uri(format!("/api/jobs/{}/cancel", alice_job.id))
.header("Authorization", "Bearer tok-bob")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::NOT_FOUND,
"bob should not be able to cancel alice's job"
);
}
}
// ═══════════════════════════════════════════════════════════════════════
// Routines Isolation Tests
// ═══════════════════════════════════════════════════════════════════════
#[cfg(feature = "libsql")]
mod routines_isolation {
use super::*;
use crate::channels::web::handlers::routines::{
routines_delete_handler, routines_detail_handler, routines_list_handler,
routines_summary_handler, routines_toggle_handler,
};
// RoutineStore methods are accessed through the Database supertrait.
fn routines_router(state: Arc<GatewayState>, auth: MultiAuthState) -> Router {
Router::new()
.route("/api/routines", get(routines_list_handler))
.route("/api/routines/summary", get(routines_summary_handler))
.route("/api/routines/{id}", get(routines_detail_handler))
.route("/api/routines/{id}/toggle", post(routines_toggle_handler))
.route("/api/routines/{id}", delete(routines_delete_handler))
.layer(middleware::from_fn_with_state(
crate::channels::web::auth::CombinedAuthState::from(auth),
auth_middleware,
))
.with_state(state)
}
#[tokio::test]
async fn test_routines_isolation() {
let (db, _dir) = test_db().await;
// Create routines for alice and bob.
let alice_routine = make_routine("alice", "alice-daily");
let bob_routine = make_routine("bob", "bob-daily");
db.create_routine(&alice_routine).await.unwrap();
db.create_routine(&bob_routine).await.unwrap();
let state = build_state(Some(db), None);
let auth = two_user_auth();
let app = routines_router(state, auth);
// Alice sees only her routine in the list.
let req = Request::builder()
.uri("/api/routines")
.header("Authorization", "Bearer tok-alice")
.body(Body::empty())
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body: serde_json::Value =
serde_json::from_slice(&axum::body::to_bytes(resp.into_body(), 8192).await.unwrap())
.unwrap();
let routines = body["routines"].as_array().unwrap();
assert_eq!(routines.len(), 1, "alice should see only her routines");
assert_eq!(routines[0]["name"], "alice-daily");
// Bob sees only his routine.
let req = Request::builder()
.uri("/api/routines")
.header("Authorization", "Bearer tok-bob")
.body(Body::empty())
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body: serde_json::Value =
serde_json::from_slice(&axum::body::to_bytes(resp.into_body(), 8192).await.unwrap())
.unwrap();
let routines = body["routines"].as_array().unwrap();
assert_eq!(routines.len(), 1, "bob should see only his routines");
assert_eq!(routines[0]["name"], "bob-daily");
// Bob cannot view alice's routine detail.
let req = Request::builder()
.uri(format!("/api/routines/{}", alice_routine.id))
.header("Authorization", "Bearer tok-bob")
.body(Body::empty())
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::NOT_FOUND,
"bob should not see alice's routine detail"
);
// Bob cannot toggle alice's routine.
let req = Request::builder()
.method(Method::POST)
.uri(format!("/api/routines/{}/toggle", alice_routine.id))
.header("Authorization", "Bearer tok-bob")
.body(Body::empty())
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::NOT_FOUND,
"bob should not toggle alice's routine"
);
// Bob cannot delete alice's routine.
let req = Request::builder()
.method(Method::DELETE)
.uri(format!("/api/routines/{}", alice_routine.id))
.header("Authorization", "Bearer tok-bob")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::NOT_FOUND,
"bob should not delete alice's routine"
);
}
}
// ═══════════════════════════════════════════════════════════════════════
// Handler Auth Enforcement Tests
// ═══════════════════════════════════════════════════════════════════════
mod auth_enforcement {
use super::*;
/// Dummy handler that extracts `AuthenticatedUser` — if the auth middleware
/// rejects the request, this handler is never reached.
async fn authed_handler(AuthenticatedUser(_user): AuthenticatedUser) -> &'static str {
"ok"
}
/// Build a router with the real auth middleware and dummy handlers at all
/// the paths we want to verify require authentication.
fn auth_test_router(auth: MultiAuthState) -> Router {
let state = build_state(None, None);
Router::new()
// Routines
.route("/api/routines", get(authed_handler))
.route("/api/routines/summary", get(authed_handler))
.route("/api/routines/{id}", get(authed_handler))
.route("/api/routines/{id}/toggle", post(authed_handler))
.route("/api/routines/{id}", delete(authed_handler))
// Skills
.route("/api/skills", get(authed_handler))
.route("/api/skills/search", post(authed_handler))
.route("/api/skills/install", post(authed_handler))
.route("/api/skills/{name}", delete(authed_handler))
// Logs
.route("/api/logs/events", get(authed_handler))
.route("/api/logs/level", get(authed_handler).put(authed_handler))
// Gateway status
.route("/api/gateway/status", get(authed_handler))
.layer(middleware::from_fn_with_state(
crate::channels::web::auth::CombinedAuthState::from(auth),
auth_middleware,
))
.with_state(state)
}
/// Send a request without auth and assert it returns UNAUTHORIZED.
async fn assert_requires_auth(app: &Router, method: Method, uri: &str) {
let req = Request::builder()
.method(method.clone())
.uri(uri)
.body(Body::empty())
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::UNAUTHORIZED,
"{} {} should require auth",
method,
uri
);
}
/// Send a request with a valid token and assert it succeeds.
async fn assert_passes_with_token(app: &Router, method: Method, uri: &str, token: &str) {
let req = Request::builder()
.method(method.clone())
.uri(uri)
.header("Authorization", format!("Bearer {token}"))
.body(Body::empty())
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::OK,
"{} {} should pass with valid token",
method,
uri
);
}
#[tokio::test]
async fn test_routines_handlers_require_auth() {
let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string());
let app = auth_test_router(auth);
let id = Uuid::new_v4();
assert_requires_auth(&app, Method::GET, "/api/routines").await;
assert_requires_auth(&app, Method::GET, "/api/routines/summary").await;
assert_requires_auth(&app, Method::GET, &format!("/api/routines/{id}")).await;
assert_requires_auth(&app, Method::POST, &format!("/api/routines/{id}/toggle")).await;
assert_requires_auth(&app, Method::DELETE, &format!("/api/routines/{id}")).await;
}
#[tokio::test]
async fn test_skills_handlers_require_auth() {
let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string());
let app = auth_test_router(auth);
assert_requires_auth(&app, Method::GET, "/api/skills").await;
assert_requires_auth(&app, Method::POST, "/api/skills/search").await;
assert_requires_auth(&app, Method::POST, "/api/skills/install").await;
assert_requires_auth(&app, Method::DELETE, "/api/skills/test-skill").await;
}
#[tokio::test]
async fn test_logs_handlers_require_auth() {
let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string());
let app = auth_test_router(auth);
assert_requires_auth(&app, Method::GET, "/api/logs/events").await;
assert_requires_auth(&app, Method::GET, "/api/logs/level").await;
assert_requires_auth(&app, Method::PUT, "/api/logs/level").await;
}
#[tokio::test]
async fn test_gateway_status_requires_auth() {
let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string());
let app = auth_test_router(auth);
assert_requires_auth(&app, Method::GET, "/api/gateway/status").await;
}
#[tokio::test]
async fn test_valid_token_passes_all_endpoints() {
let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string());
let app = auth_test_router(auth);
let id = Uuid::new_v4();
assert_passes_with_token(&app, Method::GET, "/api/routines", "secret-tok").await;
assert_passes_with_token(&app, Method::GET, "/api/skills", "secret-tok").await;
assert_passes_with_token(&app, Method::GET, "/api/logs/events", "secret-tok").await;
assert_passes_with_token(&app, Method::GET, "/api/gateway/status", "secret-tok").await;
assert_passes_with_token(
&app,
Method::GET,
&format!("/api/routines/{id}"),
"secret-tok",
)
.await;
}
#[tokio::test]
async fn test_wrong_token_rejected_on_all_endpoints() {
let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string());
let app = auth_test_router(auth);
// Wrong token should be rejected.
let req = Request::builder()
.uri("/api/routines")
.header("Authorization", "Bearer wrong-tok")
.body(Body::empty())
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
let req = Request::builder()
.uri("/api/gateway/status")
.header("Authorization", "Bearer wrong-tok")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
}
}
// ═══════════════════════════════════════════════════════════════════════
// Admin Endpoint Role Enforcement Tests
// ═══════════════════════════════════════════════════════════════════════
mod admin_role_enforcement {
use super::*;
use crate::channels::web::handlers::users::{
users_activate_handler, users_detail_handler, users_list_handler, users_suspend_handler,
users_update_handler,
};
use axum::routing::patch;
/// Build a router with admin user endpoints behind multi-user auth.
/// Uses a member-role token and an admin-role token.
fn admin_router() -> Router {
let mut tokens = HashMap::new();
tokens.insert(
"tok-admin".to_string(),
UserIdentity {
user_id: "admin-user".to_string(),
role: "admin".to_string(),
workspace_read_scopes: vec![],
},
);
tokens.insert(
"tok-member".to_string(),
UserIdentity {
user_id: "member-user".to_string(),
role: "member".to_string(),
workspace_read_scopes: vec![],
},
);
let auth = MultiAuthState::multi(tokens);
let state = build_state(None, None);
Router::new()
.route("/api/admin/users", get(users_list_handler))
.route("/api/admin/users/{id}", get(users_detail_handler))
.route("/api/admin/users/{id}", patch(users_update_handler))
.route("/api/admin/users/{id}/suspend", post(users_suspend_handler))
.route(
"/api/admin/users/{id}/activate",
post(users_activate_handler),
)
.layer(middleware::from_fn_with_state(
crate::channels::web::auth::CombinedAuthState::from(auth),
auth_middleware,
))
.with_state(state)
}
/// Assert a request returns FORBIDDEN for a member token.
async fn assert_forbidden_for_member(app: &Router, method: Method, uri: &str) {
let req = Request::builder()
.method(method)
.uri(uri)
.header("Authorization", "Bearer tok-member")
.body(Body::empty())
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::FORBIDDEN,
"expected 403 for member on {}",
uri
);
}
#[tokio::test]
async fn test_admin_user_endpoints_reject_member_role() {
let app = admin_router();
assert_forbidden_for_member(&app, Method::GET, "/api/admin/users").await;
assert_forbidden_for_member(&app, Method::GET, "/api/admin/users/some-id").await;
assert_forbidden_for_member(&app, Method::POST, "/api/admin/users/some-id/suspend").await;
assert_forbidden_for_member(&app, Method::POST, "/api/admin/users/some-id/activate").await;
}
#[tokio::test]
async fn test_admin_user_endpoints_accept_admin_role() {
let app = admin_router();
// Admin token should pass auth (will get 503 since no DB, but not 403).
let req = Request::builder()
.uri("/api/admin/users")
.header("Authorization", "Bearer tok-admin")
.body(Body::empty())
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_ne!(
resp.status(),
StatusCode::FORBIDDEN,
"admin should not get 403"
);
}
}
// ═══════════════════════════════════════════════════════════════════════
// DbAuthenticator Cache Bounded Tests
// ═══════════════════════════════════════════════════════════════════════
mod db_auth_cache {
use super::*;
use std::time::Instant;
#[tokio::test]
async fn test_cache_bounded_by_max_entries() {
// Access the internal cache and verify LRU eviction.
// We can't easily test through `authenticate()` since it hits the DB,
// so we test the LRU cache directly.
let cap = std::num::NonZeroUsize::new(4).unwrap(); // safety: test-only, 4 is non-zero
let cache: lru::LruCache<[u8; 32], (UserIdentity, Instant)> = lru::LruCache::new(cap);
let cache = Arc::new(tokio::sync::RwLock::new(cache));
{
let mut c = cache.write().await;
for i in 0..10u8 {
let mut hash = [0u8; 32];
hash[0] = i;
c.put(
hash,
(
UserIdentity {
user_id: format!("user-{i}"),
role: "member".to_string(),
workspace_read_scopes: vec![],
},
Instant::now(),
),
);
}
// Cache must be bounded at capacity, not grown to 10.
assert_eq!(c.len(), 4, "cache should be bounded to capacity"); // safety: test assertion
}
}
}
+206 -33
View File
@@ -63,9 +63,6 @@ pub struct TurnInfo {
pub started_at: String, pub started_at: String,
pub completed_at: Option<String>, pub completed_at: Option<String>,
pub tool_calls: Vec<ToolCallInfo>, pub tool_calls: Vec<ToolCallInfo>,
/// Agent's reasoning narrative for this turn.
#[serde(skip_serializing_if = "Option::is_none")]
pub narrative: Option<String>,
} }
#[derive(Debug, Serialize)] #[derive(Debug, Serialize)]
@@ -77,9 +74,6 @@ pub struct ToolCallInfo {
pub result_preview: Option<String>, pub result_preview: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
pub error: Option<String>, pub error: Option<String>,
/// Agent's reasoning for choosing this tool.
#[serde(skip_serializing_if = "Option::is_none")]
pub rationale: Option<String>,
} }
#[derive(Debug, Serialize)] #[derive(Debug, Serialize)]
@@ -120,9 +114,165 @@ pub struct ApprovalRequest {
pub thread_id: Option<String>, pub thread_id: Option<String>,
} }
// --- App Event (re-exported from ironclaw_common) --- // --- SSE Event Types ---
pub use ironclaw_common::{AppEvent, ToolDecisionDto}; #[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<String>,
},
#[serde(rename = "tool_started")]
ToolStarted {
name: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "tool_completed")]
ToolCompleted {
name: String,
success: bool,
#[serde(skip_serializing_if = "Option::is_none")]
error: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
parameters: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "tool_result")]
ToolResult {
name: String,
preview: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "stream_chunk")]
StreamChunk {
content: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "status")]
Status {
message: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[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<String>,
/// 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<String>,
#[serde(skip_serializing_if = "Option::is_none")]
auth_url: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
setup_url: Option<String>,
},
#[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<String>,
},
#[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<String>,
#[serde(skip_serializing_if = "Option::is_none")]
fallback_deliverable: Option<serde_json::Value>,
},
/// An image was generated by a tool.
#[serde(rename = "image_generated")]
ImageGenerated {
data_url: String,
#[serde(skip_serializing_if = "Option::is_none")]
path: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
/// Suggested follow-up messages for the user.
#[serde(rename = "suggestions")]
Suggestions {
suggestions: Vec<String>,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
/// 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<String>,
},
/// 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<String>,
},
}
// --- Memory --- // --- Memory ---
@@ -634,9 +784,32 @@ pub enum WsServerMessage {
} }
impl WsServerMessage { impl WsServerMessage {
/// Create a WsServerMessage from an AppEvent. /// Create a WsServerMessage from an SseEvent.
pub fn from_app_event(event: &AppEvent) -> Self { pub fn from_sse_event(event: &SseEvent) -> Self {
let event_type = event.event_type(); 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",
};
let data = serde_json::to_value(event).unwrap_or(serde_json::Value::Null); let data = serde_json::to_value(event).unwrap_or(serde_json::Value::Null);
WsServerMessage::Event { WsServerMessage::Event {
event_type: event_type.to_string(), event_type: event_type.to_string(),
@@ -928,12 +1101,12 @@ mod tests {
} }
#[test] #[test]
fn test_ws_server_from_app_event_response() { fn test_ws_server_from_sse_response() {
let event = AppEvent::Response { let sse = SseEvent::Response {
content: "hello".to_string(), content: "hello".to_string(),
thread_id: "t1".to_string(), thread_id: "t1".to_string(),
}; };
let ws = WsServerMessage::from_app_event(&event); let ws = WsServerMessage::from_sse_event(&sse);
match ws { match ws {
WsServerMessage::Event { event_type, data } => { WsServerMessage::Event { event_type, data } => {
assert_eq!(event_type, "response"); assert_eq!(event_type, "response");
@@ -945,12 +1118,12 @@ mod tests {
} }
#[test] #[test]
fn test_ws_server_from_app_event_thinking() { fn test_ws_server_from_sse_thinking() {
let event = AppEvent::Thinking { let sse = SseEvent::Thinking {
message: "reasoning...".to_string(), message: "reasoning...".to_string(),
thread_id: None, thread_id: None,
}; };
let ws = WsServerMessage::from_app_event(&event); let ws = WsServerMessage::from_sse_event(&sse);
match ws { match ws {
WsServerMessage::Event { event_type, data } => { WsServerMessage::Event { event_type, data } => {
assert_eq!(event_type, "thinking"); assert_eq!(event_type, "thinking");
@@ -961,8 +1134,8 @@ mod tests {
} }
#[test] #[test]
fn test_ws_server_from_app_event_approval_needed() { fn test_ws_server_from_sse_approval_needed() {
let event = AppEvent::ApprovalNeeded { let sse = SseEvent::ApprovalNeeded {
request_id: "r1".to_string(), request_id: "r1".to_string(),
tool_name: "shell".to_string(), tool_name: "shell".to_string(),
description: "Run ls".to_string(), description: "Run ls".to_string(),
@@ -970,7 +1143,7 @@ mod tests {
thread_id: Some("t1".to_string()), thread_id: Some("t1".to_string()),
allow_always: true, allow_always: true,
}; };
let ws = WsServerMessage::from_app_event(&event); let ws = WsServerMessage::from_sse_event(&sse);
match ws { match ws {
WsServerMessage::Event { event_type, data } => { WsServerMessage::Event { event_type, data } => {
assert_eq!(event_type, "approval_needed"); assert_eq!(event_type, "approval_needed");
@@ -982,9 +1155,9 @@ mod tests {
} }
#[test] #[test]
fn test_ws_server_from_app_event_heartbeat() { fn test_ws_server_from_sse_heartbeat() {
let event = AppEvent::Heartbeat; let sse = SseEvent::Heartbeat;
let ws = WsServerMessage::from_app_event(&event); let ws = WsServerMessage::from_sse_event(&sse);
match ws { match ws {
WsServerMessage::Event { event_type, .. } => { WsServerMessage::Event { event_type, .. } => {
assert_eq!(event_type, "heartbeat"); assert_eq!(event_type, "heartbeat");
@@ -1024,8 +1197,8 @@ mod tests {
} }
#[test] #[test]
fn test_app_event_auth_required_serialize() { fn test_sse_auth_required_serialize() {
let event = AppEvent::AuthRequired { let event = SseEvent::AuthRequired {
extension_name: "notion".to_string(), extension_name: "notion".to_string(),
instructions: Some("Get your token from...".to_string()), instructions: Some("Get your token from...".to_string()),
auth_url: None, auth_url: None,
@@ -1041,8 +1214,8 @@ mod tests {
} }
#[test] #[test]
fn test_app_event_auth_completed_serialize() { fn test_sse_auth_completed_serialize() {
let event = AppEvent::AuthCompleted { let event = SseEvent::AuthCompleted {
extension_name: "notion".to_string(), extension_name: "notion".to_string(),
success: true, success: true,
message: "notion authenticated (3 tools loaded)".to_string(), message: "notion authenticated (3 tools loaded)".to_string(),
@@ -1055,14 +1228,14 @@ mod tests {
} }
#[test] #[test]
fn test_ws_server_from_app_event_auth_required() { fn test_ws_server_from_sse_auth_required() {
let event = AppEvent::AuthRequired { let sse = SseEvent::AuthRequired {
extension_name: "openai".to_string(), extension_name: "openai".to_string(),
instructions: Some("Enter API key".to_string()), instructions: Some("Enter API key".to_string()),
auth_url: None, auth_url: None,
setup_url: None, setup_url: None,
}; };
let ws = WsServerMessage::from_app_event(&event); let ws = WsServerMessage::from_sse_event(&sse);
match ws { match ws {
WsServerMessage::Event { event_type, data } => { WsServerMessage::Event { event_type, data } => {
assert_eq!(event_type, "auth_required"); assert_eq!(event_type, "auth_required");
@@ -1073,13 +1246,13 @@ mod tests {
} }
#[test] #[test]
fn test_ws_server_from_app_event_auth_completed() { fn test_ws_server_from_sse_auth_completed() {
let event = AppEvent::AuthCompleted { let sse = SseEvent::AuthCompleted {
extension_name: "slack".to_string(), extension_name: "slack".to_string(),
success: false, success: false,
message: "Invalid token".to_string(), message: "Invalid token".to_string(),
}; };
let ws = WsServerMessage::from_app_event(&event); let ws = WsServerMessage::from_sse_event(&sse);
match ws { match ws {
WsServerMessage::Event { event_type, data } => { WsServerMessage::Event { event_type, data } => {
assert_eq!(event_type, "auth_completed"); assert_eq!(event_type, "auth_completed");
+115 -86
View File
@@ -2,21 +2,28 @@
use crate::channels::web::types::{ToolCallInfo, TurnInfo}; use crate::channels::web::types::{ToolCallInfo, TurnInfo};
pub use ironclaw_common::truncate_preview; /// Truncate a string to at most `max_bytes` bytes at a char boundary, appending "...".
///
/// If the input is wrapped in `<tool_output …>…</tool_output>` 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]);
/// Parse tool call summary JSON objects into `ToolCallInfo` structs. // Re-close <tool_output> if truncation cut through the closing tag.
fn parse_tool_call_infos(calls: &[serde_json::Value]) -> Vec<ToolCallInfo> { if s.starts_with("<tool_output") && !result.ends_with("</tool_output>") {
calls result.push_str("\n</tool_output>");
.iter() }
.map(|c| ToolCallInfo {
name: c["name"].as_str().unwrap_or("unknown").to_string(), result
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). /// Build TurnInfo pairs from flat DB messages (user/tool_calls/assistant triples).
@@ -42,7 +49,6 @@ pub fn build_turns_from_db_messages(
started_at: msg.created_at.to_rfc3339(), started_at: msg.created_at.to_rfc3339(),
completed_at: None, completed_at: None,
tool_calls: Vec::new(), tool_calls: Vec::new(),
narrative: None,
}; };
// Check if next message is a tool_calls record // Check if next message is a tool_calls record
@@ -50,28 +56,18 @@ pub fn build_turns_from_db_messages(
&& next.role == "tool_calls" && next.role == "tool_calls"
{ {
let tc_msg = iter.next().expect("peeked"); let tc_msg = iter.next().expect("peeked");
// Parse tool_calls JSON — supports two formats: match serde_json::from_str::<Vec<serde_json::Value>>(&tc_msg.content) {
// safety: no byte-index slicing; comment describes JSON shape Ok(calls) => {
match serde_json::from_str::<serde_json::Value>(&tc_msg.content) { turn.tool_calls = calls
Ok(serde_json::Value::Array(calls)) => { .iter()
// Old format: plain array .map(|c| ToolCallInfo {
turn.tool_calls = parse_tool_call_infos(&calls); name: c["name"].as_str().unwrap_or("unknown").to_string(),
} has_result: c.get("result_preview").is_some(),
Ok(serde_json::Value::Object(obj)) => { has_error: c.get("error").is_some(),
// New wrapped format with narrative result_preview: c["result_preview"].as_str().map(String::from),
turn.narrative = obj error: c["error"].as_str().map(String::from),
.get("narrative") })
.and_then(|v| v.as_str()) .collect();
.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) => { Err(e) => {
tracing::warn!( tracing::warn!(
@@ -109,7 +105,6 @@ pub fn build_turns_from_db_messages(
started_at: msg.created_at.to_rfc3339(), started_at: msg.created_at.to_rfc3339(),
completed_at: Some(msg.created_at.to_rfc3339()), completed_at: Some(msg.created_at.to_rfc3339()),
tool_calls: Vec::new(), tool_calls: Vec::new(),
narrative: None,
}); });
turn_number += 1; turn_number += 1;
} }
@@ -123,6 +118,88 @@ mod tests {
use super::*; use super::*;
use uuid::Uuid; 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 = "<tool_output name=\"search\">\nSome very long content here\n</tool_output>";
// Truncate so it cuts before the closing tag
let result = truncate_preview(s, 60);
assert!(result.ends_with("</tool_output>"));
assert!(result.contains("..."));
}
#[test]
fn test_truncate_preview_no_extra_close_when_intact() {
let s = "<tool_output name=\"echo\">\nshort\n</tool_output>";
// 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("</tool_output>").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("</tool_output>"));
}
// ---- build_turns_from_db_messages tests ---- // ---- build_turns_from_db_messages tests ----
fn make_msg(role: &str, content: &str, offset_ms: i64) -> crate::history::ConversationMessage { fn make_msg(role: &str, content: &str, offset_ms: i64) -> crate::history::ConversationMessage {
@@ -228,52 +305,4 @@ mod tests {
assert!(turns[0].tool_calls.is_empty()); assert!(turns[0].tool_calls.is_empty());
assert_eq!(turns[0].state, "Completed"); 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);
}
} }
+17 -29
View File
@@ -62,11 +62,7 @@ impl Default for WsConnectionTracker {
/// ///
/// When either task ends (client disconnect or broadcast closed), both are /// When either task ends (client disconnect or broadcast closed), both are
/// cleaned up. /// cleaned up.
pub async fn handle_ws_connection( pub async fn handle_ws_connection(socket: WebSocket, state: Arc<GatewayState>) {
socket: WebSocket,
state: Arc<GatewayState>,
user: crate::channels::web::auth::UserIdentity,
) {
let (mut ws_sink, mut ws_stream) = socket.split(); let (mut ws_sink, mut ws_stream) = socket.split();
// Track connection // Track connection
@@ -75,9 +71,9 @@ pub async fn handle_ws_connection(
} }
let tracker_for_drop = state.ws_tracker.clone(); let tracker_for_drop = state.ws_tracker.clone();
// Subscribe to broadcast events (same source as SSE), scoped to this user. // Subscribe to broadcast events (same source as SSE).
// Reject if we've hit the connection limit. // Reject if we've hit the connection limit.
let Some(raw_stream) = state.sse.subscribe_raw(Some(user.user_id.clone())) else { let Some(raw_stream) = state.sse.subscribe_raw() else {
tracing::warn!("WebSocket rejected: too many connections"); tracing::warn!("WebSocket rejected: too many connections");
// Decrement the WS tracker we already incremented above. // Decrement the WS tracker we already incremented above.
if let Some(ref tracker) = tracker_for_drop { if let Some(ref tracker) = tracker_for_drop {
@@ -97,7 +93,7 @@ pub async fn handle_ws_connection(
let msg = tokio::select! { let msg = tokio::select! {
event = event_stream.next() => { event = event_stream.next() => {
match event { match event {
Some(app_event) => WsServerMessage::from_app_event(&app_event), Some(sse_event) => WsServerMessage::from_sse_event(&sse_event),
None => break, // Broadcast channel closed None => break, // Broadcast channel closed
} }
} }
@@ -121,7 +117,7 @@ pub async fn handle_ws_connection(
}); });
// Receiver task: read client frames and route to agent // Receiver task: read client frames and route to agent
let user_id = user.user_id; let user_id = state.user_id.clone();
while let Some(Ok(frame)) = ws_stream.next().await { while let Some(Ok(frame)) = ws_stream.next().await {
match frame { match frame {
Message::Text(text) => { Message::Text(text) => {
@@ -267,15 +263,11 @@ async fn handle_client_message(
token, token,
} => { } => {
if let Some(ref ext_mgr) = state.extension_manager { if let Some(ref ext_mgr) = state.extension_manager {
match ext_mgr match ext_mgr.configure_token(&extension_name, &token).await {
.configure_token(&extension_name, &token, user_id)
.await
{
Ok(result) => { Ok(result) => {
if result.verification.is_some() { if result.verification.is_some() {
state.sse.broadcast_for_user( state.sse.broadcast(
user_id, crate::channels::web::types::SseEvent::AuthRequired {
crate::channels::web::types::AppEvent::AuthRequired {
extension_name: extension_name.clone(), extension_name: extension_name.clone(),
instructions: Some(result.message), instructions: Some(result.message),
auth_url: None, auth_url: None,
@@ -283,10 +275,9 @@ async fn handle_client_message(
}, },
); );
} else { } else {
crate::channels::web::server::clear_auth_mode(state, user_id).await; crate::channels::web::server::clear_auth_mode(state).await;
state.sse.broadcast_for_user( state.sse.broadcast(
user_id, crate::channels::web::types::SseEvent::AuthCompleted {
crate::channels::web::types::AppEvent::AuthCompleted {
extension_name, extension_name,
success: true, success: true,
message: result.message, message: result.message,
@@ -297,9 +288,8 @@ async fn handle_client_message(
Err(e) => { Err(e) => {
let msg = format!("Auth failed: {}", e); let msg = format!("Auth failed: {}", e);
if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) { if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) {
state.sse.broadcast_for_user( state.sse.broadcast(
user_id, crate::channels::web::types::SseEvent::AuthRequired {
crate::channels::web::types::AppEvent::AuthRequired {
extension_name: extension_name.clone(), extension_name: extension_name.clone(),
instructions: Some(msg.clone()), instructions: Some(msg.clone()),
auth_url: None, auth_url: None,
@@ -321,7 +311,7 @@ async fn handle_client_message(
} }
} }
WsClientMessage::AuthCancel { .. } => { WsClientMessage::AuthCancel { .. } => {
crate::channels::web::server::clear_auth_mode(state, user_id).await; crate::channels::web::server::clear_auth_mode(state).await;
} }
WsClientMessage::Ping => { WsClientMessage::Ping => {
let _ = direct_tx.send(WsServerMessage::Pong).await; let _ = direct_tx.send(WsServerMessage::Pong).await;
@@ -508,9 +498,8 @@ mod tests {
GatewayState { GatewayState {
msg_tx: tokio::sync::RwLock::new(msg_tx), msg_tx: tokio::sync::RwLock::new(msg_tx),
sse: Arc::new(SseManager::new()), sse: SseManager::new(),
workspace: None, workspace: None,
workspace_pool: None,
session_manager: None, session_manager: None,
log_broadcaster: None, log_broadcaster: None,
log_level_handle: None, log_level_handle: None,
@@ -520,13 +509,13 @@ mod tests {
job_manager: None, job_manager: None,
prompt_queue: None, prompt_queue: None,
scheduler: None, scheduler: None,
owner_id: "test".to_string(), user_id: "test".to_string(),
shutdown_tx: tokio::sync::RwLock::new(None), shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())), ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: None, llm_provider: None,
skill_registry: None, skill_registry: None,
skill_catalog: None, skill_catalog: None,
chat_rate_limiter: crate::channels::web::server::PerUserRateLimiter::new(30, 60), chat_rate_limiter: crate::channels::web::server::RateLimiter::new(30, 60),
oauth_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60), oauth_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60),
webhook_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60), webhook_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60),
registry_entries: Vec::new(), registry_entries: Vec::new(),
@@ -534,7 +523,6 @@ mod tests {
routine_engine: Arc::new(tokio::sync::RwLock::new(None)), routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
startup_time: std::time::Instant::now(), startup_time: std::time::Instant::now(),
active_config: crate::channels::web::server::ActiveConfigSnapshot::default(), active_config: crate::channels::web::server::ActiveConfigSnapshot::default(),
secrets_store: None,
} }
} }
} }
+1 -12
View File
@@ -25,7 +25,6 @@ pub mod import;
mod logs; mod logs;
mod mcp; mod mcp;
pub mod memory; pub mod memory;
mod models;
pub mod oauth_defaults; pub mod oauth_defaults;
mod pairing; mod pairing;
mod registry; mod registry;
@@ -46,7 +45,6 @@ pub use logs::{LogsCommand, run_logs_command};
pub use mcp::{McpCommand, run_mcp_command}; pub use mcp::{McpCommand, run_mcp_command};
pub use memory::MemoryCommand; pub use memory::MemoryCommand;
pub use memory::run_memory_command_with_db; pub use memory::run_memory_command_with_db;
pub use models::{ModelsCommand, run_models_command};
pub use pairing::{PairingCommand, run_pairing_command, run_pairing_command_with_store}; pub use pairing::{PairingCommand, run_pairing_command, run_pairing_command_with_store};
pub use registry::{RegistryCommand, run_registry_command}; pub use registry::{RegistryCommand, run_registry_command};
pub use routines::{RoutinesCommand, run_routines_command}; pub use routines::{RoutinesCommand, run_routines_command};
@@ -219,14 +217,6 @@ pub enum Command {
)] )]
Hooks(HooksCommand), Hooks(HooksCommand),
/// Manage LLM providers and models
#[command(
subcommand,
about = "Manage LLM providers and models",
long_about = "List providers, view current configuration, and set active provider/model.\nExamples:\n ironclaw models list\n ironclaw models list openai --verbose\n ironclaw models status\n ironclaw models set gpt-4o\n ironclaw models set-provider anthropic --model claude-sonnet-4-6-20250514"
)]
Models(ModelsCommand),
/// Probe external dependencies and validate configuration /// Probe external dependencies and validate configuration
#[command( #[command(
about = "Run diagnostics", about = "Run diagnostics",
@@ -352,8 +342,7 @@ pub async fn run_routines_cli(
.await .await
.map_err(|e| anyhow::anyhow!("{e:#}"))?; .map_err(|e| anyhow::anyhow!("{e:#}"))?;
let user_id = let user_id = std::env::var("GATEWAY_USER_ID").unwrap_or_else(|_| "default".to_string());
std::env::var("IRONCLAW_OWNER_ID").unwrap_or_else(|_| "default".to_string());
run_routines_command(routines_cmd.clone(), db, &user_id).await run_routines_command(routines_cmd.clone(), db, &user_id).await
} }
-864
View File
@@ -1,864 +0,0 @@
//! Models management CLI commands.
//!
//! Provides subcommands for listing providers, viewing current model
//! configuration, and setting the active provider/model. Settings are
//! persisted to both `config.toml` and `~/.ironclaw/.env` so changes
//! take effect immediately (no DB connection required).
use clap::Subcommand;
use std::path::Path;
use crate::llm::registry::ProviderRegistry;
use crate::settings::Settings;
#[derive(Subcommand, Debug, Clone)]
pub enum ModelsCommand {
/// List providers (or available models for a specific provider)
List {
/// Show only a specific provider (by ID or alias)
provider: Option<String>,
/// Show detailed information (env vars, base URL, protocol)
#[arg(short, long)]
verbose: bool,
/// Output as JSON
#[arg(long)]
json: bool,
},
/// Show current model configuration
Status {
/// Output as JSON
#[arg(long)]
json: bool,
},
/// Set the default model
Set {
/// Model name (e.g., "gpt-5-mini", "claude-sonnet-4-6-20250514")
model: String,
},
/// Set the LLM provider
SetProvider {
/// Provider ID or alias (e.g., "openai", "anthropic", "ollama")
provider: String,
/// Also set the model (defaults to provider's default model)
#[arg(long)]
model: Option<String>,
},
}
/// Run the models CLI subcommand.
pub async fn run_models_command(
cmd: ModelsCommand,
config_path: Option<&Path>,
) -> anyhow::Result<()> {
match cmd {
ModelsCommand::List {
provider,
verbose,
json,
} => {
if let Some(ref id) = provider {
cmd_show_provider(id, verbose, json, config_path).await
} else {
cmd_list_providers(verbose, json, config_path).await
}
}
ModelsCommand::Status { json } => cmd_status(json, config_path),
ModelsCommand::Set { model } => cmd_set_model(&model, config_path),
ModelsCommand::SetProvider { provider, model } => {
cmd_set_provider(&provider, model.as_deref(), config_path)
}
}
}
// ─── Shared helpers ───────────────────────────────────────────────
/// Resolve the currently active backend and model from env + settings.
fn resolve_active(config_path: Option<&Path>) -> (String, String) {
let settings = load_settings(config_path);
resolve_active_from_settings(&settings)
}
/// Resolve active backend + model from a pre-loaded Settings.
fn resolve_active_from_settings(settings: &Settings) -> (String, String) {
let backend = std::env::var("LLM_BACKEND")
.ok()
.or_else(|| settings.llm_backend.clone())
.unwrap_or_else(|| "nearai".to_string());
let registry = ProviderRegistry::load();
let canonical_backend = registry
.find(&backend)
.map(|d| d.id.clone())
.unwrap_or_else(|| backend.clone());
let model = if canonical_backend == "nearai" {
std::env::var("NEARAI_MODEL")
.ok()
.or_else(|| settings.selected_model.clone())
.unwrap_or_else(|| "qwen2.5-72b-instruct:free".to_string())
} else if let Some(def) = registry.find(&canonical_backend) {
std::env::var(&def.model_env)
.ok()
.or_else(|| settings.selected_model.clone())
.unwrap_or_else(|| def.default_model.clone())
} else {
settings
.selected_model
.clone()
.unwrap_or_else(|| "unknown".to_string())
};
(canonical_backend, model)
}
fn load_settings(config_path: Option<&Path>) -> Settings {
if let Some(path) = config_path {
Settings::load_toml(path).ok().flatten().unwrap_or_default()
} else {
let toml_path = config_toml_path();
if toml_path.exists() {
Settings::load_toml(&toml_path)
.ok()
.flatten()
.unwrap_or_default()
} else {
Settings::load()
}
}
}
fn save_settings(settings: &Settings, config_path: Option<&Path>) -> anyhow::Result<()> {
let path = config_path
.map(|p| p.to_path_buf())
.unwrap_or_else(config_toml_path);
settings
.save_toml(&path)
.map_err(|e| anyhow::anyhow!("{}", e))?;
Ok(())
}
fn config_toml_path() -> std::path::PathBuf {
crate::bootstrap::ironclaw_base_dir().join("config.toml")
}
/// Try to fetch the live model list from a provider.
///
/// Best-effort: returns `None` if config loading, provider creation, or the
/// `list_models()` call fails (missing API key, network error, etc.).
async fn try_fetch_models(provider_id: &str, config_path: Option<&Path>) -> Option<Vec<String>> {
let config = crate::config::Config::from_env_with_toml(config_path)
.await
.ok()?;
// Override backend to the requested provider so create_llm_provider
// constructs the right one.
let mut llm_config = config.llm.clone();
llm_config.backend = provider_id.to_string();
// For registry providers, resolve the RegistryProviderConfig if not
// already set for this backend.
if provider_id != "nearai" && provider_id != "bedrock" {
let registry = ProviderRegistry::load();
if let Some(def) = registry.find(provider_id)
&& llm_config
.provider
.as_ref()
.is_none_or(|p| p.provider_id != def.id)
{
// Build a minimal RegistryProviderConfig from env + registry
let api_key = def
.api_key_env
.as_ref()
.and_then(|env| std::env::var(env).ok());
if def.api_key_required && api_key.is_none() {
return None;
}
let base_url = def.default_base_url.clone().unwrap_or_default();
llm_config.provider = Some(crate::llm::RegistryProviderConfig {
protocol: def.protocol,
provider_id: def.id.clone(),
model: def.default_model.clone(),
api_key: api_key.map(secrecy::SecretString::from),
base_url,
extra_headers: Vec::new(),
oauth_token: None,
is_codex_chatgpt: false,
refresh_token: None,
auth_path: None,
cache_retention: Default::default(),
unsupported_params: def.unsupported_params.clone(),
});
}
}
let session = crate::llm::create_session_manager(config.llm.session.clone()).await;
let provider = crate::llm::create_llm_provider(&llm_config, session)
.await
.ok()?;
provider.list_models().await.ok().filter(|m| !m.is_empty())
}
/// Print available models section (text output).
fn print_model_list(models: &Option<Vec<String>>, active_model: Option<&String>) {
match models {
Some(models) => {
println!("\n Available models ({}):", models.len());
for m in models {
let marker = active_model
.filter(|a| a.as_str() == m)
.map(|_| " (active)")
.unwrap_or("");
println!(" {}{}", m, marker);
}
}
None => {
println!(
"\n Could not fetch model list (missing credentials or provider unavailable)."
);
}
}
}
/// Also update `~/.ironclaw/.env` so changes take effect immediately.
///
/// Skipped when `config_path` is `Some` (custom `--config`), because the user
/// is explicitly targeting a different config file and we must not pollute the
/// default profile's `.env`.
fn sync_to_dotenv(config_path: Option<&Path>, vars: &[(&str, &str)]) {
if config_path.is_some() {
return;
}
if let Err(e) = crate::bootstrap::upsert_bootstrap_vars(vars) {
eprintln!("Warning: failed to update .env: {}", e);
}
}
// ─── status ───────────────────────────────────────────────────────
fn cmd_status(json: bool, config_path: Option<&Path>) -> anyhow::Result<()> {
let settings = load_settings(config_path);
let (backend, model) = resolve_active_from_settings(&settings);
let registry = ProviderRegistry::load();
let fallback = std::env::var("NEARAI_FALLBACK_MODEL").ok();
let cheap = std::env::var("NEARAI_CHEAP_MODEL").ok();
let description = if backend == "nearai" {
"NEAR AI inference (default)".to_string()
} else {
registry
.find(&backend)
.map(|d| d.description.clone())
.unwrap_or_default()
};
if json {
let v = serde_json::json!({
"provider": backend,
"model": model,
"description": description,
"fallback_model": fallback,
"cheap_model": cheap,
});
println!(
"{}",
serde_json::to_string_pretty(&v).unwrap_or_else(|_| "{}".to_string())
);
return Ok(());
}
println!("Provider: {} ({})", backend, description);
println!("Model: {}", model);
if let Some(ref fb) = fallback {
println!("Fallback: {}", fb);
}
if let Some(ref ch) = cheap {
println!("Cheap: {}", ch);
}
Ok(())
}
// ─── set ──────────────────────────────────────────────────────────
fn cmd_set_model(model: &str, config_path: Option<&Path>) -> anyhow::Result<()> {
let trimmed = model.trim();
if trimmed.is_empty() {
anyhow::bail!("Model name cannot be empty");
}
let mut settings = load_settings(config_path);
let registry = ProviderRegistry::load();
// Warn if model name doesn't match any known provider's default model
let known_model = registry.all().iter().any(|d| d.default_model == trimmed)
|| trimmed.contains("qwen") // nearai models
|| trimmed.contains("llama")
|| trimmed.contains("gpt")
|| trimmed.contains("claude")
|| trimmed.contains("gemini")
|| trimmed.contains("mistral");
if !known_model {
eprintln!(
"Warning: '{}' is not a recognized model name. Proceeding anyway.",
trimmed
);
}
settings.selected_model = Some(trimmed.to_string());
save_settings(&settings, config_path)?;
let backend = std::env::var("LLM_BACKEND")
.ok()
.or_else(|| settings.llm_backend.clone())
.unwrap_or_else(|| "nearai".to_string());
// Also write to .env so the change takes effect immediately
let model_env = if backend == "nearai" {
"NEARAI_MODEL".to_string()
} else {
registry
.find(&backend)
.map(|d| d.model_env.clone())
.unwrap_or_default()
};
if !model_env.is_empty() {
sync_to_dotenv(config_path, &[(&model_env, trimmed)]);
}
println!("Model set to '{}' (provider: {})", trimmed, backend);
println!(
"Saved to {}",
config_path
.map(|p| p.display().to_string())
.unwrap_or_else(|| config_toml_path().display().to_string())
);
Ok(())
}
// ─── set-provider ─────────────────────────────────────────────────
fn cmd_set_provider(
provider: &str,
model: Option<&str>,
config_path: Option<&Path>,
) -> anyhow::Result<()> {
let registry = ProviderRegistry::load();
// Validate and normalize provider
let canonical_id = if provider == "nearai" || provider == "near_ai" || provider == "near" {
"nearai".to_string()
} else {
let def = registry.find(provider).ok_or_else(|| {
let known: Vec<&str> = std::iter::once("nearai")
.chain(registry.all().iter().map(|d| d.id.as_str()))
.collect();
anyhow::anyhow!(
"Unknown provider '{}'. Known providers: {}",
provider,
known.join(", ")
)
})?;
def.id.clone()
};
// Resolve model: explicit > provider default
let resolved_model = if let Some(m) = model {
m.to_string()
} else if canonical_id == "nearai" {
"qwen2.5-72b-instruct:free".to_string()
} else if let Some(def) = registry.find(&canonical_id) {
def.default_model.clone()
} else {
"default".to_string()
};
let mut settings = load_settings(config_path);
settings.llm_backend = Some(canonical_id.clone());
settings.selected_model = Some(resolved_model.clone());
save_settings(&settings, config_path)?;
// Also write to .env so the change takes effect immediately
let model_env = if canonical_id == "nearai" {
"NEARAI_MODEL".to_string()
} else {
registry
.find(&canonical_id)
.map(|d| d.model_env.clone())
.unwrap_or_default()
};
let mut vars: Vec<(&str, &str)> = vec![("LLM_BACKEND", &canonical_id)];
if !model_env.is_empty() {
vars.push((&model_env, &resolved_model));
}
sync_to_dotenv(config_path, &vars);
println!(
"Provider set to '{}', model set to '{}'",
canonical_id, resolved_model
);
println!(
"Saved to {}",
config_path
.map(|p| p.display().to_string())
.unwrap_or_else(|| config_toml_path().display().to_string())
);
Ok(())
}
// ─── list ─────────────────────────────────────────────────────────
/// List all providers with their default models.
async fn cmd_list_providers(
verbose: bool,
json: bool,
config_path: Option<&Path>,
) -> anyhow::Result<()> {
let registry = ProviderRegistry::load();
let (active_backend, active_model) = resolve_active(config_path);
if json {
let mut entries: Vec<serde_json::Value> = Vec::new();
// NEAR AI (not in registry)
let nearai_active = active_backend == "nearai";
entries.push(serde_json::json!({
"id": "nearai",
"description": "NEAR AI inference (default)",
"default_model": "qwen2.5-72b-instruct:free",
"active": nearai_active,
"active_model": if nearai_active { Some(&active_model) } else { None },
}));
for def in registry.all() {
let is_active = active_backend == def.id;
let mut v = serde_json::json!({
"id": def.id,
"description": def.description,
"default_model": def.default_model,
"protocol": format!("{:?}", def.protocol),
"active": is_active,
});
if is_active {
v["active_model"] = serde_json::json!(active_model);
}
if verbose {
v["aliases"] = serde_json::json!(def.aliases);
v["model_env"] = serde_json::json!(def.model_env);
v["api_key_env"] = serde_json::json!(def.api_key_env);
v["api_key_required"] = serde_json::json!(def.api_key_required);
if let Some(ref url) = def.default_base_url {
v["base_url"] = serde_json::json!(url);
}
if let Some(ref setup) = def.setup {
v["can_list_models"] = serde_json::json!(setup.can_list_models());
}
}
entries.push(v);
}
println!(
"{}",
serde_json::to_string_pretty(&entries).unwrap_or_else(|_| "[]".to_string())
);
return Ok(());
}
let providers = registry.all();
println!("Active: {} (model: {})\n", active_backend, active_model);
println!(
"{} provider(s) available:\n",
providers.len() + 1 // +1 for NEAR AI
);
// NEAR AI (not in registry)
let nearai_marker = if active_backend == "nearai" { " *" } else { "" };
if verbose {
println!(" nearai{}", nearai_marker);
println!(" Description: NEAR AI inference (default)");
println!(" Default model: qwen2.5-72b-instruct:free");
println!(" Model env: NEARAI_MODEL");
if active_backend == "nearai" {
println!(" Active model: {}", active_model);
}
println!();
} else {
println!(
" {:<22} {:<40} NEAR AI inference (default)",
format!("nearai{nearai_marker}"),
"qwen2.5-72b-instruct:free"
);
}
for def in providers {
let is_active = active_backend == def.id;
let marker = if is_active { " *" } else { "" };
if verbose {
println!(" {}{}", def.id, marker);
println!(" Description: {}", def.description);
println!(" Default model: {}", def.default_model);
println!(" Protocol: {:?}", def.protocol);
println!(" Model env: {}", def.model_env);
if let Some(ref env) = def.api_key_env {
println!(
" API key env: {} ({})",
env,
if def.api_key_required {
"required"
} else {
"optional"
}
);
}
if let Some(ref url) = def.default_base_url {
println!(" Base URL: {}", url);
}
if !def.aliases.is_empty() {
println!(" Aliases: {}", def.aliases.join(", "));
}
if is_active {
println!(" Active model: {}", active_model);
}
println!();
} else {
let model_display = if is_active {
active_model.clone()
} else {
def.default_model.clone()
};
println!(
" {:<22} {:<40} {}",
format!("{}{marker}", def.id),
model_display,
def.description,
);
}
}
if !verbose {
println!();
println!("* = active provider. Use --verbose for details.");
}
Ok(())
}
/// Show details for a specific provider.
async fn cmd_show_provider(
id: &str,
verbose: bool,
json: bool,
config_path: Option<&Path>,
) -> anyhow::Result<()> {
let registry = ProviderRegistry::load();
let (active_backend, active_model) = resolve_active(config_path);
// Resolve canonical ID for model fetching
let canonical_id = if id == "nearai" || id == "near_ai" || id == "near" {
"nearai".to_string()
} else {
registry
.find(id)
.map(|d| d.id.clone())
.unwrap_or_else(|| id.to_string())
};
// Try to fetch live model list from the provider
let live_models = try_fetch_models(&canonical_id, config_path).await;
// Check NEAR AI first (not in registry)
if id == "nearai" || id == "near_ai" || id == "near" {
let is_active = active_backend == "nearai";
if json {
let mut v = serde_json::json!({
"id": "nearai",
"description": "NEAR AI inference (default)",
"default_model": "qwen2.5-72b-instruct:free",
"model_env": "NEARAI_MODEL",
"active": is_active,
});
if is_active {
v["active_model"] = serde_json::json!(active_model);
}
if let Some(ref models) = live_models {
v["available_models"] = serde_json::json!(models);
}
println!(
"{}",
serde_json::to_string_pretty(&v).unwrap_or_else(|_| "{}".to_string())
);
} else {
println!("Provider: nearai");
println!(" Description: NEAR AI inference (default)");
println!(" Default model: qwen2.5-72b-instruct:free");
println!(" Model env: NEARAI_MODEL");
println!(" Active: {}", if is_active { "yes" } else { "no" });
if is_active {
println!(" Active model: {}", active_model);
}
print_model_list(&live_models, is_active.then_some(&active_model));
}
return Ok(());
}
let def = registry.find(id).ok_or_else(|| {
let known: Vec<&str> = std::iter::once("nearai")
.chain(registry.all().iter().map(|d| d.id.as_str()))
.collect();
anyhow::anyhow!(
"Unknown provider '{}'. Known providers: {}",
id,
known.join(", ")
)
})?;
let is_active = active_backend == def.id;
if json {
let mut v = serde_json::json!({
"id": def.id,
"description": def.description,
"protocol": format!("{:?}", def.protocol),
"default_model": def.default_model,
"model_env": def.model_env,
"api_key_env": def.api_key_env,
"api_key_required": def.api_key_required,
"aliases": def.aliases,
"active": is_active,
});
if let Some(ref url) = def.default_base_url {
v["base_url"] = serde_json::json!(url);
}
if let Some(ref setup) = def.setup {
v["can_list_models"] = serde_json::json!(setup.can_list_models());
v["display_name"] = serde_json::json!(setup.display_name());
}
if is_active {
v["active_model"] = serde_json::json!(active_model);
}
if verbose && !def.unsupported_params.is_empty() {
v["unsupported_params"] = serde_json::json!(def.unsupported_params);
}
if let Some(ref models) = live_models {
v["available_models"] = serde_json::json!(models);
}
println!(
"{}",
serde_json::to_string_pretty(&v).unwrap_or_else(|_| "{}".to_string())
);
return Ok(());
}
println!("Provider: {}", def.id);
println!(" Description: {}", def.description);
println!(" Protocol: {:?}", def.protocol);
println!(" Default model: {}", def.default_model);
println!(" Model env: {}", def.model_env);
if let Some(ref env) = def.api_key_env {
println!(
" API key env: {} ({})",
env,
if def.api_key_required {
"required"
} else {
"optional"
}
);
}
if let Some(ref url) = def.default_base_url {
println!(" Base URL: {}", url);
}
if !def.aliases.is_empty() {
println!(" Aliases: {}", def.aliases.join(", "));
}
if let Some(ref setup) = def.setup {
println!(
" List models: {}",
if setup.can_list_models() {
"supported"
} else {
"not supported"
}
);
println!(" Display name: {}", setup.display_name());
}
if !def.unsupported_params.is_empty() {
println!(" Unsupported: {}", def.unsupported_params.join(", "));
}
println!(" Active: {}", if is_active { "yes" } else { "no" });
if is_active {
println!(" Active model: {}", active_model);
}
print_model_list(&live_models, is_active.then_some(&active_model));
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn resolve_active_defaults_to_nearai() {
let settings = Settings::default();
assert!(settings.llm_backend.is_none());
assert!(settings.selected_model.is_none());
}
#[test]
fn registry_loads_all_providers() {
let registry = ProviderRegistry::load();
let all = registry.all();
assert!(
all.len() >= 10,
"should have at least 10 built-in providers, got {}",
all.len()
);
}
#[test]
fn registry_find_by_alias() {
let registry = ProviderRegistry::load();
let def = registry
.find("claude")
.expect("claude alias should resolve");
assert_eq!(def.id, "anthropic");
}
#[test]
fn all_providers_have_description() {
let registry = ProviderRegistry::load();
for def in registry.all() {
assert!(
!def.description.is_empty(),
"provider {} should have a description",
def.id
);
}
}
#[test]
fn set_model_persists_to_toml() {
let dir = tempfile::tempdir().expect("create temp dir");
let toml_path = dir.path().join("config.toml");
cmd_set_model("gpt-5-mini", Some(&toml_path)).expect("set model");
let settings = Settings::load_toml(&toml_path)
.expect("read toml")
.expect("should have settings");
assert_eq!(settings.selected_model.as_deref(), Some("gpt-5-mini"));
}
#[test]
fn set_provider_validates_unknown() {
let dir = tempfile::tempdir().expect("create temp dir");
let toml_path = dir.path().join("config.toml");
let result = cmd_set_provider("nonexistent_provider", None, Some(&toml_path));
assert!(result.is_err());
let err = result.unwrap_err().to_string();
assert!(
err.contains("Unknown provider"),
"should mention unknown provider: {}",
err
);
}
#[test]
fn set_provider_persists_to_toml() {
let dir = tempfile::tempdir().expect("create temp dir");
let toml_path = dir.path().join("config.toml");
cmd_set_provider("groq", None, Some(&toml_path)).expect("set provider");
let settings = Settings::load_toml(&toml_path)
.expect("read toml")
.expect("should have settings");
assert_eq!(settings.llm_backend.as_deref(), Some("groq"));
assert_eq!(
settings.selected_model.as_deref(),
Some("llama-3.3-70b-versatile")
);
}
#[test]
fn set_provider_with_custom_model() {
let dir = tempfile::tempdir().expect("create temp dir");
let toml_path = dir.path().join("config.toml");
cmd_set_provider("anthropic", Some("claude-opus-4-6"), Some(&toml_path))
.expect("set provider with model");
let settings = Settings::load_toml(&toml_path)
.expect("read toml")
.expect("should have settings");
assert_eq!(settings.llm_backend.as_deref(), Some("anthropic"));
assert_eq!(settings.selected_model.as_deref(), Some("claude-opus-4-6"));
}
#[test]
fn custom_config_does_not_pollute_default_dotenv() {
let dir = tempfile::tempdir().expect("create temp dir");
let toml_path = dir.path().join("config.toml");
// With a custom config path, sync_to_dotenv should be a no-op
// (it returns early when config_path is Some).
// We verify by checking that cmd_set_provider succeeds without
// trying to write to the default ~/.ironclaw/.env.
cmd_set_provider("groq", None, Some(&toml_path)).expect("set provider with custom config");
let settings = Settings::load_toml(&toml_path)
.expect("read toml")
.expect("should have settings");
assert_eq!(settings.llm_backend.as_deref(), Some("groq"));
// The key assertion is that no error was thrown trying to write
// to the default .env — sync_to_dotenv skipped it.
}
#[test]
fn set_model_rejects_empty_name() {
let dir = tempfile::tempdir().expect("create temp dir");
let toml_path = dir.path().join("config.toml");
let result = cmd_set_model("", Some(&toml_path));
assert!(result.is_err());
assert!(
result.unwrap_err().to_string().contains("cannot be empty"),
"should reject empty model name"
);
let result2 = cmd_set_model(" ", Some(&toml_path));
assert!(result2.is_err());
}
#[test]
fn set_provider_normalizes_alias() {
let dir = tempfile::tempdir().expect("create temp dir");
let toml_path = dir.path().join("config.toml");
cmd_set_provider("claude", None, Some(&toml_path)).expect("set via alias");
let settings = Settings::load_toml(&toml_path)
.expect("read toml")
.expect("should have settings");
assert_eq!(
settings.llm_backend.as_deref(),
Some("anthropic"),
"alias should be normalized to canonical ID"
);
}
}
+26 -402
View File
@@ -62,30 +62,6 @@ 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<String>,
builtin: Option<&OAuthCredentials>,
exchange_proxy_configured: bool,
) -> Option<String> {
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 ────────────────────────────────────────────── // ── Shared callback server ──────────────────────────────────────────────
// Core OAuth callback infrastructure is defined in `crate::llm::oauth_helpers` // Core OAuth callback infrastructure is defined in `crate::llm::oauth_helpers`
@@ -471,8 +447,8 @@ pub struct PendingOAuthFlow {
pub user_id: String, pub user_id: String,
/// Secrets store reference for token persistence. /// Secrets store reference for token persistence.
pub secrets: Arc<dyn SecretsStore + Send + Sync>, pub secrets: Arc<dyn SecretsStore + Send + Sync>,
/// SSE broadcast manager for notifying the web UI. /// SSE broadcast sender for notifying the web UI.
pub sse_manager: Option<Arc<crate::channels::web::sse::SseManager>>, pub sse_sender: Option<tokio::sync::broadcast::Sender<crate::channels::web::types::SseEvent>>,
/// Gateway auth token for authenticating with the platform token exchange proxy. /// Gateway auth token for authenticating with the platform token exchange proxy.
pub gateway_token: Option<String>, pub gateway_token: Option<String>,
/// Additional form params for the token exchange request. /// Additional form params for the token exchange request.
@@ -685,48 +661,6 @@ pub struct ProxyTokenExchangeRequest<'a> {
pub extra_token_params: &'a HashMap<String, String>, pub extra_token_params: &'a HashMap<String, String>,
} }
pub struct ProxyRefreshTokenRequest<'a> {
pub proxy_url: &'a str,
pub gateway_token: &'a str,
pub token_url: &'a str,
pub client_id: &'a str,
pub client_secret: Option<&'a str>,
pub refresh_token: &'a str,
pub provider: Option<&'a str>,
}
fn oauth_token_response_from_json(
token_data: serde_json::Value,
access_token_field: &str,
) -> Result<OAuthTokenResponse, OAuthCallbackError> {
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. /// 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 the gateway auth token (Bearer header). The caller may
@@ -748,7 +682,6 @@ pub async fn exchange_via_proxy(
let client = reqwest::Client::builder() let client = reqwest::Client::builder()
.timeout(Duration::from_secs(60)) .timeout(Duration::from_secs(60))
.redirect(reqwest::redirect::Policy::none())
.build() .build()
.map_err(|e| OAuthCallbackError::Io(format!("Failed to build HTTP client: {}", e)))?; .map_err(|e| OAuthCallbackError::Io(format!("Failed to build HTTP client: {}", e)))?;
let mut params = vec![ let mut params = vec![
@@ -791,350 +724,41 @@ pub async fn exchange_via_proxy(
.json() .json()
.await .await
.map_err(|e| OAuthCallbackError::Io(format!("Failed to parse proxy response: {}", e)))?; .map_err(|e| OAuthCallbackError::Io(format!("Failed to parse proxy response: {}", e)))?;
oauth_token_response_from_json(token_data, request.access_token_field)
}
/// Refresh an OAuth access token via the platform's token refresh proxy. let access_token = token_data
/// .get(request.access_token_field)
/// Authenticated via the gateway auth token (Bearer header). The caller may .and_then(|v| v.as_str())
/// either rely on proxy-side secret lookup or forward a `client_secret` when .ok_or_else(|| {
/// the provider requires it. let fields: Vec<&str> = token_data
pub async fn refresh_token_via_proxy( .as_object()
request: ProxyRefreshTokenRequest<'_>, .map(|o| o.keys().map(|k| k.as_str()).collect())
) -> Result<OAuthTokenResponse, OAuthCallbackError> { .unwrap_or_default();
if request.gateway_token.is_empty() { OAuthCallbackError::Io(format!(
return Err(OAuthCallbackError::Io( "No '{}' field in proxy response (fields present: {:?})",
"Gateway auth token is required for proxy token refresh".to_string(), request.access_token_field, fields
)); ))
} })?
.to_string();
let refresh_url = format!("{}/oauth/refresh", request.proxy_url.trim_end_matches('/')); let refresh_token = token_data
let client = reqwest::Client::builder() .get("refresh_token")
.timeout(Duration::from_secs(15)) .and_then(|v| v.as_str())
.redirect(reqwest::redirect::Policy::none()) .map(String::from);
.build() let expires_in = token_data.get("expires_in").and_then(|v| v.as_u64());
.map_err(|e| OAuthCallbackError::Io(format!("Failed to build HTTP client: {}", e)))?;
let mut params = vec![ Ok(OAuthTokenResponse {
("refresh_token", request.refresh_token.to_string()), access_token,
("token_url", request.token_url.to_string()), refresh_token,
("client_id", request.client_id.to_string()), expires_in,
]; })
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(&params)
.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)] #[cfg(test)]
mod tests { 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::{ use crate::cli::oauth_defaults::{
builtin_credentials, callback_host, callback_url, is_loopback_host, landing_html, builtin_credentials, callback_host, callback_url, is_loopback_host, landing_html,
}; };
use crate::config::helpers::lock_env; 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<String>,
form: HashMap<String, String>,
}
#[derive(Clone)]
struct MockProxyState {
requests: Arc<Mutex<Vec<RecordedProxyRequest>>>,
exchange_redirect_target: String,
refresh_redirect_target: String,
}
struct MockProxyServer {
addr: SocketAddr,
requests: Arc<Mutex<Vec<RecordedProxyRequest>>>,
shutdown_tx: Option<oneshot::Sender<()>>,
server_task: Option<tokio::task::JoinHandle<()>>,
}
impl MockProxyServer {
async fn start() -> Self {
async fn exchange_handler(
State(state): State<MockProxyState>,
headers: HeaderMap,
Form(form): Form<HashMap<String, String>>,
) -> Json<serde_json::Value> {
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<MockProxyState>,
headers: HeaderMap,
Form(form): Form<HashMap<String, String>>,
) -> Json<serde_json::Value> {
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<MockProxyState>) -> Redirect {
Redirect::temporary(&state.exchange_redirect_target)
}
async fn refresh_redirect_handler(State(state): State<MockProxyState>) -> 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<RecordedProxyRequest> {
self.requests.lock().await.clone()
}
async fn shutdown(mut self) {
if let Some(tx) = self.shutdown_tx.take() {
let _ = tx.send(());
}
if let Some(task) = self.server_task.take() {
let _ = task.await;
}
}
}
impl Drop for MockProxyServer {
fn drop(&mut self) {
if let Some(tx) = self.shutdown_tx.take() {
let _ = tx.send(());
}
if let Some(task) = self.server_task.take() {
task.abort();
}
}
}
#[test]
fn test_hosted_proxy_client_secret_suppresses_builtin_secret() {
let builtin = builtin_credentials("google_oauth_token").expect("google builtin creds");
let client_secret = Some(builtin.client_secret.to_string());
let result = super::hosted_proxy_client_secret(&client_secret, Some(&builtin), true);
assert_eq!(result, None);
}
#[test]
fn test_hosted_proxy_client_secret_preserves_explicit_secret() {
let builtin = builtin_credentials("google_oauth_token").expect("google builtin creds");
let client_secret = Some("hosted-server-secret".to_string());
let result = super::hosted_proxy_client_secret(&client_secret, Some(&builtin), true);
assert_eq!(result, client_secret);
}
#[tokio::test]
async fn test_refresh_token_via_proxy_sends_auth_and_form() {
let server = MockProxyServer::start().await;
let response = super::refresh_token_via_proxy(super::ProxyRefreshTokenRequest {
proxy_url: &server.base_url(),
gateway_token: "gateway-test-token",
token_url: "https://oauth2.googleapis.com/token",
client_id: TEST_OAUTH_CLIENT_ID,
client_secret: Some(TEST_OAUTH_CLIENT_SECRET),
refresh_token: "refresh-token-123",
provider: Some("google"),
})
.await
.expect("proxy refresh succeeds");
assert_eq!(response.access_token, "proxy-access-token");
assert_eq!(
response.refresh_token.as_deref(),
Some("proxy-refresh-token")
);
assert_eq!(response.expires_in, Some(7200));
let requests = server.requests().await;
assert_eq!(requests.len(), 1);
assert_eq!(
requests[0].authorization.as_deref(),
Some("Bearer gateway-test-token")
);
assert_eq!(
requests[0].form.get("token_url").map(String::as_str),
Some("https://oauth2.googleapis.com/token")
);
assert_eq!(
requests[0].form.get("client_id").map(String::as_str),
Some(TEST_OAUTH_CLIENT_ID)
);
assert_eq!(
requests[0].form.get("client_secret").map(String::as_str),
Some(TEST_OAUTH_CLIENT_SECRET)
);
assert_eq!(
requests[0].form.get("refresh_token").map(String::as_str),
Some("refresh-token-123")
);
assert_eq!(
requests[0].form.get("provider").map(String::as_str),
Some("google")
);
server.shutdown().await;
}
#[tokio::test]
async fn test_exchange_via_proxy_does_not_follow_redirects() {
let server = MockProxyServer::start().await;
let error = match super::exchange_via_proxy(super::ProxyTokenExchangeRequest {
proxy_url: &server.redirecting_base_url(),
gateway_token: "gateway-test-token",
code: "auth-code-123",
redirect_uri: "http://localhost:3000/oauth/callback",
token_url: "https://oauth2.googleapis.com/token",
client_id: TEST_OAUTH_CLIENT_ID,
client_secret: Some(TEST_OAUTH_CLIENT_SECRET),
access_token_field: "access_token",
code_verifier: Some("code-verifier-123"),
extra_token_params: &HashMap::new(),
})
.await
{
Ok(_) => panic!("redirected proxy exchange should fail"),
Err(error) => error,
};
assert!(error.to_string().contains("307"));
assert!(server.requests().await.is_empty());
server.shutdown().await;
}
#[tokio::test]
async fn test_refresh_token_via_proxy_does_not_follow_redirects() {
let server = MockProxyServer::start().await;
let error = match super::refresh_token_via_proxy(super::ProxyRefreshTokenRequest {
proxy_url: &server.redirecting_base_url(),
gateway_token: "gateway-test-token",
token_url: "https://oauth2.googleapis.com/token",
client_id: TEST_OAUTH_CLIENT_ID,
client_secret: Some(TEST_OAUTH_CLIENT_SECRET),
refresh_token: "refresh-token-123",
provider: Some("google"),
})
.await
{
Ok(_) => panic!("redirected proxy refresh should fail"),
Err(error) => error,
};
assert!(error.to_string().contains("307"));
assert!(server.requests().await.is_empty());
server.shutdown().await;
}
#[test] #[test]
fn test_is_loopback_host() { fn test_is_loopback_host() {
+10 -66
View File
@@ -10,7 +10,7 @@ use clap::Subcommand;
use uuid::Uuid; use uuid::Uuid;
use crate::agent::routine::{ use crate::agent::routine::{
NotifyConfig, Routine, RoutineAction, RoutineGuardrails, RunStatus, Trigger, next_cron_fire, NotifyConfig, Routine, RoutineAction, RoutineGuardrails, Trigger, next_cron_fire,
}; };
use crate::db::Database; use crate::db::Database;
@@ -251,26 +251,15 @@ async fn list(
); );
println!("{}", "-".repeat(130)); println!("{}", "-".repeat(130));
// Fetch last-run status for all routines in a single batch query
let routine_ids: Vec<Uuid> = filtered.iter().map(|r| r.id).collect();
let last_run_results = db
.batch_get_last_run_status(&routine_ids)
.await
.unwrap_or_default();
for r in &filtered { for r in &filtered {
let last_run_status = last_run_results.get(&r.id).copied(); let status = if r.enabled {
if r.consecutive_failures > 0 {
let status = if !r.enabled { format!("err({})", r.consecutive_failures)
"disabled".to_string() } else {
} else if last_run_status == Some(RunStatus::Running) { "active".to_string()
"running".to_string() }
} else if r.consecutive_failures > 0 {
format!("err({})", r.consecutive_failures)
} else if last_run_status == Some(RunStatus::Attention) {
"attention".to_string()
} else { } else {
"active".to_string() "disabled".to_string()
}; };
let next_fire = r let next_fire = r
@@ -351,8 +340,8 @@ async fn create(
prompt: prompt.to_string(), prompt: prompt.to_string(),
context_paths: Vec::new(), context_paths: Vec::new(),
max_tokens: 4096, max_tokens: 4096,
use_tools: true, use_tools: false,
max_tool_rounds: 3, max_tool_rounds: 0,
}, },
guardrails: RoutineGuardrails { guardrails: RoutineGuardrails {
cooldown: std::time::Duration::from_secs(cooldown_secs), cooldown: std::time::Duration::from_secs(cooldown_secs),
@@ -696,7 +685,6 @@ fn truncate(s: &str, max_chars: usize) -> String {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use crate::agent::routine::RoutineAction;
#[test] #[test]
fn format_relative_future() { fn format_relative_future() {
@@ -755,48 +743,4 @@ mod tests {
assert!(notify.on_failure); // safety: test-only assertion assert!(notify.on_failure); // safety: test-only assertion
assert!(!notify.on_success); // safety: test-only assertion assert!(!notify.on_success); // safety: test-only assertion
} }
#[cfg(feature = "libsql")]
#[tokio::test]
async fn cli_create_defaults_lightweight_routines_to_tools_enabled() {
let harness = crate::testing::TestHarnessBuilder::new().build().await;
let db = harness.db.clone();
run_routines_command(
RoutinesCommand::Create {
name: "cli-digest".to_string(),
schedule: "0 0 9 * * *".to_string(),
prompt: "Prepare the morning digest.".to_string(),
description: "CLI created routine".to_string(),
timezone: Some("UTC".to_string()),
cooldown: 300,
notify_channel: None,
},
db.clone(),
"user1",
)
.await
.expect("create routine");
let routine = db
.get_routine_by_name("user1", "cli-digest")
.await
.expect("get routine by name")
.expect("cli-digest should exist");
match routine.action {
RoutineAction::Lightweight {
use_tools,
max_tool_rounds,
..
} => {
assert!(
use_tools,
"CLI-created lightweight routines should default to tools"
);
assert_eq!(max_tool_rounds, 3);
}
other => panic!("expected lightweight action, got {other:?}"),
}
}
} }
@@ -20,7 +20,6 @@ Commands:
service Manage OS service service Manage OS service
skills Manage skills skills Manage skills
hooks Manage lifecycle hooks hooks Manage lifecycle hooks
models Manage LLM providers and models
doctor Run diagnostics doctor Run diagnostics
logs View and manage gateway logs logs View and manage gateway logs
status Show system status status Show system status
@@ -20,7 +20,6 @@ Commands:
service Manage OS service service Manage OS service
skills Manage skills skills Manage skills
hooks Manage lifecycle hooks hooks Manage lifecycle hooks
models Manage LLM providers and models
doctor Run diagnostics doctor Run diagnostics
logs View and manage gateway logs logs View and manage gateway logs
status Show system status status Show system status
@@ -23,7 +23,6 @@ Commands:
service Manage OS service service Manage OS service
skills Manage skills skills Manage skills
hooks Manage lifecycle hooks hooks Manage lifecycle hooks
models Manage LLM providers and models
doctor Run diagnostics doctor Run diagnostics
logs View and manage gateway logs logs View and manage gateway logs
status Show system status status Show system status
@@ -23,7 +23,6 @@ Commands:
service Manage OS service service Manage OS service
skills Manage skills skills Manage skills
hooks Manage lifecycle hooks hooks Manage lifecycle hooks
models Manage LLM providers and models
doctor Run diagnostics doctor Run diagnostics
logs View and manage gateway logs logs View and manage gateway logs
status Show system status status Show system status
+16 -270
View File
@@ -2,7 +2,6 @@
//! //!
//! Commands for installing, listing, removing, and authenticating WASM tools. //! Commands for installing, listing, removing, and authenticating WASM tools.
use std::collections::{HashMap, HashSet};
use std::io::Write; use std::io::Write;
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
use std::sync::Arc; use std::sync::Arc;
@@ -80,10 +79,6 @@ pub enum ToolCommand {
/// Directory to look for tool (default: ~/.ironclaw/tools/) /// Directory to look for tool (default: ~/.ironclaw/tools/)
#[arg(short, long)] #[arg(short, long)]
dir: Option<PathBuf>, dir: Option<PathBuf>,
/// User ID for checking credential status (default: "default")
#[arg(short, long, default_value = "default")]
user: String,
}, },
/// Configure authentication for a tool /// Configure authentication for a tool
@@ -129,11 +124,7 @@ pub async fn run_tool_command(cmd: ToolCommand) -> anyhow::Result<()> {
} => install_tool(path, name, capabilities, target, release, skip_build, force).await, } => install_tool(path, name, capabilities, target, release, skip_build, force).await,
ToolCommand::List { dir, verbose } => list_tools(dir, verbose).await, ToolCommand::List { dir, verbose } => list_tools(dir, verbose).await,
ToolCommand::Remove { name, dir } => remove_tool(name, dir).await, ToolCommand::Remove { name, dir } => remove_tool(name, dir).await,
ToolCommand::Info { ToolCommand::Info { name_or_path, dir } => show_tool_info(name_or_path, dir).await,
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::Auth { name, dir, user } => auth_tool(name, dir, user).await,
ToolCommand::Setup { name, dir, user } => setup_tool(name, dir, user).await, ToolCommand::Setup { name, dir, user } => setup_tool(name, dir, user).await,
} }
@@ -397,11 +388,7 @@ async fn remove_tool(name: String, dir: Option<PathBuf>) -> anyhow::Result<()> {
} }
/// Show information about a tool. /// Show information about a tool.
async fn show_tool_info( async fn show_tool_info(name_or_path: String, dir: Option<PathBuf>) -> anyhow::Result<()> {
name_or_path: String,
dir: Option<PathBuf>,
user_id: String,
) -> anyhow::Result<()> {
let wasm_path = if name_or_path.ends_with(".wasm") { let wasm_path = if name_or_path.ends_with(".wasm") {
PathBuf::from(&name_or_path) PathBuf::from(&name_or_path)
} else { } else {
@@ -436,37 +423,7 @@ async fn show_tool_info(
println!("\nCapabilities ({}):", caps_path.display()); println!("\nCapabilities ({}):", caps_path.display());
let content = fs::read_to_string(&caps_path).await?; let content = fs::read_to_string(&caps_path).await?;
match CapabilitiesFile::from_json(&content) { match CapabilitiesFile::from_json(&content) {
Ok(caps) => { Ok(caps) => print_capabilities_detail(&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), Err(e) => println!(" Error parsing: {}", e),
} }
} else { } else {
@@ -519,89 +476,8 @@ 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<String>,
/// Injection location (from http.credentials).
location: Option<String>,
}
/// Collected auth secrets and the set of secret names they cover.
struct CollectedAuthSecrets {
secrets: Vec<AuthSecretInfo>,
/// Secret names present in `secrets`, for filtering the Secrets capability section.
seen_names: HashSet<String>,
}
/// 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<AuthSecretInfo> = Vec::new();
let mut seen: HashMap<String, usize> = 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. /// Print detailed capabilities.
async fn print_capabilities_detail( fn print_capabilities_detail(caps: &CapabilitiesFile) {
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 { if let Some(ref http) = caps.http {
println!(" HTTP:"); println!(" HTTP:");
for endpoint in &http.allowlist { for endpoint in &http.allowlist {
@@ -614,6 +490,13 @@ async fn print_capabilities_detail(
println!(" {} {} {}", methods, endpoint.host, path); 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 { if let Some(ref rate) = http.rate_limit {
println!( println!(
" Rate limit: {}/min, {}/hour", " Rate limit: {}/min, {}/hour",
@@ -622,24 +505,12 @@ async fn print_capabilities_detail(
} }
} }
// Filter secrets already covered by the auth section (always rendered when non-empty).
if let Some(ref secrets) = caps.secrets if let Some(ref secrets) = caps.secrets
&& !secrets.allowed_names.is_empty() && !secrets.allowed_names.is_empty()
{ {
let extra: Vec<_> = if collected.secrets.is_empty() { println!(" Secrets (existence check only):");
secrets.allowed_names.iter().collect() for name in &secrets.allowed_names {
} else { println!(" {}", name);
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);
}
} }
} }
@@ -660,38 +531,6 @@ async fn print_capabilities_detail(
println!(" {}", prefix); 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. /// Validate a tool name to prevent path traversal.
@@ -838,7 +677,8 @@ async fn combine_provider_scopes(
secret_name: &str, secret_name: &str,
base_oauth: &crate::tools::wasm::OAuthConfigSchema, base_oauth: &crate::tools::wasm::OAuthConfigSchema,
) -> crate::tools::wasm::OAuthConfigSchema { ) -> crate::tools::wasm::OAuthConfigSchema {
let mut all_scopes: HashSet<String> = base_oauth.scopes.iter().cloned().collect(); let mut all_scopes: std::collections::HashSet<String> =
base_oauth.scopes.iter().cloned().collect();
if let Ok(mut entries) = tokio::fs::read_dir(tools_dir).await { if let Ok(mut entries) = tokio::fs::read_dir(tools_dir).await {
while let Ok(Some(entry)) = entries.next_entry().await { while let Ok(Some(entry)) = entries.next_entry().await {
@@ -1287,8 +1127,6 @@ async fn setup_tool(name: String, dir: Option<PathBuf>, user_id: String) -> anyh
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use crate::secrets::{CreateSecretParams, SecretsStore};
use crate::testing::credentials::test_secrets_store;
#[test] #[test]
fn test_format_size() { fn test_format_size() {
@@ -1305,96 +1143,4 @@ mod tests {
assert!(dir.to_string_lossy().contains(".ironclaw")); assert!(dir.to_string_lossy().contains(".ironclaw"));
assert!(dir.to_string_lossy().contains("tools")); 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());
}
} }
-24
View File
@@ -23,26 +23,14 @@ pub struct AgentConfig {
pub max_cost_per_day_cents: Option<u64>, pub max_cost_per_day_cents: Option<u64>,
/// Maximum LLM/tool actions per hour. None = unlimited. /// Maximum LLM/tool actions per hour. None = unlimited.
pub max_actions_per_hour: Option<u64>, pub max_actions_per_hour: Option<u64>,
/// Maximum daily LLM spend per user in cents. None = unlimited.
pub max_cost_per_user_per_day_cents: Option<u64>,
/// Maximum tool-call iterations per agentic loop invocation. Default 50. /// Maximum tool-call iterations per agentic loop invocation. Default 50.
pub max_tool_iterations: usize, pub max_tool_iterations: usize,
/// When true, skip tool approval checks entirely. For benchmarks/CI. /// When true, skip tool approval checks entirely. For benchmarks/CI.
pub auto_approve_tools: bool, pub auto_approve_tools: bool,
/// Default timezone for new sessions (IANA name, e.g. "America/New_York"). /// Default timezone for new sessions (IANA name, e.g. "America/New_York").
pub default_timezone: String, pub default_timezone: String,
/// Maximum concurrent jobs per user. None = use global max_parallel_jobs.
pub max_jobs_per_user: Option<usize>,
/// Maximum tokens per job (0 = unlimited). /// Maximum tokens per job (0 = unlimited).
pub max_tokens_per_job: u64, pub max_tokens_per_job: u64,
/// Whether the deployment is multi-tenant (multiple users sharing one
/// instance). Detected at runtime after DB initialization, not from config.
/// See app.rs startup logic.
pub multi_tenant: bool,
/// Maximum concurrent LLM calls per user. None = use default (4).
pub max_llm_concurrent_per_user: Option<usize>,
/// Maximum concurrent jobs per user. None = use default (3).
pub max_jobs_concurrent_per_user: Option<usize>,
} }
impl AgentConfig { impl AgentConfig {
@@ -61,15 +49,10 @@ impl AgentConfig {
allow_local_tools: true, allow_local_tools: true,
max_cost_per_day_cents: None, max_cost_per_day_cents: None,
max_actions_per_hour: None, max_actions_per_hour: None,
max_cost_per_user_per_day_cents: None,
max_tool_iterations: 10, max_tool_iterations: 10,
auto_approve_tools: true, auto_approve_tools: true,
default_timezone: "UTC".to_string(), default_timezone: "UTC".to_string(),
max_jobs_per_user: None,
max_tokens_per_job: 0, max_tokens_per_job: 0,
multi_tenant: false,
max_llm_concurrent_per_user: None,
max_jobs_concurrent_per_user: None,
} }
} }
@@ -104,7 +87,6 @@ impl AgentConfig {
allow_local_tools: parse_bool_env("ALLOW_LOCAL_TOOLS", false)?, allow_local_tools: parse_bool_env("ALLOW_LOCAL_TOOLS", false)?,
max_cost_per_day_cents: parse_option_env("MAX_COST_PER_DAY_CENTS")?, 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_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( max_tool_iterations: parse_optional_env(
"AGENT_MAX_TOOL_ITERATIONS", "AGENT_MAX_TOOL_ITERATIONS",
settings.agent.max_tool_iterations, settings.agent.max_tool_iterations,
@@ -126,16 +108,10 @@ impl AgentConfig {
} }
tz tz
}, },
max_jobs_per_user: parse_option_env("MAX_JOBS_PER_USER")?,
max_tokens_per_job: parse_optional_env( max_tokens_per_job: parse_optional_env(
"AGENT_MAX_TOKENS_PER_JOB", "AGENT_MAX_TOKENS_PER_JOB",
settings.agent.max_tokens_per_job, settings.agent.max_tokens_per_job,
)?, )?,
// Multi-tenant mode is detected at runtime after DB initialization,
// not from config. See app.rs startup logic.
multi_tenant: false,
max_llm_concurrent_per_user: parse_option_env("TENANT_MAX_LLM_CONCURRENT")?,
max_jobs_concurrent_per_user: parse_option_env("TENANT_MAX_JOBS_CONCURRENT")?,
}) })
} }
} }
+11 -91
View File
@@ -1,11 +1,12 @@
use std::collections::HashMap; use std::collections::HashMap;
use std::path::PathBuf; use std::path::PathBuf;
use secrecy::SecretString;
use crate::bootstrap::ironclaw_base_dir; use crate::bootstrap::ironclaw_base_dir;
use crate::config::helpers::{optional_env, parse_bool_env, parse_optional_env}; use crate::config::helpers::{optional_env, parse_bool_env, parse_optional_env};
use crate::error::ConfigError; use crate::error::ConfigError;
use crate::settings::Settings; use crate::settings::Settings;
use secrecy::SecretString;
/// Channel configurations. /// Channel configurations.
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
@@ -43,14 +44,7 @@ pub struct GatewayConfig {
pub port: u16, pub port: u16,
/// Bearer token for authentication. Random hex generated at startup if unset. /// Bearer token for authentication. Random hex generated at startup if unset.
pub auth_token: Option<String>, pub auth_token: Option<String>,
/// Additional user scopes for workspace reads. pub user_id: String,
///
/// When set, the workspace will be able to read (search, read, list) from
/// these additional user scopes while writes remain isolated to `user_id`.
/// Parsed from `WORKSPACE_READ_SCOPES` (comma-separated).
pub workspace_read_scopes: Vec<String>,
/// Memory layer definitions (JSON in env var, or from external config).
pub memory_layers: Vec<crate::workspace::layer::MemoryLayer>,
} }
/// Signal channel configuration (signal-cli daemon HTTP/JSON-RPC). /// Signal channel configuration (signal-cli daemon HTTP/JSON-RPC).
@@ -117,83 +111,10 @@ impl ChannelsConfig {
let gateway_enabled = parse_bool_env("GATEWAY_ENABLED", cs.gateway_enabled)?; let gateway_enabled = parse_bool_env("GATEWAY_ENABLED", cs.gateway_enabled)?;
let gateway = if gateway_enabled { let gateway = if gateway_enabled {
let memory_layers: Vec<crate::workspace::layer::MemoryLayer> = let user_id = optional_env("GATEWAY_USER_ID")?
match optional_env("MEMORY_LAYERS")? { .or_else(|| cs.gateway_user_id.clone())
Some(json_str) => { .unwrap_or_else(|| owner_id.to_string());
serde_json::from_str(&json_str).map_err(|e| ConfigError::InvalidValue {
key: "MEMORY_LAYERS".to_string(),
message: format!("must be valid JSON array of layer objects: {e}"),
})?
}
None => crate::workspace::layer::MemoryLayer::default_for_user(owner_id),
};
// Validate layer names and scopes
for layer in &memory_layers {
if layer.name.trim().is_empty() {
return Err(ConfigError::InvalidValue {
key: "MEMORY_LAYERS".to_string(),
message: "layer name must not be empty".to_string(),
});
}
if layer.name.len() > 64 {
return Err(ConfigError::InvalidValue {
key: "MEMORY_LAYERS".to_string(),
message: format!("layer name '{}' exceeds 64 characters", layer.name),
});
}
if !layer
.name
.chars()
.all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '-')
{
return Err(ConfigError::InvalidValue {
key: "MEMORY_LAYERS".to_string(),
message: format!(
"layer name '{}' contains invalid characters \
(allowed: a-z, A-Z, 0-9, _, -)",
layer.name
),
});
}
if layer.scope.trim().is_empty() {
return Err(ConfigError::InvalidValue {
key: "MEMORY_LAYERS".to_string(),
message: format!("layer '{}' has an empty scope", layer.name),
});
}
}
// Check for duplicate layer names
{
let mut seen = std::collections::HashSet::new();
for layer in &memory_layers {
if !seen.insert(&layer.name) {
return Err(ConfigError::InvalidValue {
key: "MEMORY_LAYERS".to_string(),
message: format!("duplicate layer name '{}'", layer.name),
});
}
}
}
let workspace_read_scopes: Vec<String> = optional_env("WORKSPACE_READ_SCOPES")?
.map(|s| {
s.split(',')
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.collect()
})
.unwrap_or_default();
for scope in &workspace_read_scopes {
if scope.len() > 128 {
return Err(ConfigError::InvalidValue {
key: "WORKSPACE_READ_SCOPES".to_string(),
message: format!("scope '{}...' exceeds 128 characters", &scope[..32]),
});
}
}
Some(GatewayConfig { Some(GatewayConfig {
host: optional_env("GATEWAY_HOST")? host: optional_env("GATEWAY_HOST")?
.or_else(|| cs.gateway_host.clone()) .or_else(|| cs.gateway_host.clone())
@@ -204,8 +125,7 @@ impl ChannelsConfig {
)?, )?,
auth_token: optional_env("GATEWAY_AUTH_TOKEN")? auth_token: optional_env("GATEWAY_AUTH_TOKEN")?
.or_else(|| cs.gateway_auth_token.clone()), .or_else(|| cs.gateway_auth_token.clone()),
workspace_read_scopes, user_id,
memory_layers,
}) })
} else { } else {
None None
@@ -360,12 +280,12 @@ mod tests {
host: "127.0.0.1".to_string(), host: "127.0.0.1".to_string(),
port: 3000, port: 3000,
auth_token: Some("tok-abc".to_string()), auth_token: Some("tok-abc".to_string()),
workspace_read_scopes: vec![], user_id: "default".to_string(),
memory_layers: vec![],
}; };
assert_eq!(cfg.host, "127.0.0.1"); assert_eq!(cfg.host, "127.0.0.1");
assert_eq!(cfg.port, 3000); assert_eq!(cfg.port, 3000);
assert_eq!(cfg.auth_token.as_deref(), Some("tok-abc")); assert_eq!(cfg.auth_token.as_deref(), Some("tok-abc"));
assert_eq!(cfg.user_id, "default");
} }
#[test] #[test]
@@ -374,8 +294,7 @@ mod tests {
host: "0.0.0.0".to_string(), host: "0.0.0.0".to_string(),
port: 3001, port: 3001,
auth_token: None, auth_token: None,
workspace_read_scopes: vec![], user_id: "anon".to_string(),
memory_layers: vec![],
}; };
assert!(cfg.auth_token.is_none()); assert!(cfg.auth_token.is_none());
} }
@@ -502,6 +421,7 @@ mod tests {
assert_eq!(gateway.host, "127.0.0.3"); assert_eq!(gateway.host, "127.0.0.3");
assert_eq!(gateway.port, 9191); assert_eq!(gateway.port, 9191);
assert_eq!(gateway.auth_token.as_deref(), Some("tok")); assert_eq!(gateway.auth_token.as_deref(), Some("tok"));
assert_eq!(gateway.user_id, "owner-scope");
let signal = cfg.signal.expect("signal config"); let signal = cfg.signal.expect("signal config");
assert_eq!(signal.account, "+15551234567"); assert_eq!(signal.account, "+15551234567");
-5
View File
@@ -21,9 +21,6 @@ pub struct HeartbeatConfig {
pub quiet_hours_end: Option<u32>, pub quiet_hours_end: Option<u32>,
/// Timezone for fire_at and quiet hours evaluation (IANA name). /// Timezone for fire_at and quiet hours evaluation (IANA name).
pub timezone: Option<String>, pub timezone: Option<String>,
/// When true, cycle through all users with routines. Set explicitly via
/// HEARTBEAT_MULTI_TENANT or detected at runtime after DB initialization.
pub multi_tenant: bool,
} }
impl Default for HeartbeatConfig { impl Default for HeartbeatConfig {
@@ -37,7 +34,6 @@ impl Default for HeartbeatConfig {
quiet_hours_start: None, quiet_hours_start: None,
quiet_hours_end: None, quiet_hours_end: None,
timezone: None, timezone: None,
multi_tenant: false,
} }
} }
} }
@@ -105,7 +101,6 @@ impl HeartbeatConfig {
} }
tz tz
}, },
multi_tenant: parse_bool_env("HEARTBEAT_MULTI_TENANT", false)?,
}) })
} }
} }
+4 -12
View File
@@ -406,7 +406,7 @@ impl LlmConfig {
// Resolve extra headers // Resolve extra headers
let extra_headers = if let Some(env_var) = extra_headers_env { let extra_headers = if let Some(env_var) = extra_headers_env {
optional_env(env_var)? optional_env(env_var)?
.map(|val| parse_extra_headers_with_key(&val, env_var)) .map(|val| parse_extra_headers(&val))
.transpose()? .transpose()?
.unwrap_or_default() .unwrap_or_default()
} else { } else {
@@ -475,10 +475,7 @@ impl LlmConfig {
/// ///
/// Format: `Key1:Value1,Key2:Value2` (colon-separated, not `=`, because /// Format: `Key1:Value1,Key2:Value2` (colon-separated, not `=`, because
/// header values often contain `=`). /// header values often contain `=`).
fn parse_extra_headers_with_key( fn parse_extra_headers(val: &str) -> Result<Vec<(String, String)>, ConfigError> {
val: &str,
env_var_name: &str,
) -> Result<Vec<(String, String)>, ConfigError> {
if val.trim().is_empty() { if val.trim().is_empty() {
return Ok(Vec::new()); return Ok(Vec::new());
} }
@@ -491,14 +488,14 @@ fn parse_extra_headers_with_key(
} }
let Some((key, value)) = pair.split_once(':') else { let Some((key, value)) = pair.split_once(':') else {
return Err(ConfigError::InvalidValue { return Err(ConfigError::InvalidValue {
key: env_var_name.to_string(), key: "LLM_EXTRA_HEADERS".to_string(),
message: format!("malformed header entry '{}', expected Key:Value", pair), message: format!("malformed header entry '{}', expected Key:Value", pair),
}); });
}; };
let key = key.trim(); let key = key.trim();
if key.is_empty() { if key.is_empty() {
return Err(ConfigError::InvalidValue { return Err(ConfigError::InvalidValue {
key: env_var_name.to_string(), key: "LLM_EXTRA_HEADERS".to_string(),
message: format!("empty header name in entry '{}'", pair), message: format!("empty header name in entry '{}'", pair),
}); });
} }
@@ -539,11 +536,6 @@ mod tests {
use crate::settings::Settings; use crate::settings::Settings;
use crate::testing::credentials::*; use crate::testing::credentials::*;
/// Convenience wrapper for tests — uses "TEST_HEADERS" as the env var name.
fn parse_extra_headers(val: &str) -> Result<Vec<(String, String)>, ConfigError> {
parse_extra_headers_with_key(val, "TEST_HEADERS")
}
/// Clear all openai-compatible-related env vars. /// Clear all openai-compatible-related env vars.
fn clear_openai_compatible_env() { fn clear_openai_compatible_env() {
// SAFETY: Only called under ENV_MUTEX in tests. // SAFETY: Only called under ENV_MUTEX in tests.
+7 -5
View File
@@ -312,11 +312,13 @@ impl Config {
let tunnel = TunnelConfig::resolve(settings)?; let tunnel = TunnelConfig::resolve(settings)?;
let channels = ChannelsConfig::resolve(settings, &owner_id)?; let channels = ChannelsConfig::resolve(settings, &owner_id)?;
// Resolve the startup workspace against the durable owner scope. The // Resolve workspace config using the gateway user_id for default layers.
// gateway may expose a distinct sender identity, but the base runtime let workspace_user_id = channels
// workspace stays owner-scoped and per-user gateway workspaces are .gateway
// handled separately by WorkspacePool. .as_ref()
let workspace = WorkspaceConfig::resolve(&owner_id)?; .map(|gw| gw.user_id.as_str())
.unwrap_or("default");
let workspace = WorkspaceConfig::resolve(workspace_user_id)?;
Ok(Self { Ok(Self {
owner_id: owner_id.clone(), owner_id: owner_id.clone(),
-69
View File
@@ -230,49 +230,6 @@ impl JobStore for LibSqlBackend {
Ok(jobs) Ok(jobs)
} }
async fn list_agent_jobs_for_user(
&self,
user_id: &str,
) -> Result<Vec<AgentJobRecord>, DatabaseError> {
let conn = self.connect().await?;
let mut rows = conn
.query(
r#"
SELECT id, title, status, user_id, failure_reason,
created_at, started_at, completed_at
FROM agent_jobs WHERE source = 'direct' AND user_id = ?1
ORDER BY created_at DESC
"#,
params![user_id],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
let mut jobs = Vec::new();
while let Some(row) = rows
.next()
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
{
let id_str = get_text(&row, 0);
let Ok(id) = id_str.parse() else {
tracing::warn!("Skipping agent job with invalid UUID: {}", id_str);
continue;
};
jobs.push(AgentJobRecord {
id,
title: get_text(&row, 1),
status: get_text(&row, 2),
user_id: get_text(&row, 3),
failure_reason: get_opt_text(&row, 4),
created_at: get_ts(&row, 5),
started_at: get_opt_ts(&row, 6),
completed_at: get_opt_ts(&row, 7),
});
}
Ok(jobs)
}
async fn get_agent_job_failure_reason( async fn get_agent_job_failure_reason(
&self, &self,
id: Uuid, id: Uuid,
@@ -320,32 +277,6 @@ impl JobStore for LibSqlBackend {
Ok(summary) Ok(summary)
} }
async fn agent_job_summary_for_user(
&self,
user_id: &str,
) -> Result<AgentJobSummary, DatabaseError> {
let conn = self.connect().await?;
let mut rows = conn
.query(
"SELECT status, COUNT(*) as cnt FROM agent_jobs WHERE source = 'direct' AND user_id = ?1 GROUP BY status",
params![user_id],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
let mut summary = AgentJobSummary::default();
while let Some(row) = rows
.next()
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
{
let status = get_text(&row, 0);
let count = get_i64(&row, 1) as usize;
summary.add_count(&status, count);
}
Ok(summary)
}
async fn save_action(&self, job_id: Uuid, action: &ActionRecord) -> Result<(), DatabaseError> { async fn save_action(&self, job_id: Uuid, action: &ActionRecord) -> Result<(), DatabaseError> {
let conn = self.connect().await?; let conn = self.connect().await?;
let duration_ms = action.duration.as_millis() as i64; let duration_ms = action.duration.as_millis() as i64;
-1
View File
@@ -12,7 +12,6 @@ mod routines;
mod sandbox; mod sandbox;
mod settings; mod settings;
mod tool_failures; mod tool_failures;
mod users;
mod workspace; mod workspace;
use std::path::Path; use std::path::Path;
+3 -68
View File
@@ -462,56 +462,6 @@ impl RoutineStore for LibSqlBackend {
Ok(counts) Ok(counts)
} }
async fn batch_get_last_run_status(
&self,
routine_ids: &[Uuid],
) -> Result<HashMap<Uuid, RunStatus>, DatabaseError> {
if routine_ids.is_empty() {
return Ok(HashMap::new());
}
let conn = self.connect().await?;
// SQLite doesn't support ANY($1), so we query all latest runs and filter in memory.
// Uses a subquery to pick only the most recent run per routine.
let mut rows = conn
.query(
"SELECT routine_id, status FROM routine_runs r1
WHERE started_at = (
SELECT MAX(started_at) FROM routine_runs r2
WHERE r2.routine_id = r1.routine_id
)
GROUP BY routine_id",
params![],
)
.await
.map_err(|e| {
DatabaseError::Query(format!("Failed to batch get last run status: {}", e))
})?;
let routine_id_set: HashSet<Uuid> = routine_ids.iter().copied().collect();
let mut statuses = HashMap::new();
while let Some(row) = rows
.next()
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
{
let id_str: String = get_text(&row, 0);
let id = Uuid::parse_str(&id_str)
.map_err(|e| DatabaseError::Query(format!("Invalid routine UUID: {}", e)))?;
if routine_id_set.contains(&id) {
let status_str: String = get_text(&row, 1);
if let std::result::Result::Ok(status) = status_str.parse::<RunStatus>() {
statuses.insert(id, status);
}
}
}
Ok(statuses)
}
async fn link_routine_run_to_job( async fn link_routine_run_to_job(
&self, &self,
run_id: Uuid, run_id: Uuid,
@@ -530,24 +480,10 @@ impl RoutineStore for LibSqlBackend {
async fn get_webhook_routine_by_path( async fn get_webhook_routine_by_path(
&self, &self,
path: &str, path: &str,
user_id: Option<&str>,
) -> Result<Option<Routine>, DatabaseError> { ) -> Result<Option<Routine>, DatabaseError> {
let conn = self.connect().await?; let conn = self.connect().await?;
let mut rows = if let Some(uid) = user_id { let mut rows = conn
conn.query( .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!( &format!(
"SELECT {} FROM routines WHERE enabled = 1 AND trigger_type = 'webhook' \ "SELECT {} FROM routines WHERE enabled = 1 AND trigger_type = 'webhook' \
AND (json_extract(trigger_config, '$.path') = ?1 \ AND (json_extract(trigger_config, '$.path') = ?1 \
@@ -557,8 +493,7 @@ impl RoutineStore for LibSqlBackend {
params![path], params![path],
) )
.await .await
.map_err(|e| DatabaseError::Query(e.to_string()))? .map_err(|e| DatabaseError::Query(e.to_string()))?;
};
match rows match rows
.next() .next()
-735
View File
@@ -1,735 +0,0 @@
//! UserStore implementation for LibSqlBackend.
use async_trait::async_trait;
use chrono::{DateTime, Utc};
use libsql::params;
use uuid::Uuid;
use super::{fmt_opt_ts, fmt_ts, get_opt_text, get_opt_ts, get_text, get_ts, opt_text};
use crate::db::libsql::LibSqlBackend;
use crate::db::{ApiTokenRecord, DatabaseError, UserRecord, UserStore};
fn row_to_user(row: &libsql::Row) -> Result<UserRecord, DatabaseError> {
let metadata_str = get_text(row, 9);
let metadata: serde_json::Value = serde_json::from_str(&metadata_str)
.map_err(|e| DatabaseError::Serialization(e.to_string()))?;
Ok(UserRecord {
id: get_text(row, 0),
email: get_opt_text(row, 1),
display_name: get_text(row, 2),
status: get_text(row, 3),
role: get_text(row, 4),
created_at: get_ts(row, 5),
updated_at: get_ts(row, 6),
last_login_at: get_opt_ts(row, 7),
created_by: get_opt_text(row, 8),
metadata,
})
}
fn row_to_api_token(row: &libsql::Row) -> Result<ApiTokenRecord, DatabaseError> {
let id_str = get_text(row, 0);
let id: Uuid = id_str
.parse()
.map_err(|e| DatabaseError::Serialization(format!("invalid UUID: {e}")))?;
Ok(ApiTokenRecord {
id,
user_id: get_text(row, 1),
name: get_text(row, 2),
token_prefix: get_text(row, 3),
expires_at: get_opt_ts(row, 4),
last_used_at: get_opt_ts(row, 5),
created_at: get_ts(row, 6),
revoked_at: get_opt_ts(row, 7),
})
}
#[async_trait]
impl UserStore for LibSqlBackend {
async fn create_user(&self, user: &UserRecord) -> Result<(), DatabaseError> {
let conn = self.connect().await?;
let metadata_json = serde_json::to_string(&user.metadata)
.map_err(|e| DatabaseError::Serialization(e.to_string()))?;
conn.execute(
r#"
INSERT INTO users (id, email, display_name, status, role, created_at, updated_at, last_login_at, created_by, metadata)
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10)
"#,
params![
user.id.as_str(),
opt_text(user.email.as_deref()),
user.display_name.as_str(),
user.status.as_str(),
user.role.as_str(),
fmt_ts(&user.created_at),
fmt_ts(&user.updated_at),
fmt_opt_ts(&user.last_login_at),
opt_text(user.created_by.as_deref()),
metadata_json,
],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
Ok(())
}
async fn get_user(&self, id: &str) -> Result<Option<UserRecord>, DatabaseError> {
let conn = self.connect().await?;
let mut rows = conn
.query(
r#"
SELECT id, email, display_name, status, role, created_at, updated_at,
last_login_at, created_by, metadata
FROM users WHERE id = ?1
"#,
params![id],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
match rows
.next()
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
{
Some(row) => Ok(Some(row_to_user(&row)?)),
None => Ok(None),
}
}
async fn get_user_by_email(&self, email: &str) -> Result<Option<UserRecord>, DatabaseError> {
let conn = self.connect().await?;
let mut rows = conn
.query(
r#"
SELECT id, email, display_name, status, role, created_at, updated_at,
last_login_at, created_by, metadata
FROM users WHERE email = ?1
"#,
params![email],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
match rows
.next()
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
{
Some(row) => Ok(Some(row_to_user(&row)?)),
None => Ok(None),
}
}
async fn list_users(&self, status: Option<&str>) -> Result<Vec<UserRecord>, DatabaseError> {
let conn = self.connect().await?;
let mut users = Vec::new();
let mut rows = if let Some(status) = status {
conn.query(
r#"
SELECT id, email, display_name, status, role, created_at, updated_at,
last_login_at, created_by, metadata
FROM users WHERE status = ?1
ORDER BY created_at DESC
"#,
params![status],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
} else {
conn.query(
r#"
SELECT id, email, display_name, status, role, created_at, updated_at,
last_login_at, created_by, metadata
FROM users
ORDER BY created_at DESC
"#,
(),
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
};
while let Some(row) = rows
.next()
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
{
users.push(row_to_user(&row)?);
}
Ok(users)
}
async fn update_user_status(&self, id: &str, status: &str) -> Result<(), DatabaseError> {
let conn = self.connect().await?;
let now = fmt_ts(&Utc::now());
conn.execute(
"UPDATE users SET status = ?2, updated_at = ?3 WHERE id = ?1",
params![id, status, now],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
Ok(())
}
async fn update_user_profile(
&self,
id: &str,
display_name: &str,
metadata: &serde_json::Value,
) -> Result<(), DatabaseError> {
let conn = self.connect().await?;
let now = fmt_ts(&Utc::now());
let metadata_json = serde_json::to_string(metadata)
.map_err(|e| DatabaseError::Serialization(e.to_string()))?;
conn.execute(
"UPDATE users SET display_name = ?2, metadata = ?3, updated_at = ?4 WHERE id = ?1",
params![id, display_name, metadata_json, now],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
Ok(())
}
async fn record_login(&self, id: &str) -> Result<(), DatabaseError> {
let conn = self.connect().await?;
let now = fmt_ts(&Utc::now());
conn.execute(
"UPDATE users SET last_login_at = ?2, updated_at = ?2 WHERE id = ?1",
params![id, now],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
Ok(())
}
async fn create_api_token(
&self,
user_id: &str,
name: &str,
token_hash: &[u8; 32],
token_prefix: &str,
expires_at: Option<DateTime<Utc>>,
) -> Result<ApiTokenRecord, DatabaseError> {
let conn = self.connect().await?;
let id = Uuid::new_v4();
let now = Utc::now();
conn.execute(
r#"
INSERT INTO api_tokens (id, user_id, token_hash, token_prefix, name, expires_at, created_at)
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7)
"#,
params![
id.to_string(),
user_id,
libsql::Value::Blob(token_hash.to_vec()),
token_prefix,
name,
fmt_opt_ts(&expires_at),
fmt_ts(&now),
],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
Ok(ApiTokenRecord {
id,
user_id: user_id.to_string(),
name: name.to_string(),
token_prefix: token_prefix.to_string(),
expires_at,
last_used_at: None,
created_at: now,
revoked_at: None,
})
}
async fn list_api_tokens(&self, user_id: &str) -> Result<Vec<ApiTokenRecord>, DatabaseError> {
let conn = self.connect().await?;
let mut rows = conn
.query(
r#"
SELECT id, user_id, name, token_prefix, expires_at, last_used_at, created_at, revoked_at
FROM api_tokens WHERE user_id = ?1
ORDER BY created_at DESC
"#,
params![user_id],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
let mut tokens = Vec::new();
while let Some(row) = rows
.next()
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
{
tokens.push(row_to_api_token(&row)?);
}
Ok(tokens)
}
async fn revoke_api_token(&self, token_id: Uuid, user_id: &str) -> Result<bool, DatabaseError> {
let conn = self.connect().await?;
let now = fmt_ts(&Utc::now());
let rows_affected = conn
.execute(
r#"
UPDATE api_tokens SET revoked_at = ?3
WHERE id = ?1 AND user_id = ?2 AND revoked_at IS NULL
"#,
params![token_id.to_string(), user_id, now],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
Ok(rows_affected > 0)
}
async fn authenticate_token(
&self,
token_hash: &[u8; 32],
) -> Result<Option<(ApiTokenRecord, UserRecord)>, DatabaseError> {
let conn = self.connect().await?;
let now = fmt_ts(&Utc::now());
let mut rows = conn
.query(
r#"
SELECT
t.id, t.user_id, t.name, t.token_prefix, t.expires_at,
t.last_used_at, t.created_at, t.revoked_at,
u.id, u.email, u.display_name, u.status, u.role, u.created_at,
u.updated_at, u.last_login_at, u.created_by, u.metadata
FROM api_tokens t
JOIN users u ON u.id = t.user_id
WHERE t.token_hash = ?1
AND t.revoked_at IS NULL
AND (t.expires_at IS NULL OR t.expires_at > ?2)
AND u.status = 'active'
"#,
params![libsql::Value::Blob(token_hash.to_vec()), now],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
match rows
.next()
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
{
Some(row) => {
let id_str = get_text(&row, 0);
let token_id: Uuid = id_str
.parse()
.map_err(|e| DatabaseError::Serialization(format!("invalid UUID: {e}")))?;
let token = ApiTokenRecord {
id: token_id,
user_id: get_text(&row, 1),
name: get_text(&row, 2),
token_prefix: get_text(&row, 3),
expires_at: get_opt_ts(&row, 4),
last_used_at: get_opt_ts(&row, 5),
created_at: get_ts(&row, 6),
revoked_at: get_opt_ts(&row, 7),
};
let metadata_str = get_text(&row, 17);
let metadata: serde_json::Value = serde_json::from_str(&metadata_str)
.map_err(|e| DatabaseError::Serialization(e.to_string()))?;
let user = UserRecord {
id: get_text(&row, 8),
email: get_opt_text(&row, 9),
display_name: get_text(&row, 10),
status: get_text(&row, 11),
role: get_text(&row, 12),
created_at: get_ts(&row, 13),
updated_at: get_ts(&row, 14),
last_login_at: get_opt_ts(&row, 15),
created_by: get_opt_text(&row, 16),
metadata,
};
Ok(Some((token, user)))
}
None => Ok(None),
}
}
async fn record_token_usage(&self, token_id: Uuid) -> Result<(), DatabaseError> {
let conn = self.connect().await?;
let now = fmt_ts(&Utc::now());
conn.execute(
"UPDATE api_tokens SET last_used_at = ?2 WHERE id = ?1",
params![token_id.to_string(), now],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
Ok(())
}
async fn has_any_users(&self) -> Result<bool, DatabaseError> {
let conn = self.connect().await?;
let mut rows = conn
.query("SELECT 1 FROM users LIMIT 1", ())
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
let has_users = rows
.next()
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
.is_some();
Ok(has_users)
}
async fn delete_user(&self, id: &str) -> Result<bool, DatabaseError> {
let conn = self.connect().await?;
// Delete from child tables first to avoid FK violations.
// agent_jobs cascades to job_actions, llm_calls, estimation_snapshots
// conversations cascades to conversation_messages
// memory_documents cascades to memory_chunks
// routines cascades to routine_runs
for table in &[
"settings",
"heartbeat_state",
"tool_rate_limit_state",
"secret_usage_log",
"leak_detection_events",
"secrets",
"wasm_tools",
"routines",
"memory_documents",
"conversations",
"api_tokens",
] {
conn.execute(
&format!("DELETE FROM {} WHERE user_id = ?1", table),
params![id],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
}
// job_events references agent_jobs(id) without CASCADE — delete via subquery.
conn.execute(
"DELETE FROM job_events WHERE job_id IN (SELECT id FROM agent_jobs WHERE user_id = ?1)",
params![id],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
conn.execute("DELETE FROM agent_jobs WHERE user_id = ?1", params![id])
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
// Nullify self-referencing created_by before deleting the user
conn.execute(
"UPDATE users SET created_by = NULL WHERE created_by = ?1",
params![id],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
let rows = conn
.execute("DELETE FROM users WHERE id = ?1", params![id])
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
Ok(rows > 0)
}
async fn user_usage_stats(
&self,
user_id: Option<&str>,
since: DateTime<Utc>,
) -> Result<Vec<crate::db::UserUsageStats>, DatabaseError> {
let conn = self.connect().await?;
let since_str = fmt_ts(&since);
let mut rows = if let Some(uid) = user_id {
conn.query(
r#"
SELECT j.user_id, l.model, COUNT(*) as call_count,
COALESCE(SUM(l.input_tokens), 0) as input_tokens,
COALESCE(SUM(l.output_tokens), 0) as output_tokens,
COALESCE(SUM(l.cost), 0) as total_cost
FROM llm_calls l
JOIN agent_jobs j ON l.job_id = j.id
WHERE l.created_at >= ?1
AND j.user_id = ?2
GROUP BY j.user_id, l.model
ORDER BY total_cost DESC
"#,
params![since_str, uid],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
} else {
conn.query(
r#"
SELECT j.user_id, l.model, COUNT(*) as call_count,
COALESCE(SUM(l.input_tokens), 0) as input_tokens,
COALESCE(SUM(l.output_tokens), 0) as output_tokens,
COALESCE(SUM(l.cost), 0) as total_cost
FROM llm_calls l
JOIN agent_jobs j ON l.job_id = j.id
WHERE l.created_at >= ?1
GROUP BY j.user_id, l.model
ORDER BY total_cost DESC
"#,
params![since_str],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
};
let mut stats = Vec::new();
while let Some(row) = rows
.next()
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
{
let cost_str = get_text(&row, 5);
let total_cost = rust_decimal::Decimal::from_str_exact(&cost_str).unwrap_or_default();
stats.push(crate::db::UserUsageStats {
user_id: get_text(&row, 0),
model: get_text(&row, 1),
call_count: row
.get::<i64>(2)
.map_err(|e| DatabaseError::Query(e.to_string()))?,
input_tokens: row
.get::<i64>(3)
.map_err(|e| DatabaseError::Query(e.to_string()))?,
output_tokens: row
.get::<i64>(4)
.map_err(|e| DatabaseError::Query(e.to_string()))?,
total_cost,
});
}
Ok(stats)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::db::libsql::LibSqlBackend;
use crate::db::{Database, UserStore};
use sha2::{Digest, Sha256};
fn hash(s: &str) -> [u8; 32] {
let mut h = Sha256::new();
h.update(s.as_bytes());
h.finalize().into()
}
async fn setup() -> (LibSqlBackend, tempfile::TempDir) {
let dir = tempfile::tempdir().unwrap();
let db_path = dir.path().join("test_users.db");
let db = LibSqlBackend::new_local(&db_path).await.unwrap();
db.run_migrations().await.unwrap();
(db, dir) // keep dir alive so the DB file isn't deleted
}
fn test_user(id: &str) -> UserRecord {
UserRecord {
id: id.to_string(),
email: Some(format!("{}@test.com", id)),
display_name: id.to_string(),
status: "active".to_string(),
role: "member".to_string(),
created_at: Utc::now(),
updated_at: Utc::now(),
last_login_at: None,
created_by: None,
metadata: serde_json::json!({}),
}
}
#[tokio::test]
async fn test_has_any_users_empty() {
let (db, _dir) = setup().await;
assert!(!db.has_any_users().await.unwrap());
}
#[tokio::test]
async fn test_create_and_get_user() {
let (db, _dir) = setup().await;
let user = test_user("alice");
db.create_user(&user).await.unwrap();
assert!(db.has_any_users().await.unwrap());
let found = db.get_user("alice").await.unwrap().unwrap();
assert_eq!(found.id, "alice");
assert_eq!(found.email, Some("[email protected]".to_string()));
assert_eq!(found.status, "active");
}
#[tokio::test]
async fn test_get_user_by_email() {
let (db, _dir) = setup().await;
db.create_user(&test_user("bob")).await.unwrap();
let found = db.get_user_by_email("[email protected]").await.unwrap();
assert!(found.is_some());
assert_eq!(found.unwrap().id, "bob");
assert!(
db.get_user_by_email("[email protected]")
.await
.unwrap()
.is_none()
);
}
#[tokio::test]
async fn test_list_users_with_status_filter() {
let (db, _dir) = setup().await;
db.create_user(&test_user("alice")).await.unwrap();
db.create_user(&test_user("bob")).await.unwrap();
db.update_user_status("bob", "suspended").await.unwrap();
let all = db.list_users(None).await.unwrap();
assert_eq!(all.len(), 2);
let active = db.list_users(Some("active")).await.unwrap();
assert_eq!(active.len(), 1);
assert_eq!(active[0].id, "alice");
let suspended = db.list_users(Some("suspended")).await.unwrap();
assert_eq!(suspended.len(), 1);
assert_eq!(suspended[0].id, "bob");
}
#[tokio::test]
async fn test_update_user_profile() {
let (db, _dir) = setup().await;
db.create_user(&test_user("alice")).await.unwrap();
let meta = serde_json::json!({"role": "admin"});
db.update_user_profile("alice", "Alice Smith", &meta)
.await
.unwrap();
let user = db.get_user("alice").await.unwrap().unwrap();
assert_eq!(user.display_name, "Alice Smith");
assert_eq!(user.metadata["role"], "admin");
}
#[tokio::test]
async fn test_token_lifecycle_create_authenticate_revoke() {
let (db, _dir) = setup().await;
db.create_user(&test_user("alice")).await.unwrap();
// Create token
let token_hash = hash("secret-token-123");
let record = db
.create_api_token("alice", "laptop", &token_hash, "secret-t", None)
.await
.unwrap();
assert_eq!(record.user_id, "alice");
assert_eq!(record.name, "laptop");
assert_eq!(record.token_prefix, "secret-t");
// Authenticate
let (tok, user) = db.authenticate_token(&token_hash).await.unwrap().unwrap();
assert_eq!(tok.id, record.id);
assert_eq!(user.id, "alice");
// List tokens
let tokens = db.list_api_tokens("alice").await.unwrap();
assert_eq!(tokens.len(), 1);
// Revoke
assert!(db.revoke_api_token(record.id, "alice").await.unwrap());
// Auth should fail after revoke
assert!(db.authenticate_token(&token_hash).await.unwrap().is_none());
}
#[tokio::test]
async fn test_token_auth_fails_for_suspended_user() {
let (db, _dir) = setup().await;
db.create_user(&test_user("alice")).await.unwrap();
let token_hash = hash("token-abc");
db.create_api_token("alice", "test", &token_hash, "token-ab", None)
.await
.unwrap();
// Auth works while active
assert!(db.authenticate_token(&token_hash).await.unwrap().is_some());
// Suspend user
db.update_user_status("alice", "suspended").await.unwrap();
// Auth should fail
assert!(db.authenticate_token(&token_hash).await.unwrap().is_none());
}
#[tokio::test]
async fn test_token_revoke_wrong_user_returns_false() {
let (db, _dir) = setup().await;
db.create_user(&test_user("alice")).await.unwrap();
db.create_user(&test_user("bob")).await.unwrap();
let token_hash = hash("alice-token");
let record = db
.create_api_token("alice", "test", &token_hash, "alice-to", None)
.await
.unwrap();
// Bob can't revoke Alice's token
assert!(!db.revoke_api_token(record.id, "bob").await.unwrap());
// Alice can
assert!(db.revoke_api_token(record.id, "alice").await.unwrap());
}
#[tokio::test]
async fn test_record_login_and_token_usage() {
let (db, _dir) = setup().await;
db.create_user(&test_user("alice")).await.unwrap();
let token_hash = hash("tok");
let record = db
.create_api_token("alice", "test", &token_hash, "tok", None)
.await
.unwrap();
// Record usage
db.record_token_usage(record.id).await.unwrap();
db.record_login("alice").await.unwrap();
// Verify timestamps updated
let user = db.get_user("alice").await.unwrap().unwrap();
assert!(user.last_login_at.is_some());
let tokens = db.list_api_tokens("alice").await.unwrap();
assert!(tokens[0].last_used_at.is_some());
}
#[tokio::test]
async fn test_delete_user_removes_api_tokens() {
let (db, _dir) = setup().await;
db.create_user(&test_user("alice")).await.unwrap();
let token_hash = hash("alice-tok");
db.create_api_token("alice", "primary", &token_hash, "alice-to", None)
.await
.unwrap();
// Verify token exists before deletion.
let tokens = db.list_api_tokens("alice").await.unwrap();
assert_eq!(tokens.len(), 1);
// Delete user — should also remove their api_tokens.
assert!(db.delete_user("alice").await.unwrap());
// api_tokens must be gone (not orphaned).
let tokens = db.list_api_tokens("alice").await.unwrap();
assert!(
tokens.is_empty(),
"expected api_tokens to be deleted with user, found {}",
tokens.len()
);
}
}

Some files were not shown because too many files have changed in this diff Show More