Compare commits

..
Author SHA1 Message Date
serrrfiratandSisyphus 8e48f36f1b style: apply rustfmt to error-path regressions
Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent)

Co-authored-by: Sisyphus <[email protected]>
2026-03-25 09:33:28 +03:00
serrrfiratandSisyphus 91fe89b078 fix: wrap preflight tool rejection errors for llm safety
Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent)

Co-authored-by: Sisyphus <[email protected]>
2026-03-25 09:26:04 +03:00
serrrfiratandSisyphus e9d56dfcc7 fix: sanitize tool error results before llm injection
Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent)

Co-authored-by: Sisyphus <[email protected]>
2026-03-25 09:15:59 +03:00
[email protected]andClaude Opus 4.6 6dfe246288 fix: address PR review — cost model attribution, heartbeat concurrency, pruning
Fixes from review comments on #1614:

- Cost tracking now uses the override model name (not active_model_name)
  when a per-user model override is active, for accurate attribution.
- Multi-user heartbeat runs per-user checks concurrently via JoinSet
  instead of sequentially, preventing one slow user from blocking others.
- Per-user failure counts tracked independently; users exceeding
  max_failures are skipped (matching single-user semantics).
- per_user_daily_cost HashMap pruned on day rollover to prevent
  unbounded growth in long-lived deployments.
- Doc comment fixed: says "routines" not "active routines".

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-24 00:06:16 -07:00
[email protected]andClaude Opus 4.6 9ff4af5734 fix: heartbeat hygiene, /model multi-tenant guard, RigAdapter model override
Three follow-up fixes for multi-tenant isolation:

1. Multi-user heartbeat now runs memory hygiene per user before each
   heartbeat check, matching single-user heartbeat behavior.

2. /model command in multi-tenant mode only persists to per-user
   settings (selected_model) without calling set_model() on the shared
   LlmProvider. The per-request model_override in the dispatcher reads
   from the same setting. Added multi_tenant flag to AgentConfig
   (auto-detected from GATEWAY_USER_TOKENS).

3. RigAdapter now supports per-request model overrides by injecting the
   model name into rig-core's additional_params. OpenAI/Anthropic/Ollama
   API servers use last-key-wins for duplicate JSON keys, so the override
   takes effect via serde's flatten serialization order.

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-23 23:29:26 -07:00
[email protected]andClaude Opus 4.6 af5daca0d9 fix: use selected_model setting key to match /model command persistence
The dispatcher was reading "preferred_model" but the /model command
(merged from staging) persists to "selected_model". Since set_setting
is already per-user scoped, using the same key makes /model work as
the per-user model override in multi-tenant mode.

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-23 23:11:34 -07:00
[email protected] 8db3638a42 Merge remote-tracking branch 'origin/staging' into feat/multi-tenant-isolation-phases-2-4 2026-03-23 23:09:17 -07:00
[email protected]andClaude Opus 4.6 9d7cdc0cf1 feat: complete multi-tenant isolation — per-user budgets, model selection, heartbeat cycling
Finishes the remaining isolation work from phases 2–4 of #59:

Phase 2 (DB scoping): Fix /status and /list commands to use _for_user
DB variants instead of global queries that leaked cross-user job data.

Phase 3 (Runtime isolation): Per-user workspace in routine engine's
spawn_fire so lightweight routines run in the correct user context.
Per-user daily cost tracking in CostGuard with configurable budget via
MAX_COST_PER_USER_PER_DAY_CENTS. Multi-user heartbeat that cycles
through all users with routines, auto-detected from GATEWAY_USER_TOKENS.

Phase 4 (Provider/tools): Per-user model selection via preferred_model
setting — looked up from SettingsStore on first iteration, threaded
through ReasoningContext.model_override to CompletionRequest. Works
with providers that support per-request model overrides (NearAI).

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-23 22:46:59 -07:00
153 changed files with 1874 additions and 15114 deletions
+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
+5 -14
View File
@@ -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.52.0",
] ]
[[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",
@@ -3486,17 +3485,9 @@ dependencies = [
"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",
@@ -5481,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.52.0",
] ]
[[package]] [[package]]
@@ -6388,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.52.0",
] ]
[[package]] [[package]]
+3 -6
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"
@@ -100,11 +100,8 @@ 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"
-7
View File
@@ -44,7 +44,6 @@ version = "0.1.0"
dependencies = [ dependencies = [
"serde", "serde",
"serde_json", "serde_json",
"subtle",
"wit-bindgen", "wit-bindgen",
] ]
@@ -209,12 +208,6 @@ dependencies = [
"smallvec", "smallvec",
] ]
[[package]]
name = "subtle"
version = "2.6.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292"
[[package]] [[package]]
name = "syn" name = "syn"
version = "2.0.117" version = "2.0.117"
-1
View File
@@ -15,7 +15,6 @@ wit-bindgen = "0.36"
# Serialization # Serialization
serde = { version = "1.0", features = ["derive"] } serde = { version = "1.0", features = ["derive"] }
serde_json = "1.0" serde_json = "1.0"
subtle = "2.6"
# Exclude from parent workspace (this is a standalone WASM component) # Exclude from parent workspace (this is a standalone WASM component)
+2 -4
View File
@@ -27,7 +27,7 @@
{ {
"name": "feishu_verification_token", "name": "feishu_verification_token",
"prompt": "Enter your Feishu/Lark Verification Token (from Event Subscription webhook settings)", "prompt": "Enter your Feishu/Lark Verification Token (from Event Subscription webhook settings)",
"optional": false "optional": true
} }
], ],
"setup_url": "https://open.feishu.cn/app" "setup_url": "https://open.feishu.cn/app"
@@ -63,15 +63,13 @@
}, },
"webhook": { "webhook": {
"secret_header": "X-Feishu-Verification-Token", "secret_header": "X-Feishu-Verification-Token",
"secret_name": "feishu_verification_token", "secret_name": "feishu_verification_token"
"managed_by_host": false
} }
} }
}, },
"config": { "config": {
"app_id": null, "app_id": null,
"app_secret": null, "app_secret": null,
"verification_token": null,
"api_base": "https://open.feishu.cn", "api_base": "https://open.feishu.cn",
"owner_id": null, "owner_id": null,
"dm_policy": "pairing", "dm_policy": "pairing",
+2 -120
View File
@@ -23,8 +23,7 @@
//! - App credentials (app_id, app_secret) are injected by the host into //! - App credentials (app_id, app_secret) are injected by the host into
//! the config JSON during startup for token exchange //! the config JSON during startup for token exchange
//! - Bearer token for API calls is obtained via token exchange and cached //! - Bearer token for API calls is obtained via token exchange and cached
//! - Webhook requests must be authenticated by the host or by a matching //! - Verification token validated by host for webhook requests
//! Feishu verification token in the request body
// Generate bindings from the WIT file // Generate bindings from the WIT file
wit_bindgen::generate!({ wit_bindgen::generate!({
@@ -33,7 +32,6 @@ wit_bindgen::generate!({
}); });
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use subtle::ConstantTimeEq;
// Re-export generated types // Re-export generated types
use exports::near::agent::channel::{ use exports::near::agent::channel::{
@@ -52,7 +50,6 @@ const ALLOW_FROM_PATH: &str = "allow_from";
const API_BASE_PATH: &str = "api_base"; const API_BASE_PATH: &str = "api_base";
const APP_ID_PATH: &str = "app_id"; const APP_ID_PATH: &str = "app_id";
const APP_SECRET_PATH: &str = "app_secret"; const APP_SECRET_PATH: &str = "app_secret";
const VERIFICATION_TOKEN_PATH: &str = "verification_token";
const TOKEN_PATH: &str = "tenant_access_token"; const TOKEN_PATH: &str = "tenant_access_token";
const TOKEN_EXPIRY_PATH: &str = "token_expiry"; const TOKEN_EXPIRY_PATH: &str = "token_expiry";
@@ -105,10 +102,6 @@ struct FeishuEventHeader {
/// Tenant key. /// Tenant key.
#[serde(default)] #[serde(default)]
tenant_key: Option<String>, tenant_key: Option<String>,
/// Verification token for v2 event payloads.
#[serde(default)]
token: Option<String>,
} }
/// Message receive event payload (im.message.receive_v1). /// Message receive event payload (im.message.receive_v1).
@@ -258,9 +251,6 @@ struct FeishuConfig {
/// Feishu App Secret (for token exchange). /// Feishu App Secret (for token exchange).
app_secret: Option<String>, app_secret: Option<String>,
/// Feishu Event Subscription verification token.
verification_token: Option<String>,
/// API base URL. Defaults to "https://open.feishu.cn" (use /// API base URL. Defaults to "https://open.feishu.cn" (use
/// "https://open.larksuite.com" for Lark international). /// "https://open.larksuite.com" for Lark international).
#[serde(default = "default_api_base")] #[serde(default = "default_api_base")]
@@ -310,9 +300,6 @@ impl Guest for FeishuChannel {
if let Some(ref app_secret) = config.app_secret { if let Some(ref app_secret) = config.app_secret {
let _ = channel_host::workspace_write(APP_SECRET_PATH, app_secret); let _ = channel_host::workspace_write(APP_SECRET_PATH, app_secret);
} }
if let Some(ref verification_token) = config.verification_token {
let _ = channel_host::workspace_write(VERIFICATION_TOKEN_PATH, verification_token);
}
if let Some(owner_id) = &config.owner_id { if let Some(owner_id) = &config.owner_id {
let _ = channel_host::workspace_write(OWNER_ID_PATH, owner_id); let _ = channel_host::workspace_write(OWNER_ID_PATH, owner_id);
@@ -389,23 +376,6 @@ impl Guest for FeishuChannel {
} }
}; };
let configured_token =
channel_host::workspace_read(VERIFICATION_TOKEN_PATH).filter(|token| !token.is_empty());
if !is_authenticated_webhook(
req.secret_validated,
configured_token.as_deref(),
request_verification_token(&event),
) {
channel_host::log(
channel_host::LogLevel::Warn,
"Rejecting unauthenticated Feishu webhook request",
);
return json_response(
401,
serde_json::json!({"error": "Webhook authentication failed"}),
);
}
// Handle URL verification challenge (initial webhook setup). // Handle URL verification challenge (initial webhook setup).
if event.event_type.as_deref() == Some("url_verification") { if event.event_type.as_deref() == Some("url_verification") {
if let Some(challenge) = &event.challenge { if let Some(challenge) = &event.challenge {
@@ -869,31 +839,6 @@ fn json_response(status: u16, body: serde_json::Value) -> OutgoingHttpResponse {
} }
} }
fn is_authenticated_webhook(
secret_validated: bool,
configured_token: Option<&str>,
request_token: Option<&str>,
) -> bool {
if secret_validated {
return true;
}
match (configured_token, request_token) {
(Some(expected), Some(provided)) => {
bool::from(expected.as_bytes().ct_eq(provided.as_bytes()))
}
_ => false,
}
}
fn request_verification_token(event: &FeishuEvent) -> Option<&str> {
event
.header
.as_ref()
.and_then(|header| header.token.as_deref())
.or(event.token.as_deref())
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
@@ -917,10 +862,7 @@ mod tests {
fn parse_token_response_rejects_missing_token() { fn parse_token_response_rejects_missing_token() {
let json = r#"{"code": 0, "msg": "ok", "expire": 7200}"#; let json = r#"{"code": 0, "msg": "ok", "expire": 7200}"#;
let result: Result<TenantAccessTokenResponse, _> = serde_json::from_str(json); let result: Result<TenantAccessTokenResponse, _> = serde_json::from_str(json);
assert!( assert!(result.is_err(), "should fail when tenant_access_token is missing");
result.is_err(),
"should fail when tenant_access_token is missing"
);
} }
#[test] #[test]
@@ -952,64 +894,4 @@ mod tests {
assert_eq!(resp.code, 10003); assert_eq!(resp.code, 10003);
assert!(resp.tenant_access_token.is_empty()); assert!(resp.tenant_access_token.is_empty());
} }
#[test]
fn webhook_auth_requires_host_auth_or_matching_verification_token() {
assert!(
!is_authenticated_webhook(false, None, Some("token")),
"requests without any configured verification mechanism must be rejected"
);
assert!(
!is_authenticated_webhook(false, Some("expected"), None),
"requests missing the Feishu token must be rejected when host auth did not pass"
);
assert!(
!is_authenticated_webhook(false, Some("expected"), Some("wrong")),
"requests with the wrong Feishu token must be rejected"
);
assert!(
is_authenticated_webhook(false, Some("expected"), Some("expected")),
"matching Feishu verification token should authenticate the request"
);
assert!(
is_authenticated_webhook(true, None, None),
"host-authenticated requests should still be accepted"
);
assert!(
is_authenticated_webhook(true, Some("expected"), Some("wrong")),
"host authentication should take precedence over body token checks"
);
}
#[test]
fn request_verification_token_prefers_v2_header_token() {
let event: FeishuEvent = serde_json::from_str(
r#"{
"schema": "2.0",
"header": {
"event_id": "evt_123",
"event_type": "im.message.receive_v1",
"token": "header-token"
},
"event": {}
}"#,
)
.unwrap();
assert_eq!(request_verification_token(&event), Some("header-token"));
}
#[test]
fn request_verification_token_falls_back_to_top_level_token() {
let event: FeishuEvent = serde_json::from_str(
r#"{
"type": "url_verification",
"challenge": "abc",
"token": "top-level-token"
}"#,
)
.unwrap();
assert_eq!(request_verification_token(&event), Some("top-level-token"));
}
} }
-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
+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
+38 -197
View File
@@ -16,7 +16,6 @@ use crate::agent::context_monitor::ContextMonitor;
use crate::agent::heartbeat::{spawn_heartbeat, spawn_multi_user_heartbeat}; use crate::agent::heartbeat::{spawn_heartbeat, spawn_multi_user_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>,
@@ -182,8 +172,6 @@ pub struct AgentDeps {
/// Resolved LLM backend identifier (e.g., "nearai", "openai", "groq"). /// Resolved LLM backend identifier (e.g., "nearai", "openai", "groq").
/// Used by `/model` persistence to determine which env var to update. /// Used by `/model` persistence to determine which env var to update.
pub llm_backend: String, 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,10 +234,7 @@ 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(),
}, },
); );
@@ -330,50 +315,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 +400,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()));
@@ -597,52 +538,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
);
}
} }
} }
} }
@@ -656,13 +575,13 @@ impl Agent {
.unwrap_or_default(); .unwrap_or_default();
if config.multi_tenant { if config.multi_tenant {
if let Some(admin) = self.admin_store() { if let Some(store) = self.store() {
Some(spawn_multi_user_heartbeat( Some(spawn_multi_user_heartbeat(
config, config,
hygiene, hygiene,
self.cheap_llm().clone(), self.cheap_llm().clone(),
Some(notify_tx), Some(notify_tx),
admin, Arc::clone(store),
)) ))
} else { } else {
tracing::warn!("Multi-tenant heartbeat requires a database store"); tracing::warn!("Multi-tenant heartbeat requires a database store");
@@ -675,7 +594,7 @@ impl Agent {
workspace.clone(), workspace.clone(),
self.cheap_llm().clone(), self.cheap_llm().clone(),
Some(notify_tx), Some(notify_tx),
self.admin_store(), self.store().map(Arc::clone),
)) ))
} }
} else { } else {
@@ -699,7 +618,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 +1071,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 +1155,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 +1165,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 +1240,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 +1265,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 +1279,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 +1324,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 +1340,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 +1502,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"
);
}
} }
+69 -172
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,22 +155,21 @@ 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;
let mut failed = 0; let mut failed = 0;
let mut stuck = 0; let mut stuck = 0;
if let Ok(s) = store.agent_job_summary().await { if let Ok(s) = store.agent_job_summary_for_user(user_id).await {
total += s.total; total += s.total;
in_progress += s.in_progress; in_progress += s.in_progress;
completed += s.completed; completed += s.completed;
failed += s.failed; failed += s.failed;
stuck += s.stuck; stuck += s.stuck;
} }
if let Ok(s) = store.sandbox_job_summary().await { if let Ok(s) = store.sandbox_job_summary_for_user(user_id).await {
total += s.total; total += s.total;
in_progress += s.running; in_progress += s.running;
completed += s.completed; completed += s.completed;
@@ -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,20 +221,19 @@ 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_for_user(user_id).await {
let agent_jobs = match store.list_agent_jobs().await {
Ok(jobs) => jobs, Ok(jobs) => jobs,
Err(e) => { Err(e) => {
tracing::warn!("Failed to list agent jobs: {}", e); tracing::warn!("Failed to list agent jobs: {}", e);
Vec::new() Vec::new()
} }
}; };
let sandbox_jobs = match store.list_sandbox_jobs().await { let sandbox_jobs = match store.list_sandbox_jobs_for_user(user_id).await {
Ok(jobs) => jobs, Ok(jobs) => jobs,
Err(e) => { Err(e) => {
tracing::warn!("Failed to list sandbox jobs: {}", e); tracing::warn!("Failed to list sandbox jobs: {}", 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",
@@ -761,12 +664,12 @@ impl Agent {
} }
if self.config.multi_tenant { if self.config.multi_tenant {
// Multi-tenant: only persist to per-user DB settings. // Multi-tenant: only persist to per-user settings.
// Do NOT call set_model() on the shared provider — that // Do NOT call set_model() on the shared provider — that
// would change the default for all users. The per-request // would change the default for all users. The per-request
// model_override in the dispatcher reads from the same // model_override in the dispatcher reads from the same
// "selected_model" setting and applies it per-user. // "selected_model" setting and applies it per-user.
self.persist_selected_model(tenant, requested).await; self.persist_selected_model(requested).await;
Ok(SubmissionResult::response(format!( Ok(SubmissionResult::response(format!(
"Model preference set to: {} (per-user)", "Model preference set to: {} (per-user)",
requested requested
@@ -775,7 +678,7 @@ impl Agent {
match self.llm().set_model(requested) { match self.llm().set_model(requested) {
Ok(()) => { Ok(()) => {
// Persist the model choice so it survives restarts. // Persist the model choice so it survives restarts.
self.persist_selected_model(tenant, requested).await; self.persist_selected_model(requested).await;
Ok(SubmissionResult::response(format!( Ok(SubmissionResult::response(format!(
"Switched model to: {}", "Switched model to: {}",
requested requested
@@ -927,14 +830,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,50 +845,38 @@ 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) {
/// The DB setting is the primary persistence layer. For LLM settings the // 1. Persist to DB if available.
/// resolution priority is `DB > env > TOML > default`, so writing to DB if let Some(store) = self.store() {
/// is sufficient for the change to survive restarts. The `.env` and TOML
/// files are only updated as a courtesy when they already contain a model
/// var, to avoid user confusion.
///
/// In multi-tenant mode, only the per-user DB setting is written — global
/// .env and TOML files are shared across users and must not be mutated.
async fn persist_selected_model(&self, tenant: &crate::tenant::TenantCtx, model: &str) {
// 1. Persist to DB if available (per-user scoped via TenantScope).
if let Some(store) = tenant.store() {
let value = serde_json::Value::String(model.to_string()); 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 { } else {
tracing::debug!( tracing::debug!("Persisted selected_model to DB: {}", model);
user_id = tenant.user_id(),
"Persisted selected_model to DB: {}",
model
);
} }
} else { } else {
tracing::warn!("No database store available — model choice will not persist to DB"); 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 .env and TOML config file (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. Best-effort update of .env and TOML if they already contain a
// model var. DB is authoritative (DB > env > TOML), but keeping
// these in sync avoids confusion when users inspect the files.
let model_owned = model.to_string(); let model_owned = model.to_string();
let backend = self.deps.llm_backend.clone(); 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 || {
// 3a. Update the backend-specific model env var in ~/.ironclaw/.env // 2a. Update the backend-specific model env var in ~/.ironclaw/.env.
// only if the var already exists (don't inject new vars). //
// 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 registry = crate::llm::ProviderRegistry::load();
let model_env = registry.model_env_var(&backend); let model_env = registry.model_env_var(&backend);
let env_var_prefix = format!("{}=", model_env); 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_path = crate::bootstrap::ironclaw_env_path();
let env_has_var = std::fs::read_to_string(&env_path) let env_has_var = std::fs::read_to_string(&env_path)
.ok() .ok()
@@ -1007,8 +894,10 @@ impl Agent {
} }
} }
// 3b. Update TOML config file if it already exists. // 2b. Update (or create) the TOML config file.
// Don't create a new one — DB persistence is sufficient. //
// 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,7 +907,15 @@ impl Agent {
} }
} }
Ok(None) => { Ok(None) => {
// No config file on disk; DB persistence is sufficient. // No config file yet — create one so the model choice
// survives restarts even when the DB is unavailable.
let settings = crate::settings::Settings {
selected_model: Some(model_owned),
..Default::default()
};
if let Err(e) = settings.save_toml(&toml_path) {
tracing::warn!("Failed to create config.toml for model persistence: {}", e);
}
} }
Err(e) => { Err(e) => {
tracing::warn!("Failed to load config.toml for model persistence: {}", e); tracing::warn!("Failed to load config.toml for model persistence: {}", e);
+29 -121
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 {
@@ -342,7 +332,12 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
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 (global + per-user)
if let Err(limit) = self.tenant.check_cost_allowed().await { if let Err(limit) = self
.agent
.cost_guard()
.check_allowed_for_user(&self.message.user_id)
.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(),
@@ -353,10 +348,12 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
// Apply per-user model override from settings (first iteration only // Apply per-user model override from settings (first iteration only
// to avoid repeated DB lookups within the same agentic loop). // to avoid repeated DB lookups within the same agentic loop).
// Uses "selected_model" — the same key the /model command persists to // Uses "selected_model" — the same key the /model command persists to
// via SettingsStore (per-user scoped via TenantScope). // via SettingsStore (per-user scoped).
if iteration == 0 if iteration == 0
&& let Some(store) = self.tenant.store() && let Some(store) = self.agent.store()
&& let Ok(Some(value)) = store.get_setting("selected_model").await && let Ok(Some(value)) = store
.get_setting(&self.message.user_id, "selected_model")
.await
&& let Some(model) = value.as_str() && let Some(model) = value.as_str()
{ {
let model = model.trim(); let model = model.trim();
@@ -400,22 +397,18 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
}; };
// Record cost and track token usage (global + per-user). // Record cost and track token usage (global + per-user).
// When a model override is active, use the override name for attribution // Use the override model name if set so cost attribution is accurate.
// and let CostGuard look up pricing via costs::model_cost() instead of let model_name = reason_ctx
// using the default provider's cost_per_token (which reflects the wrong model). .model_override
let (model_name, cost_per_token) = if let Some(ref ovr) = reason_ctx.model_override { .clone()
(ovr.clone(), None) .unwrap_or_else(|| self.agent.llm().active_model_name());
} else {
(
self.agent.llm().active_model_name(),
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
.record_llm_call( .cost_guard()
.record_llm_call_for_user(
&self.message.user_id,
&model_name, &model_name,
output.usage.input_tokens, output.usage.input_tokens,
output.usage.output_tokens, output.usage.output_tokens,
@@ -423,7 +416,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!(
@@ -454,19 +447,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
@@ -487,41 +467,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());
@@ -537,23 +482,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()),
);
} }
} }
} }
@@ -825,7 +755,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, result_content.clone()); turn.record_tool_error(result_content.clone());
} }
} }
reason_ctx.messages.push(tool_message); reason_ctx.messages.push(tool_message);
@@ -944,19 +874,16 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
&tool_result, &tool_result,
); );
// 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),
);
} }
} }
} }
@@ -1340,7 +1267,6 @@ mod tests {
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(), llm_backend: "nearai".to_string(),
tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)),
}; };
Agent::new( Agent::new(
@@ -1362,8 +1288,6 @@ mod tests {
default_timezone: "UTC".to_string(), default_timezone: "UTC".to_string(),
max_tokens_per_job: 0, max_tokens_per_job: 0,
multi_tenant: false, 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 +1497,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 +1687,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 +1779,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 +1817,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 +1947,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 +2100,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,
@@ -2221,7 +2137,6 @@ mod tests {
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(), llm_backend: "nearai".to_string(),
tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)),
}; };
Agent::new( Agent::new(
@@ -2243,8 +2158,6 @@ mod tests {
default_timezone: "UTC".to_string(), default_timezone: "UTC".to_string(),
max_tokens_per_job: 0, max_tokens_per_job: 0,
multi_tenant: false, 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()),
@@ -2279,14 +2192,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,7 +2260,6 @@ mod tests {
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(), llm_backend: "nearai".to_string(),
tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)),
}; };
Agent::new( Agent::new(
@@ -2370,8 +2281,6 @@ mod tests {
default_timezone: "UTC".to_string(), default_timezone: "UTC".to_string(),
max_tokens_per_job: 0, max_tokens_per_job: 0,
multi_tenant: false, 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()),
@@ -2391,14 +2300,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;
+42 -66
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;
@@ -182,7 +182,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 +211,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 +497,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 {
@@ -521,7 +521,7 @@ pub fn spawn_multi_user_heartbeat(
hygiene_config: HygieneConfig, hygiene_config: HygieneConfig,
llm: Arc<dyn LlmProvider>, llm: Arc<dyn LlmProvider>,
response_tx: Option<mpsc::Sender<OutgoingResponse>>, response_tx: Option<mpsc::Sender<OutgoingResponse>>,
store: AdminScope, store: Arc<dyn Database>,
) -> tokio::task::JoinHandle<()> { ) -> tokio::task::JoinHandle<()> {
tokio::spawn(async move { tokio::spawn(async move {
if !config.enabled { if !config.enabled {
@@ -574,9 +574,8 @@ pub fn spawn_multi_user_heartbeat(
} }
}; };
// Run user heartbeats concurrently so one slow LLM call doesn't // Run all user heartbeats concurrently so one slow LLM call
// block others. Cap concurrency to avoid flooding the LLM provider. // doesn't block others.
const MAX_CONCURRENT_HEARTBEATS: usize = 8;
let mut join_set = tokio::task::JoinSet::new(); let mut join_set = tokio::task::JoinSet::new();
for user_id in &user_ids { for user_id in &user_ids {
@@ -586,7 +585,7 @@ pub fn spawn_multi_user_heartbeat(
continue; continue;
} }
let workspace = Arc::new(Workspace::new_with_db(user_id, Arc::clone(store.db()))); let workspace = Arc::new(Workspace::new_with_db(user_id, store.clone()));
// Run memory hygiene per user (same as single-user heartbeat). // Run memory hygiene per user (same as single-user heartbeat).
let hygiene_ws = Arc::clone(&workspace); let hygiene_ws = Arc::clone(&workspace);
@@ -605,26 +604,19 @@ pub fn spawn_multi_user_heartbeat(
} }
}); });
// Drain completed tasks to stay within the concurrency cap.
while join_set.len() >= MAX_CONCURRENT_HEARTBEATS {
if let Some(join_result) = join_set.join_next().await {
collect_heartbeat_result(join_result, &mut user_failures, &config);
}
}
let uid = user_id.clone(); let uid = user_id.clone();
let cfg = config.clone(); let cfg = config.clone();
let hyg = hygiene_config.clone(); let hyg = hygiene_config.clone();
let llm_clone = llm.clone(); let llm_clone = llm.clone();
let tx = response_tx.clone(); let tx = response_tx.clone();
let admin = store.clone(); let st = store.clone();
join_set.spawn(async move { join_set.spawn(async move {
let mut runner = HeartbeatRunner::new(cfg, hyg, workspace, llm_clone); let mut runner = HeartbeatRunner::new(cfg, hyg, workspace, llm_clone);
if let Some(tx) = tx { if let Some(tx) = tx {
runner = runner.with_response_channel(tx); runner = runner.with_response_channel(tx);
} }
runner = runner.with_store(admin); runner = runner.with_store(st);
let result = runner.check_heartbeat().await; let result = runner.check_heartbeat().await;
if let HeartbeatResult::NeedsAttention(msg) = &result { if let HeartbeatResult::NeedsAttention(msg) = &result {
@@ -634,57 +626,41 @@ pub fn spawn_multi_user_heartbeat(
}); });
} }
// Collect remaining results and update failure counts // Collect results and update failure counts
while let Some(join_result) = join_set.join_next().await { while let Some(Ok((uid, result))) = join_set.join_next().await {
collect_heartbeat_result(join_result, &mut user_failures, &config); 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
);
}
}
}
} }
} }
}) })
} }
/// 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::*;
@@ -903,7 +879,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;
} }
+24 -24
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, String, 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, String, 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>>,
@@ -74,7 +74,7 @@ pub fn spawn_job_monitor_with_context(
} }
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, String, 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,7 +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, _user_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" {
@@ -229,7 +229,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, String, 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();
@@ -240,7 +240,7 @@ mod tests {
.send(( .send((
job_id, job_id,
"test-user".to_string(), "test-user".to_string(),
AppEvent::JobMessage { SseEvent::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 +262,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, String, 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();
@@ -274,7 +274,7 @@ mod tests {
.send(( .send((
other_job_id, other_job_id,
"test-user".to_string(), "test-user".to_string(),
AppEvent::JobMessage { SseEvent::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 +293,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, String, 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();
@@ -304,7 +304,7 @@ mod tests {
.send(( .send((
job_id, job_id,
"test-user".to_string(), "test-user".to_string(),
AppEvent::JobResult { SseEvent::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 +329,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, String, 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();
@@ -340,7 +340,7 @@ mod tests {
.send(( .send((
job_id, job_id,
"test-user".to_string(), "test-user".to_string(),
AppEvent::JobToolUse { SseEvent::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"}),
@@ -353,7 +353,7 @@ mod tests {
.send(( .send((
job_id, job_id,
"test-user".to_string(), "test-user".to_string(),
AppEvent::JobMessage { SseEvent::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 +409,7 @@ mod tests {
.await .await
.unwrap(); .unwrap();
let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (event_tx, _) = broadcast::channel::<(Uuid, String, 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(
@@ -425,7 +425,7 @@ mod tests {
.send(( .send((
job_id, job_id,
"test-user".to_string(), "test-user".to_string(),
AppEvent::JobResult { SseEvent::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 +458,7 @@ mod tests {
.await .await
.unwrap(); .unwrap();
let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (event_tx, _) = broadcast::channel::<(Uuid, String, 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(
@@ -474,7 +474,7 @@ mod tests {
.send(( .send((
job_id, job_id,
"test-user".to_string(), "test-user".to_string(),
AppEvent::JobResult { SseEvent::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 +507,14 @@ mod tests {
.await .await
.unwrap(); .unwrap();
let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16); let (event_tx, _) = broadcast::channel::<(Uuid, String, 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(), "test-user".to_string(),
AppEvent::JobResult { SseEvent::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,
+37 -270
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,
@@ -926,10 +827,7 @@ impl RoutineEngine {
let routine_workspace = if routine.user_id == self.workspace.user_id() { let routine_workspace = if routine.user_id == self.workspace.user_id() {
self.workspace.clone() self.workspace.clone()
} else { } else {
Arc::new(Workspace::new_with_db( Arc::new(Workspace::new_with_db(&routine.user_id, self.store.clone()))
&routine.user_id,
Arc::clone(self.store.db()),
))
}; };
let engine = EngineContext { let engine = EngineContext {
@@ -954,7 +852,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 +887,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 +898,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 +970,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>,
@@ -1613,10 +1511,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 +1790,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 +1797,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 +1863,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 +2061,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![
+3 -5
View File
@@ -11,12 +11,12 @@ use uuid::Uuid;
use crate::agent::task::{Task, TaskContext, TaskOutput}; use crate::agent::task::{Task, TaskContext, TaskOutput};
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 +52,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,7 +64,7 @@ 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 manager for live job event streaming.
sse_tx: Option<Arc<crate::channels::web::sse::SseManager>>, sse_tx: Option<Arc<crate::channels::web::sse::SseManager>>,
@@ -786,8 +786,6 @@ mod tests {
default_timezone: "UTC".to_string(), default_timezone: "UTC".to_string(),
max_tokens_per_job, max_tokens_per_job,
multi_tenant: false, 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 {
+17 -99
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);
@@ -1128,12 +1104,9 @@ impl Agent {
&& 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(&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),
);
} }
} }
} }
@@ -1385,12 +1358,9 @@ impl Agent {
&& let Some(turn) = thread.last_turn_mut() && 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),
);
} }
} }
} }
@@ -1474,13 +1444,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 +1459,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 +1471,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(
@@ -1853,20 +1816,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 +1839,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();
@@ -1907,10 +1853,7 @@ fn rebuild_chat_messages_from_db(
let name = c["name"].as_str().unwrap_or("unknown").to_string(); let name = c["name"].as_str().unwrap_or("unknown").to_string();
let content = if let Some(err) = c.get("error").and_then(|v| v.as_str()) let content = if let Some(err) = c.get("error").and_then(|v| v.as_str())
{ {
// Both wrapped (new) and legacy (plain) errors pass format!("Error: {}", err)
// through as-is. Legacy errors are already descriptive
// (e.g. "Tool 'http' failed: timeout"), so no prefix needed.
err.to_string()
} else if let Some(res) = c.get("result").and_then(|v| v.as_str()) { } else if let Some(res) = c.get("result").and_then(|v| v.as_str()) {
res.to_string() res.to_string()
} else if let Some(preview) = } else if let Some(preview) =
@@ -1996,38 +1939,13 @@ mod tests {
assert_eq!(result[3].role, crate::llm::Role::Tool); assert_eq!(result[3].role, crate::llm::Role::Tool);
assert_eq!(result[3].tool_call_id, Some("call_1".to_string())); assert_eq!(result[3].tool_call_id, Some("call_1".to_string()));
assert!(result[3].content.contains("timeout")); assert!(result[3].content.contains("Error: timeout"));
// final assistant // final assistant
assert_eq!(result[4].role, crate::llm::Role::Assistant); assert_eq!(result[4].role, crate::llm::Role::Assistant);
assert_eq!(result[4].content, "I found some results."); assert_eq!(result[4].content, "I found some results.");
} }
#[test]
fn test_rebuild_chat_messages_preserves_wrapped_tool_error() {
let wrapped_error =
"<tool_output name=\"http\">\nTool 'http' failed: timeout\n</tool_output>";
let tool_json = serde_json::json!([
{
"name": "http",
"call_id": "call_1",
"parameters": {"url": "https://example.com"},
"error": wrapped_error
}
]);
let messages = vec![
make_db_msg("user", "Fetch example"),
make_db_msg("tool_calls", &tool_json.to_string()),
];
let result = rebuild_chat_messages_from_db(&messages);
assert_eq!(result.len(), 3);
assert_eq!(result[2].role, crate::llm::Role::Tool);
assert_eq!(result[2].tool_call_id, Some("call_1".to_string()));
assert_eq!(result[2].content, wrapped_error);
}
#[test] #[test]
fn test_rebuild_chat_messages_legacy_tool_calls_skipped() { fn test_rebuild_chat_messages_legacy_tool_calls_skipped() {
// Legacy format: no call_id field // Legacy format: no call_id field
+10 -21
View File
@@ -229,35 +229,18 @@ impl AppBuilder {
let store = crate::secrets::create_secrets_store(crypto, handles); let store = crate::secrets::create_secrets_store(crypto, handles);
if let Some(ref secrets) = store { if let Some(ref secrets) = store {
// Migrate any plaintext API keys from the settings table to the
// encrypted secrets store. Idempotent — safe to run on every startup.
if let Some(ref db) = self.db {
crate::config::migrate_plaintext_llm_keys(
db.as_ref(),
secrets.as_ref(),
&self.config.owner_id,
)
.await;
}
// Inject LLM API keys from encrypted storage // Inject LLM API keys from encrypted storage
crate::config::inject_llm_keys_from_secrets(secrets.as_ref(), &self.config.owner_id) crate::config::inject_llm_keys_from_secrets(secrets.as_ref(), &self.config.owner_id)
.await; .await;
// Re-resolve only the LLM config with newly available keys, // Re-resolve only the LLM config with newly available keys.
// including keys hydrated from the secrets store. let store: Option<&(dyn crate::db::SettingsStore + Sync)> =
let settings_store: Option<&(dyn crate::db::SettingsStore + Sync)> =
self.db.as_ref().map(|db| db.as_ref() as _); self.db.as_ref().map(|db| db.as_ref() as _);
let toml_path = self.toml_path.as_deref(); let toml_path = self.toml_path.as_deref();
let owner_id = self.config.owner_id.clone(); let owner_id = self.config.owner_id.clone();
if let Err(e) = self if let Err(e) = self
.config .config
.re_resolve_llm_with_secrets( .re_resolve_llm(store, &owner_id, toml_path)
settings_store,
&owner_id,
toml_path,
Some(secrets.as_ref()),
)
.await .await
{ {
tracing::warn!("Failed to re-resolve LLM config after secret injection: {e}"); tracing::warn!("Failed to re-resolve LLM config after secret injection: {e}");
@@ -329,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,
-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;
+6 -61
View File
@@ -122,32 +122,18 @@ impl RelayClient {
/// instance_url in chat-api. IronClaw only passes an optional CSRF nonce /// instance_url in chat-api. IronClaw only passes an optional CSRF nonce
/// for validating the callback — no URLs. /// for validating the callback — no URLs.
pub async fn initiate_oauth(&self, state_nonce: Option<&str>) -> Result<String, RelayError> { pub async fn initiate_oauth(&self, state_nonce: Option<&str>) -> Result<String, RelayError> {
let url = format!("{}/oauth/slack/auth", self.base_url);
tracing::trace!(relay_url = %url, "RelayClient::initiate_oauth: sending request");
let mut query: Vec<(&str, &str)> = vec![]; let mut query: Vec<(&str, &str)> = vec![];
if let Some(nonce) = state_nonce { if let Some(nonce) = state_nonce {
query.push(("state_nonce", nonce)); query.push(("state_nonce", nonce));
} }
let resp = self let resp = self
.http .http
.get(&url) .get(format!("{}/oauth/slack/auth", self.base_url))
.bearer_auth(self.api_key.expose_secret()) .bearer_auth(self.api_key.expose_secret())
.query(&query) .query(&query)
.send() .send()
.await .await
.map_err(|e| { .map_err(|e| RelayError::Network(e.to_string()))?;
tracing::warn!(
relay_url = %url,
error = %e,
"RelayClient::initiate_oauth: network request failed"
);
RelayError::Network(e.to_string())
})?;
tracing::trace!(
relay_url = %url,
status = %resp.status(),
"RelayClient::initiate_oauth: received response"
);
let status = resp.status(); let status = resp.status();
if status.is_redirection() { if status.is_redirection() {
@@ -238,39 +224,20 @@ impl RelayClient {
method: &str, method: &str,
body: serde_json::Value, body: serde_json::Value,
) -> Result<serde_json::Value, RelayError> { ) -> Result<serde_json::Value, RelayError> {
let url = format!("{}/proxy/{}/{}", self.base_url, provider, method);
tracing::trace!(
relay_url = %url,
provider = %provider,
method = %method,
"RelayClient::proxy_provider: sending request"
);
let query: Vec<(&str, &str)> = vec![("team_id", team_id)]; let query: Vec<(&str, &str)> = vec![("team_id", team_id)];
let resp = self let resp = self
.http .http
.post(&url) .post(format!("{}/proxy/{}/{}", self.base_url, provider, method))
.bearer_auth(self.api_key.expose_secret()) .bearer_auth(self.api_key.expose_secret())
.query(&query) .query(&query)
.json(&body) .json(&body)
.send() .send()
.await .await
.map_err(|e| { .map_err(|e| RelayError::Network(e.to_string()))?;
tracing::warn!(
relay_url = %url,
error = %e,
"RelayClient::proxy_provider: network request failed"
);
RelayError::Network(e.to_string())
})?;
if !resp.status().is_success() { if !resp.status().is_success() {
let status = resp.status().as_u16(); let status = resp.status().as_u16();
let body = resp.text().await.unwrap_or_default(); let body = resp.text().await.unwrap_or_default();
tracing::warn!(
relay_url = %url,
status = status,
"RelayClient::proxy_provider: channel-relay returned error"
);
return Err(RelayError::Api { return Err(RelayError::Api {
status, status,
message: body, message: body,
@@ -288,45 +255,23 @@ impl RelayClient {
/// 32-byte secret. Called once at activation time; the result is cached in the /// 32-byte secret. Called once at activation time; the result is cached in the
/// extension manager so subsequent calls to `relay_signing_secret()` use it. /// extension manager so subsequent calls to `relay_signing_secret()` use it.
pub async fn get_signing_secret(&self, team_id: &str) -> Result<Vec<u8>, RelayError> { pub async fn get_signing_secret(&self, team_id: &str) -> Result<Vec<u8>, RelayError> {
let url = format!("{}/relay/signing-secret", self.base_url);
tracing::trace!(
relay_url = %url,
"RelayClient::get_signing_secret: fetching signing secret"
);
let resp = self let resp = self
.http .http
.get(&url) .get(format!("{}/relay/signing-secret", self.base_url))
.bearer_auth(self.api_key.expose_secret()) .bearer_auth(self.api_key.expose_secret())
.query(&[("team_id", team_id)]) .query(&[("team_id", team_id)])
.send() .send()
.await .await
.map_err(|e| { .map_err(|e| RelayError::Network(e.to_string()))?;
tracing::warn!(
relay_url = %url,
error = %e,
"RelayClient::get_signing_secret: network request failed"
);
RelayError::Network(e.to_string())
})?;
if !resp.status().is_success() { if !resp.status().is_success() {
let status = resp.status().as_u16(); let status = resp.status().as_u16();
let body = resp.text().await.unwrap_or_default(); let body = resp.text().await.unwrap_or_default();
tracing::warn!(
relay_url = %url,
status = status,
body = %body,
"RelayClient::get_signing_secret: channel-relay returned error"
);
return Err(RelayError::Api { return Err(RelayError::Api {
status, status,
message: body, message: body,
}); });
} }
tracing::trace!(
relay_url = %url,
"RelayClient::get_signing_secret: received successful response"
);
let body: serde_json::Value = resp let body: serde_json::Value = resp
.json() .json()
+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"
); );
} }
-8
View File
@@ -317,14 +317,6 @@ impl LoadedChannel {
.map(|f| f.webhook_secret_name()) .map(|f| f.webhook_secret_name())
.unwrap_or_else(|| format!("{}_webhook_secret", self.channel.channel_name())) .unwrap_or_else(|| format!("{}_webhook_secret", self.channel.channel_name()))
} }
/// Whether the host should enforce generic webhook-secret validation.
pub fn webhook_secret_managed_by_host(&self) -> bool {
self.capabilities_file
.as_ref()
.map(|f| f.webhook_secret_managed_by_host())
.unwrap_or(true)
}
} }
/// Results from loading multiple channels. /// Results from loading multiple channels.
-40
View File
@@ -185,19 +185,6 @@ impl ChannelCapabilitiesFile {
.and_then(|w| w.secret_name.clone()) .and_then(|w| w.secret_name.clone())
.unwrap_or_else(|| format!("{}_webhook_secret", self.name)) .unwrap_or_else(|| format!("{}_webhook_secret", self.name))
} }
/// Whether the host should enforce generic webhook-secret validation.
///
/// Defaults to true. Channels can opt out when they validate the shared
/// secret themselves using provider-specific request body fields.
pub fn webhook_secret_managed_by_host(&self) -> bool {
self.capabilities
.channel
.as_ref()
.and_then(|c| c.webhook.as_ref())
.and_then(|w| w.managed_by_host)
.unwrap_or(true)
}
} }
/// Schema for channel capabilities. /// Schema for channel capabilities.
@@ -315,14 +302,6 @@ pub struct WebhookSchema {
/// Secret name in secrets store for HMAC-SHA256 signing (Slack-style). /// Secret name in secrets store for HMAC-SHA256 signing (Slack-style).
#[serde(default)] #[serde(default)]
pub hmac_secret_name: Option<String>, pub hmac_secret_name: Option<String>,
/// Whether the host/router should enforce generic webhook-secret
/// validation before the channel sees the request.
///
/// Default: true. Set to false when the provider sends the shared secret
/// in a provider-specific request field rather than the configured header.
#[serde(default)]
pub managed_by_host: Option<bool>,
} }
/// Setup configuration schema. /// Setup configuration schema.
@@ -632,25 +611,6 @@ mod tests {
Some("X-Telegram-Bot-Api-Secret-Token") Some("X-Telegram-Bot-Api-Secret-Token")
); );
assert_eq!(file.webhook_secret_name(), "telegram_webhook_secret"); assert_eq!(file.webhook_secret_name(), "telegram_webhook_secret");
assert!(file.webhook_secret_managed_by_host());
}
#[test]
fn test_webhook_schema_can_disable_host_managed_secret_validation() {
let json = r#"{
"name": "feishu",
"capabilities": {
"channel": {
"webhook": {
"secret_name": "feishu_verification_token",
"managed_by_host": false
}
}
}
}"#;
let file = ChannelCapabilitiesFile::from_json(json).unwrap();
assert!(!file.webhook_secret_managed_by_host());
} }
#[test] #[test]
+5 -12
View File
@@ -139,18 +139,13 @@ async fn register_channel(
}; };
let secret_header = loaded.webhook_secret_header().map(|s| s.to_string()); let secret_header = loaded.webhook_secret_header().map(|s| s.to_string());
let host_webhook_secret = if loaded.webhook_secret_managed_by_host() {
webhook_secret.clone()
} else {
None
};
let webhook_path = format!("/webhook/{}", channel_name); let webhook_path = format!("/webhook/{}", channel_name);
let endpoints = vec![RegisteredEndpoint { let endpoints = vec![RegisteredEndpoint {
channel_name: channel_name.clone(), channel_name: channel_name.clone(),
path: webhook_path, path: webhook_path,
methods: vec!["POST".to_string()], methods: vec!["POST".to_string()],
require_secret: host_webhook_secret.is_some(), require_secret: webhook_secret.is_some(),
}]; }];
let channel_arc = Arc::new(loaded.channel.with_owner_actor_id(owner_actor_id.clone())); let channel_arc = Arc::new(loaded.channel.with_owner_actor_id(owner_actor_id.clone()));
@@ -210,7 +205,7 @@ async fn register_channel(
tracing::info!( tracing::info!(
channel = %channel_name, channel = %channel_name,
has_webhook_secret = host_webhook_secret.is_some(), has_webhook_secret = webhook_secret.is_some(),
secret_header = ?secret_header, secret_header = ?secret_header,
"Registering channel with router" "Registering channel with router"
); );
@@ -219,7 +214,7 @@ async fn register_channel(
.register( .register(
Arc::clone(&channel_arc), Arc::clone(&channel_arc),
endpoints, endpoints,
host_webhook_secret.clone(), webhook_secret.clone(),
secret_header, secret_header,
) )
.await; .await;
@@ -397,9 +392,8 @@ pub async fn inject_channel_credentials(
/// placeholders in URLs and headers, so this function fills config fields /// placeholders in URLs and headers, so this function fills config fields
/// that map to secret names. /// that map to secret names.
/// ///
/// Mapping: for a channel named "feishu", secrets `feishu_app_id`, /// Mapping: for a channel named "feishu", secrets `feishu_app_id` and
/// `feishu_app_secret`, and `feishu_verification_token` are injected as config /// `feishu_app_secret` are injected as config keys `app_id` and `app_secret`.
/// keys `app_id`, `app_secret`, and `verification_token`.
async fn inject_channel_secrets_into_config( async fn inject_channel_secrets_into_config(
channel_name: &str, channel_name: &str,
secrets_store: &Option<Arc<dyn SecretsStore + Send + Sync>>, secrets_store: &Option<Arc<dyn SecretsStore + Send + Sync>>,
@@ -410,7 +404,6 @@ async fn inject_channel_secrets_into_config(
"feishu" => &[ "feishu" => &[
("app_id", "feishu_app_id"), ("app_id", "feishu_app_id"),
("app_secret", "feishu_app_secret"), ("app_secret", "feishu_app_secret"),
("verification_token", "feishu_verification_token"),
], ],
_ => return, _ => return,
}; };
-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,
}
}
}) })
} }
+6 -10
View File
@@ -15,9 +15,7 @@ use crate::channels::IncomingMessage;
use crate::channels::web::auth::AuthenticatedUser; 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::{ use crate::channels::web::util::{build_turns_from_db_messages, truncate_preview};
build_turns_from_db_messages, tool_error_for_display, truncate_preview,
};
pub async fn chat_send_handler( pub async fn chat_send_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
@@ -177,7 +175,7 @@ pub async fn chat_auth_token_handler(
if result.verification.is_some() { if result.verification.is_some() {
state.sse.broadcast_for_user( state.sse.broadcast_for_user(
&user.user_id, &user.user_id,
AppEvent::AuthRequired { SseEvent::AuthRequired {
extension_name: req.extension_name.clone(), extension_name: req.extension_name.clone(),
instructions: Some(result.message), instructions: Some(result.message),
auth_url: None, auth_url: None,
@@ -189,7 +187,7 @@ pub async fn chat_auth_token_handler(
state.sse.broadcast_for_user( state.sse.broadcast_for_user(
&user.user_id, &user.user_id,
AppEvent::AuthCompleted { SseEvent::AuthCompleted {
extension_name: req.extension_name.clone(), extension_name: req.extension_name.clone(),
success: true, success: true,
message: result.message, message: result.message,
@@ -204,7 +202,7 @@ pub async fn chat_auth_token_handler(
if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) { if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) {
state.sse.broadcast_for_user( state.sse.broadcast_for_user(
&user.user_id, &user.user_id,
AppEvent::AuthRequired { SseEvent::AuthRequired {
extension_name: req.extension_name.clone(), extension_name: req.extension_name.clone(),
instructions: Some(msg.clone()), instructions: Some(msg.clone()),
auth_url: None, auth_url: None,
@@ -399,11 +397,9 @@ pub async fn chat_history_handler(
}; };
truncate_preview(&s, 500) truncate_preview(&s, 500)
}), }),
error: tc.error.as_deref().map(tool_error_for_display), error: tc.error.clone(),
rationale: tc.rationale.clone(),
}) })
.collect(), .collect(),
narrative: t.narrative.clone(),
}) })
.collect(); .collect();
@@ -535,7 +531,7 @@ pub async fn chat_threads_handler(
// Fallback: in-memory only (no assistant thread without DB) // Fallback: in-memory only (no assistant thread without DB)
let sess = session.lock().await; let sess = session.lock().await;
let mut sorted_threads: Vec<_> = sess.threads.values().collect(); let mut sorted_threads: Vec<_> = sess.threads.values().collect();
sorted_threads.sort_by_key(|t| std::cmp::Reverse(t.updated_at)); sorted_threads.sort_by(|a, b| b.updated_at.cmp(&a.updated_at));
let threads: Vec<ThreadInfo> = sorted_threads let threads: Vec<ThreadInfo> = sorted_threads
.into_iter() .into_iter()
.map(|t| ThreadInfo { .map(|t| ThreadInfo {
+2 -2
View File
@@ -114,7 +114,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,
@@ -324,7 +324,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,
+8 -833
View File
@@ -7,15 +7,10 @@ use axum::{
extract::{Path, State}, extract::{Path, State},
http::StatusCode, http::StatusCode,
}; };
use secrecy::SecretString;
use crate::channels::web::auth::AuthenticatedUser; 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::secrets::{CreateSecretParams, SecretsStore};
/// Sentinel value the frontend sends to mean "key is unchanged, don't touch it".
const API_KEY_UNCHANGED: &str = "••••••••";
pub async fn settings_list_handler( pub async fn settings_list_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
@@ -30,34 +25,12 @@ pub async fn settings_list_handler(
StatusCode::INTERNAL_SERVER_ERROR StatusCode::INTERNAL_SERVER_ERROR
})?; })?;
// Build a map of sensitive keys so we can annotate and mask them.
let sensitive_keys = ["llm_builtin_overrides", "llm_custom_providers"];
let mut sensitive_map: std::collections::HashMap<String, serde_json::Value> = rows
.iter()
.filter(|r| sensitive_keys.contains(&r.key.as_str()))
.map(|r| (r.key.clone(), r.value.clone()))
.collect();
if !sensitive_map.is_empty() {
annotate_secret_key_presence(&state, &user.user_id, &mut sensitive_map).await;
mask_settings_api_keys(&mut sensitive_map);
}
let settings = rows let settings = rows
.into_iter() .into_iter()
.map(|r| { .map(|r| SettingResponse {
let value = if sensitive_keys.contains(&r.key.as_str()) { key: r.key,
sensitive_map value: r.value,
.get(&r.key) updated_at: r.updated_at.to_rfc3339(),
.cloned()
.unwrap_or(r.value.clone())
} else {
r.value
};
SettingResponse {
key: r.key,
value,
updated_at: r.updated_at.to_rfc3339(),
}
}) })
.collect(); .collect();
@@ -82,22 +55,9 @@ pub async fn settings_get_handler(
})? })?
.ok_or(StatusCode::NOT_FOUND)?; .ok_or(StatusCode::NOT_FOUND)?;
// Mask any plaintext API keys that may exist from legacy data.
let value = if matches!(
key.as_str(),
"llm_builtin_overrides" | "llm_custom_providers"
) {
let mut map = std::collections::HashMap::from([(key.clone(), row.value.clone())]);
annotate_secret_key_presence(&state, &user.user_id, &mut map).await;
mask_settings_api_keys(&mut map);
map.remove(&key).unwrap_or(row.value)
} else {
row.value
};
Ok(Json(SettingResponse { Ok(Json(SettingResponse {
key: row.key, key: row.key,
value, value: row.value,
updated_at: row.updated_at.to_rfc3339(), updated_at: row.updated_at.to_rfc3339(),
})) }))
} }
@@ -112,27 +72,8 @@ pub async fn settings_set_handler(
.store .store
.as_ref() .as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?; .ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
// Guard: cannot remove a custom provider that is currently active.
if key == "llm_custom_providers" {
guard_active_provider_not_removed(store, &user.user_id, &body.value).await?;
validate_custom_providers(&body.value)?;
}
// Extract API keys from LLM settings and vault them in the secrets store.
// The sanitized value has api_key fields removed (stored encrypted instead).
let sanitized_value = match key.as_str() {
"llm_builtin_overrides" => {
extract_builtin_override_keys(&state, &user.user_id, &body.value).await?
}
"llm_custom_providers" => {
extract_custom_provider_keys(&state, &user.user_id, &body.value).await?
}
_ => body.value.clone(),
};
store store
.set_setting(&user.user_id, &key, &sanitized_value) .set_setting(&user.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);
@@ -142,110 +83,6 @@ pub async fn settings_set_handler(
Ok(StatusCode::NO_CONTENT) Ok(StatusCode::NO_CONTENT)
} }
const VALID_ADAPTERS: &[&str] = &["open_ai_completions", "anthropic", "ollama"];
/// Valid provider ID: lowercase alphanumeric and hyphens, 1-64 chars.
fn is_valid_provider_id(id: &str) -> bool {
!id.is_empty()
&& id.len() <= 64
&& id
.bytes()
.all(|b| b.is_ascii_lowercase() || b.is_ascii_digit() || b == b'-')
}
/// Returns `Err(422)` if any provider has an invalid ID or unrecognised adapter.
fn validate_custom_providers(value: &serde_json::Value) -> Result<(), StatusCode> {
let providers = match value.as_array() {
Some(arr) => arr,
None => return Ok(()),
};
for p in providers {
let id = p.get("id").and_then(|v| v.as_str()).unwrap_or("");
if !is_valid_provider_id(id) {
tracing::warn!(
id = %id,
"Rejected custom provider with invalid ID (must be lowercase alphanumeric/hyphens, 1-64 chars)"
);
return Err(StatusCode::UNPROCESSABLE_ENTITY);
}
}
validate_custom_providers_adapters(value)
}
/// Returns `Err(422)` if any provider in the incoming list has an unrecognised adapter.
fn validate_custom_providers_adapters(value: &serde_json::Value) -> Result<(), StatusCode> {
let providers = match value.as_array() {
Some(arr) => arr,
None => return Ok(()),
};
for p in providers {
let adapter = p.get("adapter").and_then(|v| v.as_str()).unwrap_or("");
if adapter.is_empty() {
tracing::warn!("Rejected custom provider with missing adapter field");
return Err(StatusCode::UNPROCESSABLE_ENTITY);
}
if !VALID_ADAPTERS.contains(&adapter) {
tracing::warn!(adapter = %adapter, "Rejected unknown LLM adapter");
return Err(StatusCode::UNPROCESSABLE_ENTITY);
}
}
Ok(())
}
/// Returns `Err(409)` if the active `llm_backend` is a custom provider that
/// would be removed by the incoming update to `llm_custom_providers`.
async fn guard_active_provider_not_removed(
store: &Arc<dyn crate::db::Database>,
user_id: &str,
new_value: &serde_json::Value,
) -> Result<(), StatusCode> {
// Get the currently active backend.
let active_backend = match store.get_setting(user_id, "llm_backend").await {
Ok(Some(v)) => match v.as_str() {
Some(s) if !s.is_empty() => s.to_string(),
_ => return Ok(()),
},
_ => return Ok(()),
};
// Parse the incoming provider list.
let new_providers: Vec<serde_json::Value> = match new_value.as_array() {
Some(arr) => arr.clone(),
None => return Ok(()),
};
// Check whether the active backend exists in the OLD custom providers list.
let old_providers_value = match store.get_setting(user_id, "llm_custom_providers").await {
Ok(Some(v)) => v,
_ => return Ok(()),
};
let old_providers: Vec<serde_json::Value> = match old_providers_value.as_array() {
Some(arr) => arr.clone(),
None => return Ok(()),
};
let active_was_custom = old_providers
.iter()
.any(|p| p.get("id").and_then(|v| v.as_str()) == Some(&active_backend));
if !active_was_custom {
return Ok(());
}
// Reject if the active provider is absent from the new list.
let still_present = new_providers
.iter()
.any(|p| p.get("id").and_then(|v| v.as_str()) == Some(&active_backend));
if !still_present {
tracing::warn!(
active_backend = %active_backend,
"Rejected attempt to delete the active custom LLM provider"
);
return Err(StatusCode::CONFLICT);
}
Ok(())
}
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, AuthenticatedUser(user): AuthenticatedUser,
@@ -255,14 +92,6 @@ pub async fn settings_delete_handler(
.store .store
.as_ref() .as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?; .ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
// Guard: deleting llm_custom_providers is equivalent to setting it to [].
// Reject if the active backend is a custom provider that would be removed.
if key == "llm_custom_providers" {
guard_active_provider_not_removed(store, &user.user_id, &serde_json::Value::Array(vec![]))
.await?;
}
store store
.delete_setting(&user.user_id, &key) .delete_setting(&user.user_id, &key)
.await .await
@@ -282,16 +111,11 @@ pub async fn settings_export_handler(
.store .store
.as_ref() .as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?; .ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
let mut settings = store.get_all_settings(&user.user_id).await.map_err(|e| { let settings = store.get_all_settings(&user.user_id).await.map_err(|e| {
tracing::error!("Failed to export settings: {}", e); tracing::error!("Failed to export settings: {}", e);
StatusCode::INTERNAL_SERVER_ERROR StatusCode::INTERNAL_SERVER_ERROR
})?; })?;
// Indicate key presence from secrets store without exposing values.
annotate_secret_key_presence(&state, &user.user_id, &mut settings).await;
mask_settings_api_keys(&mut settings);
Ok(Json(SettingsExportResponse { settings })) Ok(Json(SettingsExportResponse { settings }))
} }
@@ -304,21 +128,8 @@ pub async fn settings_import_handler(
.store .store
.as_ref() .as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?; .ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
// Vault any API keys present in the imported settings, same as the
// individual SET handler does, so plaintext keys never reach the DB.
let mut sanitized = body.settings.clone();
if let Some(v) = sanitized.get("llm_builtin_overrides").cloned() {
let clean = extract_builtin_override_keys(&state, &user.user_id, &v).await?;
sanitized.insert("llm_builtin_overrides".to_string(), clean);
}
if let Some(v) = sanitized.get("llm_custom_providers").cloned() {
let clean = extract_custom_provider_keys(&state, &user.user_id, &v).await?;
sanitized.insert("llm_custom_providers".to_string(), clean);
}
store store
.set_all_settings(&user.user_id, &sanitized) .set_all_settings(&user.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);
@@ -327,639 +138,3 @@ pub async fn settings_import_handler(
Ok(StatusCode::NO_CONTENT) Ok(StatusCode::NO_CONTENT)
} }
// ---------------------------------------------------------------------------
// LLM API key vaulting helpers
// ---------------------------------------------------------------------------
/// Canonical secret name for a built-in provider's API key.
fn builtin_secret_name(provider_id: &str) -> String {
format!("llm_builtin_{}_api_key", provider_id)
}
/// Canonical secret name for a custom provider's API key.
fn custom_secret_name(provider_id: &str) -> String {
format!("llm_custom_{}_api_key", provider_id)
}
/// Returns true if the `api_key` value is a real key (not sentinel/empty).
fn is_real_api_key(key: &str) -> bool {
!key.is_empty() && key != API_KEY_UNCHANGED
}
/// Require the secrets store when real API keys are present.
/// Returns `Ok(None)` when no secrets store and no real keys (passthrough).
fn require_secrets_store(
state: &GatewayState,
has_real_keys: bool,
) -> Result<Option<&Arc<dyn SecretsStore + Send + Sync>>, StatusCode> {
match state.secrets_store.as_ref() {
Some(s) => Ok(Some(s)),
None if has_real_keys => {
tracing::error!("Cannot store API keys: secrets store is not available");
Err(StatusCode::SERVICE_UNAVAILABLE)
}
None => Ok(None),
}
}
/// Extract API keys from builtin overrides, store in secrets, return sanitized JSON.
async fn extract_builtin_override_keys(
state: &GatewayState,
user_id: &str,
value: &serde_json::Value,
) -> Result<serde_json::Value, StatusCode> {
let obj = match value.as_object() {
Some(o) => o,
None => return Ok(value.clone()),
};
let has_real_keys = obj.values().any(|v| {
v.get("api_key")
.and_then(|k| k.as_str())
.is_some_and(is_real_api_key)
});
let secrets = match require_secrets_store(state, has_real_keys)? {
Some(s) => s,
None => return Ok(value.clone()),
};
let mut sanitized = obj.clone();
for (provider_id, override_val) in obj {
if let Some(api_key) = override_val.get("api_key").and_then(|v| v.as_str()) {
if !is_real_api_key(api_key) {
// Unchanged or empty — remove from settings, keep existing secret.
if let Some(o) = sanitized
.get_mut(provider_id)
.and_then(|v| v.as_object_mut())
{
o.remove("api_key");
}
continue;
}
vault_secret(
secrets.as_ref(),
user_id,
&builtin_secret_name(provider_id),
api_key,
provider_id,
)
.await?;
if let Some(o) = sanitized
.get_mut(provider_id)
.and_then(|v| v.as_object_mut())
{
o.remove("api_key");
}
}
}
Ok(serde_json::Value::Object(sanitized))
}
/// Extract API keys from custom providers, store in secrets, return sanitized JSON.
async fn extract_custom_provider_keys(
state: &GatewayState,
user_id: &str,
value: &serde_json::Value,
) -> Result<serde_json::Value, StatusCode> {
let arr = match value.as_array() {
Some(a) => a,
None => return Ok(value.clone()),
};
let has_real_keys = arr.iter().any(|v| {
v.get("api_key")
.and_then(|k| k.as_str())
.is_some_and(is_real_api_key)
});
let secrets = match require_secrets_store(state, has_real_keys)? {
Some(s) => s,
None => return Ok(value.clone()),
};
let mut sanitized = arr.clone();
for (idx, provider_val) in arr.iter().enumerate() {
let provider_id = provider_val
.get("id")
.and_then(|v| v.as_str())
.unwrap_or("");
if provider_id.is_empty() {
continue;
}
if let Some(api_key) = provider_val.get("api_key").and_then(|v| v.as_str()) {
if !is_real_api_key(api_key) {
if let Some(o) = sanitized[idx].as_object_mut() {
o.remove("api_key");
}
continue;
}
vault_secret(
secrets.as_ref(),
user_id,
&custom_secret_name(provider_id),
api_key,
provider_id,
)
.await?;
if let Some(o) = sanitized[idx].as_object_mut() {
o.remove("api_key");
}
}
}
Ok(serde_json::Value::Array(sanitized))
}
/// Encrypt and store an API key in the secrets store.
async fn vault_secret(
secrets: &(dyn SecretsStore + Send + Sync),
user_id: &str,
secret_name: &str,
api_key: &str,
provider_id: &str,
) -> Result<(), StatusCode> {
secrets
.create(
user_id,
CreateSecretParams {
name: secret_name.to_string(),
value: SecretString::from(api_key.to_string()),
provider: Some(provider_id.to_string()),
expires_at: None,
},
)
.await
.map_err(|e| {
tracing::error!(
"Failed to store secret '{}' for provider '{}': {}",
secret_name,
provider_id,
e
);
StatusCode::INTERNAL_SERVER_ERROR
})?;
Ok(())
}
/// Mask plaintext API keys in settings values before returning to the frontend.
///
/// Any `api_key` field still present in the settings JSON (legacy plaintext)
/// is replaced with the sentinel so the frontend shows "key configured".
fn mask_settings_api_keys(settings: &mut std::collections::HashMap<String, serde_json::Value>) {
if let Some(obj) = settings
.get_mut("llm_builtin_overrides")
.and_then(|v| v.as_object_mut())
{
for override_val in obj.values_mut() {
if let Some(o) = override_val.as_object_mut()
&& o.contains_key("api_key")
{
o.insert(
"api_key".to_string(),
serde_json::Value::String(API_KEY_UNCHANGED.to_string()),
);
}
}
}
if let Some(arr) = settings
.get_mut("llm_custom_providers")
.and_then(|v| v.as_array_mut())
{
for provider_val in arr.iter_mut() {
if let Some(o) = provider_val.as_object_mut()
&& o.contains_key("api_key")
{
o.insert(
"api_key".to_string(),
serde_json::Value::String(API_KEY_UNCHANGED.to_string()),
);
}
}
}
}
/// Check the secrets store for vaulted API keys and annotate the settings map.
///
/// For builtin overrides and custom providers whose API key was stripped from
/// settings (stored in secrets), this adds `api_key: "••••••••"` so the
/// frontend knows a key is configured without seeing the actual value.
async fn annotate_secret_key_presence(
state: &GatewayState,
user_id: &str,
settings: &mut std::collections::HashMap<String, serde_json::Value>,
) {
let secrets = match state.secrets_store.as_ref() {
Some(s) => s,
None => return,
};
// Annotate builtin overrides
if let Some(obj) = settings
.get_mut("llm_builtin_overrides")
.and_then(|v| v.as_object_mut())
{
let provider_ids: Vec<String> = obj.keys().cloned().collect();
for provider_id in provider_ids {
let has_key_in_settings = obj
.get(&provider_id)
.and_then(|v| v.get("api_key"))
.is_some();
if has_key_in_settings {
continue; // Will be masked by mask_settings_api_keys
}
let secret_name = builtin_secret_name(&provider_id);
if secrets.exists(user_id, &secret_name).await.unwrap_or(false)
&& let Some(o) = obj.get_mut(&provider_id).and_then(|v| v.as_object_mut())
{
o.insert(
"api_key".to_string(),
serde_json::Value::String(API_KEY_UNCHANGED.to_string()),
);
}
}
}
// Annotate custom providers
if let Some(arr) = settings
.get_mut("llm_custom_providers")
.and_then(|v| v.as_array_mut())
{
for provider_val in arr.iter_mut() {
let provider_id = provider_val
.get("id")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
if provider_id.is_empty() {
continue;
}
let has_key_in_settings = provider_val.get("api_key").is_some();
if has_key_in_settings {
continue;
}
let secret_name = custom_secret_name(&provider_id);
if secrets.exists(user_id, &secret_name).await.unwrap_or(false)
&& let Some(o) = provider_val.as_object_mut()
{
o.insert(
"api_key".to_string(),
serde_json::Value::String(API_KEY_UNCHANGED.to_string()),
);
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashMap;
#[test]
fn test_mask_settings_api_keys_builtin_overrides() {
let mut settings = HashMap::new();
settings.insert(
"llm_builtin_overrides".to_string(),
serde_json::json!({
"openai": { "api_key": "sk-secret-123", "model": "gpt-4" },
"anthropic": { "model": "claude-3" }
}),
);
mask_settings_api_keys(&mut settings);
let overrides = settings["llm_builtin_overrides"].as_object().unwrap();
assert_eq!(
overrides["openai"]["api_key"].as_str().unwrap(),
API_KEY_UNCHANGED,
);
assert_eq!(overrides["openai"]["model"].as_str().unwrap(), "gpt-4");
assert!(overrides["anthropic"].get("api_key").is_none());
}
#[test]
fn test_mask_settings_api_keys_custom_providers() {
let mut settings = HashMap::new();
settings.insert(
"llm_custom_providers".to_string(),
serde_json::json!([
{ "id": "my-llm", "api_key": "secret-key", "adapter": "open_ai_completions" },
{ "id": "no-key", "adapter": "ollama" }
]),
);
mask_settings_api_keys(&mut settings);
let providers = settings["llm_custom_providers"].as_array().unwrap();
assert_eq!(providers[0]["api_key"].as_str().unwrap(), API_KEY_UNCHANGED,);
assert!(providers[1].get("api_key").is_none());
}
#[test]
fn test_mask_settings_no_llm_keys_is_noop() {
let mut settings = HashMap::new();
settings.insert("some_other_setting".to_string(), serde_json::json!("value"));
mask_settings_api_keys(&mut settings);
assert_eq!(settings["some_other_setting"].as_str().unwrap(), "value");
}
#[test]
fn test_builtin_secret_name_format() {
assert_eq!(builtin_secret_name("openai"), "llm_builtin_openai_api_key");
}
#[test]
fn test_custom_secret_name_format() {
assert_eq!(custom_secret_name("my-groq"), "llm_custom_my-groq_api_key");
}
fn test_secrets_store() -> Arc<dyn SecretsStore + Send + Sync> {
let crypto = Arc::new(
crate::secrets::SecretsCrypto::new(secrecy::SecretString::from(
crate::secrets::keychain::generate_master_key_hex(),
))
.unwrap(),
);
Arc::new(crate::secrets::InMemorySecretsStore::new(crypto))
}
fn test_gateway_state(secrets: Arc<dyn SecretsStore + Send + Sync>) -> GatewayState {
GatewayState {
msg_tx: tokio::sync::RwLock::new(None),
sse: Arc::new(crate::channels::web::sse::SseManager::new()),
workspace: None,
workspace_pool: None,
session_manager: None,
log_broadcaster: None,
log_level_handle: None,
extension_manager: None,
tool_registry: None,
store: None,
job_manager: None,
prompt_queue: None,
scheduler: None,
owner_id: "test".to_string(),
default_sender_id: "test".to_string(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: None,
llm_provider: None,
skill_registry: None,
skill_catalog: None,
chat_rate_limiter: crate::channels::web::server::PerUserRateLimiter::new(30, 60),
oauth_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60),
webhook_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60),
registry_entries: Vec::new(),
cost_guard: None,
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
startup_time: std::time::Instant::now(),
active_config: crate::channels::web::server::ActiveConfigSnapshot::default(),
secrets_store: Some(secrets),
}
}
#[tokio::test]
async fn test_extract_builtin_keys_vaults_and_strips() {
let secrets = test_secrets_store();
let state = test_gateway_state(Arc::clone(&secrets));
let input = serde_json::json!({
"openai": { "api_key": "sk-test-key", "model": "gpt-4" },
"anthropic": { "model": "claude-3" }
});
let result = extract_builtin_override_keys(&state, "test", &input)
.await
.unwrap();
let obj = result.as_object().unwrap();
assert!(
obj["openai"].get("api_key").is_none(),
"api_key should be stripped"
);
assert_eq!(obj["openai"]["model"].as_str().unwrap(), "gpt-4");
assert_eq!(obj["anthropic"]["model"].as_str().unwrap(), "claude-3");
let decrypted = secrets
.get_decrypted("test", "llm_builtin_openai_api_key")
.await
.unwrap();
assert_eq!(decrypted.expose(), "sk-test-key");
}
#[tokio::test]
async fn test_extract_custom_keys_vaults_and_strips() {
let secrets = test_secrets_store();
let state = test_gateway_state(Arc::clone(&secrets));
let input = serde_json::json!([
{ "id": "my-llm", "api_key": "gsk-custom-key", "adapter": "open_ai_completions" },
{ "id": "local", "adapter": "ollama" }
]);
let result = extract_custom_provider_keys(&state, "test", &input)
.await
.unwrap();
let arr = result.as_array().unwrap();
assert!(
arr[0].get("api_key").is_none(),
"api_key should be stripped"
);
assert_eq!(arr[0]["id"].as_str().unwrap(), "my-llm");
assert!(arr[1].get("api_key").is_none());
let decrypted = secrets
.get_decrypted("test", "llm_custom_my-llm_api_key")
.await
.unwrap();
assert_eq!(decrypted.expose(), "gsk-custom-key");
}
#[tokio::test]
async fn test_unchanged_sentinel_preserves_existing_secret() {
let secrets = test_secrets_store();
secrets
.create(
"test",
CreateSecretParams {
name: "llm_builtin_openai_api_key".to_string(),
value: SecretString::from("sk-original".to_string()),
provider: Some("openai".to_string()),
expires_at: None,
},
)
.await
.unwrap();
let state = test_gateway_state(Arc::clone(&secrets));
let input = serde_json::json!({
"openai": { "api_key": "••••••••", "model": "gpt-4" }
});
let result = extract_builtin_override_keys(&state, "test", &input)
.await
.unwrap();
assert!(result["openai"].get("api_key").is_none());
let decrypted = secrets
.get_decrypted("test", "llm_builtin_openai_api_key")
.await
.unwrap();
assert_eq!(decrypted.expose(), "sk-original");
}
/// When secrets store is unavailable, attempting to save a real API key
/// must fail with 503 rather than silently storing plaintext.
#[tokio::test]
async fn test_extract_builtin_keys_rejects_without_secrets_store() {
let state = GatewayState {
secrets_store: None,
..test_gateway_state(test_secrets_store())
};
let input = serde_json::json!({
"openai": { "api_key": "sk-real-key", "model": "gpt-4" }
});
let err = extract_builtin_override_keys(&state, "test", &input)
.await
.unwrap_err();
assert_eq!(err, StatusCode::SERVICE_UNAVAILABLE);
}
/// When secrets store is unavailable but no real keys are present
/// (only sentinels or no api_key at all), the call should succeed.
#[tokio::test]
async fn test_extract_builtin_keys_allows_no_keys_without_secrets_store() {
let state = GatewayState {
secrets_store: None,
..test_gateway_state(test_secrets_store())
};
let input = serde_json::json!({
"openai": { "api_key": "••••••••", "model": "gpt-4" },
"anthropic": { "model": "claude-3" }
});
let result = extract_builtin_override_keys(&state, "test", &input)
.await
.unwrap();
// Without secrets store, the value passes through unchanged (no vaulting needed).
assert!(result.as_object().is_some());
}
#[tokio::test]
async fn test_extract_custom_keys_rejects_without_secrets_store() {
let state = GatewayState {
secrets_store: None,
..test_gateway_state(test_secrets_store())
};
let input = serde_json::json!([
{ "id": "my-llm", "api_key": "gsk-real-key", "adapter": "open_ai_completions" }
]);
let err = extract_custom_provider_keys(&state, "test", &input)
.await
.unwrap_err();
assert_eq!(err, StatusCode::SERVICE_UNAVAILABLE);
}
// --- Provider ID validation tests ---
#[test]
fn test_valid_provider_ids() {
assert!(is_valid_provider_id("my-llm"));
assert!(is_valid_provider_id("openai"));
assert!(is_valid_provider_id("custom-provider-123"));
assert!(is_valid_provider_id("a"));
}
#[test]
fn test_invalid_provider_ids() {
assert!(!is_valid_provider_id(""), "empty ID");
assert!(!is_valid_provider_id("My-LLM"), "uppercase");
assert!(!is_valid_provider_id("my llm"), "spaces");
assert!(!is_valid_provider_id("my_llm"), "underscores");
assert!(!is_valid_provider_id("../../etc"), "path traversal");
assert!(!is_valid_provider_id("a.b"), "dots");
assert!(
!is_valid_provider_id(&"a".repeat(65)),
"exceeds 64 char limit"
);
}
#[test]
fn test_validate_custom_providers_rejects_bad_id() {
let input = serde_json::json!([
{ "id": "UPPER-CASE", "adapter": "open_ai_completions" }
]);
assert_eq!(
validate_custom_providers(&input).unwrap_err(),
StatusCode::UNPROCESSABLE_ENTITY,
);
}
#[test]
fn test_validate_custom_providers_accepts_valid() {
let input = serde_json::json!([
{ "id": "my-llm", "adapter": "open_ai_completions" },
{ "id": "local-ollama", "adapter": "ollama" }
]);
assert!(validate_custom_providers(&input).is_ok());
}
// --- Adapter validation tests ---
#[test]
fn test_validate_adapters_rejects_unknown() {
let input = serde_json::json!([
{ "id": "test", "adapter": "not_a_real_adapter" }
]);
assert_eq!(
validate_custom_providers_adapters(&input).unwrap_err(),
StatusCode::UNPROCESSABLE_ENTITY,
);
}
#[test]
fn test_validate_adapters_rejects_missing() {
let input = serde_json::json!([
{ "id": "test" }
]);
assert_eq!(
validate_custom_providers_adapters(&input).unwrap_err(),
StatusCode::UNPROCESSABLE_ENTITY,
);
}
#[test]
fn test_validate_adapters_accepts_all_valid() {
for adapter in VALID_ADAPTERS {
let input = serde_json::json!([
{ "id": "test", "adapter": adapter }
]);
assert!(
validate_custom_providers_adapters(&input).is_ok(),
"adapter '{}' should be accepted",
adapter
);
}
}
#[test]
fn test_validate_adapters_non_array_is_ok() {
let input = serde_json::json!("not-an-array");
assert!(validate_custom_providers_adapters(&input).is_ok());
}
}
+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 { .. }
+19 -65
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;
@@ -59,7 +58,7 @@ use self::log_layer::{LogBroadcaster, LogLevelHandle};
use self::auth::MultiAuthState; use self::auth::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 {
@@ -99,8 +98,7 @@ impl GatewayChannel {
job_manager: None, job_manager: None,
prompt_queue: None, prompt_queue: None,
scheduler: None, scheduler: None,
owner_id: config.user_id.clone(), default_user_id: config.user_id.clone(),
default_sender_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,
@@ -114,7 +112,6 @@ 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 {
@@ -124,22 +121,6 @@ impl GatewayChannel {
} }
} }
/// Rebind the single-user auth identity to the durable owner scope while
/// preserving the configured gateway sender/routing identity.
pub fn with_owner_scope(mut self, owner_id: impl Into<String>) -> Self {
let owner_id = owner_id.into();
let single_user_token = if self.config.user_tokens.is_none() {
self.auth.first_token().map(ToOwned::to_owned)
} else {
None
};
if let Some(token) = single_user_token {
self.auth = MultiAuthState::single(token, owner_id.clone());
}
self.rebuild_state(|s| s.owner_id = owner_id);
self
}
/// Create a gateway channel with a pre-built multi-user auth state. /// Create a gateway channel with a pre-built multi-user auth state.
pub fn new_multi_auth(config: GatewayConfig, auth: MultiAuthState) -> Self { pub fn new_multi_auth(config: GatewayConfig, auth: MultiAuthState) -> Self {
let state = Arc::new(GatewayState { let state = Arc::new(GatewayState {
@@ -156,8 +137,7 @@ impl GatewayChannel {
job_manager: None, job_manager: None,
prompt_queue: None, prompt_queue: None,
scheduler: None, scheduler: None,
owner_id: config.user_id.clone(), default_user_id: config.user_id.clone(),
default_sender_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,
@@ -171,7 +151,6 @@ impl GatewayChannel {
startup_time: std::time::Instant::now(), startup_time: std::time::Instant::now(),
webhook_rate_limiter: server::RateLimiter::new(10, 60), webhook_rate_limiter: server::RateLimiter::new(10, 60),
active_config: server::ActiveConfigSnapshot::default(), active_config: server::ActiveConfigSnapshot::default(),
secrets_store: None,
}); });
Self { Self {
@@ -198,8 +177,7 @@ 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(), default_user_id: self.state.default_user_id.clone(),
default_sender_id: self.state.default_sender_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(),
@@ -213,7 +191,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);
@@ -331,15 +308,6 @@ impl GatewayChannel {
self self
} }
/// Inject the secrets store for encrypting LLM API keys in settings handlers.
pub fn with_secrets_store(
mut self,
ss: Arc<dyn crate::secrets::SecretsStore + Send + Sync>,
) -> Self {
self.rebuild_state(|s| s.secrets_store = Some(ss));
self
}
/// Inject the per-user workspace pool for multi-user mode. /// Inject the per-user workspace pool for multi-user mode.
pub fn with_workspace_pool(mut self, pool: Arc<server::WorkspacePool>) -> Self { pub fn with_workspace_pool(mut self, pool: Arc<server::WorkspacePool>) -> Self {
self.rebuild_state(|s| s.workspace_pool = Some(pool)); self.rebuild_state(|s| s.workspace_pool = Some(pool));
@@ -399,7 +367,7 @@ impl Channel for GatewayChannel {
self.state.sse.broadcast_for_user( self.state.sse.broadcast_for_user(
&msg.user_id, &msg.user_id,
AppEvent::Response { SseEvent::Response {
content: response.content, content: response.content,
thread_id, thread_id,
}, },
@@ -418,11 +386,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(),
}, },
@@ -431,23 +399,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(),
}, },
@@ -455,7 +423,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,
@@ -466,7 +434,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,
@@ -480,7 +448,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,
@@ -490,39 +458,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,
@@ -558,7 +512,7 @@ impl Channel for GatewayChannel {
}; };
self.state.sse.broadcast_for_user( self.state.sse.broadcast_for_user(
user_id, user_id,
AppEvent::Response { SseEvent::Response {
content: response.content, content: response.content,
thread_id, thread_id,
}, },
-2
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(
@@ -955,7 +954,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
+143 -1083
View File
File diff suppressed because it is too large Load Diff
+42 -19
View File
@@ -11,7 +11,7 @@ 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.
@@ -25,7 +25,7 @@ const MAX_CONNECTIONS: u64 = 100;
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub(crate) struct ScopedEvent { pub(crate) struct ScopedEvent {
pub(crate) user_id: Option<String>, pub(crate) user_id: Option<String>,
pub(crate) event: AppEvent, pub(crate) event: SseEvent,
} }
/// Manages SSE broadcast to all connected browser tabs. /// Manages SSE broadcast to all connected browser tabs.
@@ -75,7 +75,7 @@ impl SseManager {
} }
/// Broadcast an event to all connected clients (global/unscoped). /// Broadcast an event to all connected clients (global/unscoped).
pub fn broadcast(&self, event: AppEvent) { pub fn broadcast(&self, event: SseEvent) {
let _ = self.tx.send(ScopedEvent { let _ = self.tx.send(ScopedEvent {
user_id: None, user_id: None,
event, event,
@@ -86,7 +86,7 @@ impl SseManager {
/// ///
/// Only subscribers for this user_id (or unscoped subscribers) will /// Only subscribers for this user_id (or unscoped subscribers) will
/// receive the event. /// receive the event.
pub fn broadcast_for_user(&self, user_id: &str, event: AppEvent) { pub fn broadcast_for_user(&self, user_id: &str, event: SseEvent) {
let _ = self.tx.send(ScopedEvent { let _ = self.tx.send(ScopedEvent {
user_id: Some(user_id.to_string()), user_id: Some(user_id.to_string()),
event, event,
@@ -108,7 +108,7 @@ impl SseManager {
pub fn subscribe_raw( pub fn subscribe_raw(
&self, &self,
user_id: Option<String>, user_id: Option<String>,
) -> Option<impl Stream<Item = AppEvent> + Send + 'static + use<>> { ) -> Option<impl Stream<Item = SseEvent> + 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);
@@ -186,7 +186,30 @@ impl SseManager {
return None; return None;
} }
}; };
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::ApprovalNeeded { .. } => "approval_needed",
SseEvent::AuthRequired { .. } => "auth_required",
SseEvent::AuthCompleted { .. } => "auth_completed",
SseEvent::Error { .. } => "error",
SseEvent::JobStarted { .. } => "job_started",
SseEvent::JobMessage { .. } => "job_message",
SseEvent::JobToolUse { .. } => "job_tool_use",
SseEvent::JobToolResult { .. } => "job_tool_result",
SseEvent::JobStatus { .. } => "job_status",
SseEvent::JobResult { .. } => "job_result",
SseEvent::Heartbeat => "heartbeat",
SseEvent::ImageGenerated { .. } => "image_generated",
SseEvent::Suggestions { .. } => "suggestions",
SseEvent::TurnCost { .. } => "turn_cost",
SseEvent::ExtensionStatus { .. } => "extension_status",
};
Some(Ok(Event::default().event(event_type).data(data))) Some(Ok(Event::default().event(event_type).data(data)))
}); });
@@ -249,7 +272,7 @@ 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]
@@ -257,14 +280,14 @@ mod tests {
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(None).expect("should 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 = stream.next().await.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"),
} }
} }
@@ -276,14 +299,14 @@ mod tests {
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"),
} }
} }
@@ -306,12 +329,12 @@ mod tests {
let mut s2 = Box::pin(manager.subscribe_raw(None).expect("should subscribe")); let mut s2 = Box::pin(manager.subscribe_raw(None).expect("should subscribe"));
assert_eq!(manager.connection_count(), 2); 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);
@@ -350,25 +373,25 @@ mod tests {
// Send event scoped to alice // Send event scoped to alice
manager.broadcast_for_user( manager.broadcast_for_user(
"alice", "alice",
AppEvent::Status { SseEvent::Status {
message: "alice only".to_string(), message: "alice only".to_string(),
thread_id: None, thread_id: None,
}, },
); );
// Send global event // Send global event
manager.broadcast(AppEvent::Heartbeat); manager.broadcast(SseEvent::Heartbeat);
// Alice gets her scoped event // Alice gets her scoped event
let e = alice.next().await.unwrap(); let e = alice.next().await.unwrap();
assert!(matches!(e, AppEvent::Status { .. })); assert!(matches!(e, SseEvent::Status { .. }));
// Alice also gets the global heartbeat // Alice also gets the global heartbeat
let e = alice.next().await.unwrap(); let e = alice.next().await.unwrap();
assert!(matches!(e, AppEvent::Heartbeat)); assert!(matches!(e, SseEvent::Heartbeat));
// Bob only gets the global heartbeat (alice's event was filtered) // Bob only gets the global heartbeat (alice's event was filtered)
let e = bob.next().await.unwrap(); // safety: test-only let e = bob.next().await.unwrap(); // safety: test-only
assert!(matches!(e, AppEvent::Heartbeat)); // safety: test assertion assert!(matches!(e, SseEvent::Heartbeat)); // safety: test assertion
} }
} }
+42 -650
View File
@@ -4265,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>'
@@ -5034,6 +5034,25 @@ function loadSettingsSubtab(subtab) {
// --- Structured Settings Definitions --- // --- Structured Settings Definitions ---
var INFERENCE_SETTINGS = [ var INFERENCE_SETTINGS = [
{
group: 'cfg.group.llm',
settings: [
{ key: 'llm_backend', label: 'cfg.llm_backend.label', description: 'cfg.llm_backend.desc',
type: 'select', options: ['nearai', 'anthropic', 'openai', 'ollama', 'openai_compatible', 'tinfoil', 'bedrock'] },
{ key: 'selected_model', label: 'cfg.selected_model.label', description: 'cfg.selected_model.desc', type: 'text' },
{ key: 'ollama_base_url', label: 'cfg.ollama_base_url.label', description: 'cfg.ollama_base_url.desc', type: 'text',
showWhen: { key: 'llm_backend', value: 'ollama' } },
{ key: 'openai_compatible_base_url', label: 'cfg.openai_compatible_base_url.label', description: 'cfg.openai_compatible_base_url.desc', type: 'text',
showWhen: { key: 'llm_backend', value: 'openai_compatible' } },
{ key: 'bedrock_region', label: 'cfg.bedrock_region.label', description: 'cfg.bedrock_region.desc', type: 'text',
showWhen: { key: 'llm_backend', value: 'bedrock' } },
{ key: 'bedrock_cross_region', label: 'cfg.bedrock_cross_region.label', description: 'cfg.bedrock_cross_region.desc',
type: 'select', options: ['us', 'eu', 'apac', 'global'],
showWhen: { key: 'llm_backend', value: 'bedrock' } },
{ key: 'bedrock_profile', label: 'cfg.bedrock_profile.label', description: 'cfg.bedrock_profile.desc', type: 'text',
showWhen: { key: 'llm_backend', value: 'bedrock' } },
]
},
{ {
group: 'cfg.group.embeddings', group: 'cfg.group.embeddings',
settings: [ settings: [
@@ -5156,64 +5175,31 @@ function loadInferenceSettings() {
Promise.all([ Promise.all([
apiFetch('/api/settings/export'), apiFetch('/api/settings/export'),
apiFetch('/api/gateway/status').catch(function() { return {}; }), apiFetch('/api/gateway/status').catch(function() { return {}; }),
apiFetch('/v1/models').catch(function() { return { data: [] }; })
]).then(function(results) { ]).then(function(results) {
var settings = results[0].settings || {}; var settings = results[0].settings || {};
var status = results[1]; var status = results[1];
var modelsData = results[2];
var activeValues = {
'llm_backend': status.llm_backend,
'selected_model': status.llm_model
};
// Inject available model IDs as suggestions for the selected_model field
var modelIds = (modelsData.data || []).map(function(m) { return m.id; }).filter(Boolean);
if (modelIds.length > 0) {
var llmGroup = INFERENCE_SETTINGS[0];
for (var i = 0; i < llmGroup.settings.length; i++) {
if (llmGroup.settings[i].key === 'selected_model') {
llmGroup.settings[i].suggestions = modelIds;
break;
}
}
}
container.innerHTML = ''; container.innerHTML = '';
renderStructuredSettingsInto(container, INFERENCE_SETTINGS, settings, activeValues);
// LLM Provider display — derived from active Model Provider
var activeBackend = settings['llm_backend'] || status.llm_backend || 'nearai';
var activeModel = settings['selected_model'] || status.llm_model || '';
var allP = (typeof BUILTIN_PROVIDERS !== 'undefined' ? BUILTIN_PROVIDERS : []);
var customP = [];
try {
var cpVal = settings['llm_custom_providers'];
customP = Array.isArray(cpVal) ? cpVal : (cpVal ? JSON.parse(cpVal) : []);
} catch (e) { customP = []; }
var provider = allP.concat(customP).find(function(p) { return p.id === activeBackend; });
var providerName = provider ? (provider.name || provider.id) : activeBackend;
if (!activeModel && provider) activeModel = provider.default_model || '';
var group = document.createElement('div');
group.className = 'settings-group';
var title = document.createElement('div');
title.className = 'settings-group-title';
title.textContent = I18n.t('cfg.group.llm');
group.appendChild(title);
var notice = document.createElement('div');
notice.className = 'config-notice';
notice.id = 'llm-restart-notice';
var restartNoticeEl = document.getElementById('config-restart-notice');
notice.style.display = (restartNoticeEl && restartNoticeEl.style.display !== 'none') ? 'flex' : 'none';
notice.innerHTML = '<span>\u26A0</span><span>' + escapeHtml(I18n.t('config.restartNotice')) + '</span>';
group.appendChild(notice);
var backendRow = document.createElement('div');
backendRow.className = 'settings-row';
backendRow.innerHTML =
'<div class="settings-label-wrap"><label class="settings-label">' + escapeHtml(I18n.t('cfg.llm_backend.label')) + '</label>' +
'<div class="settings-description">' + escapeHtml(I18n.t('cfg.llm_backend.desc')) + '</div></div>' +
'<div class="settings-display-value">' + escapeHtml(providerName) + '</div>';
group.appendChild(backendRow);
var modelRow = document.createElement('div');
modelRow.className = 'settings-row';
modelRow.innerHTML =
'<div class="settings-label-wrap"><label class="settings-label">' + escapeHtml(I18n.t('cfg.selected_model.label')) + '</label>' +
'<div class="settings-description">' + escapeHtml(I18n.t('cfg.selected_model.desc')) + '</div></div>' +
'<div class="settings-display-value">' + escapeHtml(activeModel || '\u2014') + '</div>';
group.appendChild(modelRow);
container.appendChild(group);
// Remaining editable settings (embeddings, etc.)
renderStructuredSettingsInto(container, INFERENCE_SETTINGS, settings, {});
loadConfig();
}).catch(function(err) { }).catch(function(err) {
container.innerHTML = '<div class="empty-state">' + I18n.t('common.loadFailed') + ': ' container.innerHTML = '<div class="empty-state">' + I18n.t('common.loadFailed') + ': '
+ escapeHtml(err.message) + '</div>'; + escapeHtml(err.message) + '</div>';
loadConfig();
}); });
} }
@@ -5457,7 +5443,8 @@ function renderStructuredSettingsRow(def, value, activeValue) {
return row; return row;
} }
var RESTART_REQUIRED_KEYS = ['embeddings.enabled', 'embeddings.provider', 'embeddings.model', var RESTART_REQUIRED_KEYS = ['llm_backend', 'selected_model', 'ollama_base_url', 'openai_compatible_base_url',
'bedrock_region', 'bedrock_cross_region', 'bedrock_profile', 'embeddings.enabled', 'embeddings.provider', 'embeddings.model',
'agent.auto_approve_tools', 'tunnel.provider', 'tunnel.public_url', 'gateway.rate_limit', 'gateway.max_connections']; 'agent.auto_approve_tools', 'tunnel.provider', 'tunnel.public_url', 'gateway.rate_limit', 'gateway.max_connections'];
var _settingsSavedTimers = {}; var _settingsSavedTimers = {};
@@ -6042,18 +6029,6 @@ document.addEventListener('click', function(e) {
case 'switch-language': case 'switch-language':
if (typeof switchLanguage === 'function') switchLanguage(el.dataset.lang); if (typeof switchLanguage === 'function') switchLanguage(el.dataset.lang);
break; break;
case 'set-active-provider':
setActiveProvider(el.dataset.id);
break;
case 'delete-custom-provider':
deleteCustomProvider(el.dataset.id);
break;
case 'edit-custom-provider':
editCustomProvider(el.dataset.id);
break;
case 'configure-builtin-provider':
configureBuiltinProvider(el.dataset.id);
break;
} }
}); });
@@ -6095,9 +6070,6 @@ document.addEventListener('keydown', function(e) {
if (e.key === 'Escape' && document.getElementById('confirm-modal').style.display === 'flex') { if (e.key === 'Escape' && document.getElementById('confirm-modal').style.display === 'flex') {
closeConfirmModal(); closeConfirmModal();
} }
if (e.key === 'Escape' && document.getElementById('provider-dialog').style.display === 'flex') {
resetProviderForm();
}
}); });
// --- Settings Import/Export --- // --- Settings Import/Export ---
@@ -6185,583 +6157,3 @@ document.getElementById('settings-search-input').addEventListener('input', funct
activePanel.appendChild(empty); activePanel.appendChild(empty);
} }
}); });
// --- Config Tab ---
// Like apiFetch but for endpoints that return 204 No Content
function apiFetchVoid(path, options) {
const opts = options || {};
opts.headers = opts.headers || {};
opts.headers['Authorization'] = 'Bearer ' + token;
if (opts.body && typeof opts.body === 'object') {
opts.headers['Content-Type'] = 'application/json';
opts.body = JSON.stringify(opts.body);
}
return fetch(path, opts).then((res) => {
if (!res.ok) {
return res.text().then((body) => { throw new Error(body || (res.status + ' ' + res.statusText)); });
}
});
}
// BUILTIN_PROVIDERS and ADAPTER_LABELS are defined in /providers.js
let _customProviders = [];
let _activeLlmBackend = '';
let _selectedModel = '';
let _builtinOverrides = {};
let _editingProviderId = null;
let _configuringBuiltinId = null;
let _configLoaded = false;
let _envDefaults = {};
function loadConfig() {
const list = document.getElementById('providers-list');
list.innerHTML = '<div class="empty-state">' + I18n.t('common.loading') + '</div>';
Promise.all([
apiFetch('/api/settings/export'),
apiFetch('/api/llm/env_defaults').catch(() => ({})),
]).then(([d, envDefs]) => {
const s = (d && d.settings) ? d.settings : {};
_activeLlmBackend = s['llm_backend'] ? String(s['llm_backend']) : 'nearai';
_selectedModel = s['selected_model'] ? String(s['selected_model']) : '';
try {
const val = s['llm_custom_providers'];
_customProviders = Array.isArray(val) ? val : (val ? JSON.parse(val) : []);
} catch (e) {
_customProviders = [];
}
try {
const val = s['llm_builtin_overrides'];
_builtinOverrides = (val && typeof val === 'object' && !Array.isArray(val)) ? val : {};
} catch (e) {
_builtinOverrides = {};
}
_envDefaults = (envDefs && typeof envDefs === 'object') ? envDefs : {};
_configLoaded = true;
renderProviders();
}).catch(() => {
_activeLlmBackend = 'nearai';
_selectedModel = '';
_customProviders = [];
_builtinOverrides = {};
_envDefaults = {};
_configLoaded = true;
renderProviders();
});
}
function scrollToProviders() {
const section = document.getElementById('providers-section');
if (section) section.scrollIntoView({ behavior: 'smooth', block: 'start' });
}
function renderProviders() {
const list = document.getElementById('providers-list');
const allProviders = [...BUILTIN_PROVIDERS, ..._customProviders].sort((a, b) => {
if (a.id === _activeLlmBackend) return -1;
if (b.id === _activeLlmBackend) return 1;
return 0;
});
if (allProviders.length === 0) {
list.innerHTML = '<div class="empty-state">No providers</div>';
return;
}
list.innerHTML = allProviders.map((p) => {
const isActive = p.id === _activeLlmBackend;
const adapterLabel = ADAPTER_LABELS[p.adapter] || p.adapter;
const activeBadge = isActive
? '<span class="provider-badge provider-badge-active">' + I18n.t('status.active') + '</span>'
: '';
const builtinBadge = p.builtin
? '<span class="provider-badge provider-badge-builtin">' + I18n.t('config.builtin') + '</span>'
: '';
const deleteBtn = !p.builtin && !isActive
? '<button class="provider-action-btn provider-delete-btn" data-action="delete-custom-provider" data-id="' + escapeHtml(p.id) + '">' + I18n.t('common.delete') + '</button>'
: '';
const editBtn = !p.builtin
? '<button class="provider-action-btn" data-action="edit-custom-provider" data-id="' + escapeHtml(p.id) + '">' + I18n.t('common.edit') + '</button>'
: '';
// Show Configure for built-in providers that support it (not bedrock — uses AWS credential chain)
const configureBtn = p.builtin && p.id !== 'bedrock'
? '<button class="provider-action-btn" data-action="configure-builtin-provider" data-id="' + escapeHtml(p.id) + '">' + I18n.t('config.configureProvider') + '</button>'
: '';
const useBtn = !isActive
? '<button class="provider-action-btn" data-action="set-active-provider" data-id="' + escapeHtml(p.id) + '">' + I18n.t('config.useProvider') + '</button>'
: '';
const envDef = _envDefaults[p.id] || {};
const overrideBaseUrl = p.builtin && _builtinOverrides[p.id] ? (_builtinOverrides[p.id].base_url || '') : '';
const effectiveBaseUrl = overrideBaseUrl || envDef.base_url || p.base_url;
const baseUrlText = effectiveBaseUrl
? '<span class="provider-url">' + escapeHtml(effectiveBaseUrl) + '</span>'
: '';
// Show configured model: for active provider use _selectedModel, for others check _builtinOverrides then env defaults
const overrideModel = p.builtin && _builtinOverrides[p.id] ? (_builtinOverrides[p.id].model || '') : '';
const displayModel = isActive
? (_selectedModel || envDef.model || '')
: (overrideModel || envDef.model || '');
const modelText = displayModel
? '<span class="provider-current-model">' + escapeHtml(I18n.t('config.currentModel', { model: displayModel })) + '</span>'
: '';
return '<div class="provider-card' + (isActive ? ' provider-card-active' : '') + '">'
+ '<div class="provider-card-header">'
+ '<span class="provider-name">' + escapeHtml(p.name || p.id) + '</span>'
+ '<span class="provider-id-label">' + escapeHtml(p.id) + '</span>'
+ activeBadge + builtinBadge
+ '</div>'
+ '<div class="provider-card-meta">'
+ '<span class="provider-adapter">' + escapeHtml(adapterLabel) + '</span>'
+ baseUrlText
+ modelText
+ '</div>'
+ '<div class="provider-card-actions">'
+ useBtn + configureBtn + editBtn + deleteBtn
+ '</div>'
+ '</div>';
}).join('');
}
function setActiveProvider(id) {
const provider = [...BUILTIN_PROVIDERS, ..._customProviders].find((p) => p.id === id);
// Restore the last-configured model for this provider, falling back to the provider's default
const restoredModel =
(_builtinOverrides[id] && _builtinOverrides[id].model) ||
(provider && provider.default_model) ||
null;
const defaultModel = restoredModel;
const modelUpdate = () => defaultModel
? apiFetchVoid('/api/settings/selected_model', { method: 'PUT', body: { value: defaultModel } })
: apiFetchVoid('/api/settings/selected_model', { method: 'DELETE' });
apiFetchVoid('/api/settings/llm_backend', { method: 'PUT', body: { value: id } })
.then(() => modelUpdate())
.then(() => {
_activeLlmBackend = id;
_selectedModel = defaultModel || '';
renderProviders();
loadInferenceSettings();
scrollToProviders();
document.getElementById('config-restart-notice').style.display = 'flex';
var llmNotice = document.getElementById('llm-restart-notice');
if (llmNotice) llmNotice.style.display = 'flex';
showToast(I18n.t('config.providerActivated', { name: id }));
})
.catch((e) => showToast(I18n.t('error.unknown') + ': ' + e.message, 'error'));
}
function deleteCustomProvider(id) {
if (id === _activeLlmBackend) {
showToast(I18n.t('config.cannotDeleteActiveProvider'), 'error');
return;
}
if (!confirm(I18n.t('config.confirmDeleteProvider', { id }))) return;
const originalProviders = _customProviders;
_customProviders = _customProviders.filter((p) => p.id !== id);
saveCustomProviders().then(() => {
renderProviders();
showToast(I18n.t('config.providerDeleted'));
}).catch((e) => {
_customProviders = originalProviders;
showToast(I18n.t('error.unknown') + ': ' + e.message, 'error');
});
}
function saveCustomProviders() {
return apiFetchVoid('/api/settings/llm_custom_providers', { method: 'PUT', body: { value: _customProviders } });
}
function editCustomProvider(id) {
const p = _customProviders.find((p) => p.id === id);
if (!p) return;
_editingProviderId = id;
const titleEl = document.getElementById('provider-form-title');
titleEl.textContent = I18n.t('config.editProvider');
titleEl.removeAttribute('data-i18n');
document.getElementById('provider-name').value = p.name || '';
const idField = document.getElementById('provider-id');
idField.value = p.id;
idField.readOnly = true;
idField.style.opacity = '0.6';
document.getElementById('provider-adapter').value = p.adapter || 'open_ai_completions';
document.getElementById('provider-base-url').value = p.base_url || '';
const editApiKeyInput = document.getElementById('provider-api-key');
if (p.api_key === '••••••••') {
editApiKeyInput.value = '';
editApiKeyInput.placeholder = 'Key configured (leave blank to keep)';
} else {
editApiKeyInput.value = '';
editApiKeyInput.placeholder = 'Enter API key';
}
document.getElementById('provider-model').value = p.default_model || '';
openProviderDialog(true);
document.getElementById('provider-name').focus();
}
function configureBuiltinProvider(id) {
const p = BUILTIN_PROVIDERS.find((p) => p.id === id);
if (!p) return;
_configuringBuiltinId = id;
const titleEl = document.getElementById('provider-form-title');
titleEl.textContent = I18n.t('config.configureProvider') + ': ' + (p.name || id);
titleEl.removeAttribute('data-i18n');
// Hide name/id/adapter rows; show base-url as editable
document.getElementById('provider-name-row').style.display = 'none';
document.getElementById('provider-id-row').style.display = 'none';
document.getElementById('provider-adapter-row').style.display = 'none';
const baseUrlInput = document.getElementById('provider-base-url');
const override = _builtinOverrides[id] || {};
const envDef = _envDefaults[id] || {};
// Priority: db override > env > hardcoded default
const effectiveBaseUrl = override.base_url || envDef.base_url || p.base_url;
document.getElementById('provider-base-url-row').style.display = '';
baseUrlInput.value = effectiveBaseUrl || '';
baseUrlInput.readOnly = false;
baseUrlInput.style.opacity = '';
baseUrlInput.placeholder = p.base_url || '';
document.getElementById('provider-api-key-row').style.display = p.api_key_required !== false ? '' : 'none';
document.getElementById('fetch-models-btn').style.display = p.can_list_models ? '' : 'none';
const apiKeyInput = document.getElementById('provider-api-key');
const hasDbKey = override.api_key === '••••••••';
const hasEnvKey = envDef.has_api_key === true;
apiKeyInput.value = '';
if (hasDbKey) {
apiKeyInput.placeholder = 'Key configured (leave blank to keep)';
} else if (hasEnvKey) {
apiKeyInput.placeholder = 'Key set via environment variable';
} else {
apiKeyInput.placeholder = 'Enter API key';
}
document.getElementById('provider-model').value = override.model || envDef.model || p.default_model || '';
openProviderDialog(true);
document.getElementById('provider-model').focus();
}
// Add provider form
document.getElementById('add-provider-btn').addEventListener('click', () => {
openProviderDialog(false);
});
document.getElementById('cancel-provider-btn').addEventListener('click', () => {
resetProviderForm();
});
document.getElementById('cancel-provider-footer-btn').addEventListener('click', () => {
resetProviderForm();
});
document.getElementById('provider-dialog-overlay').addEventListener('click', () => {
resetProviderForm();
});
function openProviderDialog(isEdit) {
if (!isEdit) {
// Add mode: ensure all rows visible
['provider-name-row', 'provider-id-row', 'provider-adapter-row',
'provider-base-url-row', 'provider-api-key-row'].forEach((id) => {
document.getElementById(id).style.display = '';
});
document.getElementById('fetch-models-btn').style.display = '';
}
document.getElementById('provider-dialog').style.display = 'flex';
if (!isEdit) {
document.getElementById('provider-name').focus();
}
}
document.getElementById('test-provider-btn').addEventListener('click', () => {
let adapter = document.getElementById('provider-adapter').value;
let baseUrl = document.getElementById('provider-base-url').value.trim();
const apiKey = document.getElementById('provider-api-key').value.trim();
const model = document.getElementById('provider-model').value.trim();
// For built-in providers, use the hardcoded adapter from BUILTIN_PROVIDERS.
// base_url comes from the form which already reflects: env > hardcoded default.
if (_configuringBuiltinId) {
const p = BUILTIN_PROVIDERS.find((x) => x.id === _configuringBuiltinId);
if (p) {
adapter = p.adapter;
if (!baseUrl) baseUrl = p.base_url;
}
}
const btn = document.getElementById('test-provider-btn');
const result = document.getElementById('test-connection-result');
btn.disabled = true;
btn.textContent = I18n.t('config.testing');
result.style.display = 'none';
result.className = 'test-connection-result';
// Resolve provider_id so the backend can look up vaulted API keys.
const providerId = _configuringBuiltinId || document.getElementById('provider-id').value.trim();
if (!model) {
result.textContent = I18n.t('config.modelRequired') || 'Model is required for connection test';
result.className = 'test-connection-result test-fail';
result.style.display = '';
btn.disabled = false;
btn.textContent = I18n.t('config.testConnection');
return;
}
apiFetch('/api/llm/test_connection', {
method: 'POST',
body: {
adapter, base_url: baseUrl,
api_key: apiKey || undefined,
model,
provider_id: providerId || undefined,
provider_type: _configuringBuiltinId ? 'builtin' : 'custom',
},
})
.then((data) => {
result.textContent = data.message;
result.className = 'test-connection-result ' + (data.ok ? 'test-ok' : 'test-fail');
result.style.display = '';
})
.catch((e) => {
result.textContent = e.message;
result.className = 'test-connection-result test-fail';
result.style.display = '';
})
.finally(() => {
btn.disabled = false;
btn.textContent = I18n.t('config.testConnection');
});
});
document.getElementById('save-provider-btn').addEventListener('click', () => {
// Built-in configure mode: save api_key + model to llm_builtin_overrides
if (_configuringBuiltinId) {
const apiKey = document.getElementById('provider-api-key').value.trim();
const model = document.getElementById('provider-model').value.trim();
const baseUrl = document.getElementById('provider-base-url').value.trim();
const id = _configuringBuiltinId;
const prevOverride = _builtinOverrides[id] || {};
const hadKey = prevOverride.api_key === '••••••••';
const override = {};
if (apiKey) {
override.api_key = apiKey; // New key entered — backend will encrypt it
} else if (hadKey) {
override.api_key = '••••••••'; // Sentinel: keep existing encrypted key
}
// If neither — key is cleared (no key configured)
if (model) override.model = model;
if (baseUrl) override.base_url = baseUrl;
const prev = _builtinOverrides[id];
_builtinOverrides[id] = override;
const isActive = id === _activeLlmBackend;
const modelUpdate = () => {
if (!isActive) return Promise.resolve();
if (model) {
return apiFetchVoid('/api/settings/selected_model', { method: 'PUT', body: { value: model } });
}
return apiFetchVoid('/api/settings/selected_model', { method: 'DELETE' });
};
apiFetchVoid('/api/settings/llm_builtin_overrides', { method: 'PUT', body: { value: _builtinOverrides } })
.then(() => modelUpdate())
.then(() => {
if (isActive) _selectedModel = model;
renderProviders();
if (isActive) loadInferenceSettings();
resetProviderForm();
scrollToProviders();
if (isActive) {
document.getElementById('config-restart-notice').style.display = 'flex';
var llmNotice = document.getElementById('llm-restart-notice');
if (llmNotice) llmNotice.style.display = 'flex';
}
showToast(I18n.t('config.providerConfigured', { name: id }));
})
.catch((e) => {
if (prev !== undefined) { _builtinOverrides[id] = prev; } else { delete _builtinOverrides[id]; }
showToast(I18n.t('error.unknown') + ': ' + e.message, 'error');
});
return;
}
const name = document.getElementById('provider-name').value.trim();
const id = document.getElementById('provider-id').value.trim();
const adapter = document.getElementById('provider-adapter').value;
const baseUrl = document.getElementById('provider-base-url').value.trim();
const apiKey = document.getElementById('provider-api-key').value.trim();
const model = document.getElementById('provider-model').value.trim();
if (!id || !name) {
showToast(I18n.t('config.providerFieldsRequired'), 'error');
return;
}
if (_editingProviderId) {
// Update existing provider
const idx = _customProviders.findIndex((p) => p.id === _editingProviderId);
if (idx === -1) return;
const original = _customProviders[idx];
const hadCustomKey = original.api_key === '••••••••';
let effectiveApiKey;
if (apiKey) {
effectiveApiKey = apiKey; // New key — backend will encrypt it
} else if (hadCustomKey) {
effectiveApiKey = '••••••••'; // Sentinel: keep existing encrypted key
} else {
effectiveApiKey = undefined; // No key
}
_customProviders[idx] = { ...original, name, adapter, base_url: baseUrl, default_model: model || undefined, api_key: effectiveApiKey };
const isActive = _editingProviderId === _activeLlmBackend;
const modelUpdate = () => {
if (!isActive) return Promise.resolve();
if (model) {
return apiFetchVoid('/api/settings/selected_model', { method: 'PUT', body: { value: model } });
}
return apiFetchVoid('/api/settings/selected_model', { method: 'DELETE' });
};
saveCustomProviders().then(() => modelUpdate()).then(() => {
if (isActive) _selectedModel = model;
renderProviders();
if (isActive) loadInferenceSettings();
resetProviderForm();
scrollToProviders();
if (isActive) {
document.getElementById('config-restart-notice').style.display = 'flex';
var llmNotice = document.getElementById('llm-restart-notice');
if (llmNotice) llmNotice.style.display = 'flex';
}
showToast(I18n.t('config.providerUpdated', { name }));
}).catch((e) => {
_customProviders[idx] = original;
showToast(I18n.t('error.unknown') + ': ' + e.message, 'error');
});
return;
}
if (!/^[a-z0-9-]+$/.test(id)) {
showToast(I18n.t('config.providerIdInvalid'), 'error');
return;
}
const allIds = [...BUILTIN_PROVIDERS.map((p) => p.id), ..._customProviders.map((p) => p.id)];
if (allIds.includes(id)) {
showToast(I18n.t('config.providerIdTaken', { id }), 'error');
return;
}
const newProvider = { id, name, adapter, base_url: baseUrl, default_model: model, api_key: apiKey || undefined, builtin: false };
_customProviders.push(newProvider);
saveCustomProviders().then(() => {
renderProviders();
resetProviderForm();
scrollToProviders();
showToast(I18n.t('config.providerAdded', { name }));
}).catch((e) => {
_customProviders.pop();
showToast(I18n.t('error.unknown') + ': ' + e.message, 'error');
});
});
function resetProviderForm() {
_editingProviderId = null;
_configuringBuiltinId = null;
document.getElementById('provider-dialog').style.display = 'none';
// Restore all hidden rows and buttons
['provider-name-row', 'provider-id-row', 'provider-adapter-row',
'provider-base-url-row', 'provider-api-key-row'].forEach((id) => {
document.getElementById(id).style.display = '';
});
document.getElementById('fetch-models-btn').style.display = '';
const titleEl = document.getElementById('provider-form-title');
titleEl.setAttribute('data-i18n', 'config.newProvider');
titleEl.textContent = I18n.t('config.newProvider');
const idField = document.getElementById('provider-id');
idField.readOnly = false;
idField.style.opacity = '';
delete idField.dataset.edited;
const baseUrlField = document.getElementById('provider-base-url');
baseUrlField.readOnly = false;
baseUrlField.style.opacity = '';
['provider-name', 'provider-id', 'provider-base-url', 'provider-api-key', 'provider-model'].forEach((id) => {
document.getElementById(id).value = '';
});
document.getElementById('provider-adapter').selectedIndex = 0;
const sel = document.getElementById('provider-model-select');
sel.innerHTML = '';
sel.style.display = 'none';
document.getElementById('test-connection-result').style.display = 'none';
}
document.getElementById('provider-model-select').addEventListener('change', (e) => {
document.getElementById('provider-model').value = e.target.value;
});
document.getElementById('fetch-models-btn').addEventListener('click', () => {
let adapter = document.getElementById('provider-adapter').value;
let baseUrl = document.getElementById('provider-base-url').value.trim();
const apiKey = document.getElementById('provider-api-key').value.trim();
// For built-in providers, use the hardcoded adapter from BUILTIN_PROVIDERS.
// base_url comes from the form which already reflects: env > hardcoded default.
if (_configuringBuiltinId) {
const p = BUILTIN_PROVIDERS.find((x) => x.id === _configuringBuiltinId);
if (p) {
adapter = p.adapter;
if (!baseUrl) baseUrl = p.base_url;
}
}
if (!baseUrl) {
showToast(I18n.t('config.providerBaseUrlRequired'), 'error');
return;
}
const btn = document.getElementById('fetch-models-btn');
btn.disabled = true;
btn.textContent = I18n.t('config.fetchingModels');
// Resolve provider_id so the backend can look up vaulted API keys.
const providerId = _configuringBuiltinId || document.getElementById('provider-id').value.trim();
apiFetch('/api/llm/list_models', {
method: 'POST',
body: {
adapter, base_url: baseUrl,
api_key: apiKey || undefined,
provider_id: providerId || undefined,
provider_type: _configuringBuiltinId ? 'builtin' : 'custom',
},
})
.then((data) => {
const select = document.getElementById('provider-model-select');
if (data.ok && data.models && data.models.length > 0) {
const currentModel = document.getElementById('provider-model').value;
select.innerHTML = data.models
.map((m) => `<option value="${escapeHtml(m)}"${m === currentModel ? ' selected' : ''}>${escapeHtml(m)}</option>`)
.join('');
select.style.display = '';
btn.style.display = 'none';
showToast(I18n.t('config.modelsFetched', { count: data.models.length }));
} else {
showToast(data.message || I18n.t('config.modelsFetchFailed'), 'error');
}
})
.catch((e) => showToast(e.message, 'error'))
.finally(() => {
btn.disabled = false;
btn.textContent = I18n.t('config.fetchModels');
});
});
// Auto-fill provider ID from name
document.getElementById('provider-name').addEventListener('input', (e) => {
const idField = document.getElementById('provider-id');
if (!idField.dataset.edited) {
idField.value = e.target.value.toLowerCase().replace(/[^a-z0-9]+/g, '-').replace(/^-|-$/g, '');
}
});
document.getElementById('provider-id').addEventListener('input', (e) => {
e.target.dataset.edited = e.target.value ? '1' : '';
});
+3 -14
View File
@@ -35,27 +35,16 @@ function switchLanguage(lang) {
if (I18n.setLanguage(lang)) { if (I18n.setLanguage(lang)) {
// Update slash commands // Update slash commands
updateSlashCommands(); updateSlashCommands();
// Update language menu active state // Update language menu active state
updateLanguageMenu(); updateLanguageMenu();
// Re-render dynamically built sections that use I18n.t()
if (typeof renderProviders === 'function' && typeof _configLoaded !== 'undefined' && _configLoaded) {
renderProviders();
}
if (typeof loadInferenceSettings === 'function') {
var inferencePanel = document.getElementById('settings-inference');
if (inferencePanel && inferencePanel.classList.contains('active')) {
loadInferenceSettings();
}
}
// Close menu // Close menu
const menu = document.getElementById('language-menu'); const menu = document.getElementById('language-menu');
if (menu) { if (menu) {
menu.style.display = 'none'; menu.style.display = 'none';
} }
// Show toast notification // Show toast notification
showToast(I18n.t('language.switch') + ': ' + (lang === 'zh-CN' ? '简体中文' : 'English')); showToast(I18n.t('language.switch') + ': ' + (lang === 'zh-CN' ? '简体中文' : 'English'));
} }
-41
View File
@@ -38,14 +38,12 @@ I18n.register('en', {
'tab.settings': 'Settings', 'tab.settings': 'Settings',
'tab.extensions': 'Extensions', 'tab.extensions': 'Extensions',
'tab.skills': 'Skills', 'tab.skills': 'Skills',
'tab.config': 'Config',
'tab.logs': 'Logs', 'tab.logs': 'Logs',
'settings.inference': 'Inference', 'settings.inference': 'Inference',
'settings.agent': 'Agent', 'settings.agent': 'Agent',
'settings.channels': 'Channels', 'settings.channels': 'Channels',
'settings.networking': 'Networking', 'settings.networking': 'Networking',
'settings.mcp': 'MCP', 'settings.mcp': 'MCP',
'settings.providers': 'Providers',
// Status // Status
'status.connected': 'Connected', 'status.connected': 'Connected',
@@ -352,45 +350,6 @@ I18n.register('en', {
'ext.removed': 'Removed {name}', 'ext.removed': 'Removed {name}',
'ext.installFailed': 'Install failed: {message}', 'ext.installFailed': 'Install failed: {message}',
// Config Tab — Model Providers
'config.modelProviders': 'Model Providers',
'config.addProvider': '+ Add Provider',
'config.newProvider': 'New Provider',
'config.restartNotice': 'Changes take effect after restart.',
'config.builtin': 'built-in',
'config.useProvider': 'Use',
'config.configureProvider': 'Configure',
'config.providerConfigured': 'Provider "{name}" configured (restart to apply)',
'config.currentModel': 'Model: {model}',
'config.providerName': 'Display Name',
'config.providerNamePlaceholder': 'My Provider',
'config.providerId': 'Provider ID',
'config.providerIdPlaceholder': 'my-provider',
'config.providerIdHint': 'Lowercase letters, numbers, hyphens',
'config.providerAdapter': 'API Adapter',
'config.adapterOpenAI': 'OpenAI Compatible',
'config.adapterAnthropic': 'Anthropic',
'config.adapterOllama': 'Ollama',
'config.providerBaseUrl': 'Base URL',
'config.providerApiKey': 'API Key',
'config.providerModel': 'Default Model',
'config.providerActivated': 'Switched to {name} (restart to apply)',
'config.providerAdded': 'Added provider "{name}" (restart to apply)',
'config.providerUpdated': 'Provider "{name}" updated (restart to apply)',
'config.editProvider': 'Edit Provider',
'config.providerDeleted': 'Provider deleted',
'config.confirmDeleteProvider': 'Delete provider "{id}"?',
'config.cannotDeleteActiveProvider': 'Cannot delete the active provider. Switch to another provider first.',
'config.testConnection': 'Test',
'config.testing': 'Testing…',
'config.fetchModels': 'Fetch available models',
'config.modelsFetched': '{count} model(s) loaded — type to filter',
'config.modelsFetchFailed': 'Failed to fetch models',
'config.providerBaseUrlRequired': 'Base URL is required to fetch models',
'config.providerFieldsRequired': 'Display name and Provider ID are required',
'config.providerIdInvalid': 'Provider ID: use only lowercase letters, numbers, hyphens',
'config.providerIdTaken': 'Provider ID "{id}" is already taken',
// Configure // Configure
'config.title': 'Configure {name}', 'config.title': 'Configure {name}',
'config.telegramOwnerHint': 'After saving, IronClaw will show a one-time code. Send `/start CODE` to your bot in Telegram and IronClaw will finish setup automatically.', 'config.telegramOwnerHint': 'After saving, IronClaw will show a one-time code. Send `/start CODE` to your bot in Telegram and IronClaw will finish setup automatically.',
-41
View File
@@ -38,14 +38,12 @@ I18n.register('zh-CN', {
'tab.settings': '设置', 'tab.settings': '设置',
'tab.extensions': '扩展', 'tab.extensions': '扩展',
'tab.skills': '技能', 'tab.skills': '技能',
'tab.config': '配置',
'tab.logs': '日志', 'tab.logs': '日志',
'settings.inference': '推理', 'settings.inference': '推理',
'settings.agent': '代理', 'settings.agent': '代理',
'settings.channels': '频道', 'settings.channels': '频道',
'settings.networking': '网络', 'settings.networking': '网络',
'settings.mcp': 'MCP', 'settings.mcp': 'MCP',
'settings.providers': '模型提供商',
// 状态 // 状态
'status.connected': '已连接', 'status.connected': '已连接',
@@ -352,45 +350,6 @@ I18n.register('zh-CN', {
'ext.removed': '已移除 {name}', 'ext.removed': '已移除 {name}',
'ext.installFailed': '安装失败: {message}', 'ext.installFailed': '安装失败: {message}',
// 配置页 — 模型提供商
'config.modelProviders': '模型提供商',
'config.addProvider': '+ 添加提供商',
'config.newProvider': '新建提供商',
'config.restartNotice': '更改将在重启后生效。',
'config.builtin': '内置',
'config.useProvider': '使用',
'config.configureProvider': '配置',
'config.providerConfigured': '提供商 "{name}" 已配置(重启后生效)',
'config.currentModel': '模型:{model}',
'config.providerName': '显示名称',
'config.providerNamePlaceholder': '我的提供商',
'config.providerId': '提供商 ID',
'config.providerIdPlaceholder': 'my-provider',
'config.providerIdHint': '小写字母、数字、连字符',
'config.providerAdapter': 'API 适配器',
'config.adapterOpenAI': 'OpenAI 兼容',
'config.adapterAnthropic': 'Anthropic',
'config.adapterOllama': 'Ollama',
'config.providerBaseUrl': '基础 URL',
'config.providerApiKey': 'API 密钥',
'config.providerModel': '默认模型',
'config.providerActivated': '已切换到 {name}(重启后生效)',
'config.providerAdded': '已添加提供商 "{name}"(重启后生效)',
'config.providerUpdated': '提供商 "{name}" 已更新(重启后生效)',
'config.editProvider': '编辑提供商',
'config.providerDeleted': '提供商已删除',
'config.confirmDeleteProvider': '确定删除提供商 "{id}"',
'config.cannotDeleteActiveProvider': '无法删除当前正在使用的提供商,请先切换到其他提供商。',
'config.testConnection': '测试',
'config.testing': '测试中…',
'config.fetchModels': '获取可用模型',
'config.modelsFetched': '已加载 {count} 个模型,可输入过滤',
'config.modelsFetchFailed': '获取模型列表失败',
'config.providerBaseUrlRequired': '请先填写 Base URL',
'config.providerFieldsRequired': '显示名称和提供商 ID 为必填项',
'config.providerIdInvalid': '提供商 ID 只能包含小写字母、数字和连字符',
'config.providerIdTaken': '提供商 ID "{id}" 已被占用',
// 配置 // 配置
'config.title': '配置 {name}', 'config.title': '配置 {name}',
'config.telegramOwnerHint': '保存后,IronClaw 会显示一次性验证码。将 `/start CODE` 发送给你的 Telegram 机器人,IronClaw 会自动完成设置。', 'config.telegramOwnerHint': '保存后,IronClaw 会显示一次性验证码。将 `/start CODE` 发送给你的 Telegram 机器人,IronClaw 会自动完成设置。',
+2 -70
View File
@@ -45,58 +45,6 @@
</div> </div>
</div> </div>
<!-- Provider Add/Edit Dialog -->
<div id="provider-dialog" class="provider-dialog" style="display:none">
<div class="provider-dialog-overlay" id="provider-dialog-overlay"></div>
<div class="provider-dialog-content">
<div class="provider-dialog-header">
<h2 id="provider-form-title" data-i18n="config.newProvider">New Provider</h2>
<button class="provider-dialog-close" id="cancel-provider-btn" title="Close">×</button>
</div>
<div class="provider-dialog-body">
<div class="config-form">
<div class="config-form-row" id="provider-name-row">
<label data-i18n="config.providerName">Display Name</label>
<input type="text" id="provider-name" data-i18n="config.providerNamePlaceholder" data-i18n-attr="placeholder" placeholder="My Provider">
</div>
<div class="config-form-row" id="provider-id-row">
<label data-i18n="config.providerId">Provider ID</label>
<input type="text" id="provider-id" data-i18n="config.providerIdPlaceholder" data-i18n-attr="placeholder" placeholder="my-provider">
<span class="config-form-hint" data-i18n="config.providerIdHint">Lowercase letters, numbers, hyphens</span>
</div>
<div class="config-form-row" id="provider-adapter-row">
<label data-i18n="config.providerAdapter">API Adapter</label>
<select id="provider-adapter">
<option value="open_ai_completions" data-i18n="config.adapterOpenAI">OpenAI Compatible</option>
<option value="anthropic" data-i18n="config.adapterAnthropic">Anthropic</option>
<option value="ollama" data-i18n="config.adapterOllama">Ollama</option>
</select>
</div>
<div class="config-form-row" id="provider-base-url-row">
<label data-i18n="config.providerBaseUrl">Base URL</label>
<input type="text" id="provider-base-url" placeholder="https://api.example.com/v1">
</div>
<div class="config-form-row" id="provider-api-key-row">
<label data-i18n="config.providerApiKey">API Key</label>
<input type="password" id="provider-api-key" placeholder="sk-...">
</div>
<div class="config-form-row">
<label data-i18n="config.providerModel">Default Model</label>
<input type="text" id="provider-model" placeholder="gpt-4o">
<button id="fetch-models-btn" class="btn-fetch-models" type="button" data-i18n="config.fetchModels">↻ Fetch available models</button>
<select id="provider-model-select" style="display:none"></select>
</div>
<div id="test-connection-result" class="test-connection-result" style="display:none"></div>
</div>
</div>
<div class="provider-dialog-footer">
<button id="save-provider-btn" data-i18n="common.save">Save</button>
<button id="test-provider-btn" class="btn-secondary" data-i18n="config.testConnection">Test</button>
<button id="cancel-provider-footer-btn" class="btn-secondary" data-i18n="common.cancel">Cancel</button>
</div>
</div>
</div>
<!-- Restart Confirmation Modal --> <!-- Restart Confirmation Modal -->
<div id="restart-confirm-modal" class="restart-modal" style="display: none;"> <div id="restart-confirm-modal" class="restart-modal" style="display: none;">
<div class="restart-modal-overlay" id="restart-overlay"></div> <div class="restart-modal-overlay" id="restart-overlay"></div>
@@ -357,23 +305,8 @@
<button id="settings-import-btn" class="settings-toolbar-btn" data-i18n="settings.import">Import</button> <button id="settings-import-btn" class="settings-toolbar-btn" data-i18n="settings.import">Import</button>
</div> </div>
<div class="settings-subpanel active" id="settings-inference"> <div class="settings-subpanel active" id="settings-inference">
<div class="extensions-container"> <div class="extensions-container" id="settings-inference-content">
<div id="settings-inference-content"> <div class="empty-state" data-i18n="common.loading">Loading settings...</div>
<div class="empty-state" data-i18n="common.loading">Loading settings...</div>
</div>
<div class="extensions-section" id="providers-section">
<div class="config-section-header">
<h3 data-i18n="config.modelProviders">Model Providers</h3>
<button id="add-provider-btn" class="btn-add-provider" data-i18n="config.addProvider">+ Add Provider</button>
</div>
<div class="config-notice" id="config-restart-notice" style="display:none">
<span></span>
<span data-i18n="config.restartNotice">Changes take effect after restart.</span>
</div>
<div id="providers-list" class="providers-list">
<div class="empty-state" data-i18n="common.loading">Loading...</div>
</div>
</div>
</div> </div>
</div> </div>
<div class="settings-subpanel" id="settings-agent"> <div class="settings-subpanel" id="settings-agent">
@@ -475,7 +408,6 @@
</div> </div>
<div id="toasts"></div> <div id="toasts"></div>
<script src="/providers.js"></script>
<script src="/app.js"></script> <script src="/app.js"></script>
<script src="/i18n-app.js"></script> <script src="/i18n-app.js"></script>
</body> </body>
-37
View File
@@ -1,37 +0,0 @@
// Built-in LLM provider definitions.
// Generated from providers.json + nearai/bedrock (handled separately in llm.rs)
// Fields: id, name, adapter, base_url, builtin, default_model, api_key_required, can_list_models
// nearai/bedrock use special auth flows — no Configure button (api_key_required=false, can_list_models=false)
const BUILTIN_PROVIDERS = [
{ id: 'nearai', name: 'NEAR AI', adapter: 'nearai', base_url: 'https://cloud-api.near.ai/v1', builtin: true, default_model: 'zai-org/GLM-5-FP8', api_key_required: true, can_list_models: true },
{ id: 'openai', name: 'OpenAI', adapter: 'open_ai_completions', base_url: 'https://api.openai.com/v1', builtin: true, default_model: 'gpt-4o-mini', api_key_required: true, can_list_models: true },
{ id: 'anthropic', name: 'Anthropic', adapter: 'anthropic', base_url: 'https://api.anthropic.com', builtin: true, default_model: 'claude-sonnet-4-20250514', api_key_required: true, can_list_models: true },
{ id: 'ollama', name: 'Ollama', adapter: 'ollama', base_url: 'http://localhost:11434', builtin: true, default_model: 'llama3', api_key_required: false, can_list_models: true },
{ id: 'openai_compatible', name: 'OpenAI Compatible', adapter: 'open_ai_completions', base_url: '', builtin: true, default_model: 'default', api_key_required: false, can_list_models: false },
{ id: 'gemini', name: 'Google Gemini', adapter: 'open_ai_completions', base_url: 'https://generativelanguage.googleapis.com/v1beta/openai', builtin: true, default_model: 'gemini-2.5-flash', api_key_required: true, can_list_models: true },
{ id: 'groq', name: 'Groq', adapter: 'open_ai_completions', base_url: 'https://api.groq.com/openai/v1', builtin: true, default_model: 'llama-3.3-70b-versatile', api_key_required: true, can_list_models: true },
{ id: 'openrouter', name: 'OpenRouter', adapter: 'open_ai_completions', base_url: 'https://openrouter.ai/api/v1', builtin: true, default_model: 'openai/gpt-4o', api_key_required: true, can_list_models: false },
{ id: 'deepseek', name: 'DeepSeek', adapter: 'open_ai_completions', base_url: 'https://api.deepseek.com/v1', builtin: true, default_model: 'deepseek-chat', api_key_required: true, can_list_models: false },
{ id: 'mistral', name: 'Mistral', adapter: 'open_ai_completions', base_url: 'https://api.mistral.ai/v1', builtin: true, default_model: 'mistral-large-latest', api_key_required: true, can_list_models: true },
{ id: 'tinfoil', name: 'Tinfoil', adapter: 'open_ai_completions', base_url: 'https://inference.tinfoil.sh/v1', builtin: true, default_model: 'kimi-k2-5', api_key_required: true, can_list_models: false },
{ id: 'nvidia', name: 'NVIDIA NIM', adapter: 'open_ai_completions', base_url: 'https://integrate.api.nvidia.com/v1', builtin: true, default_model: 'meta/llama-3.3-70b-instruct', api_key_required: true, can_list_models: true },
{ id: 'together', name: 'Together AI', adapter: 'open_ai_completions', base_url: 'https://api.together.xyz/v1', builtin: true, default_model: 'meta-llama/Llama-3-70b-chat-hf', api_key_required: true, can_list_models: false },
{ id: 'fireworks', name: 'Fireworks AI', adapter: 'open_ai_completions', base_url: 'https://api.fireworks.ai/inference/v1', builtin: true, default_model: 'accounts/fireworks/models/llama-v3p1-70b-instruct', api_key_required: true, can_list_models: false },
{ id: 'cerebras', name: 'Cerebras', adapter: 'open_ai_completions', base_url: 'https://api.cerebras.ai/v1', builtin: true, default_model: 'llama-3.3-70b', api_key_required: true, can_list_models: false },
{ id: 'sambanova', name: 'SambaNova', adapter: 'open_ai_completions', base_url: 'https://api.sambanova.ai/v1', builtin: true, default_model: 'Meta-Llama-3.1-70B-Instruct', api_key_required: true, can_list_models: false },
{ id: 'zai', name: 'Z.AI', adapter: 'open_ai_completions', base_url: 'https://api.z.ai/api/paas/v4', builtin: true, default_model: 'glm-5', api_key_required: true, can_list_models: false },
{ id: 'venice', name: 'Venice.ai', adapter: 'open_ai_completions', base_url: 'https://api.venice.ai/api/v1', builtin: true, default_model: 'llama-3.3-70b', api_key_required: true, can_list_models: false },
{ id: 'minimax', name: 'MiniMax', adapter: 'open_ai_completions', base_url: 'https://api.minimax.io/v1', builtin: true, default_model: 'MiniMax-M2.5', api_key_required: true, can_list_models: false },
{ id: 'ionet', name: 'io.net', adapter: 'open_ai_completions', base_url: 'https://api.intelligence.io.solutions/api/v1', builtin: true, default_model: 'deepseek-coder-v2-instruct', api_key_required: true, can_list_models: true },
{ id: 'cloudflare', name: 'Cloudflare AI', adapter: 'open_ai_completions', base_url: '', builtin: true, default_model: '@cf/meta/llama-3.3-70b-instruct-fp8-fast', api_key_required: true, can_list_models: false },
{ id: 'yandex', name: 'Yandex AI Studio', adapter: 'open_ai_completions', base_url: 'https://ai.api.cloud.yandex.net/v1', builtin: true, default_model: 'yandexgpt-lite', api_key_required: true, can_list_models: true },
{ id: 'bedrock', name: 'AWS Bedrock', adapter: 'bedrock', base_url: '', builtin: true, default_model: 'anthropic.claude-3-sonnet-20240229-v1:0', api_key_required: false, can_list_models: false },
];
const ADAPTER_LABELS = {
open_ai_completions: 'OpenAI Compatible',
anthropic: 'Anthropic',
ollama: 'Ollama',
bedrock: 'AWS Bedrock',
nearai: 'NEAR AI',
};
-420
View File
@@ -2801,22 +2801,10 @@ body {
padding: var(--space-4); padding: var(--space-4);
} }
#settings-inference > .extensions-container {
display: flex;
flex-direction: column;
}
.extensions-section { .extensions-section {
margin-bottom: 24px; margin-bottom: 24px;
} }
#providers-section {
flex: 1;
min-height: 0;
display: flex;
flex-direction: column;
}
.extensions-section h3 { .extensions-section h3 {
font-size: var(--text-xs); font-size: var(--text-xs);
font-weight: 600; font-weight: 600;
@@ -4605,12 +4593,6 @@ mark {
min-width: 180px; min-width: 180px;
} }
.settings-display-value {
font-size: var(--text-sm);
color: var(--text);
font-family: 'IBM Plex Mono', monospace;
}
.settings-input { .settings-input {
padding: 6px 10px; padding: 6px 10px;
background: var(--bg); background: var(--bg);
@@ -5447,405 +5429,3 @@ body.theme-transition *:not(svg):not(path):not(line):not(circle):not(rect) {
--text-muted: #a1a1aa; --text-muted: #a1a1aa;
} }
} }
/* --- Config Tab --- */
.config-section-header {
display: flex;
align-items: center;
justify-content: space-between;
margin-bottom: 12px;
}
.config-section-header h3 {
margin-bottom: 0;
}
.btn-add-provider {
padding: 5px 14px;
background: var(--accent);
color: #09090b;
border: none;
border-radius: var(--radius);
cursor: pointer;
font-size: 13px;
font-weight: 600;
transition: background 0.2s, transform 0.2s;
}
.btn-add-provider:hover {
background: var(--accent-hover);
transform: translateY(-1px);
}
.config-notice {
display: flex;
align-items: center;
gap: 8px;
padding: 8px 12px;
background: rgba(245, 166, 35, 0.1);
border: 1px solid rgba(245, 166, 35, 0.3);
border-radius: var(--radius);
color: var(--warning);
font-size: 13px;
margin-bottom: 12px;
}
.providers-list {
display: flex;
flex-direction: column;
gap: 8px;
min-height: 420px;
overflow-y: auto;
}
.provider-card {
background: var(--bg-secondary);
border: 1px solid var(--border);
border-radius: var(--radius-lg);
padding: 12px 14px;
display: flex;
flex-direction: column;
gap: 6px;
transition: border-color 0.2s;
}
.provider-card:hover {
border-color: rgba(255, 255, 255, 0.15);
}
.provider-card-active {
border-color: var(--accent);
}
.provider-card-header {
display: flex;
align-items: center;
gap: 8px;
flex-wrap: wrap;
}
.provider-name {
font-weight: 600;
font-size: 14px;
color: var(--text);
}
.provider-id-label {
font-size: 11px;
color: var(--text-secondary);
font-family: var(--font-mono);
}
.provider-badge {
font-size: 10px;
padding: 2px 7px;
border-radius: 20px;
font-weight: 600;
letter-spacing: 0.02em;
}
.provider-badge-active {
background: rgba(52, 211, 153, 0.15);
color: var(--accent);
}
.provider-badge-builtin {
background: rgba(161, 161, 170, 0.12);
color: var(--text-secondary);
}
.provider-card-meta {
display: flex;
align-items: center;
gap: 10px;
flex-wrap: wrap;
}
.provider-adapter {
font-size: 12px;
color: var(--text-secondary);
}
.provider-url {
font-size: 11px;
color: var(--text-secondary);
font-family: var(--font-mono);
opacity: 0.7;
}
.provider-current-model {
font-size: 11px;
color: var(--accent);
font-family: var(--font-mono);
font-weight: 500;
}
.provider-card-actions {
display: flex;
gap: 6px;
margin-top: 2px;
}
.provider-action-btn {
padding: 4px 12px;
background: var(--bg-tertiary);
border: 1px solid var(--border);
border-radius: var(--radius);
color: var(--text-secondary);
cursor: pointer;
font-size: 12px;
transition: color 0.2s, border-color 0.2s, background 0.2s;
}
.provider-action-btn:hover {
color: var(--text);
border-color: rgba(255, 255, 255, 0.2);
background: var(--bg);
}
.provider-delete-btn:hover {
color: var(--danger);
border-color: var(--danger);
}
/* Config form */
.provider-dialog {
position: fixed;
top: 0;
left: 0;
right: 0;
bottom: 0;
z-index: 9999;
display: flex;
align-items: center;
justify-content: center;
}
.provider-dialog-overlay {
position: absolute;
top: 0;
left: 0;
right: 0;
bottom: 0;
background: rgba(0, 0, 0, 0.5);
backdrop-filter: blur(4px);
}
.provider-dialog-content {
position: relative;
z-index: 10000;
background: var(--bg-secondary);
border: 1px solid var(--border);
border-radius: var(--radius-lg);
box-shadow: 0 25px 50px -12px rgba(0, 0, 0, 0.4);
width: 100%;
max-width: 480px;
margin: 0 1rem;
display: flex;
flex-direction: column;
max-height: 90vh;
}
.provider-dialog-header {
display: flex;
align-items: center;
justify-content: space-between;
padding: 14px 18px;
border-bottom: 1px solid var(--border);
flex-shrink: 0;
}
.provider-dialog-header h2 {
font-size: 14px;
font-weight: 600;
color: var(--text);
margin: 0;
}
.provider-dialog-close {
color: var(--text-secondary);
font-size: 18px;
line-height: 1;
padding: 2px 6px;
background: transparent;
border: none;
border-radius: var(--radius);
cursor: pointer;
transition: color 0.15s, background 0.15s;
}
.provider-dialog-close:hover {
color: var(--text);
background: var(--bg-hover);
}
.provider-dialog-body {
padding: 18px;
overflow-y: auto;
flex: 1;
}
.provider-dialog-footer {
display: flex;
gap: 8px;
padding: 14px 18px;
border-top: 1px solid var(--border);
flex-shrink: 0;
}
.provider-dialog-footer button {
padding: 6px 18px;
border-radius: var(--radius);
font-size: 13px;
font-weight: 600;
cursor: pointer;
transition: background 0.2s, transform 0.2s;
}
.provider-dialog-footer button:first-child {
background: var(--accent);
color: #09090b;
border: none;
}
.provider-dialog-footer button:first-child:hover {
background: var(--accent-hover);
transform: translateY(-1px);
}
.provider-dialog-footer .btn-secondary {
background: transparent;
color: var(--text-secondary);
border: 1px solid var(--border);
}
.provider-dialog-footer .btn-secondary:hover {
color: var(--text);
border-color: rgba(255, 255, 255, 0.2);
}
.config-form {
display: flex;
flex-direction: column;
gap: 12px;
}
.config-form-row {
display: flex;
flex-direction: column;
gap: 4px;
}
.config-form-row label {
font-size: 12px;
font-weight: 500;
color: var(--text-secondary);
}
.config-form-row input,
.config-form-row select {
padding: 7px 10px;
background: var(--bg);
border: 1px solid var(--border);
border-radius: var(--radius);
color: var(--text);
font-size: 13px;
}
.config-form-row input:focus,
.config-form-row select:focus {
outline: none;
border-color: var(--accent);
box-shadow: 0 0 0 3px rgba(52, 211, 153, 0.1);
}
.config-form-hint {
font-size: 11px;
color: var(--text-secondary);
opacity: 0.7;
}
.config-form-actions {
display: flex;
gap: 8px;
margin-top: 4px;
}
.config-form-actions button {
padding: 6px 18px;
border-radius: var(--radius);
font-size: 13px;
font-weight: 600;
cursor: pointer;
transition: background 0.2s, transform 0.2s;
}
.config-form-actions button:first-child {
background: var(--accent);
color: #09090b;
border: none;
}
.config-form-actions button:first-child:hover {
background: var(--accent-hover);
transform: translateY(-1px);
}
.config-form-actions .btn-secondary {
background: transparent;
color: var(--text-secondary);
border: 1px solid var(--border);
}
.config-form-actions .btn-secondary:hover {
color: var(--text);
border-color: rgba(255, 255, 255, 0.2);
}
.btn-fetch-models {
display: inline-flex;
align-items: center;
gap: 5px;
margin-top: 6px;
padding: 5px 11px;
background: transparent;
border: 1px solid var(--border);
border-radius: var(--radius);
color: var(--text-secondary);
cursor: pointer;
font-size: 12px;
transition: color 0.15s, border-color 0.15s, background 0.15s;
}
.btn-fetch-models:hover {
color: var(--text);
border-color: var(--accent);
background: color-mix(in srgb, var(--accent) 8%, transparent);
}
.btn-fetch-models:disabled {
opacity: 0.5;
cursor: not-allowed;
}
.test-connection-result {
margin-top: 8px;
padding: 6px 12px;
border-radius: var(--radius);
font-size: 13px;
}
.test-connection-result.test-ok {
background: rgba(74, 222, 128, 0.12);
color: #4ade80;
border: 1px solid rgba(74, 222, 128, 0.3);
}
.test-connection-result.test-fail {
background: rgba(248, 113, 113, 0.12);
color: #f87171;
border: 1px solid rgba(248, 113, 113, 0.3);
}
+1 -3
View File
@@ -76,8 +76,7 @@ impl TestGatewayBuilder {
store: None, store: None,
job_manager: None, job_manager: None,
prompt_queue: None, prompt_queue: None,
owner_id: self.user_id.clone(), default_user_id: self.user_id,
default_sender_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,
@@ -92,7 +91,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,
}) })
} }
+1 -39
View File
@@ -16,7 +16,6 @@ use axum::routing::{delete, get, post};
use tower::ServiceExt; use tower::ServiceExt;
use uuid::Uuid; use uuid::Uuid;
use crate::channels::web::GatewayChannel;
use crate::channels::web::auth::{ use crate::channels::web::auth::{
AuthenticatedUser, MultiAuthState, UserIdentity, auth_middleware, AuthenticatedUser, MultiAuthState, UserIdentity, auth_middleware,
}; };
@@ -24,7 +23,6 @@ use crate::channels::web::server::{
ActiveConfigSnapshot, GatewayState, PerUserRateLimiter, PromptQueue, RateLimiter, WorkspacePool, ActiveConfigSnapshot, GatewayState, PerUserRateLimiter, PromptQueue, RateLimiter, WorkspacePool,
}; };
use crate::channels::web::sse::SseManager; use crate::channels::web::sse::SseManager;
use crate::config::GatewayConfig;
// ── Helpers ──────────────────────────────────────────────────────────── // ── Helpers ────────────────────────────────────────────────────────────
@@ -66,8 +64,7 @@ fn build_state(
store, store,
job_manager: None, job_manager: None,
prompt_queue, prompt_queue,
owner_id: "test".to_string(), default_user_id: "test".to_string(),
default_sender_id: "test".to_string(),
shutdown_tx: tokio::sync::RwLock::new(None), shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: None, ws_tracker: None,
llm_provider: None, llm_provider: None,
@@ -82,44 +79,9 @@ fn build_state(
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: ActiveConfigSnapshot::default(), active_config: ActiveConfigSnapshot::default(),
secrets_store: None,
}) })
} }
fn gateway_config() -> GatewayConfig {
GatewayConfig {
host: "127.0.0.1".to_string(),
port: 3000,
auth_token: Some("gateway-auth".to_string()),
user_id: "gateway-sender".to_string(),
workspace_read_scopes: Vec::new(),
memory_layers: Vec::new(),
user_tokens: None,
}
}
#[test]
fn with_owner_scope_updates_gateway_owner_scope_in_multi_user_mode() {
let mut gateway = GatewayChannel::new(gateway_config());
gateway.auth = two_user_auth();
gateway.config.user_tokens = Some(HashMap::new());
let gateway = gateway.with_owner_scope("owner-scope");
assert_eq!(gateway.state.owner_id, "owner-scope");
assert_eq!(gateway.state.default_sender_id, "gateway-sender");
let alice = gateway
.auth
.authenticate("tok-alice")
.expect("alice token should remain valid");
let bob = gateway
.auth
.authenticate("tok-bob")
.expect("bob token should remain valid");
assert_eq!(alice.user_id, "alice");
assert_eq!(bob.user_id, "bob");
}
/// Create a libSQL-backed test database in a temporary directory. /// Create a libSQL-backed test database in a temporary directory.
/// ///
/// Returns the database and a `TempDir` guard — the database file is /// Returns the database and a `TempDir` guard — the database file is
+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");
+114 -113
View File
@@ -2,26 +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]);
/// Convert stored tool errors into plain text suitable for UI display. // Re-close <tool_output> if truncation cut through the closing tag.
pub fn tool_error_for_display(error: &str) -> String { if s.starts_with("<tool_output") && !result.ends_with("</tool_output>") {
ironclaw_safety::SafetyLayer::unwrap_tool_output(error).unwrap_or_else(|| error.to_string()) result.push_str("\n</tool_output>");
} }
/// Parse tool call summary JSON objects into `ToolCallInfo` structs. result
fn parse_tool_call_infos(calls: &[serde_json::Value]) -> Vec<ToolCallInfo> {
calls
.iter()
.map(|c| ToolCallInfo {
name: c["name"].as_str().unwrap_or("unknown").to_string(),
has_result: c.get("result_preview").is_some_and(|v| !v.is_null()),
has_error: c.get("error").is_some_and(|v| !v.is_null()),
result_preview: c["result_preview"].as_str().map(String::from),
error: c["error"].as_str().map(tool_error_for_display),
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).
@@ -47,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
@@ -55,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!(
@@ -114,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;
} }
@@ -128,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 {
@@ -186,29 +258,6 @@ mod tests {
assert_eq!(turns[0].response.as_deref(), Some("Done")); assert_eq!(turns[0].response.as_deref(), Some("Done"));
} }
#[test]
fn test_build_turns_unwrap_wrapped_tool_error_for_display() {
let tc_json = serde_json::json!([
{
"name": "http",
"error": "<tool_output name=\"http\">\nTool 'http' failed: timeout\n</tool_output>"
}
]);
let messages = vec![
make_msg("user", "Run it", 0),
make_msg("tool_calls", &tc_json.to_string(), 500),
];
let turns = build_turns_from_db_messages(&messages);
assert_eq!(turns.len(), 1);
assert_eq!(turns[0].tool_calls.len(), 1);
assert_eq!(
turns[0].tool_calls[0].error.as_deref(),
Some("Tool 'http' failed: timeout")
);
}
#[test] #[test]
fn test_build_turns_malformed_tool_calls() { fn test_build_turns_malformed_tool_calls() {
let messages = vec![ let messages = vec![
@@ -256,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);
}
} }
+5 -7
View File
@@ -97,7 +97,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
} }
} }
@@ -275,7 +275,7 @@ async fn handle_client_message(
if result.verification.is_some() { if result.verification.is_some() {
state.sse.broadcast_for_user( state.sse.broadcast_for_user(
user_id, user_id,
crate::channels::web::types::AppEvent::AuthRequired { crate::channels::web::types::SseEvent::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,
@@ -286,7 +286,7 @@ async fn handle_client_message(
crate::channels::web::server::clear_auth_mode(state, user_id).await; crate::channels::web::server::clear_auth_mode(state, user_id).await;
state.sse.broadcast_for_user( state.sse.broadcast_for_user(
user_id, user_id,
crate::channels::web::types::AppEvent::AuthCompleted { crate::channels::web::types::SseEvent::AuthCompleted {
extension_name, extension_name,
success: true, success: true,
message: result.message, message: result.message,
@@ -299,7 +299,7 @@ async fn handle_client_message(
if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) { if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) {
state.sse.broadcast_for_user( state.sse.broadcast_for_user(
user_id, user_id,
crate::channels::web::types::AppEvent::AuthRequired { crate::channels::web::types::SseEvent::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,
@@ -520,8 +520,7 @@ mod tests {
job_manager: None, job_manager: None,
prompt_queue: None, prompt_queue: None,
scheduler: None, scheduler: None,
owner_id: "test".to_string(), default_user_id: "test".to_string(),
default_sender_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,
@@ -535,7 +534,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,
} }
} }
} }
+5 -7
View File
@@ -80,7 +80,7 @@ pub async fn run_doctor_command() -> anyhow::Result<()> {
check( check(
"Routines config", "Routines config",
check_routines_config(&settings), check_routines_config(),
&mut passed, &mut passed,
&mut failed, &mut failed,
&mut skipped, &mut skipped,
@@ -434,8 +434,8 @@ fn check_embeddings(settings: &Settings) -> CheckResult {
// ── Routines config ───────────────────────────────────────── // ── Routines config ─────────────────────────────────────────
fn check_routines_config(settings: &Settings) -> CheckResult { fn check_routines_config() -> CheckResult {
match crate::config::RoutineConfig::resolve(settings) { match crate::config::RoutineConfig::resolve() {
Ok(config) => { Ok(config) => {
if config.enabled { if config.enabled {
CheckResult::Pass(format!( CheckResult::Pass(format!(
@@ -737,8 +737,7 @@ mod tests {
#[test] #[test]
fn check_routines_config_does_not_panic() { fn check_routines_config_does_not_panic() {
let settings = Settings::default(); let result = check_routines_config();
let result = check_routines_config(&settings);
match result { match result {
CheckResult::Pass(_) | CheckResult::Fail(_) | CheckResult::Skip(_) => {} CheckResult::Pass(_) | CheckResult::Fail(_) | CheckResult::Skip(_) => {}
} }
@@ -867,8 +866,7 @@ mod tests {
unsafe { unsafe {
std::env::remove_var("ROUTINES_ENABLED"); std::env::remove_var("ROUTINES_ENABLED");
} }
let settings = Settings::default(); match check_routines_config() {
match check_routines_config(&settings) {
CheckResult::Pass(msg) => { CheckResult::Pass(msg) => {
assert!( assert!(
msg.contains("enabled"), msg.contains("enabled"),
+27 -582
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`
@@ -473,8 +449,7 @@ pub struct PendingOAuthFlow {
pub secrets: Arc<dyn SecretsStore + Send + Sync>, pub secrets: Arc<dyn SecretsStore + Send + Sync>,
/// SSE broadcast manager for notifying the web UI. /// SSE broadcast manager for notifying the web UI.
pub sse_manager: Option<Arc<crate::channels::web::sse::SseManager>>, pub sse_manager: Option<Arc<crate::channels::web::sse::SseManager>>,
/// OAuth proxy auth token for authenticating with the hosted token exchange proxy. /// Gateway auth token for authenticating with the platform token exchange proxy.
/// Kept as `gateway_token` for public API compatibility.
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.
/// Used for provider-specific requirements such as RFC 8707 `resource`. /// Used for provider-specific requirements such as RFC 8707 `resource`.
@@ -497,12 +472,6 @@ impl std::fmt::Debug for PendingOAuthFlow {
} }
} }
impl PendingOAuthFlow {
pub fn oauth_proxy_auth_token(&self) -> Option<&str> {
self.gateway_token.as_deref()
}
}
/// Thread-safe registry of pending OAuth flows, keyed by CSRF `state` parameter. /// Thread-safe registry of pending OAuth flows, keyed by CSRF `state` parameter.
pub type PendingOAuthRegistry = Arc<RwLock<HashMap<String, PendingOAuthFlow>>>; pub type PendingOAuthRegistry = Arc<RwLock<HashMap<String, PendingOAuthFlow>>>;
@@ -536,22 +505,6 @@ pub fn exchange_proxy_url() -> Option<String> {
.filter(|url| !url.is_empty()) .filter(|url| !url.is_empty())
} }
/// Returns the configured OAuth proxy auth token, if any.
///
/// New hosted infra can inject a dedicated shared proxy secret via
/// `IRONCLAW_OAUTH_PROXY_AUTH_TOKEN`. Existing hosted instances continue to
/// work by falling back to `GATEWAY_AUTH_TOKEN`.
pub fn oauth_proxy_auth_token() -> Option<String> {
fn normalized_env_value(key: &str) -> Option<String> {
crate::config::helpers::env_or_override(key)
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
}
normalized_env_value("IRONCLAW_OAUTH_PROXY_AUTH_TOKEN")
.or_else(|| normalized_env_value("GATEWAY_AUTH_TOKEN"))
}
/// Maximum age for pending OAuth flows (5 minutes, matching TCP listener timeout). /// Maximum age for pending OAuth flows (5 minutes, matching TCP listener timeout).
pub const OAUTH_FLOW_EXPIRY: Duration = Duration::from_secs(300); pub const OAUTH_FLOW_EXPIRY: Duration = Duration::from_secs(300);
@@ -697,8 +650,6 @@ pub fn strip_instance_prefix(state: &str) -> &str {
pub struct ProxyTokenExchangeRequest<'a> { pub struct ProxyTokenExchangeRequest<'a> {
pub proxy_url: &'a str, pub proxy_url: &'a str,
/// OAuth proxy auth token.
/// Kept as `gateway_token` for public API compatibility.
pub gateway_token: &'a str, pub gateway_token: &'a str,
pub token_url: &'a str, pub token_url: &'a str,
pub client_id: &'a str, pub client_id: &'a str,
@@ -710,53 +661,9 @@ 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,
/// OAuth proxy auth token.
/// Kept as `gateway_token` for public API compatibility.
pub gateway_token: &'a str,
pub token_url: &'a str,
pub client_id: &'a str,
pub client_secret: Option<&'a str>,
pub refresh_token: &'a str,
pub provider: Option<&'a str>,
}
fn oauth_token_response_from_json(
token_data: serde_json::Value,
access_token_field: &str,
) -> Result<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 an OAuth proxy auth token (Bearer header). The caller may /// Authenticated via the gateway auth token (Bearer header). The caller may
/// either rely on proxy-side secret lookup or forward a `client_secret` when /// either rely on proxy-side secret lookup or forward a `client_secret` when
/// the provider requires it. /// the provider requires it.
/// ///
@@ -768,14 +675,13 @@ pub async fn exchange_via_proxy(
) -> Result<OAuthTokenResponse, OAuthCallbackError> { ) -> Result<OAuthTokenResponse, OAuthCallbackError> {
if request.gateway_token.is_empty() { if request.gateway_token.is_empty() {
return Err(OAuthCallbackError::Io( return Err(OAuthCallbackError::Io(
"OAuth proxy auth token is required for proxy token exchange".to_string(), "Gateway auth token is required for proxy token exchange".to_string(),
)); ));
} }
let exchange_url = format!("{}/oauth/exchange", request.proxy_url.trim_end_matches('/')); let exchange_url = format!("{}/oauth/exchange", request.proxy_url.trim_end_matches('/'));
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![
@@ -818,454 +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 an OAuth proxy 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: {:?})",
"OAuth proxy 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();
}
}
}
struct EnvVarGuard {
key: &'static str,
original: Option<String>,
}
impl Drop for EnvVarGuard {
fn drop(&mut self) {
// SAFETY: Under ENV_MUTEX, no concurrent env access.
unsafe {
if let Some(ref value) = self.original {
std::env::set_var(self.key, value);
} else {
std::env::remove_var(self.key);
}
}
}
}
fn set_env_var(key: &'static str, value: Option<&str>) -> EnvVarGuard {
let original = std::env::var(key).ok();
// SAFETY: Under ENV_MUTEX, no concurrent env access.
unsafe {
if let Some(value) = value {
std::env::set_var(key, value);
} else {
std::env::remove_var(key);
}
}
EnvVarGuard { key, original }
}
#[test]
fn test_hosted_proxy_client_secret_suppresses_builtin_secret() {
let builtin = builtin_credentials("google_oauth_token").expect("google builtin creds");
let client_secret = Some(builtin.client_secret.to_string());
let result = super::hosted_proxy_client_secret(&client_secret, Some(&builtin), true);
assert_eq!(result, None);
}
#[test]
fn test_hosted_proxy_client_secret_preserves_explicit_secret() {
let builtin = builtin_credentials("google_oauth_token").expect("google builtin creds");
let client_secret = Some("hosted-server-secret".to_string());
let result = super::hosted_proxy_client_secret(&client_secret, Some(&builtin), true);
assert_eq!(result, client_secret);
}
#[tokio::test]
async fn test_exchange_via_proxy_sends_auth_and_form() {
let server = MockProxyServer::start().await;
let mut extra_token_params = HashMap::new();
extra_token_params.insert("resource".to_string(), "https://mcp.notion.com".to_string());
let response = super::exchange_via_proxy(super::ProxyTokenExchangeRequest {
proxy_url: &server.base_url(),
gateway_token: "shared-oauth-proxy-secret",
code: "auth-code-123",
redirect_uri: "https://oauth.example.com/oauth/callback",
token_url: "https://oauth2.googleapis.com/token",
client_id: TEST_OAUTH_CLIENT_ID,
client_secret: Some(TEST_OAUTH_CLIENT_SECRET),
access_token_field: "access_token",
code_verifier: Some("code-verifier-123"),
extra_token_params: &extra_token_params,
})
.await
.expect("proxy exchange succeeds");
assert_eq!(response.access_token, "proxy-access-token");
assert_eq!(
response.refresh_token.as_deref(),
Some("proxy-refresh-token")
);
assert_eq!(response.expires_in, Some(7200));
let requests = server.requests().await;
assert_eq!(requests.len(), 1);
assert_eq!(
requests[0].authorization.as_deref(),
Some("Bearer shared-oauth-proxy-secret")
);
assert_eq!(
requests[0].form.get("code").map(String::as_str),
Some("auth-code-123")
);
assert_eq!(
requests[0].form.get("redirect_uri").map(String::as_str),
Some("https://oauth.example.com/oauth/callback")
);
assert_eq!(
requests[0].form.get("token_url").map(String::as_str),
Some("https://oauth2.googleapis.com/token")
);
assert_eq!(
requests[0].form.get("client_id").map(String::as_str),
Some(TEST_OAUTH_CLIENT_ID)
);
assert_eq!(
requests[0].form.get("client_secret").map(String::as_str),
Some(TEST_OAUTH_CLIENT_SECRET)
);
assert_eq!(
requests[0]
.form
.get("access_token_field")
.map(String::as_str),
Some("access_token")
);
assert_eq!(
requests[0].form.get("code_verifier").map(String::as_str),
Some("code-verifier-123")
);
assert_eq!(
requests[0].form.get("resource").map(String::as_str),
Some("https://mcp.notion.com")
);
server.shutdown().await;
}
#[tokio::test]
async fn test_refresh_token_via_proxy_sends_auth_and_form() {
let server = MockProxyServer::start().await;
let response = super::refresh_token_via_proxy(super::ProxyRefreshTokenRequest {
proxy_url: &server.base_url(),
gateway_token: "gateway-test-token",
token_url: "https://oauth2.googleapis.com/token",
client_id: TEST_OAUTH_CLIENT_ID,
client_secret: Some(TEST_OAUTH_CLIENT_SECRET),
refresh_token: "refresh-token-123",
provider: Some("google"),
})
.await
.expect("proxy refresh succeeds");
assert_eq!(response.access_token, "proxy-access-token");
assert_eq!(
response.refresh_token.as_deref(),
Some("proxy-refresh-token")
);
assert_eq!(response.expires_in, Some(7200));
let requests = server.requests().await;
assert_eq!(requests.len(), 1);
assert_eq!(
requests[0].authorization.as_deref(),
Some("Bearer gateway-test-token")
);
assert_eq!(
requests[0].form.get("token_url").map(String::as_str),
Some("https://oauth2.googleapis.com/token")
);
assert_eq!(
requests[0].form.get("client_id").map(String::as_str),
Some(TEST_OAUTH_CLIENT_ID)
);
assert_eq!(
requests[0].form.get("client_secret").map(String::as_str),
Some(TEST_OAUTH_CLIENT_SECRET)
);
assert_eq!(
requests[0].form.get("refresh_token").map(String::as_str),
Some("refresh-token-123")
);
assert_eq!(
requests[0].form.get("provider").map(String::as_str),
Some("google")
);
server.shutdown().await;
}
#[tokio::test]
async fn test_exchange_via_proxy_does_not_follow_redirects() {
let server = MockProxyServer::start().await;
let error = match super::exchange_via_proxy(super::ProxyTokenExchangeRequest {
proxy_url: &server.redirecting_base_url(),
gateway_token: "gateway-test-token",
code: "auth-code-123",
redirect_uri: "http://localhost:3000/oauth/callback",
token_url: "https://oauth2.googleapis.com/token",
client_id: TEST_OAUTH_CLIENT_ID,
client_secret: Some(TEST_OAUTH_CLIENT_SECRET),
access_token_field: "access_token",
code_verifier: Some("code-verifier-123"),
extra_token_params: &HashMap::new(),
})
.await
{
Ok(_) => panic!("redirected proxy exchange should fail"),
Err(error) => error,
};
assert!(error.to_string().contains("307"));
assert!(server.requests().await.is_empty());
server.shutdown().await;
}
#[tokio::test]
async fn test_refresh_token_via_proxy_does_not_follow_redirects() {
let server = MockProxyServer::start().await;
let error = match super::refresh_token_via_proxy(super::ProxyRefreshTokenRequest {
proxy_url: &server.redirecting_base_url(),
gateway_token: "gateway-test-token",
token_url: "https://oauth2.googleapis.com/token",
client_id: TEST_OAUTH_CLIENT_ID,
client_secret: Some(TEST_OAUTH_CLIENT_SECRET),
refresh_token: "refresh-token-123",
provider: Some("google"),
})
.await
{
Ok(_) => panic!("redirected proxy refresh should fail"),
Err(error) => error,
};
assert!(error.to_string().contains("307"));
assert!(server.requests().await.is_empty());
server.shutdown().await;
}
#[test] #[test]
fn test_is_loopback_host() { fn test_is_loopback_host() {
@@ -1666,54 +1159,6 @@ mod tests {
} }
} }
#[test]
fn test_oauth_proxy_auth_token_prefers_dedicated_env() {
let _guard = lock_env();
let _proxy_guard = set_env_var(
"IRONCLAW_OAUTH_PROXY_AUTH_TOKEN",
Some("shared-proxy-secret"),
);
let _gateway_guard = set_env_var("GATEWAY_AUTH_TOKEN", Some("gateway-token"));
assert_eq!(
crate::cli::oauth_defaults::oauth_proxy_auth_token().as_deref(),
Some("shared-proxy-secret")
);
}
#[test]
fn test_oauth_proxy_auth_token_falls_back_to_gateway_token() {
let _guard = lock_env();
let _proxy_guard = set_env_var("IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", None);
let _gateway_guard = set_env_var("GATEWAY_AUTH_TOKEN", Some("gateway-token"));
assert_eq!(
crate::cli::oauth_defaults::oauth_proxy_auth_token().as_deref(),
Some("gateway-token")
);
}
#[test]
fn test_oauth_proxy_auth_token_whitespace_dedicated_env_falls_back_to_gateway_token() {
let _guard = lock_env();
let _proxy_guard = set_env_var("IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", Some(" "));
let _gateway_guard = set_env_var("GATEWAY_AUTH_TOKEN", Some("gateway-token"));
assert_eq!(
crate::cli::oauth_defaults::oauth_proxy_auth_token().as_deref(),
Some("gateway-token")
);
}
#[test]
fn test_oauth_proxy_auth_token_returns_none_when_unset() {
let _guard = lock_env();
let _proxy_guard = set_env_var("IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", None);
let _gateway_guard = set_env_var("GATEWAY_AUTH_TOKEN", None);
assert_eq!(crate::cli::oauth_defaults::oauth_proxy_auth_token(), None);
}
#[test] #[test]
fn test_strip_instance_prefix_with_colon() { fn test_strip_instance_prefix_with_colon() {
use crate::cli::oauth_defaults::strip_instance_prefix; use crate::cli::oauth_defaults::strip_instance_prefix;
+8 -19
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
+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());
}
} }
+27 -52
View File
@@ -1,8 +1,6 @@
use std::time::Duration; use std::time::Duration;
use crate::config::helpers::{ use crate::config::helpers::{optional_env, parse_bool_env, parse_option_env, parse_optional_env};
db_first_bool, db_first_or_default, optional_env, parse_bool_env, parse_option_env,
};
use crate::error::ConfigError; use crate::error::ConfigError;
use crate::settings::Settings; use crate::settings::Settings;
@@ -38,10 +36,6 @@ pub struct AgentConfig {
/// Whether the deployment is multi-tenant (multiple users sharing one /// Whether the deployment is multi-tenant (multiple users sharing one
/// instance). Auto-detected from GATEWAY_USER_TOKENS presence. /// instance). Auto-detected from GATEWAY_USER_TOKENS presence.
pub multi_tenant: bool, 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 {
@@ -66,70 +60,53 @@ impl AgentConfig {
default_timezone: "UTC".to_string(), default_timezone: "UTC".to_string(),
max_tokens_per_job: 0, max_tokens_per_job: 0,
multi_tenant: false, multi_tenant: false,
max_llm_concurrent_per_user: None,
max_jobs_concurrent_per_user: None,
} }
} }
pub(crate) fn resolve(settings: &Settings) -> Result<Self, ConfigError> { pub(crate) fn resolve(settings: &Settings) -> Result<Self, ConfigError> {
let defaults = crate::settings::AgentSettings::default();
Ok(Self { Ok(Self {
name: db_first_or_default(&settings.agent.name, &defaults.name, "AGENT_NAME")?, name: parse_optional_env("AGENT_NAME", settings.agent.name.clone())?,
max_parallel_jobs: db_first_or_default( max_parallel_jobs: parse_optional_env(
&(settings.agent.max_parallel_jobs as usize),
&(defaults.max_parallel_jobs as usize),
"AGENT_MAX_PARALLEL_JOBS", "AGENT_MAX_PARALLEL_JOBS",
settings.agent.max_parallel_jobs as usize,
)?, )?,
job_timeout: Duration::from_secs(db_first_or_default( job_timeout: Duration::from_secs(parse_optional_env(
&settings.agent.job_timeout_secs,
&defaults.job_timeout_secs,
"AGENT_JOB_TIMEOUT_SECS", "AGENT_JOB_TIMEOUT_SECS",
settings.agent.job_timeout_secs,
)?), )?),
stuck_threshold: Duration::from_secs(db_first_or_default( stuck_threshold: Duration::from_secs(parse_optional_env(
&settings.agent.stuck_threshold_secs,
&defaults.stuck_threshold_secs,
"AGENT_STUCK_THRESHOLD_SECS", "AGENT_STUCK_THRESHOLD_SECS",
settings.agent.stuck_threshold_secs,
)?), )?),
repair_check_interval: Duration::from_secs(db_first_or_default( repair_check_interval: Duration::from_secs(parse_optional_env(
&settings.agent.repair_check_interval_secs,
&defaults.repair_check_interval_secs,
"SELF_REPAIR_CHECK_INTERVAL_SECS", "SELF_REPAIR_CHECK_INTERVAL_SECS",
settings.agent.repair_check_interval_secs,
)?), )?),
max_repair_attempts: db_first_or_default( max_repair_attempts: parse_optional_env(
&settings.agent.max_repair_attempts,
&defaults.max_repair_attempts,
"SELF_REPAIR_MAX_ATTEMPTS", "SELF_REPAIR_MAX_ATTEMPTS",
settings.agent.max_repair_attempts,
)?, )?,
use_planning: db_first_bool( use_planning: parse_bool_env("AGENT_USE_PLANNING", settings.agent.use_planning)?,
settings.agent.use_planning, session_idle_timeout: Duration::from_secs(parse_optional_env(
defaults.use_planning,
"AGENT_USE_PLANNING",
)?,
session_idle_timeout: Duration::from_secs(db_first_or_default(
&settings.agent.session_idle_timeout_secs,
&defaults.session_idle_timeout_secs,
"SESSION_IDLE_TIMEOUT_SECS", "SESSION_IDLE_TIMEOUT_SECS",
settings.agent.session_idle_timeout_secs,
)?), )?),
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_cost_per_user_per_day_cents: parse_option_env("MAX_COST_PER_USER_PER_DAY_CENTS")?,
max_tool_iterations: db_first_or_default( max_tool_iterations: parse_optional_env(
&settings.agent.max_tool_iterations,
&defaults.max_tool_iterations,
"AGENT_MAX_TOOL_ITERATIONS", "AGENT_MAX_TOOL_ITERATIONS",
settings.agent.max_tool_iterations,
)?, )?,
auto_approve_tools: db_first_bool( auto_approve_tools: parse_bool_env(
settings.agent.auto_approve_tools,
defaults.auto_approve_tools,
"AGENT_AUTO_APPROVE_TOOLS", "AGENT_AUTO_APPROVE_TOOLS",
settings.agent.auto_approve_tools,
)?, )?,
default_timezone: { default_timezone: {
let tz: String = db_first_or_default( let tz: String = parse_optional_env(
&settings.agent.default_timezone,
&defaults.default_timezone,
"DEFAULT_TIMEZONE", "DEFAULT_TIMEZONE",
settings.agent.default_timezone.clone(),
)?; )?;
if crate::timezone::parse_timezone(&tz).is_none() { if crate::timezone::parse_timezone(&tz).is_none() {
return Err(ConfigError::InvalidValue { return Err(ConfigError::InvalidValue {
@@ -139,16 +116,14 @@ impl AgentConfig {
} }
tz tz
}, },
max_tokens_per_job: db_first_or_default( max_tokens_per_job: parse_optional_env(
&settings.agent.max_tokens_per_job,
&defaults.max_tokens_per_job,
"AGENT_MAX_TOKENS_PER_JOB", "AGENT_MAX_TOKENS_PER_JOB",
settings.agent.max_tokens_per_job,
)?,
multi_tenant: parse_bool_env(
"MULTI_TENANT",
optional_env("GATEWAY_USER_TOKENS")?.is_some(),
)?, )?,
// Auto-detected from GATEWAY_USER_TOKENS presence. Not a separate
// knob — multi-tenant mode is always implied by configuring user tokens.
multi_tenant: optional_env("GATEWAY_USER_TOKENS")?.is_some(),
max_llm_concurrent_per_user: parse_option_env("TENANT_MAX_LLM_CONCURRENT")?,
max_jobs_concurrent_per_user: parse_option_env("TENANT_MAX_JOBS_CONCURRENT")?,
}) })
} }
} }
+10 -41
View File
@@ -1,7 +1,7 @@
use std::path::PathBuf; use std::path::PathBuf;
use std::time::Duration; use std::time::Duration;
use crate::config::helpers::{db_first_bool, db_first_or_default, optional_env}; use crate::config::helpers::{optional_env, parse_bool_env, parse_optional_env};
use crate::error::ConfigError; use crate::error::ConfigError;
/// Builder mode configuration. /// Builder mode configuration.
@@ -34,29 +34,14 @@ impl Default for BuilderModeConfig {
impl BuilderModeConfig { impl BuilderModeConfig {
pub(crate) fn resolve(settings: &crate::settings::Settings) -> Result<Self, ConfigError> { pub(crate) fn resolve(settings: &crate::settings::Settings) -> Result<Self, ConfigError> {
let bs = &settings.builder; let bs = &settings.builder;
let defaults = crate::settings::BuilderSettings::default();
Ok(Self { Ok(Self {
enabled: db_first_bool(bs.enabled, defaults.enabled, "BUILDER_ENABLED")?, enabled: parse_bool_env("BUILDER_ENABLED", bs.enabled)?,
build_dir: if let Some(ref dir) = bs.build_dir { build_dir: optional_env("BUILDER_DIR")?
Some(dir.clone()) .map(PathBuf::from)
} else { .or_else(|| bs.build_dir.clone()),
optional_env("BUILDER_DIR")?.map(PathBuf::from) max_iterations: parse_optional_env("BUILDER_MAX_ITERATIONS", bs.max_iterations)?,
}, timeout_secs: parse_optional_env("BUILDER_TIMEOUT_SECS", bs.timeout_secs)?,
max_iterations: db_first_or_default( auto_register: parse_bool_env("BUILDER_AUTO_REGISTER", bs.auto_register)?,
&bs.max_iterations,
&defaults.max_iterations,
"BUILDER_MAX_ITERATIONS",
)?,
timeout_secs: db_first_or_default(
&bs.timeout_secs,
&defaults.timeout_secs,
"BUILDER_TIMEOUT_SECS",
)?,
auto_register: db_first_bool(
bs.auto_register,
defaults.auto_register,
"BUILDER_AUTO_REGISTER",
)?,
}) })
} }
@@ -94,7 +79,7 @@ mod tests {
} }
#[test] #[test]
fn db_settings_override_env() { fn env_overrides_settings() {
let _guard = lock_env(); let _guard = lock_env();
let mut settings = Settings::default(); let mut settings = Settings::default();
settings.builder.timeout_secs = 123; settings.builder.timeout_secs = 123;
@@ -104,22 +89,6 @@ mod tests {
let cfg = BuilderModeConfig::resolve(&settings).expect("resolve"); let cfg = BuilderModeConfig::resolve(&settings).expect("resolve");
unsafe { std::env::remove_var("BUILDER_TIMEOUT_SECS") }; unsafe { std::env::remove_var("BUILDER_TIMEOUT_SECS") };
assert_eq!(cfg.timeout_secs, 123, "DB setting should win over env"); assert_eq!(cfg.timeout_secs, 3);
}
#[test]
fn env_used_when_no_db_setting() {
let _guard = lock_env();
let settings = Settings::default();
// SAFETY: Under ENV_MUTEX, no concurrent env access.
unsafe { std::env::set_var("BUILDER_TIMEOUT_SECS", "42") };
let cfg = BuilderModeConfig::resolve(&settings).expect("resolve");
unsafe { std::env::remove_var("BUILDER_TIMEOUT_SECS") };
assert_eq!(
cfg.timeout_secs, 42,
"env should be used when DB has the default value"
);
} }
} }
+52 -86
View File
@@ -5,11 +5,9 @@ use secrecy::SecretString;
use serde::Deserialize; use serde::Deserialize;
use crate::bootstrap::ironclaw_base_dir; use crate::bootstrap::ironclaw_base_dir;
use crate::config::helpers::{ use crate::config::helpers::{optional_env, parse_bool_env, parse_optional_env};
db_first_bool, db_first_optional_string, db_first_or_default, optional_env, parse_optional_env,
};
use crate::error::ConfigError; use crate::error::ConfigError;
use crate::settings::{ChannelSettings, Settings}; use crate::settings::Settings;
/// Channel configurations. /// Channel configurations.
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
@@ -116,24 +114,15 @@ pub struct SignalConfig {
impl ChannelsConfig { impl ChannelsConfig {
pub(crate) fn resolve(settings: &Settings, owner_id: &str) -> Result<Self, ConfigError> { pub(crate) fn resolve(settings: &Settings, owner_id: &str) -> Result<Self, ConfigError> {
let cs = &settings.channels; let cs = &settings.channels;
let defaults = ChannelSettings::default();
let http_enabled_by_env = let http_enabled_by_env =
optional_env("HTTP_PORT")?.is_some() || optional_env("HTTP_HOST")?.is_some(); optional_env("HTTP_PORT")?.is_some() || optional_env("HTTP_HOST")?.is_some();
let http_enabled_by_db = let http = if http_enabled_by_env || cs.http_enabled {
db_first_bool(cs.http_enabled, defaults.http_enabled, "HTTP_ENABLED")?;
let http = if http_enabled_by_env || http_enabled_by_db {
Some(HttpConfig { Some(HttpConfig {
host: db_first_optional_string(&cs.http_host, "HTTP_HOST")? host: optional_env("HTTP_HOST")?
.or_else(|| cs.http_host.clone())
.unwrap_or_else(|| "0.0.0.0".to_string()), .unwrap_or_else(|| "0.0.0.0".to_string()),
port: { port: parse_optional_env("HTTP_PORT", cs.http_port.unwrap_or(8080))?,
// defaults.http_port is None, so any Some(..) is an explicit DB override.
if let Some(ref db_port) = cs.http_port {
db_first_or_default(db_port, &8080, "HTTP_PORT")?
} else {
parse_optional_env("HTTP_PORT", 8080)?
}
},
webhook_secret: optional_env("HTTP_WEBHOOK_SECRET")?.map(SecretString::from), webhook_secret: optional_env("HTTP_WEBHOOK_SECRET")?.map(SecretString::from),
user_id: owner_id.to_string(), user_id: owner_id.to_string(),
}) })
@@ -141,13 +130,10 @@ impl ChannelsConfig {
None None
}; };
let gateway_enabled = db_first_bool( let gateway_enabled = parse_bool_env("GATEWAY_ENABLED", cs.gateway_enabled)?;
cs.gateway_enabled,
defaults.gateway_enabled,
"GATEWAY_ENABLED",
)?;
let gateway = if gateway_enabled { let gateway = if gateway_enabled {
let user_id = db_first_optional_string(&cs.gateway_user_id, "GATEWAY_USER_ID")? let user_id = optional_env("GATEWAY_USER_ID")?
.or_else(|| cs.gateway_user_id.clone())
.unwrap_or_else(|| owner_id.to_string()); .unwrap_or_else(|| owner_id.to_string());
let memory_layers: Vec<crate::workspace::layer::MemoryLayer> = let memory_layers: Vec<crate::workspace::layer::MemoryLayer> =
@@ -263,16 +249,13 @@ impl ChannelsConfig {
} }
} }
Some(GatewayConfig { Some(GatewayConfig {
host: db_first_optional_string(&cs.gateway_host, "GATEWAY_HOST")? host: optional_env("GATEWAY_HOST")?
.or_else(|| cs.gateway_host.clone())
.unwrap_or_else(|| "127.0.0.1".to_string()), .unwrap_or_else(|| "127.0.0.1".to_string()),
port: { port: parse_optional_env(
// defaults.gateway_port is None, so any Some(..) is an explicit DB override. "GATEWAY_PORT",
if let Some(ref db_port) = cs.gateway_port { cs.gateway_port.unwrap_or(DEFAULT_GATEWAY_PORT),
db_first_or_default(db_port, &DEFAULT_GATEWAY_PORT, "GATEWAY_PORT")? )?,
} else {
parse_optional_env("GATEWAY_PORT", DEFAULT_GATEWAY_PORT)?
}
},
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()),
user_id, user_id,
@@ -284,22 +267,16 @@ impl ChannelsConfig {
None None
}; };
let signal_enabled = let signal_url = optional_env("SIGNAL_HTTP_URL")?.or_else(|| cs.signal_http_url.clone());
db_first_bool(cs.signal_enabled, defaults.signal_enabled, "SIGNAL_ENABLED")?; let signal = if let Some(http_url) = signal_url {
let signal_url = db_first_optional_string(&cs.signal_http_url, "SIGNAL_HTTP_URL")?; let account = optional_env("SIGNAL_ACCOUNT")?
let signal = if signal_enabled || signal_url.is_some() { .or_else(|| cs.signal_account.clone())
let http_url = signal_url.ok_or(ConfigError::InvalidValue { .ok_or(ConfigError::InvalidValue {
key: "SIGNAL_HTTP_URL".to_string(),
message: "SIGNAL_HTTP_URL is required when Signal is enabled".to_string(),
})?;
let account = db_first_optional_string(&cs.signal_account, "SIGNAL_ACCOUNT")?.ok_or(
ConfigError::InvalidValue {
key: "SIGNAL_ACCOUNT".to_string(), key: "SIGNAL_ACCOUNT".to_string(),
message: "SIGNAL_ACCOUNT is required when SIGNAL_HTTP_URL is set".to_string(), message: "SIGNAL_ACCOUNT is required when SIGNAL_HTTP_URL is set".to_string(),
}, })?;
)?;
let allow_from = let allow_from =
match db_first_optional_string(&cs.signal_allow_from, "SIGNAL_ALLOW_FROM")? { match optional_env("SIGNAL_ALLOW_FROM")?.or_else(|| cs.signal_allow_from.clone()) {
None => vec![account.clone()], None => vec![account.clone()],
Some(s) => s Some(s) => s
.split(',') .split(',')
@@ -307,39 +284,36 @@ impl ChannelsConfig {
.filter(|s| !s.is_empty()) .filter(|s| !s.is_empty())
.collect(), .collect(),
}; };
let dm_policy = db_first_optional_string(&cs.signal_dm_policy, "SIGNAL_DM_POLICY")? let dm_policy = optional_env("SIGNAL_DM_POLICY")?
.or_else(|| cs.signal_dm_policy.clone())
.unwrap_or_else(|| "pairing".to_string()); .unwrap_or_else(|| "pairing".to_string());
let group_policy = let group_policy = optional_env("SIGNAL_GROUP_POLICY")?
db_first_optional_string(&cs.signal_group_policy, "SIGNAL_GROUP_POLICY")? .or_else(|| cs.signal_group_policy.clone())
.unwrap_or_else(|| "allowlist".to_string()); .unwrap_or_else(|| "allowlist".to_string());
Some(SignalConfig { Some(SignalConfig {
http_url, http_url,
account, account,
allow_from, allow_from,
allow_from_groups: db_first_optional_string( allow_from_groups: optional_env("SIGNAL_ALLOW_FROM_GROUPS")?
&cs.signal_allow_from_groups, .or_else(|| cs.signal_allow_from_groups.clone())
"SIGNAL_ALLOW_FROM_GROUPS", .map(|s| {
)? s.split(',')
.map(|s| { .map(|e| e.trim().to_string())
s.split(',') .filter(|s| !s.is_empty())
.map(|e| e.trim().to_string()) .collect()
.filter(|s| !s.is_empty()) })
.collect() .unwrap_or_default(),
})
.unwrap_or_default(),
dm_policy, dm_policy,
group_policy, group_policy,
group_allow_from: db_first_optional_string( group_allow_from: optional_env("SIGNAL_GROUP_ALLOW_FROM")?
&cs.signal_group_allow_from, .or_else(|| cs.signal_group_allow_from.clone())
"SIGNAL_GROUP_ALLOW_FROM", .map(|s| {
)? s.split(',')
.map(|s| { .map(|e| e.trim().to_string())
s.split(',') .filter(|s| !s.is_empty())
.map(|e| e.trim().to_string()) .collect()
.filter(|s| !s.is_empty()) })
.collect() .unwrap_or_default(),
})
.unwrap_or_default(),
ignore_attachments: optional_env("SIGNAL_IGNORE_ATTACHMENTS")? ignore_attachments: optional_env("SIGNAL_IGNORE_ATTACHMENTS")?
.map(|s| s.to_lowercase() == "true" || s == "1") .map(|s| s.to_lowercase() == "true" || s == "1")
.unwrap_or(false), .unwrap_or(false),
@@ -351,7 +325,7 @@ impl ChannelsConfig {
None None
}; };
let cli_enabled = db_first_bool(cs.cli_enabled, defaults.cli_enabled, "CLI_ENABLED")?; let cli_enabled = parse_bool_env("CLI_ENABLED", cs.cli_enabled)?;
Ok(Self { Ok(Self {
cli: CliConfig { cli: CliConfig {
@@ -360,21 +334,13 @@ impl ChannelsConfig {
http, http,
gateway, gateway,
signal, signal,
wasm_channels_dir: { wasm_channels_dir: optional_env("WASM_CHANNELS_DIR")?
// DB-first: use settings if explicitly set, else env, else default. .map(PathBuf::from)
// defaults.wasm_channels_dir is None, so any Some(..) is an explicit DB override. .or_else(|| cs.wasm_channels_dir.clone())
if let Some(ref db_dir) = cs.wasm_channels_dir { .unwrap_or_else(default_channels_dir),
db_dir.clone() wasm_channels_enabled: parse_bool_env(
} else {
optional_env("WASM_CHANNELS_DIR")?
.map(PathBuf::from)
.unwrap_or_else(default_channels_dir)
}
},
wasm_channels_enabled: db_first_bool(
cs.wasm_channels_enabled,
defaults.wasm_channels_enabled,
"WASM_CHANNELS_ENABLED", "WASM_CHANNELS_ENABLED",
cs.wasm_channels_enabled,
)?, )?,
wasm_channel_owner_ids: { wasm_channel_owner_ids: {
let mut ids = cs.wasm_channel_owner_ids.clone(); let mut ids = cs.wasm_channel_owner_ids.clone();
+16 -85
View File
@@ -2,9 +2,7 @@ use std::sync::Arc;
use secrecy::{ExposeSecret, SecretString}; use secrecy::{ExposeSecret, SecretString};
use crate::config::helpers::{ use crate::config::helpers::{optional_env, parse_bool_env, parse_optional_env, validate_base_url};
db_first_bool, db_first_or_default, optional_env, parse_optional_env, validate_base_url,
};
use crate::error::ConfigError; use crate::error::ConfigError;
use crate::llm::SessionManager; use crate::llm::SessionManager;
use crate::settings::Settings; use crate::settings::Settings;
@@ -73,41 +71,22 @@ pub(crate) fn default_dimension_for_model(model: &str) -> usize {
impl EmbeddingsConfig { impl EmbeddingsConfig {
pub(crate) fn resolve(settings: &Settings) -> Result<Self, ConfigError> { pub(crate) fn resolve(settings: &Settings) -> Result<Self, ConfigError> {
let defaults = crate::settings::EmbeddingsSettings::default();
let openai_api_key = optional_env("OPENAI_API_KEY")?.map(SecretString::from); let openai_api_key = optional_env("OPENAI_API_KEY")?.map(SecretString::from);
let provider = db_first_or_default( let provider = optional_env("EMBEDDING_PROVIDER")?
&settings.embeddings.provider, .unwrap_or_else(|| settings.embeddings.provider.clone());
&defaults.provider,
"EMBEDDING_PROVIDER",
)?;
let model = db_first_or_default( let model =
&settings.embeddings.model, optional_env("EMBEDDING_MODEL")?.unwrap_or_else(|| settings.embeddings.model.clone());
&defaults.model,
"EMBEDDING_MODEL",
)?;
// ollama_base_url lives on the top-level Settings, not the embeddings let ollama_base_url = optional_env("OLLAMA_BASE_URL")?
// sub-struct. Use a manual DB > env > default chain. .or_else(|| settings.ollama_base_url.clone())
let default_ollama_url = "http://localhost:11434".to_string(); .unwrap_or_else(|| "http://localhost:11434".to_string());
let ollama_base_url = settings
.ollama_base_url
.as_ref()
.filter(|s| !s.is_empty())
.cloned()
.or_else(|| optional_env("OLLAMA_BASE_URL").ok().flatten())
.unwrap_or(default_ollama_url);
let dim_default = default_dimension_for_model(&model); let dimension =
let dimension = db_first_or_default(&dim_default, &dim_default, "EMBEDDING_DIMENSION")?; parse_optional_env("EMBEDDING_DIMENSION", default_dimension_for_model(&model))?;
let enabled = db_first_bool( let enabled = parse_bool_env("EMBEDDING_ENABLED", settings.embeddings.enabled)?;
settings.embeddings.enabled,
defaults.enabled,
"EMBEDDING_ENABLED",
)?;
let openai_base_url = optional_env("EMBEDDING_BASE_URL")?; let openai_base_url = optional_env("EMBEDDING_BASE_URL")?;
@@ -228,11 +207,9 @@ mod tests {
std::env::remove_var("EMBEDDING_ENABLED"); std::env::remove_var("EMBEDDING_ENABLED");
std::env::remove_var("EMBEDDING_PROVIDER"); std::env::remove_var("EMBEDDING_PROVIDER");
std::env::remove_var("EMBEDDING_MODEL"); std::env::remove_var("EMBEDDING_MODEL");
std::env::remove_var("EMBEDDING_DIMENSION");
std::env::remove_var("OPENAI_API_KEY"); std::env::remove_var("OPENAI_API_KEY");
std::env::remove_var("EMBEDDING_BASE_URL"); std::env::remove_var("EMBEDDING_BASE_URL");
std::env::remove_var("EMBEDDING_CACHE_SIZE"); std::env::remove_var("EMBEDDING_CACHE_SIZE");
std::env::remove_var("OLLAMA_BASE_URL");
} }
} }
@@ -287,21 +264,18 @@ mod tests {
} }
#[test] #[test]
fn db_settings_override_env() { fn embeddings_env_override_takes_precedence() {
let _guard = lock_env(); let _guard = lock_env();
clear_embedding_env(); clear_embedding_env();
// SAFETY: Under ENV_MUTEX. // SAFETY: Under ENV_MUTEX.
unsafe { unsafe {
std::env::set_var("EMBEDDING_ENABLED", "false"); std::env::set_var("EMBEDDING_ENABLED", "true");
std::env::set_var("EMBEDDING_PROVIDER", "ollama");
std::env::set_var("EMBEDDING_MODEL", "all-minilm");
} }
let settings = Settings { let settings = Settings {
embeddings: EmbeddingsSettings { embeddings: EmbeddingsSettings {
enabled: true, enabled: false,
provider: "openai".to_string(), ..Default::default()
model: "text-embedding-3-large".to_string(),
}, },
..Default::default() ..Default::default()
}; };
@@ -309,55 +283,12 @@ mod tests {
let config = EmbeddingsConfig::resolve(&settings).expect("resolve should succeed"); let config = EmbeddingsConfig::resolve(&settings).expect("resolve should succeed");
assert!( assert!(
config.enabled, config.enabled,
"DB enabled=true should win over env EMBEDDING_ENABLED=false" "EMBEDDING_ENABLED=true env var should override settings"
);
assert_eq!(config.provider, "openai", "DB provider should win over env");
assert_eq!(
config.model, "text-embedding-3-large",
"DB model should win over env"
); );
// SAFETY: Under ENV_MUTEX. // SAFETY: Under ENV_MUTEX.
unsafe { unsafe {
std::env::remove_var("EMBEDDING_ENABLED"); std::env::remove_var("EMBEDDING_ENABLED");
std::env::remove_var("EMBEDDING_PROVIDER");
std::env::remove_var("EMBEDDING_MODEL");
}
}
#[test]
fn env_used_when_no_db_setting() {
let _guard = lock_env();
clear_embedding_env();
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::set_var("EMBEDDING_ENABLED", "true");
std::env::set_var("EMBEDDING_PROVIDER", "ollama");
std::env::set_var("EMBEDDING_MODEL", "nomic-embed-text");
}
// Settings left at defaults — no explicit DB/TOML override
let settings = Settings::default();
let config = EmbeddingsConfig::resolve(&settings).expect("resolve should succeed");
assert!(
config.enabled,
"env EMBEDDING_ENABLED should be used when settings at default"
);
assert_eq!(
config.provider, "ollama",
"env EMBEDDING_PROVIDER should be used when settings at default"
);
assert_eq!(
config.model, "nomic-embed-text",
"env EMBEDDING_MODEL should be used when settings at default"
);
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::remove_var("EMBEDDING_ENABLED");
std::env::remove_var("EMBEDDING_PROVIDER");
std::env::remove_var("EMBEDDING_MODEL");
} }
} }
+38 -175
View File
@@ -1,7 +1,4 @@
use crate::config::helpers::{ use crate::config::helpers::{optional_env, parse_bool_env, parse_option_env, parse_optional_env};
db_first_bool, db_first_optional_string, db_first_or_default, optional_env, parse_bool_env,
parse_option_env,
};
use crate::error::ConfigError; use crate::error::ConfigError;
use crate::settings::Settings; use crate::settings::Settings;
@@ -47,11 +44,8 @@ impl Default for HeartbeatConfig {
impl HeartbeatConfig { impl HeartbeatConfig {
pub(crate) fn resolve(settings: &Settings) -> Result<Self, ConfigError> { pub(crate) fn resolve(settings: &Settings) -> Result<Self, ConfigError> {
let defaults = crate::settings::HeartbeatSettings::default();
// fire_at: DB > env, then parse into NaiveTime
let fire_at_str = let fire_at_str =
db_first_optional_string(&settings.heartbeat.fire_at, "HEARTBEAT_FIRE_AT")?; optional_env("HEARTBEAT_FIRE_AT")?.or_else(|| settings.heartbeat.fire_at.clone());
let fire_at = fire_at_str let fire_at = fire_at_str
.map(|s| { .map(|s| {
chrono::NaiveTime::parse_from_str(&s, "%H:%M").map_err(|e| { chrono::NaiveTime::parse_from_str(&s, "%H:%M").map_err(|e| {
@@ -63,62 +57,44 @@ impl HeartbeatConfig {
}) })
.transpose()?; .transpose()?;
// quiet_hours: settings first, then env fallback
let quiet_hours_start = settings
.heartbeat
.quiet_hours_start
.or(parse_option_env::<u32>("HEARTBEAT_QUIET_START")?)
.map(|h| {
if h > 23 {
return Err(ConfigError::InvalidValue {
key: "HEARTBEAT_QUIET_START".into(),
message: "must be 0-23".into(),
});
}
Ok(h)
})
.transpose()?;
let quiet_hours_end = settings
.heartbeat
.quiet_hours_end
.or(parse_option_env::<u32>("HEARTBEAT_QUIET_END")?)
.map(|h| {
if h > 23 {
return Err(ConfigError::InvalidValue {
key: "HEARTBEAT_QUIET_END".into(),
message: "must be 0-23".into(),
});
}
Ok(h)
})
.transpose()?;
Ok(Self { Ok(Self {
enabled: db_first_bool( enabled: parse_bool_env("HEARTBEAT_ENABLED", settings.heartbeat.enabled)?,
settings.heartbeat.enabled, interval_secs: parse_optional_env(
defaults.enabled,
"HEARTBEAT_ENABLED",
)?,
interval_secs: db_first_or_default(
&settings.heartbeat.interval_secs,
&defaults.interval_secs,
"HEARTBEAT_INTERVAL_SECS", "HEARTBEAT_INTERVAL_SECS",
settings.heartbeat.interval_secs,
)?, )?,
notify_channel: db_first_optional_string( notify_channel: optional_env("HEARTBEAT_NOTIFY_CHANNEL")?
&settings.heartbeat.notify_channel, .or_else(|| settings.heartbeat.notify_channel.clone()),
"HEARTBEAT_NOTIFY_CHANNEL", notify_user: optional_env("HEARTBEAT_NOTIFY_USER")?
)?, .or_else(|| settings.heartbeat.notify_user.clone()),
notify_user: db_first_optional_string(
&settings.heartbeat.notify_user,
"HEARTBEAT_NOTIFY_USER",
)?,
fire_at, fire_at,
quiet_hours_start, quiet_hours_start: parse_option_env::<u32>("HEARTBEAT_QUIET_START")?
quiet_hours_end, .or(settings.heartbeat.quiet_hours_start)
.map(|h| {
if h > 23 {
return Err(ConfigError::InvalidValue {
key: "HEARTBEAT_QUIET_START".into(),
message: "must be 0-23".into(),
});
}
Ok(h)
})
.transpose()?,
quiet_hours_end: parse_option_env::<u32>("HEARTBEAT_QUIET_END")?
.or(settings.heartbeat.quiet_hours_end)
.map(|h| {
if h > 23 {
return Err(ConfigError::InvalidValue {
key: "HEARTBEAT_QUIET_END".into(),
message: "must be 0-23".into(),
});
}
Ok(h)
})
.transpose()?,
timezone: { timezone: {
let tz = let tz = optional_env("HEARTBEAT_TIMEZONE")?
db_first_optional_string(&settings.heartbeat.timezone, "HEARTBEAT_TIMEZONE")?; .or_else(|| settings.heartbeat.timezone.clone());
if let Some(ref tz_str) = tz if let Some(ref tz_str) = tz
&& crate::timezone::parse_timezone(tz_str).is_none() && crate::timezone::parse_timezone(tz_str).is_none()
{ {
@@ -130,7 +106,7 @@ impl HeartbeatConfig {
tz tz
}, },
// Auto-detect multi-tenant mode from GATEWAY_USER_TOKENS presence, // Auto-detect multi-tenant mode from GATEWAY_USER_TOKENS presence,
// or allow explicit override via HEARTBEAT_MULTI_TENANT. Stays env-only. // or allow explicit override via HEARTBEAT_MULTI_TENANT.
multi_tenant: parse_bool_env( multi_tenant: parse_bool_env(
"HEARTBEAT_MULTI_TENANT", "HEARTBEAT_MULTI_TENANT",
optional_env("GATEWAY_USER_TOKENS")?.is_some(), optional_env("GATEWAY_USER_TOKENS")?.is_some(),
@@ -142,11 +118,10 @@ impl HeartbeatConfig {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use crate::config::helpers::lock_env;
#[test] #[test]
fn test_quiet_hours_settings_have_priority() { fn test_quiet_hours_settings_fallback() {
// DB/settings values should take priority over env // When env vars are not set, settings values should be used
let mut settings = Settings::default(); let mut settings = Settings::default();
settings.heartbeat.quiet_hours_start = Some(22); settings.heartbeat.quiet_hours_start = Some(22);
settings.heartbeat.quiet_hours_end = Some(6); settings.heartbeat.quiet_hours_end = Some(6);
@@ -193,116 +168,4 @@ mod tests {
let config = HeartbeatConfig::resolve(&settings).expect("resolve"); let config = HeartbeatConfig::resolve(&settings).expect("resolve");
assert_eq!(config.timezone.as_deref(), Some("America/New_York")); assert_eq!(config.timezone.as_deref(), Some("America/New_York"));
} }
#[test]
fn test_db_first_enabled_beats_env() {
let _guard = lock_env();
// SAFETY: under ENV_MUTEX
unsafe { std::env::set_var("HEARTBEAT_ENABLED", "false") };
let mut settings = Settings::default();
settings.heartbeat.enabled = true; // DB says enabled
let config = HeartbeatConfig::resolve(&settings).expect("resolve");
assert!(config.enabled, "DB value (true) should beat env (false)");
unsafe { std::env::remove_var("HEARTBEAT_ENABLED") };
}
#[test]
fn test_db_first_interval_beats_env() {
let _guard = lock_env();
unsafe { std::env::set_var("HEARTBEAT_INTERVAL_SECS", "999") };
let mut settings = Settings::default();
settings.heartbeat.interval_secs = 600; // DB says 600
let config = HeartbeatConfig::resolve(&settings).expect("resolve");
assert_eq!(config.interval_secs, 600, "DB value should beat env");
unsafe { std::env::remove_var("HEARTBEAT_INTERVAL_SECS") };
}
#[test]
fn test_db_first_notify_channel_beats_env() {
let _guard = lock_env();
unsafe { std::env::set_var("HEARTBEAT_NOTIFY_CHANNEL", "env-channel") };
let mut settings = Settings::default();
settings.heartbeat.notify_channel = Some("db-channel".to_string());
let config = HeartbeatConfig::resolve(&settings).expect("resolve");
assert_eq!(
config.notify_channel.as_deref(),
Some("db-channel"),
"DB value should beat env"
);
unsafe { std::env::remove_var("HEARTBEAT_NOTIFY_CHANNEL") };
}
#[test]
fn test_env_fallback_when_db_at_default() {
let _guard = lock_env();
unsafe { std::env::set_var("HEARTBEAT_INTERVAL_SECS", "999") };
// Settings at default => env should win
let settings = Settings::default();
let config = HeartbeatConfig::resolve(&settings).expect("resolve");
assert_eq!(
config.interval_secs, 999,
"env should win when DB at default"
);
unsafe { std::env::remove_var("HEARTBEAT_INTERVAL_SECS") };
}
#[test]
fn test_fire_at_db_first() {
let _guard = lock_env();
unsafe { std::env::set_var("HEARTBEAT_FIRE_AT", "08:00") };
let mut settings = Settings::default();
settings.heartbeat.fire_at = Some("14:30".to_string());
let config = HeartbeatConfig::resolve(&settings).expect("resolve");
assert_eq!(
config.fire_at,
Some(chrono::NaiveTime::from_hms_opt(14, 30, 0).unwrap()),
"DB fire_at should beat env"
);
unsafe { std::env::remove_var("HEARTBEAT_FIRE_AT") };
}
#[test]
fn test_timezone_db_first() {
let _guard = lock_env();
unsafe { std::env::set_var("HEARTBEAT_TIMEZONE", "UTC") };
let mut settings = Settings::default();
settings.heartbeat.timezone = Some("America/New_York".to_string());
let config = HeartbeatConfig::resolve(&settings).expect("resolve");
assert_eq!(
config.timezone.as_deref(),
Some("America/New_York"),
"DB timezone should beat env"
);
unsafe { std::env::remove_var("HEARTBEAT_TIMEZONE") };
}
#[test]
fn test_multi_tenant_stays_env_only() {
let _guard = lock_env();
unsafe { std::env::set_var("HEARTBEAT_MULTI_TENANT", "true") };
let settings = Settings::default();
let config = HeartbeatConfig::resolve(&settings).expect("resolve");
assert!(config.multi_tenant, "multi_tenant should read from env");
unsafe { std::env::remove_var("HEARTBEAT_MULTI_TENANT") };
}
} }
-227
View File
@@ -331,93 +331,6 @@ pub(crate) fn validate_base_url(url: &str, field_name: &str) -> Result<(), Confi
Ok(()) Ok(())
} }
// ---------------------------------------------------------------------------
// DB-first resolution helpers (DB > env > default)
// ---------------------------------------------------------------------------
/// Log a warning when a DB/TOML setting shadows a set env var.
fn warn_if_db_shadows_env(env_key: &str, db_value: &dyn std::fmt::Display) {
if let Ok(env_val) = std::env::var(env_key)
&& !env_val.is_empty()
{
tracing::warn!(
db_value = %db_value,
env_value = %env_val,
"{env_key} env var is set but DB/TOML setting takes priority. \
Remove the setting from the DB to use the env var."
);
}
}
/// Resolve with DB > env > default priority for concrete settings fields.
///
/// If `settings_val != default_val`, the settings value wins (it was explicitly
/// set in DB or TOML). Otherwise falls back to `optional_env(env_key)`, then
/// `default_val`.
pub(crate) fn db_first_or_default<T>(
settings_val: &T,
default_val: &T,
env_key: &str,
) -> Result<T, ConfigError>
where
T: std::str::FromStr + Clone + PartialEq + std::fmt::Display,
T::Err: std::fmt::Display,
{
if settings_val != default_val {
warn_if_db_shadows_env(env_key, settings_val);
return Ok(settings_val.clone());
}
parse_optional_env(env_key, default_val.clone())
}
/// Resolve a bool with DB > env > default priority.
pub(crate) fn db_first_bool(
settings_val: bool,
default_val: bool,
env_key: &str,
) -> Result<bool, ConfigError> {
if settings_val != default_val {
warn_if_db_shadows_env(env_key, &settings_val);
return Ok(settings_val);
}
parse_bool_env(env_key, default_val)
}
/// Resolve an `Option<String>` with DB > env priority (no hardcoded default).
///
/// Non-empty `Some` means DB set it; `None` or empty falls back to env.
pub(crate) fn db_first_optional_string(
settings_val: &Option<String>,
env_key: &str,
) -> Result<Option<String>, ConfigError> {
if let Some(val) = settings_val
&& !val.is_empty()
{
warn_if_db_shadows_env(env_key, val);
return Ok(Some(val.clone()));
}
optional_env(env_key)
}
/// Resolve an `Option<T>` with DB > env priority (no hardcoded default).
///
/// `Some(v)` means DB set it; `None` falls back to env.
#[allow(dead_code)] // Used by Group 2 resolvers (routines, skills, etc.)
pub(crate) fn db_first_option<T>(
settings_val: &Option<T>,
env_key: &str,
) -> Result<Option<T>, ConfigError>
where
T: std::str::FromStr + Clone + std::fmt::Display,
T::Err: std::fmt::Display,
{
if let Some(val) = settings_val {
warn_if_db_shadows_env(env_key, val);
return Ok(Some(val.clone()));
}
parse_option_env(env_key)
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
@@ -606,144 +519,4 @@ mod tests {
"Expected DNS resolution failure, got: {err}" "Expected DNS resolution failure, got: {err}"
); );
} }
// --- db_first_* helper tests ---
#[test]
fn db_first_or_default_prefers_settings_over_env() {
let _guard = lock_env();
let key = "IRONCLAW_TEST_DB_FIRST_1";
// SAFETY: under ENV_MUTEX
unsafe { std::env::set_var(key, "from-env") };
let result: String =
db_first_or_default(&"from-db".to_string(), &"default".to_string(), key)
.expect("should resolve");
assert_eq!(result, "from-db", "DB value should win over env");
unsafe { std::env::remove_var(key) };
}
#[test]
fn db_first_or_default_falls_back_to_env() {
let _guard = lock_env();
let key = "IRONCLAW_TEST_DB_FIRST_2";
unsafe { std::env::set_var(key, "from-env") };
// settings_val == default_val → treated as "unset"
let result: String =
db_first_or_default(&"default".to_string(), &"default".to_string(), key)
.expect("should resolve");
assert_eq!(
result, "from-env",
"env should win when settings at default"
);
unsafe { std::env::remove_var(key) };
}
#[test]
fn db_first_or_default_uses_default_when_neither_set() {
let _guard = lock_env();
let key = "IRONCLAW_TEST_DB_FIRST_3";
unsafe { std::env::remove_var(key) };
let result: String =
db_first_or_default(&"default".to_string(), &"default".to_string(), key)
.expect("should resolve");
assert_eq!(result, "default");
}
#[test]
fn db_first_bool_prefers_settings() {
let _guard = lock_env();
let key = "IRONCLAW_TEST_DB_FIRST_BOOL_1";
unsafe { std::env::set_var(key, "false") };
let result = db_first_bool(true, false, key).expect("should resolve");
assert!(result, "DB true should win over env false");
unsafe { std::env::remove_var(key) };
}
#[test]
fn db_first_bool_falls_back_to_env() {
let _guard = lock_env();
let key = "IRONCLAW_TEST_DB_FIRST_BOOL_2";
unsafe { std::env::set_var(key, "true") };
// settings == default → falls back to env
let result = db_first_bool(false, false, key).expect("should resolve");
assert!(result, "env should win when settings at default");
unsafe { std::env::remove_var(key) };
}
#[test]
fn db_first_optional_string_prefers_settings() {
let _guard = lock_env();
let key = "IRONCLAW_TEST_DB_FIRST_OPT_1";
unsafe { std::env::set_var(key, "from-env") };
let val = Some("from-db".to_string());
let result = db_first_optional_string(&val, key).expect("should resolve");
assert_eq!(result, Some("from-db".to_string()));
unsafe { std::env::remove_var(key) };
}
#[test]
fn db_first_optional_string_falls_back_to_env() {
let _guard = lock_env();
let key = "IRONCLAW_TEST_DB_FIRST_OPT_2";
unsafe { std::env::set_var(key, "from-env") };
let result = db_first_optional_string(&None, key).expect("should resolve");
assert_eq!(result, Some("from-env".to_string()));
unsafe { std::env::remove_var(key) };
}
#[test]
fn db_first_optional_string_empty_treated_as_unset() {
let _guard = lock_env();
let key = "IRONCLAW_TEST_DB_FIRST_OPT_3";
unsafe { std::env::set_var(key, "from-env") };
let val = Some(String::new());
let result = db_first_optional_string(&val, key).expect("should resolve");
assert_eq!(
result,
Some("from-env".to_string()),
"empty string should be treated as unset"
);
unsafe { std::env::remove_var(key) };
}
#[test]
fn db_first_option_prefers_settings() {
let _guard = lock_env();
let key = "IRONCLAW_TEST_DB_FIRST_OPT_T_1";
unsafe { std::env::set_var(key, "99") };
let val: Option<u64> = Some(42);
let result = db_first_option(&val, key).expect("should resolve");
assert_eq!(result, Some(42));
unsafe { std::env::remove_var(key) };
}
#[test]
fn db_first_option_falls_back_to_env() {
let _guard = lock_env();
let key = "IRONCLAW_TEST_DB_FIRST_OPT_T_2";
unsafe { std::env::set_var(key, "99") };
let val: Option<u64> = None;
let result = db_first_option(&val, key).expect("should resolve");
assert_eq!(result, Some(99));
unsafe { std::env::remove_var(key) };
}
} }
+7 -20
View File
@@ -1,7 +1,6 @@
use crate::bootstrap::ironclaw_base_dir; use crate::bootstrap::ironclaw_base_dir;
use crate::config::helpers::{db_first_bool, db_first_or_default}; use crate::config::helpers::{parse_bool_env, parse_optional_env};
use crate::error::ConfigError; use crate::error::ConfigError;
use crate::settings::Settings;
/// Memory hygiene configuration. /// Memory hygiene configuration.
/// ///
@@ -31,27 +30,15 @@ impl Default for HygieneConfig {
} }
impl HygieneConfig { impl HygieneConfig {
pub(crate) fn resolve(settings: &Settings) -> Result<Self, ConfigError> { pub(crate) fn resolve() -> Result<Self, ConfigError> {
let defaults = crate::settings::HygieneSettings::default();
let hs = &settings.hygiene;
Ok(Self { Ok(Self {
enabled: db_first_bool(hs.enabled, defaults.enabled, "MEMORY_HYGIENE_ENABLED")?, enabled: parse_bool_env("MEMORY_HYGIENE_ENABLED", true)?,
daily_retention_days: db_first_or_default( daily_retention_days: parse_optional_env("MEMORY_HYGIENE_DAILY_RETENTION_DAYS", 30)?,
&hs.daily_retention_days, conversation_retention_days: parse_optional_env(
&defaults.daily_retention_days,
"MEMORY_HYGIENE_DAILY_RETENTION_DAYS",
)?,
conversation_retention_days: db_first_or_default(
&hs.conversation_retention_days,
&defaults.conversation_retention_days,
"MEMORY_HYGIENE_CONVERSATION_RETENTION_DAYS", "MEMORY_HYGIENE_CONVERSATION_RETENTION_DAYS",
7,
)?, )?,
cadence_hours: db_first_or_default( cadence_hours: parse_optional_env("MEMORY_HYGIENE_CADENCE_HOURS", 12)?,
&hs.cadence_hours,
&defaults.cadence_hours,
"MEMORY_HYGIENE_CADENCE_HOURS",
)?,
}) })
} }
+58 -818
View File
File diff suppressed because it is too large Load Diff
+29 -220
View File
@@ -1,15 +1,9 @@
//! Configuration for IronClaw. //! Configuration for IronClaw.
//! //!
//! Settings are loaded with priority: **DB > env > TOML > default**. //! Settings are loaded with priority: env var > database > default.
//!
//! Exceptions:
//! - Bootstrap configs (database, secrets): env-only (DB not yet available)
//! - Security-sensitive fields (allow_local_tools, allow_full_access,
//! cost limits, auth tokens): env-only
//! - API keys: env/secrets store only
//!
//! `DATABASE_URL` lives in `~/.ironclaw/.env` (loaded via dotenvy early //! `DATABASE_URL` lives in `~/.ironclaw/.env` (loaded via dotenvy early
//! in startup). //! in startup). Everything else comes from env vars, the DB settings
//! table, or auto-detection.
mod agent; mod agent;
mod builder; mod builder;
@@ -192,9 +186,8 @@ impl Config {
/// Load configuration from environment variables and the database. /// Load configuration from environment variables and the database.
/// ///
/// Priority: DB > env > TOML > default. TOML is loaded first as a /// Priority: env var > TOML config file > DB settings > default.
/// base, then DB values are merged on top. Subsystem resolvers check /// This is the primary way to load config after DB is connected.
/// DB-backed settings before env vars (except bootstrap/security fields).
pub async fn from_db( pub async fn from_db(
store: &(dyn crate::db::SettingsStore + Sync), store: &(dyn crate::db::SettingsStore + Sync),
user_id: &str, user_id: &str,
@@ -203,9 +196,6 @@ impl Config {
} }
/// Load from DB with an optional TOML config file overlay. /// Load from DB with an optional TOML config file overlay.
///
/// Priority: DB > env > TOML > default. TOML is loaded as the base,
/// then DB values are merged on top. See module docs for exceptions.
pub async fn from_db_with_toml( pub async fn from_db_with_toml(
store: &(dyn crate::db::SettingsStore + Sync), store: &(dyn crate::db::SettingsStore + Sync),
user_id: &str, user_id: &str,
@@ -214,22 +204,19 @@ impl Config {
let _ = dotenvy::dotenv(); let _ = dotenvy::dotenv();
crate::bootstrap::load_ironclaw_env(); crate::bootstrap::load_ironclaw_env();
// Start with TOML config as a base (lowest priority among the two). // Load all settings from DB into a Settings struct
let mut settings = Settings::default(); let mut db_settings = match store.get_all_settings(user_id).await {
Self::apply_toml_overlay(&mut settings, toml_path)?; Ok(map) => Settings::from_db_map(&map),
// Overlay DB settings on top so DB values win over TOML.
match store.get_all_settings(user_id).await {
Ok(map) => {
let db_settings = Settings::from_db_map(&map);
settings.merge_from(&db_settings);
}
Err(e) => { Err(e) => {
tracing::warn!("Failed to load settings from DB, using defaults: {}", e); tracing::warn!("Failed to load settings from DB, using defaults: {}", e);
Settings::default()
} }
}; };
Self::build(&settings).await // Overlay TOML config file (values win over DB settings)
Self::apply_toml_overlay(&mut db_settings, toml_path)?;
Self::build(&db_settings).await
} }
/// Load configuration from environment variables only (no database). /// Load configuration from environment variables only (no database).
@@ -304,38 +291,16 @@ impl Config {
user_id: &str, user_id: &str,
toml_path: Option<&std::path::Path>, toml_path: Option<&std::path::Path>,
) -> Result<(), ConfigError> { ) -> Result<(), ConfigError> {
self.re_resolve_llm_with_secrets(store, user_id, toml_path, None) let settings = if let Some(store) = store {
.await let mut s = match store.get_all_settings(user_id).await {
} Ok(map) => Settings::from_db_map(&map),
Err(_) => Settings::default(),
/// Re-resolve LLM config, hydrating API keys from the secrets store. };
pub async fn re_resolve_llm_with_secrets(
&mut self,
store: Option<&(dyn crate::db::SettingsStore + Sync)>,
user_id: &str,
toml_path: Option<&std::path::Path>,
secrets: Option<&(dyn crate::secrets::SecretsStore + Send + Sync)>,
) -> Result<(), ConfigError> {
let mut settings = if let Some(store) = store {
// TOML as base, then DB on top (DB wins).
let mut s = Settings::default();
Self::apply_toml_overlay(&mut s, toml_path)?; Self::apply_toml_overlay(&mut s, toml_path)?;
if let Ok(map) = store.get_all_settings(user_id).await {
let db_settings = Settings::from_db_map(&map);
s.merge_from(&db_settings);
}
s s
} else { } else {
Settings::default() Settings::default()
}; };
// Hydrate API keys from encrypted secrets store into the settings
// struct so that LlmConfig::resolve() sees them without any changes
// to its synchronous resolution logic.
if let Some(secrets) = secrets {
hydrate_llm_keys_from_secrets(&mut settings, secrets, user_id).await;
}
self.llm = LlmConfig::resolve(&settings)?; self.llm = LlmConfig::resolve(&settings)?;
Ok(()) Ok(())
} }
@@ -347,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(),
@@ -366,13 +333,13 @@ impl Config {
secrets: SecretsConfig::resolve().await?, secrets: SecretsConfig::resolve().await?,
builder: BuilderModeConfig::resolve(settings)?, builder: BuilderModeConfig::resolve(settings)?,
heartbeat: HeartbeatConfig::resolve(settings)?, heartbeat: HeartbeatConfig::resolve(settings)?,
hygiene: HygieneConfig::resolve(settings)?, hygiene: HygieneConfig::resolve()?,
routines: RoutineConfig::resolve(settings)?, routines: RoutineConfig::resolve()?,
sandbox: SandboxModeConfig::resolve(settings)?, sandbox: SandboxModeConfig::resolve(settings)?,
claude_code: ClaudeCodeConfig::resolve(settings)?, claude_code: ClaudeCodeConfig::resolve(settings)?,
skills: SkillsConfig::resolve(settings)?, skills: SkillsConfig::resolve()?,
transcription: TranscriptionConfig::resolve(settings)?, transcription: TranscriptionConfig::resolve(settings)?,
search: WorkspaceSearchConfig::resolve(settings)?, search: WorkspaceSearchConfig::resolve()?,
workspace, workspace,
observability: crate::observability::ObservabilityConfig { observability: crate::observability::ObservabilityConfig {
backend: std::env::var("OBSERVABILITY_BACKEND").unwrap_or_else(|_| "none".into()), backend: std::env::var("OBSERVABILITY_BACKEND").unwrap_or_else(|_| "none".into()),
@@ -536,161 +503,3 @@ fn inject_os_credential_store_tokens(injected: &mut HashMap<String, String>) {
tracing::debug!("Refreshed ANTHROPIC_OAUTH_TOKEN from OS credential store"); tracing::debug!("Refreshed ANTHROPIC_OAUTH_TOKEN from OS credential store");
} }
} }
/// Hydrate LLM API keys from the secrets store into the settings struct.
///
/// Called after loading settings from DB but before `LlmConfig::resolve()`.
/// Populates `api_key` fields that were stripped from settings during the
/// write path and stored encrypted in the secrets store instead.
pub async fn hydrate_llm_keys_from_secrets(
settings: &mut Settings,
secrets: &(dyn crate::secrets::SecretsStore + Send + Sync),
user_id: &str,
) {
// Hydrate builtin overrides
for (provider_id, override_val) in settings.llm_builtin_overrides.iter_mut() {
if override_val.api_key.is_some() {
continue; // Already has a key (legacy plaintext or TOML)
}
let secret_name = format!("llm_builtin_{}_api_key", provider_id);
if let Ok(decrypted) = secrets.get_decrypted(user_id, &secret_name).await {
override_val.api_key = Some(decrypted.expose().to_string());
}
}
// Hydrate custom providers
for provider in settings.llm_custom_providers.iter_mut() {
if provider.api_key.is_some() {
continue;
}
let secret_name = format!("llm_custom_{}_api_key", provider.id);
if let Ok(decrypted) = secrets.get_decrypted(user_id, &secret_name).await {
provider.api_key = Some(decrypted.expose().to_string());
}
}
}
/// Migrate plaintext API keys from the settings table to the encrypted secrets store.
///
/// Idempotent: skips keys that are already in the secrets store.
/// After migration, strips plaintext keys from the settings table.
pub async fn migrate_plaintext_llm_keys(
settings_store: &(dyn crate::db::SettingsStore + Sync),
secrets: &(dyn crate::secrets::SecretsStore + Send + Sync),
user_id: &str,
) {
let settings_map = match settings_store.get_all_settings(user_id).await {
Ok(m) => m,
Err(_) => return,
};
let mut migrated = 0u32;
// Migrate builtin overrides
if let Some(obj) = settings_map
.get("llm_builtin_overrides")
.and_then(|v| v.as_object())
{
let mut sanitized = obj.clone();
for (provider_id, override_val) in obj {
if let Some(api_key) = override_val.get("api_key").and_then(|v| v.as_str()) {
if api_key.is_empty() {
continue;
}
let secret_name = format!("llm_builtin_{}_api_key", provider_id);
if !secrets.exists(user_id, &secret_name).await.unwrap_or(false)
&& let Err(e) = secrets
.create(
user_id,
crate::secrets::CreateSecretParams {
name: secret_name.clone(),
value: secrecy::SecretString::from(api_key.to_string()),
provider: Some(provider_id.clone()),
expires_at: None,
},
)
.await
{
tracing::warn!("Failed to migrate key for builtin '{}': {}", provider_id, e);
continue;
}
if let Some(o) = sanitized
.get_mut(provider_id)
.and_then(|v| v.as_object_mut())
{
o.remove("api_key");
}
migrated += 1;
}
}
if migrated > 0 {
let _ = settings_store
.set_setting(
user_id,
"llm_builtin_overrides",
&serde_json::Value::Object(sanitized),
)
.await;
}
}
// Migrate custom providers
let before = migrated;
if let Some(arr) = settings_map
.get("llm_custom_providers")
.and_then(|v| v.as_array())
{
let mut sanitized = arr.clone();
for (idx, provider_val) in arr.iter().enumerate() {
let provider_id = provider_val
.get("id")
.and_then(|v| v.as_str())
.unwrap_or("");
if provider_id.is_empty() {
continue;
}
if let Some(api_key) = provider_val.get("api_key").and_then(|v| v.as_str()) {
if api_key.is_empty() {
continue;
}
let secret_name = format!("llm_custom_{}_api_key", provider_id);
if !secrets.exists(user_id, &secret_name).await.unwrap_or(false)
&& let Err(e) = secrets
.create(
user_id,
crate::secrets::CreateSecretParams {
name: secret_name.clone(),
value: secrecy::SecretString::from(api_key.to_string()),
provider: Some(provider_id.to_string()),
expires_at: None,
},
)
.await
{
tracing::warn!("Failed to migrate key for custom '{}': {}", provider_id, e);
continue;
}
if let Some(o) = sanitized[idx].as_object_mut() {
o.remove("api_key");
}
migrated += 1;
}
}
if migrated > before {
let _ = settings_store
.set_setting(
user_id,
"llm_custom_providers",
&serde_json::Value::Array(sanitized),
)
.await;
}
}
if migrated > 0 {
tracing::info!(
"Migrated {} plaintext LLM API key(s) to encrypted secrets store",
migrated
);
}
}
+9 -37
View File
@@ -1,6 +1,5 @@
use crate::config::helpers::{db_first_bool, db_first_or_default}; use crate::config::helpers::{parse_bool_env, parse_optional_env};
use crate::error::ConfigError; use crate::error::ConfigError;
use crate::settings::Settings;
/// Routines configuration. /// Routines configuration.
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
@@ -36,42 +35,15 @@ impl Default for RoutineConfig {
} }
impl RoutineConfig { impl RoutineConfig {
pub(crate) fn resolve(settings: &Settings) -> Result<Self, ConfigError> { pub(crate) fn resolve() -> Result<Self, ConfigError> {
let defaults = crate::settings::RoutineSettings::default(); let max_iterations: u32 = parse_optional_env("ROUTINES_LIGHTWEIGHT_MAX_ITERATIONS", 3)?;
let rs = &settings.routines;
let max_iterations: u32 = db_first_or_default(
&rs.lightweight_max_iterations,
&defaults.lightweight_max_iterations,
"ROUTINES_LIGHTWEIGHT_MAX_ITERATIONS",
)?;
Ok(Self { Ok(Self {
enabled: db_first_bool(rs.enabled, defaults.enabled, "ROUTINES_ENABLED")?, enabled: parse_bool_env("ROUTINES_ENABLED", true)?,
cron_check_interval_secs: db_first_or_default( cron_check_interval_secs: parse_optional_env("ROUTINES_CRON_INTERVAL", 15)?,
&rs.cron_check_interval_secs, max_concurrent_routines: parse_optional_env("ROUTINES_MAX_CONCURRENT", 10)?,
&defaults.cron_check_interval_secs, default_cooldown_secs: parse_optional_env("ROUTINES_DEFAULT_COOLDOWN", 300)?,
"ROUTINES_CRON_INTERVAL", max_lightweight_tokens: parse_optional_env("ROUTINES_MAX_TOKENS", 4096)?,
)?, lightweight_tools_enabled: parse_bool_env("ROUTINES_LIGHTWEIGHT_TOOLS", true)?,
max_concurrent_routines: db_first_or_default(
&rs.max_concurrent_routines,
&defaults.max_concurrent_routines,
"ROUTINES_MAX_CONCURRENT",
)?,
default_cooldown_secs: db_first_or_default(
&rs.default_cooldown_secs,
&defaults.default_cooldown_secs,
"ROUTINES_DEFAULT_COOLDOWN",
)?,
max_lightweight_tokens: db_first_or_default(
&rs.max_lightweight_tokens,
&defaults.max_lightweight_tokens,
"ROUTINES_MAX_TOKENS",
)?,
lightweight_tools_enabled: db_first_bool(
rs.lightweight_tools_enabled,
defaults.lightweight_tools_enabled,
"ROUTINES_LIGHTWEIGHT_TOOLS",
)?,
lightweight_max_iterations: max_iterations.min(5), // cap at 5 lightweight_max_iterations: max_iterations.min(5), // cap at 5
}) })
} }
+5 -31
View File
@@ -1,4 +1,4 @@
use crate::config::helpers::{db_first_bool, db_first_or_default}; use crate::config::helpers::{parse_bool_env, parse_optional_env};
use crate::error::ConfigError; use crate::error::ConfigError;
pub use ironclaw_safety::SafetyConfig; pub use ironclaw_safety::SafetyConfig;
@@ -7,17 +7,11 @@ pub(crate) fn resolve_safety_config(
settings: &crate::settings::Settings, settings: &crate::settings::Settings,
) -> Result<SafetyConfig, ConfigError> { ) -> Result<SafetyConfig, ConfigError> {
let ss = &settings.safety; let ss = &settings.safety;
let defaults = crate::settings::SafetySettings::default();
Ok(SafetyConfig { Ok(SafetyConfig {
max_output_length: db_first_or_default( max_output_length: parse_optional_env("SAFETY_MAX_OUTPUT_LENGTH", ss.max_output_length)?,
&ss.max_output_length, injection_check_enabled: parse_bool_env(
&defaults.max_output_length,
"SAFETY_MAX_OUTPUT_LENGTH",
)?,
injection_check_enabled: db_first_bool(
ss.injection_check_enabled,
defaults.injection_check_enabled,
"SAFETY_INJECTION_CHECK_ENABLED", "SAFETY_INJECTION_CHECK_ENABLED",
ss.injection_check_enabled,
)?, )?,
}) })
} }
@@ -41,10 +35,9 @@ mod tests {
} }
#[test] #[test]
fn db_settings_override_env() { fn env_overrides_settings() {
let _guard = lock_env(); let _guard = lock_env();
let mut settings = Settings::default(); let mut settings = Settings::default();
// Non-default value simulates an explicit DB/TOML setting
settings.safety.max_output_length = 42; settings.safety.max_output_length = 42;
// SAFETY: Under ENV_MUTEX, no concurrent env access. // SAFETY: Under ENV_MUTEX, no concurrent env access.
@@ -52,25 +45,6 @@ mod tests {
let cfg = resolve_safety_config(&settings).expect("resolve"); let cfg = resolve_safety_config(&settings).expect("resolve");
unsafe { std::env::remove_var("SAFETY_MAX_OUTPUT_LENGTH") }; unsafe { std::env::remove_var("SAFETY_MAX_OUTPUT_LENGTH") };
// DB value (42) wins over env value (7)
assert_eq!(cfg.max_output_length, 42);
}
#[test]
fn env_used_when_no_db_setting() {
let _guard = lock_env();
// Settings left at defaults — no explicit DB/TOML override
let settings = Settings::default();
// SAFETY: Under ENV_MUTEX, no concurrent env access.
unsafe { std::env::set_var("SAFETY_MAX_OUTPUT_LENGTH", "7") };
unsafe { std::env::set_var("SAFETY_INJECTION_CHECK_ENABLED", "false") };
let cfg = resolve_safety_config(&settings).expect("resolve");
unsafe { std::env::remove_var("SAFETY_MAX_OUTPUT_LENGTH") };
unsafe { std::env::remove_var("SAFETY_INJECTION_CHECK_ENABLED") };
// Env values win when settings are at their defaults
assert_eq!(cfg.max_output_length, 7); assert_eq!(cfg.max_output_length, 7);
assert!(!cfg.injection_check_enabled);
} }
} }
+24 -72
View File
@@ -1,7 +1,4 @@
use crate::config::helpers::{ use crate::config::helpers::{optional_env, parse_bool_env, parse_optional_env, parse_string_env};
db_first_bool, db_first_or_default, optional_env, parse_bool_env, parse_optional_env,
parse_string_env,
};
use crate::error::ConfigError; use crate::error::ConfigError;
/// Docker sandbox configuration. /// Docker sandbox configuration.
@@ -57,16 +54,16 @@ impl Default for SandboxModeConfig {
impl SandboxModeConfig { impl SandboxModeConfig {
pub(crate) fn resolve(settings: &crate::settings::Settings) -> Result<Self, ConfigError> { pub(crate) fn resolve(settings: &crate::settings::Settings) -> Result<Self, ConfigError> {
let ss = &settings.sandbox; let ss = &settings.sandbox;
let defaults = crate::settings::SandboxSettings::default();
// extra_allowed_domains: DB wins if non-empty, otherwise env, otherwise empty. let extra_domains = optional_env("SANDBOX_EXTRA_DOMAINS")?
let extra_domains = if !ss.extra_allowed_domains.is_empty() { .map(|s| s.split(',').map(|d| d.trim().to_string()).collect())
ss.extra_allowed_domains.clone() .unwrap_or_else(|| {
} else { if ss.extra_allowed_domains.is_empty() {
optional_env("SANDBOX_EXTRA_DOMAINS")? Vec::new()
.map(|s| s.split(',').map(|d| d.trim().to_string()).collect()) } else {
.unwrap_or_default() ss.extra_allowed_domains.clone()
}; }
});
// reaper/orphan fields have no Settings counterpart — env > default only. // reaper/orphan fields have no Settings counterpart — env > default only.
let reaper_interval_secs: u64 = parse_optional_env("SANDBOX_REAPER_INTERVAL_SECS", 300)?; let reaper_interval_secs: u64 = parse_optional_env("SANDBOX_REAPER_INTERVAL_SECS", 300)?;
@@ -88,31 +85,15 @@ impl SandboxModeConfig {
} }
Ok(Self { Ok(Self {
enabled: db_first_bool(ss.enabled, defaults.enabled, "SANDBOX_ENABLED")?, enabled: parse_bool_env("SANDBOX_ENABLED", ss.enabled)?,
policy: db_first_or_default(&ss.policy, &defaults.policy, "SANDBOX_POLICY")?, policy: parse_string_env("SANDBOX_POLICY", ss.policy.clone())?,
// allow_full_access has no Settings counterpart — env > default only (security). // allow_full_access has no Settings counterpart — env > default only.
allow_full_access: parse_bool_env("SANDBOX_ALLOW_FULL_ACCESS", false)?, allow_full_access: parse_bool_env("SANDBOX_ALLOW_FULL_ACCESS", false)?,
timeout_secs: db_first_or_default( timeout_secs: parse_optional_env("SANDBOX_TIMEOUT_SECS", ss.timeout_secs)?,
&ss.timeout_secs, memory_limit_mb: parse_optional_env("SANDBOX_MEMORY_LIMIT_MB", ss.memory_limit_mb)?,
&defaults.timeout_secs, cpu_shares: parse_optional_env("SANDBOX_CPU_SHARES", ss.cpu_shares)?,
"SANDBOX_TIMEOUT_SECS", image: parse_string_env("SANDBOX_IMAGE", ss.image.clone())?,
)?, auto_pull_image: parse_bool_env("SANDBOX_AUTO_PULL", ss.auto_pull_image)?,
memory_limit_mb: db_first_or_default(
&ss.memory_limit_mb,
&defaults.memory_limit_mb,
"SANDBOX_MEMORY_LIMIT_MB",
)?,
cpu_shares: db_first_or_default(
&ss.cpu_shares,
&defaults.cpu_shares,
"SANDBOX_CPU_SHARES",
)?,
image: db_first_or_default(&ss.image, &defaults.image, "SANDBOX_IMAGE")?,
auto_pull_image: db_first_bool(
ss.auto_pull_image,
defaults.auto_pull_image,
"SANDBOX_AUTO_PULL",
)?,
extra_allowed_domains: extra_domains, extra_allowed_domains: extra_domains,
reaper_interval_secs, reaper_interval_secs,
orphan_threshold_secs, orphan_threshold_secs,
@@ -283,28 +264,19 @@ impl ClaudeCodeConfig {
} }
pub(crate) fn resolve(settings: &crate::settings::Settings) -> Result<Self, ConfigError> { pub(crate) fn resolve(settings: &crate::settings::Settings) -> Result<Self, ConfigError> {
let ss = &settings.sandbox;
let defaults = Self::default(); let defaults = Self::default();
Ok(Self { Ok(Self {
enabled: db_first_bool( // Use settings.sandbox.claude_code_enabled as fallback (written by setup wizard).
ss.claude_code_enabled, enabled: parse_bool_env("CLAUDE_CODE_ENABLED", settings.sandbox.claude_code_enabled)?,
defaults.enabled,
"CLAUDE_CODE_ENABLED",
)?,
// config_dir has no Settings counterpart — env > default only.
config_dir: optional_env("CLAUDE_CONFIG_DIR")? config_dir: optional_env("CLAUDE_CONFIG_DIR")?
.map(std::path::PathBuf::from) .map(std::path::PathBuf::from)
.unwrap_or(defaults.config_dir), .unwrap_or(defaults.config_dir),
// model has no Settings counterpart — env > default only.
model: parse_string_env("CLAUDE_CODE_MODEL", defaults.model)?, model: parse_string_env("CLAUDE_CODE_MODEL", defaults.model)?,
// max_turns has no Settings counterpart — env > default only.
max_turns: parse_optional_env("CLAUDE_CODE_MAX_TURNS", defaults.max_turns)?, max_turns: parse_optional_env("CLAUDE_CODE_MAX_TURNS", defaults.max_turns)?,
// memory_limit_mb has no Settings counterpart — env > default only.
memory_limit_mb: parse_optional_env( memory_limit_mb: parse_optional_env(
"CLAUDE_CODE_MEMORY_LIMIT_MB", "CLAUDE_CODE_MEMORY_LIMIT_MB",
defaults.memory_limit_mb, defaults.memory_limit_mb,
)?, )?,
// allowed_tools has no Settings counterpart — env > default only.
allowed_tools: optional_env("CLAUDE_CODE_ALLOWED_TOOLS")? allowed_tools: optional_env("CLAUDE_CODE_ALLOWED_TOOLS")?
.map(|s| { .map(|s| {
s.split(',') s.split(',')
@@ -635,7 +607,7 @@ mod tests {
} }
#[test] #[test]
fn sandbox_db_settings_override_env() { fn sandbox_env_overrides_settings() {
let _guard = crate::config::helpers::lock_env(); let _guard = crate::config::helpers::lock_env();
let mut settings = crate::settings::Settings::default(); let mut settings = crate::settings::Settings::default();
settings.sandbox.timeout_secs = 999; settings.sandbox.timeout_secs = 999;
@@ -645,26 +617,7 @@ mod tests {
let cfg = SandboxModeConfig::resolve(&settings).expect("resolve"); let cfg = SandboxModeConfig::resolve(&settings).expect("resolve");
unsafe { std::env::remove_var("SANDBOX_TIMEOUT_SECS") }; unsafe { std::env::remove_var("SANDBOX_TIMEOUT_SECS") };
// DB value (999) wins over env (5) under DB-first priority. assert_eq!(cfg.timeout_secs, 5);
assert_eq!(cfg.timeout_secs, 999);
}
#[test]
fn sandbox_env_used_when_no_db_setting() {
let _guard = crate::config::helpers::lock_env();
// Default settings — all fields at their defaults, so DB is "unset".
let settings = crate::settings::Settings::default();
// SAFETY: Under ENV_MUTEX, no concurrent env access.
unsafe { std::env::set_var("SANDBOX_TIMEOUT_SECS", "42") };
unsafe { std::env::set_var("SANDBOX_MEMORY_LIMIT_MB", "512") };
let cfg = SandboxModeConfig::resolve(&settings).expect("resolve");
unsafe { std::env::remove_var("SANDBOX_TIMEOUT_SECS") };
unsafe { std::env::remove_var("SANDBOX_MEMORY_LIMIT_MB") };
// Env values win when settings are at their defaults.
assert_eq!(cfg.timeout_secs, 42);
assert_eq!(cfg.memory_limit_mb, 512);
} }
// ── ClaudeCodeConfig settings fallback tests ──────────────────── // ── ClaudeCodeConfig settings fallback tests ────────────────────
@@ -688,7 +641,7 @@ mod tests {
} }
#[test] #[test]
fn claude_code_db_settings_override_env() { fn claude_code_env_overrides_settings() {
let _guard = crate::config::helpers::lock_env(); let _guard = crate::config::helpers::lock_env();
let mut settings = crate::settings::Settings::default(); let mut settings = crate::settings::Settings::default();
settings.sandbox.claude_code_enabled = true; settings.sandbox.claude_code_enabled = true;
@@ -698,8 +651,7 @@ mod tests {
let cfg = ClaudeCodeConfig::resolve(&settings).expect("resolve"); let cfg = ClaudeCodeConfig::resolve(&settings).expect("resolve");
unsafe { std::env::remove_var("CLAUDE_CODE_ENABLED") }; unsafe { std::env::remove_var("CLAUDE_CODE_ENABLED") };
// DB value (true) wins over env (false) under DB-first priority. assert!(!cfg.enabled);
assert!(cfg.enabled);
} }
#[test] #[test]
+24 -77
View File
@@ -1,6 +1,5 @@
use crate::config::helpers::{db_first_or_default, parse_optional_env}; use crate::config::helpers::{optional_env, parse_optional_env};
use crate::error::ConfigError; use crate::error::ConfigError;
use crate::settings::Settings;
use crate::workspace::FusionStrategy; use crate::workspace::FusionStrategy;
/// Workspace search configuration resolved from environment variables. /// Workspace search configuration resolved from environment variables.
@@ -34,48 +33,30 @@ impl Default for WorkspaceSearchConfig {
} }
impl WorkspaceSearchConfig { impl WorkspaceSearchConfig {
pub(crate) fn resolve(settings: &Settings) -> Result<Self, ConfigError> { pub(crate) fn resolve() -> Result<Self, ConfigError> {
let defaults = crate::settings::SearchSettings::default(); let fusion_strategy = match optional_env("SEARCH_FUSION_STRATEGY")? {
let ss = &settings.search; Some(s) => match s.to_lowercase().as_str() {
"rrf" => FusionStrategy::Rrf,
// Resolve fusion_strategy string via DB-first, then parse into enum. "weighted" => FusionStrategy::WeightedScore,
let strategy_str = db_first_or_default( other => {
&ss.fusion_strategy, return Err(ConfigError::InvalidValue {
&defaults.fusion_strategy, key: "SEARCH_FUSION_STRATEGY".to_string(),
"SEARCH_FUSION_STRATEGY", message: format!("must be 'rrf' or 'weighted', got '{other}'"),
)?; });
let fusion_strategy = match strategy_str.to_lowercase().as_str() { }
"rrf" => FusionStrategy::Rrf, },
"weighted" => FusionStrategy::WeightedScore, None => FusionStrategy::default(),
other => {
return Err(ConfigError::InvalidValue {
key: "SEARCH_FUSION_STRATEGY".to_string(),
message: format!("must be 'rrf' or 'weighted', got '{other}'"),
});
}
}; };
let rrf_k = db_first_or_default(&ss.rrf_k, &defaults.rrf_k, "SEARCH_RRF_K")?; let rrf_k = parse_optional_env("SEARCH_RRF_K", 60u32)?;
// Per-strategy weight defaults: RRF uses 0.5/0.5, weighted uses 0.3/0.7 (vector-biased). // Per-strategy weight defaults: RRF uses 0.5/0.5, weighted uses 0.3/0.7 (vector-biased).
let (default_fts, default_vec) = match fusion_strategy { let (default_fts, default_vec) = match fusion_strategy {
FusionStrategy::Rrf => (0.5f32, 0.5f32), FusionStrategy::Rrf => (0.5f32, 0.5f32),
FusionStrategy::WeightedScore => (0.3f32, 0.7f32), FusionStrategy::WeightedScore => (0.3f32, 0.7f32),
}; };
let fts_weight = parse_optional_env("SEARCH_FTS_WEIGHT", default_fts)?;
// For weights, we need to check whether the settings value differs from let vector_weight = parse_optional_env("SEARCH_VECTOR_WEIGHT", default_vec)?;
// the *static* default (0.5) to detect DB overrides. If it does, use it;
// otherwise fall back to env, then per-strategy default.
let fts_weight = if (ss.fts_weight - defaults.fts_weight).abs() > f32::EPSILON {
ss.fts_weight
} else {
parse_optional_env("SEARCH_FTS_WEIGHT", default_fts)?
};
let vector_weight = if (ss.vector_weight - defaults.vector_weight).abs() > f32::EPSILON {
ss.vector_weight
} else {
parse_optional_env("SEARCH_VECTOR_WEIGHT", default_vec)?
};
if !fts_weight.is_finite() || fts_weight < 0.0 { if !fts_weight.is_finite() || fts_weight < 0.0 {
return Err(ConfigError::InvalidValue { return Err(ConfigError::InvalidValue {
@@ -128,8 +109,7 @@ mod tests {
let _guard = lock_env(); let _guard = lock_env();
clear_search_env(); clear_search_env();
let settings = Settings::default(); let config = WorkspaceSearchConfig::resolve().expect("should resolve");
let config = WorkspaceSearchConfig::resolve(&settings).expect("should resolve");
assert_eq!(config.fusion_strategy, FusionStrategy::Rrf); assert_eq!(config.fusion_strategy, FusionStrategy::Rrf);
assert_eq!(config.rrf_k, 60); assert_eq!(config.rrf_k, 60);
assert!((config.fts_weight - 0.5).abs() < 0.001); assert!((config.fts_weight - 0.5).abs() < 0.001);
@@ -137,35 +117,7 @@ mod tests {
} }
#[test] #[test]
fn db_settings_override_env() { fn env_overrides() {
let _guard = lock_env();
clear_search_env();
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::set_var("SEARCH_FUSION_STRATEGY", "rrf");
std::env::set_var("SEARCH_RRF_K", "30");
std::env::set_var("SEARCH_FTS_WEIGHT", "0.9");
std::env::set_var("SEARCH_VECTOR_WEIGHT", "0.1");
}
let mut settings = Settings::default();
settings.search.fusion_strategy = "weighted".to_string();
settings.search.rrf_k = 42;
settings.search.fts_weight = 0.4;
settings.search.vector_weight = 0.6;
let config = WorkspaceSearchConfig::resolve(&settings).expect("should resolve");
assert_eq!(config.fusion_strategy, FusionStrategy::WeightedScore);
assert_eq!(config.rrf_k, 42);
assert!((config.fts_weight - 0.4).abs() < 0.001);
assert!((config.vector_weight - 0.6).abs() < 0.001);
clear_search_env();
}
#[test]
fn env_fallback_when_settings_at_default() {
let _guard = lock_env(); let _guard = lock_env();
clear_search_env(); clear_search_env();
@@ -177,8 +129,7 @@ mod tests {
std::env::set_var("SEARCH_VECTOR_WEIGHT", "0.1"); std::env::set_var("SEARCH_VECTOR_WEIGHT", "0.1");
} }
let settings = Settings::default(); let config = WorkspaceSearchConfig::resolve().expect("should resolve");
let config = WorkspaceSearchConfig::resolve(&settings).expect("should resolve");
assert_eq!(config.fusion_strategy, FusionStrategy::WeightedScore); assert_eq!(config.fusion_strategy, FusionStrategy::WeightedScore);
assert_eq!(config.rrf_k, 30); assert_eq!(config.rrf_k, 30);
assert!((config.fts_weight - 0.9).abs() < 0.001); assert!((config.fts_weight - 0.9).abs() < 0.001);
@@ -197,8 +148,7 @@ mod tests {
std::env::set_var("SEARCH_FUSION_STRATEGY", "bm25"); std::env::set_var("SEARCH_FUSION_STRATEGY", "bm25");
} }
let settings = Settings::default(); let result = WorkspaceSearchConfig::resolve();
let result = WorkspaceSearchConfig::resolve(&settings);
assert!(result.is_err()); assert!(result.is_err());
clear_search_env(); clear_search_env();
@@ -214,8 +164,7 @@ mod tests {
std::env::set_var("SEARCH_FUSION_STRATEGY", "weighted"); std::env::set_var("SEARCH_FUSION_STRATEGY", "weighted");
} }
let settings = Settings::default(); let config = WorkspaceSearchConfig::resolve().expect("should resolve");
let config = WorkspaceSearchConfig::resolve(&settings).expect("should resolve");
assert_eq!(config.fusion_strategy, FusionStrategy::WeightedScore); assert_eq!(config.fusion_strategy, FusionStrategy::WeightedScore);
// Weighted mode should default to 0.3 FTS / 0.7 vector // Weighted mode should default to 0.3 FTS / 0.7 vector
assert!((config.fts_weight - 0.3).abs() < 0.001); assert!((config.fts_weight - 0.3).abs() < 0.001);
@@ -236,8 +185,7 @@ mod tests {
std::env::set_var("SEARCH_VECTOR_WEIGHT", "0.0"); std::env::set_var("SEARCH_VECTOR_WEIGHT", "0.0");
} }
let settings = Settings::default(); let result = WorkspaceSearchConfig::resolve();
let result = WorkspaceSearchConfig::resolve(&settings);
assert!(result.is_err()); assert!(result.is_err());
clear_search_env(); clear_search_env();
@@ -255,8 +203,7 @@ mod tests {
} }
// RRF ignores weights, so both=0 is fine // RRF ignores weights, so both=0 is fine
let settings = Settings::default(); let config = WorkspaceSearchConfig::resolve().expect("should resolve");
let config = WorkspaceSearchConfig::resolve(&settings).expect("should resolve");
assert_eq!(config.fusion_strategy, FusionStrategy::Rrf); assert_eq!(config.fusion_strategy, FusionStrategy::Rrf);
clear_search_env(); clear_search_env();
+5 -18
View File
@@ -1,9 +1,8 @@
use std::path::PathBuf; use std::path::PathBuf;
use crate::bootstrap::ironclaw_base_dir; use crate::bootstrap::ironclaw_base_dir;
use crate::config::helpers::{db_first_bool, db_first_or_default, 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;
/// Skills system configuration. /// Skills system configuration.
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
@@ -45,29 +44,17 @@ fn default_installed_skills_dir() -> PathBuf {
} }
impl SkillsConfig { impl SkillsConfig {
pub(crate) fn resolve(settings: &Settings) -> Result<Self, ConfigError> { pub(crate) fn resolve() -> Result<Self, ConfigError> {
let defaults = crate::settings::SkillsSettings::default();
let ss = &settings.skills;
Ok(Self { Ok(Self {
enabled: db_first_bool(ss.enabled, defaults.enabled, "SKILLS_ENABLED")?, enabled: parse_bool_env("SKILLS_ENABLED", true)?,
// local_dir and installed_dir are env-only (filesystem paths, no settings counterpart)
local_dir: optional_env("SKILLS_DIR")? local_dir: optional_env("SKILLS_DIR")?
.map(PathBuf::from) .map(PathBuf::from)
.unwrap_or_else(default_skills_dir), .unwrap_or_else(default_skills_dir),
installed_dir: optional_env("SKILLS_INSTALLED_DIR")? installed_dir: optional_env("SKILLS_INSTALLED_DIR")?
.map(PathBuf::from) .map(PathBuf::from)
.unwrap_or_else(default_installed_skills_dir), .unwrap_or_else(default_installed_skills_dir),
max_active_skills: db_first_or_default( max_active_skills: parse_optional_env("SKILLS_MAX_ACTIVE", 3)?,
&ss.max_active_skills, max_context_tokens: parse_optional_env("SKILLS_MAX_CONTEXT_TOKENS", 4000)?,
&defaults.max_active_skills,
"SKILLS_MAX_ACTIVE",
)?,
max_context_tokens: db_first_or_default(
&ss.max_context_tokens,
&defaults.max_context_tokens,
"SKILLS_MAX_CONTEXT_TOKENS",
)?,
}) })
} }
} }
+5 -3
View File
@@ -1,6 +1,6 @@
use secrecy::SecretString; use secrecy::SecretString;
use crate::config::helpers::{db_first_bool, optional_env, validate_base_url}; use crate::config::helpers::{optional_env, parse_bool_env, validate_base_url};
use crate::error::ConfigError; use crate::error::ConfigError;
use crate::settings::Settings; use crate::settings::Settings;
@@ -39,8 +39,10 @@ impl Default for TranscriptionConfig {
impl TranscriptionConfig { impl TranscriptionConfig {
pub(crate) fn resolve(settings: &Settings) -> Result<Self, ConfigError> { pub(crate) fn resolve(settings: &Settings) -> Result<Self, ConfigError> {
let settings_enabled = settings.transcription.as_ref().is_some_and(|t| t.enabled); let enabled = parse_bool_env(
let enabled = db_first_bool(settings_enabled, false, "TRANSCRIPTION_ENABLED")?; "TRANSCRIPTION_ENABLED",
settings.transcription.as_ref().is_some_and(|t| t.enabled),
)?;
let provider = let provider =
optional_env("TRANSCRIPTION_PROVIDER")?.unwrap_or_else(|| "openai".to_string()); optional_env("TRANSCRIPTION_PROVIDER")?.unwrap_or_else(|| "openai".to_string());
+29 -42
View File
@@ -1,14 +1,12 @@
use crate::config::helpers::{db_first_bool, db_first_optional_string}; use crate::config::helpers::optional_env;
use crate::error::ConfigError; use crate::error::ConfigError;
use crate::settings::{Settings, TunnelSettings}; use crate::settings::Settings;
/// Tunnel configuration for exposing the agent to the internet. /// Tunnel configuration for exposing the agent to the internet.
/// ///
/// Used by channels and tools that need public webhook endpoints. /// Used by channels and tools that need public webhook endpoints.
/// The tunnel URL is shared across all channels (Telegram, Slack, etc.). /// The tunnel URL is shared across all channels (Telegram, Slack, etc.).
/// ///
/// Resolution priority: DB/settings > env var > default.
///
/// Two modes: /// Two modes:
/// - **Static URL** (`TUNNEL_URL`): set the public URL directly (manual tunnel) /// - **Static URL** (`TUNNEL_URL`): set the public URL directly (manual tunnel)
/// - **Managed provider** (`TUNNEL_PROVIDER`): lifecycle-managed tunnel process /// - **Managed provider** (`TUNNEL_PROVIDER`): lifecycle-managed tunnel process
@@ -27,10 +25,8 @@ pub struct TunnelConfig {
impl TunnelConfig { impl TunnelConfig {
pub(crate) fn resolve(settings: &Settings) -> Result<Self, ConfigError> { pub(crate) fn resolve(settings: &Settings) -> Result<Self, ConfigError> {
let defaults = TunnelSettings::default(); let public_url = optional_env("TUNNEL_URL")?
.or_else(|| settings.tunnel.public_url.clone().filter(|s| !s.is_empty()));
// Priority: DB/settings > env > default.
let public_url = db_first_optional_string(&settings.tunnel.public_url, "TUNNEL_URL")?;
if let Some(ref url) = public_url if let Some(ref url) = public_url
&& !url.starts_with("https://") && !url.starts_with("https://")
@@ -42,8 +38,9 @@ impl TunnelConfig {
} }
// Resolve managed tunnel provider config. // Resolve managed tunnel provider config.
// Priority: DB/settings > env > default (none). // Priority: env var > settings > default (none).
let provider_name = db_first_optional_string(&settings.tunnel.provider, "TUNNEL_PROVIDER")? let provider_name = optional_env("TUNNEL_PROVIDER")?
.or_else(|| settings.tunnel.provider.clone())
.unwrap_or_default(); .unwrap_or_default();
let provider = if provider_name.is_empty() || provider_name == "none" { let provider = if provider_name.is_empty() || provider_name == "none" {
@@ -51,48 +48,38 @@ impl TunnelConfig {
} else { } else {
Some(crate::tunnel::TunnelProviderConfig { Some(crate::tunnel::TunnelProviderConfig {
provider: provider_name.clone(), provider: provider_name.clone(),
cloudflare: db_first_optional_string(&settings.tunnel.cf_token, "TUNNEL_CF_TOKEN")? cloudflare: optional_env("TUNNEL_CF_TOKEN")?
.or_else(|| settings.tunnel.cf_token.clone())
.map(|token| crate::tunnel::CloudflareTunnelConfig { token }), .map(|token| crate::tunnel::CloudflareTunnelConfig { token }),
tailscale: Some(crate::tunnel::TailscaleTunnelConfig { tailscale: Some(crate::tunnel::TailscaleTunnelConfig {
funnel: db_first_bool( funnel: optional_env("TUNNEL_TS_FUNNEL")?
settings.tunnel.ts_funnel, .map(|s| s == "true" || s == "1")
defaults.ts_funnel, .unwrap_or(settings.tunnel.ts_funnel),
"TUNNEL_TS_FUNNEL", hostname: optional_env("TUNNEL_TS_HOSTNAME")?
)?, .or_else(|| settings.tunnel.ts_hostname.clone()),
hostname: db_first_optional_string(
&settings.tunnel.ts_hostname,
"TUNNEL_TS_HOSTNAME",
)?,
}), }),
ngrok: { ngrok: {
let ngrok_domain = db_first_optional_string( let ngrok_domain = optional_env("TUNNEL_NGROK_DOMAIN")?
&settings.tunnel.ngrok_domain, .or_else(|| settings.tunnel.ngrok_domain.clone());
"TUNNEL_NGROK_DOMAIN", optional_env("TUNNEL_NGROK_TOKEN")?
)?; .or_else(|| settings.tunnel.ngrok_token.clone())
db_first_optional_string(&settings.tunnel.ngrok_token, "TUNNEL_NGROK_TOKEN")?
.map(|auth_token| crate::tunnel::NgrokTunnelConfig { .map(|auth_token| crate::tunnel::NgrokTunnelConfig {
auth_token, auth_token,
domain: ngrok_domain, domain: ngrok_domain,
}) })
}, },
custom: { custom: {
let health_url = db_first_optional_string( let health_url = optional_env("TUNNEL_CUSTOM_HEALTH_URL")?
&settings.tunnel.custom_health_url, .or_else(|| settings.tunnel.custom_health_url.clone());
"TUNNEL_CUSTOM_HEALTH_URL", let url_pattern = optional_env("TUNNEL_CUSTOM_URL_PATTERN")?
)?; .or_else(|| settings.tunnel.custom_url_pattern.clone());
let url_pattern = db_first_optional_string( optional_env("TUNNEL_CUSTOM_COMMAND")?
&settings.tunnel.custom_url_pattern, .or_else(|| settings.tunnel.custom_command.clone())
"TUNNEL_CUSTOM_URL_PATTERN", .map(|start_command| crate::tunnel::CustomTunnelConfig {
)?; start_command,
db_first_optional_string( health_url,
&settings.tunnel.custom_command, url_pattern,
"TUNNEL_CUSTOM_COMMAND", })
)?
.map(|start_command| crate::tunnel::CustomTunnelConfig {
start_command,
health_url,
url_pattern,
})
}, },
}) })
}; };
+17 -43
View File
@@ -2,7 +2,7 @@ use std::path::PathBuf;
use std::time::Duration; use std::time::Duration;
use crate::bootstrap::ironclaw_base_dir; use crate::bootstrap::ironclaw_base_dir;
use crate::config::helpers::{db_first_bool, db_first_or_default, optional_env}; use crate::config::helpers::{optional_env, parse_bool_env, parse_optional_env};
use crate::error::ConfigError; use crate::error::ConfigError;
/// WASM sandbox configuration. /// WASM sandbox configuration.
@@ -46,41 +46,28 @@ fn default_tools_dir() -> PathBuf {
impl WasmConfig { impl WasmConfig {
pub(crate) fn resolve(settings: &crate::settings::Settings) -> Result<Self, ConfigError> { pub(crate) fn resolve(settings: &crate::settings::Settings) -> Result<Self, ConfigError> {
let ws = &settings.wasm; let ws = &settings.wasm;
let defaults = crate::settings::WasmSettings::default();
Ok(Self { Ok(Self {
enabled: db_first_bool(ws.enabled, defaults.enabled, "WASM_ENABLED")?, enabled: parse_bool_env("WASM_ENABLED", ws.enabled)?,
tools_dir: if let Some(ref dir) = ws.tools_dir { tools_dir: optional_env("WASM_TOOLS_DIR")?
dir.clone() .map(PathBuf::from)
} else { .or_else(|| ws.tools_dir.clone())
optional_env("WASM_TOOLS_DIR")? .unwrap_or_else(default_tools_dir),
.map(PathBuf::from) default_memory_limit: parse_optional_env(
.unwrap_or_else(default_tools_dir)
},
default_memory_limit: db_first_or_default(
&ws.default_memory_limit,
&defaults.default_memory_limit,
"WASM_DEFAULT_MEMORY_LIMIT", "WASM_DEFAULT_MEMORY_LIMIT",
ws.default_memory_limit,
)?, )?,
default_timeout_secs: db_first_or_default( default_timeout_secs: parse_optional_env(
&ws.default_timeout_secs,
&defaults.default_timeout_secs,
"WASM_DEFAULT_TIMEOUT_SECS", "WASM_DEFAULT_TIMEOUT_SECS",
ws.default_timeout_secs,
)?, )?,
default_fuel_limit: db_first_or_default( default_fuel_limit: parse_optional_env(
&ws.default_fuel_limit,
&defaults.default_fuel_limit,
"WASM_DEFAULT_FUEL_LIMIT", "WASM_DEFAULT_FUEL_LIMIT",
ws.default_fuel_limit,
)?, )?,
cache_compiled: db_first_bool( cache_compiled: parse_bool_env("WASM_CACHE_COMPILED", ws.cache_compiled)?,
ws.cache_compiled, cache_dir: optional_env("WASM_CACHE_DIR")?
defaults.cache_compiled, .map(PathBuf::from)
"WASM_CACHE_COMPILED", .or_else(|| ws.cache_dir.clone()),
)?,
cache_dir: if let Some(ref dir) = ws.cache_dir {
Some(dir.clone())
} else {
optional_env("WASM_CACHE_DIR")?.map(PathBuf::from)
},
}) })
} }
@@ -124,7 +111,7 @@ mod tests {
} }
#[test] #[test]
fn db_settings_override_env() { fn env_overrides_settings() {
let _guard = lock_env(); let _guard = lock_env();
let mut settings = Settings::default(); let mut settings = Settings::default();
settings.wasm.default_fuel_limit = 42; settings.wasm.default_fuel_limit = 42;
@@ -134,19 +121,6 @@ mod tests {
let cfg = WasmConfig::resolve(&settings).expect("resolve"); let cfg = WasmConfig::resolve(&settings).expect("resolve");
unsafe { std::env::remove_var("WASM_DEFAULT_FUEL_LIMIT") }; unsafe { std::env::remove_var("WASM_DEFAULT_FUEL_LIMIT") };
assert_eq!(cfg.default_fuel_limit, 42);
}
#[test]
fn env_used_when_no_db_setting() {
let _guard = lock_env();
let settings = Settings::default();
// SAFETY: Under ENV_MUTEX, no concurrent env access.
unsafe { std::env::set_var("WASM_DEFAULT_FUEL_LIMIT", "7") };
let cfg = WasmConfig::resolve(&settings).expect("resolve");
unsafe { std::env::remove_var("WASM_DEFAULT_FUEL_LIMIT") };
assert_eq!(cfg.default_fuel_limit, 7); assert_eq!(cfg.default_fuel_limit, 7);
} }
} }
-3
View File
@@ -192,9 +192,6 @@ pub struct JobContext {
/// but subsequent tools (e.g., `json`) may need the full output. This /// but subsequent tools (e.g., `json`) may need the full output. This
/// stash stores the complete, unsanitized output so tools can reference /// stash stores the complete, unsanitized output so tools can reference
/// previous results by ID via `$tool_call_id` parameter syntax. /// previous results by ID via `$tool_call_id` parameter syntax.
///
/// Also used for cross-tool implicit state (keys prefixed with `__`) such
/// as `__routine_last_name` for fallback recovery in routine tool chains.
#[serde(skip)] #[serde(skip)]
pub tool_output_stash: Arc<tokio::sync::RwLock<HashMap<String, String>>>, pub tool_output_stash: Arc<tokio::sync::RwLock<HashMap<String, String>>>,
/// User's preferred timezone (IANA name, e.g. "America/New_York"). Defaults to "UTC". /// User's preferred timezone (IANA name, e.g. "America/New_York"). Defaults to "UTC".
+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()
-10
View File
@@ -528,15 +528,6 @@ pub trait RoutineStore: Send + Sync {
&self, &self,
routine_ids: &[Uuid], routine_ids: &[Uuid],
) -> Result<HashMap<Uuid, i64>, DatabaseError>; ) -> Result<HashMap<Uuid, i64>, DatabaseError>;
/// Fetch the last run status for multiple routines in a single query.
/// Returns a map from routine_id to its most recent RunStatus.
/// Routines with no runs are omitted from the result.
async fn batch_get_last_run_status(
&self,
routine_ids: &[Uuid],
) -> Result<HashMap<Uuid, RunStatus>, DatabaseError>;
async fn link_routine_run_to_job( async fn link_routine_run_to_job(
&self, &self,
run_id: Uuid, run_id: Uuid,
@@ -545,7 +536,6 @@ pub trait RoutineStore: Send + Sync {
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>;
/// List routine runs that were dispatched as full_job but have not yet /// List routine runs that were dispatched as full_job but have not yet
+1 -10
View File
@@ -510,14 +510,6 @@ impl RoutineStore for PgBackend {
.await .await
} }
async fn batch_get_last_run_status(
&self,
routine_ids: &[Uuid],
) -> Result<std::collections::HashMap<Uuid, crate::agent::routine::RunStatus>, DatabaseError>
{
self.store.batch_get_last_run_status(routine_ids).await
}
async fn link_routine_run_to_job( async fn link_routine_run_to_job(
&self, &self,
run_id: Uuid, run_id: Uuid,
@@ -529,9 +521,8 @@ impl RoutineStore for PgBackend {
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> {
self.store.get_webhook_routine_by_path(path, user_id).await self.store.get_webhook_routine_by_path(path).await
} }
async fn list_dispatched_routine_runs(&self) -> Result<Vec<RoutineRun>, DatabaseError> { async fn list_dispatched_routine_runs(&self) -> Result<Vec<RoutineRun>, DatabaseError> {
+110 -455
View File
@@ -53,6 +53,22 @@ struct HostedOAuthFlowStart {
flow: crate::cli::oauth_defaults::PendingOAuthFlow, flow: crate::cli::oauth_defaults::PendingOAuthFlow,
} }
fn hosted_proxy_client_secret(
client_secret: &Option<String>,
builtin: Option<&crate::cli::oauth_defaults::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(),
}
}
fn normalize_oauth_callback_path(path: &str) -> String { fn normalize_oauth_callback_path(path: &str) -> String {
let trimmed_path = path.trim_end_matches('/'); let trimmed_path = path.trim_end_matches('/');
if trimmed_path.is_empty() { if trimmed_path.is_empty() {
@@ -403,10 +419,9 @@ pub struct ExtensionManager {
/// when running in gateway mode, consumed by the web gateway's /// when running in gateway mode, consumed by the web gateway's
/// `/oauth/callback` handler. /// `/oauth/callback` handler.
pending_oauth_flows: crate::cli::oauth_defaults::PendingOAuthRegistry, pending_oauth_flows: crate::cli::oauth_defaults::PendingOAuthRegistry,
/// OAuth proxy auth token for authenticating with the hosted token exchange proxy. /// Gateway auth token for authenticating with the platform token exchange proxy.
/// Resolved once at construction from `IRONCLAW_OAUTH_PROXY_AUTH_TOKEN`, /// Read once at construction from `GATEWAY_AUTH_TOKEN` env var.
/// then `GATEWAY_AUTH_TOKEN` as a backward-compatible fallback. gateway_token: Option<String>,
oauth_proxy_auth_token: Option<String>,
/// Relay config captured at startup. Used by `auth_channel_relay` and /// Relay config captured at startup. Used by `auth_channel_relay` and
/// `activate_channel_relay` instead of re-reading env vars. /// `activate_channel_relay` instead of re-reading env vars.
relay_config: Option<crate::config::RelayConfig>, relay_config: Option<crate::config::RelayConfig>,
@@ -536,7 +551,7 @@ impl ExtensionManager {
activation_errors: RwLock::new(HashMap::new()), activation_errors: RwLock::new(HashMap::new()),
sse_manager: RwLock::new(None), sse_manager: RwLock::new(None),
pending_oauth_flows: crate::cli::oauth_defaults::new_pending_oauth_registry(), pending_oauth_flows: crate::cli::oauth_defaults::new_pending_oauth_registry(),
oauth_proxy_auth_token: crate::cli::oauth_defaults::oauth_proxy_auth_token(), gateway_token: std::env::var("GATEWAY_AUTH_TOKEN").ok(),
relay_config: crate::config::RelayConfig::from_env(), relay_config: crate::config::RelayConfig::from_env(),
relay_event_tx: Arc::new(tokio::sync::Mutex::new(None)), relay_event_tx: Arc::new(tokio::sync::Mutex::new(None)),
relay_signing_secret_cache: Arc::new(std::sync::Mutex::new(None)), relay_signing_secret_cache: Arc::new(std::sync::Mutex::new(None)),
@@ -660,66 +675,6 @@ impl ExtensionManager {
}) })
} }
/// Resolve the relay URL override for an extension from settings.
///
/// Returns `Some(url)` if a non-empty per-extension `relay_url` override is
/// set for the given extension; otherwise returns `None` and callers should
/// fall back to the env-level `RelayConfig`.
///
/// Uses `self.user_id` (owner scope) for consistency with `configure()`,
/// which also writes setting_path fields under the owner scope.
///
/// The override is validated: only `http` / `https` schemes are accepted
/// and the URL must not contain userinfo (embedded credentials). This
/// prevents a malicious override from exfiltrating the instance-wide relay
/// API key to an attacker-controlled host.
async fn effective_relay_url(&self, name: &str) -> Option<String> {
if let Some(ref store) = self.store {
let key = format!("extensions.{name}.relay_url");
if let Ok(Some(v)) = store.get_setting(&self.user_id, &key).await {
let url = v
.as_str()
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty());
if let Some(ref u) = url {
// Validate the override to prevent API-key exfiltration:
// only allow http(s) with no embedded credentials.
match url::Url::parse(u) {
Ok(parsed)
if (parsed.scheme() == "http" || parsed.scheme() == "https")
&& parsed.username().is_empty()
&& parsed.password().is_none() =>
{
tracing::trace!(
extension = %name,
relay_url_host = %parsed.host_str().unwrap_or("unknown"),
"effective_relay_url: using per-extension override from settings"
);
return url;
}
Ok(parsed) => {
tracing::warn!(
extension = %name,
scheme = %parsed.scheme(),
has_userinfo = !parsed.username().is_empty() || parsed.password().is_some(),
"effective_relay_url: rejecting override — \
only http/https without embedded credentials is allowed"
);
}
Err(e) => {
tracing::warn!(
extension = %name,
error = %e,
"effective_relay_url: rejecting override — invalid URL"
);
}
}
}
}
}
None
}
/// Get the shared relay event sender for the webhook endpoint. /// Get the shared relay event sender for the webhook endpoint.
pub fn relay_event_tx( pub fn relay_event_tx(
&self, &self,
@@ -936,67 +891,24 @@ impl ExtensionManager {
*self.relay_channel_manager.write().await = Some(channel_manager); *self.relay_channel_manager.write().await = Some(channel_manager);
} }
/// Check if a channel name corresponds to a relay extension (has stored team_id /// Check if a channel name corresponds to a relay extension (has stored stream token
/// or is tracked in the installed relay extensions set). /// or is tracked in the installed relay extensions set).
pub async fn is_relay_channel(&self, name: &str, user_id: &str) -> bool { pub async fn is_relay_channel(&self, name: &str, user_id: &str) -> bool {
// Check in-memory installed set first (supports no-store mode) // Check in-memory installed set first (supports no-store mode)
if self.installed_relay_extensions.read().await.contains(name) { if self.installed_relay_extensions.read().await.contains(name) {
return true; return true;
} }
// Check for stored team_id (persisted across restarts by the OAuth callback) // Then check for stored stream token
if let Some(ref store) = self.store { self.secrets
let key = format!("relay:{}:team_id", name); .exists(user_id, &format!("relay:{}:stream_token", name))
if let Ok(Some(v)) = store.get_setting(user_id, &key).await { .await
return v.as_str().is_some_and(|s| !s.is_empty()); .unwrap_or(false)
}
}
false
}
/// Check whether a stored `team_id` setting exists for the given relay extension.
///
/// Unlike [`is_relay_channel`], this does **not** consult the in-memory
/// `installed_relay_extensions` set — it only looks at the persistent settings
/// store. This distinction matters for `auth_channel_relay`: an extension can
/// be *installed* (present in the in-memory set) but not yet *authenticated*
/// (no OAuth completed, no team_id stored).
async fn has_stored_team_id(&self, name: &str, _user_id: &str) -> bool {
if let Some(ref store) = self.store {
let key = format!("relay:{}:team_id", name);
// Use owner scope (self.user_id) for consistency: the OAuth callback
// stores team_id under state.owner_id which maps to self.user_id.
match store.get_setting(&self.user_id, &key).await {
Ok(Some(v)) => {
let has_id = v.as_str().is_some_and(|s| !s.is_empty());
tracing::trace!(
extension = %name,
has_team_id = has_id,
"has_stored_team_id: checked store"
);
return has_id;
}
Ok(None) => {
tracing::trace!(
extension = %name,
"has_stored_team_id: no team_id setting found"
);
}
Err(e) => {
tracing::warn!(
extension = %name,
error = %e,
"has_stored_team_id: failed to read from settings store"
);
}
}
}
false
} }
/// Restore persisted relay channels after startup. /// Restore persisted relay channels after startup.
/// ///
/// Loads the persisted active channel list, filters to relay types (those with /// Loads the persisted active channel list, filters to relay types (those with
/// a stored team_id setting), and activates each via `activate_stored_relay()`. /// a stored stream token), and activates each via `activate_stored_relay()`.
/// Skips channels that are already active. /// Skips channels that are already active.
/// ///
/// Call this only after `set_relay_channel_manager()` or `set_channel_runtime()`. /// Call this only after `set_relay_channel_manager()` or `set_channel_runtime()`.
@@ -1219,7 +1131,7 @@ impl ExtensionManager {
/// Broadcast an extension status change to the web UI via SSE. /// Broadcast an extension status change to the web UI via SSE.
async fn broadcast_extension_status(&self, name: &str, status: &str, message: Option<&str>) { async fn broadcast_extension_status(&self, name: &str, status: &str, message: Option<&str>) {
if let Some(ref sse) = *self.sse_manager.read().await { if let Some(ref sse) = *self.sse_manager.read().await {
sse.broadcast(ironclaw_common::AppEvent::ExtensionStatus { sse.broadcast(crate::channels::web::types::SseEvent::ExtensionStatus {
extension_name: name.to_string(), extension_name: name.to_string(),
status: status.to_string(), status: status.to_string(),
message: message.map(|m| m.to_string()), message: message.map(|m| m.to_string()),
@@ -1516,11 +1428,9 @@ impl ExtensionManager {
if kind_filter.is_none() || kind_filter == Some(ExtensionKind::ChannelRelay) { if kind_filter.is_none() || kind_filter == Some(ExtensionKind::ChannelRelay) {
let installed = self.installed_relay_extensions.read().await; let installed = self.installed_relay_extensions.read().await;
let active_names = self.active_channel_names.read().await; let active_names = self.active_channel_names.read().await;
let errors = self.activation_errors.read().await;
for name in installed.iter() { for name in installed.iter() {
let active = active_names.contains(name); let active = active_names.contains(name);
let authenticated = self.has_stored_team_id(name, user_id).await; let has_token = self.is_relay_channel(name, user_id).await;
let activation_error = errors.get(name).cloned();
let registry_entry = self let registry_entry = self
.registry .registry
.get_with_kind(name, Some(ExtensionKind::ChannelRelay)) .get_with_kind(name, Some(ExtensionKind::ChannelRelay))
@@ -1533,13 +1443,13 @@ impl ExtensionManager {
display_name, display_name,
description, description,
url: None, url: None,
authenticated, authenticated: has_token,
active, active,
tools: Vec::new(), tools: Vec::new(),
needs_setup: false, needs_setup: false,
has_auth: true, has_auth: true,
installed: true, installed: true,
activation_error, activation_error: None,
version: None, version: None,
}); });
} }
@@ -1716,22 +1626,7 @@ impl ExtensionManager {
self.persist_active_channels(user_id).await; self.persist_active_channels(user_id).await;
self.activation_errors.write().await.remove(name); self.activation_errors.write().await.remove(name);
// Remove stored team_id setting and clean up secrets // Remove stored stream token
if let Some(ref store) = self.store
&& let Err(e) = store
.delete_setting(user_id, &format!("relay:{}:team_id", name))
.await
{
tracing::warn!(error = %e, name, "Failed to delete relay team_id setting on removal");
}
if let Err(e) = self
.secrets
.delete(user_id, &format!("relay:{}:oauth_state", name))
.await
{
tracing::warn!(error = %e, name, "Failed to delete relay oauth_state secret on removal");
}
// Clean up legacy stream_token secret from pre-webhook installs
let _ = self let _ = self
.secrets .secrets
.delete(user_id, &format!("relay:{}:stream_token", name)) .delete(user_id, &format!("relay:{}:stream_token", name))
@@ -2789,7 +2684,7 @@ impl ExtensionManager {
user_id: user_id.to_string(), user_id: user_id.to_string(),
secrets: Arc::clone(&self.secrets), secrets: Arc::clone(&self.secrets),
sse_manager: self.sse_manager.read().await.clone(), sse_manager: self.sse_manager.read().await.clone(),
gateway_token: self.oauth_proxy_auth_token.clone(), gateway_token: self.gateway_token.clone(),
token_exchange_extra_params, token_exchange_extra_params,
client_id_secret_name: if server.oauth.is_none() { client_id_secret_name: if server.oauth.is_none() {
Some(server.client_id_secret_name()) Some(server.client_id_secret_name())
@@ -3284,7 +3179,7 @@ impl ExtensionManager {
// apps. Sending the desktop secret would cause a client_id/secret // apps. Sending the desktop secret would cause a client_id/secret
// mismatch because the container's GOOGLE_OAUTH_CLIENT_ID is the web // mismatch because the container's GOOGLE_OAUTH_CLIENT_ID is the web
// app, not the desktop app. // app, not the desktop app.
let proxy_client_secret = oauth_defaults::hosted_proxy_client_secret( let proxy_client_secret = hosted_proxy_client_secret(
&client_secret, &client_secret,
builtin.as_ref(), builtin.as_ref(),
oauth_defaults::exchange_proxy_url().is_some(), oauth_defaults::exchange_proxy_url().is_some(),
@@ -3306,7 +3201,7 @@ impl ExtensionManager {
user_id: user_id.to_string(), user_id: user_id.to_string(),
secrets: Arc::clone(&self.secrets), secrets: Arc::clone(&self.secrets),
sse_manager: self.sse_manager.read().await.clone(), sse_manager: self.sse_manager.read().await.clone(),
gateway_token: self.oauth_proxy_auth_token.clone(), gateway_token: self.gateway_token.clone(),
token_exchange_extra_params: std::collections::HashMap::new(), token_exchange_extra_params: std::collections::HashMap::new(),
client_id_secret_name: None, client_id_secret_name: None,
created_at: std::time::Instant::now(), created_at: std::time::Instant::now(),
@@ -3389,7 +3284,7 @@ impl ExtensionManager {
} }
.await; .await;
// Broadcast auth result event // Broadcast SSE event
let (success, message) = match result { let (success, message) = match result {
Ok(()) => (true, format!("{} authenticated successfully", display_name)), Ok(()) => (true, format!("{} authenticated successfully", display_name)),
Err(ref e) => ( Err(ref e) => (
@@ -3415,7 +3310,7 @@ impl ExtensionManager {
} }
if let Some(ref sse) = sse_manager { if let Some(ref sse) = sse_manager {
sse.broadcast(ironclaw_common::AppEvent::AuthCompleted { sse.broadcast(crate::channels::web::types::SseEvent::AuthCompleted {
extension_name: ext_name, extension_name: ext_name,
success, success,
message, message,
@@ -4286,75 +4181,26 @@ impl ExtensionManager {
/// ///
/// For Slack: initiates OAuth flow (redirect-based). /// For Slack: initiates OAuth flow (redirect-based).
/// For Telegram: accepts a bot token, registers it with channel-relay, /// For Telegram: accepts a bot token, registers it with channel-relay,
/// and stores the team_id setting. /// and stores the returned stream token.
async fn auth_channel_relay( async fn auth_channel_relay(
&self, &self,
name: &str, name: &str,
user_id: &str, user_id: &str,
) -> Result<AuthResult, ExtensionError> { ) -> Result<AuthResult, ExtensionError> {
tracing::trace!( // Check if already authenticated (stream token exists)
extension = %name, if self.is_relay_channel(name, user_id).await {
user_id = %user_id,
"auth_channel_relay: starting"
);
// Check if already authenticated by looking for a stored team_id.
// We intentionally skip the `installed_relay_extensions` in-memory set
// here because that set only tracks *installed* extensions — an extension
// can be installed (via registry) but not yet authenticated (no OAuth
// completed). Checking just `is_relay_channel()` would short-circuit
// to "authenticated" even when no team_id exists, preventing the OAuth
// flow from being offered to the user.
if self.has_stored_team_id(name, user_id).await {
tracing::trace!(
extension = %name,
"auth_channel_relay: already authenticated (team_id in store)"
);
return Ok(AuthResult::authenticated(name, ExtensionKind::ChannelRelay)); return Ok(AuthResult::authenticated(name, ExtensionKind::ChannelRelay));
} }
tracing::trace!(
extension = %name,
"auth_channel_relay: no stored team_id, initiating OAuth"
);
// Use relay config captured at startup // Use relay config captured at startup
let relay_config = self.relay_config().map_err(|e| { let relay_config = self.relay_config()?;
tracing::warn!(
extension = %name,
error = %e,
"auth_channel_relay: relay config not available — \
CHANNEL_RELAY_URL and CHANNEL_RELAY_API_KEY must be set"
);
e
})?;
// Allow per-extension URL override from settings
let effective_url = self
.effective_relay_url(name)
.await
.unwrap_or_else(|| relay_config.url.clone());
tracing::trace!(
extension = %name,
relay_url = %effective_url,
"auth_channel_relay: creating relay client for OAuth"
);
let client = crate::channels::relay::RelayClient::new( let client = crate::channels::relay::RelayClient::new(
effective_url.clone(), relay_config.url.clone(),
relay_config.api_key.clone(), relay_config.api_key.clone(),
relay_config.request_timeout_secs, relay_config.request_timeout_secs,
) )
.map_err(|e| { .map_err(|e| ExtensionError::Config(e.to_string()))?;
tracing::warn!(
extension = %name,
relay_url = %effective_url,
error = %e,
"auth_channel_relay: failed to create relay HTTP client"
);
ExtensionError::Config(e.to_string())
})?;
// Generate CSRF nonce — IronClaw validates this on the callback to ensure // Generate CSRF nonce — IronClaw validates this on the callback to ensure
// the OAuth completion is legitimate. Channel-relay embeds it in the signed // the OAuth completion is legitimate. Channel-relay embeds it in the signed
@@ -4366,44 +4212,18 @@ impl ExtensionManager {
self.secrets self.secrets
.create(user_id, CreateSecretParams::new(&state_key, &state_nonce)) .create(user_id, CreateSecretParams::new(&state_key, &state_nonce))
.await .await
.map_err(|e| { .map_err(|e| ExtensionError::AuthFailed(format!("Failed to store OAuth state: {e}")))?;
tracing::warn!(
extension = %name,
error = %e,
"auth_channel_relay: failed to store OAuth state nonce"
);
ExtensionError::AuthFailed(format!("Failed to store OAuth state: {e}"))
})?;
// Channel-relay derives all URLs from trusted instance_url in chat-api. // Channel-relay derives all URLs from trusted instance_url in chat-api.
// We only pass the nonce for CSRF validation on the callback. // We only pass the nonce for CSRF validation on the callback.
tracing::trace!(
extension = %name,
relay_url = %effective_url,
"auth_channel_relay: calling initiate_oauth on channel-relay"
);
match client.initiate_oauth(Some(&state_nonce)).await { match client.initiate_oauth(Some(&state_nonce)).await {
Ok(auth_url) => { Ok(auth_url) => Ok(AuthResult::awaiting_authorization(
tracing::info!( name,
extension = %name, ExtensionKind::ChannelRelay,
"auth_channel_relay: OAuth URL obtained, awaiting user authorization" auth_url,
); "redirect".to_string(),
Ok(AuthResult::awaiting_authorization( )),
name, Err(e) => Err(ExtensionError::AuthFailed(e.to_string())),
ExtensionKind::ChannelRelay,
auth_url,
"redirect".to_string(),
))
}
Err(e) => {
tracing::warn!(
extension = %name,
relay_url = %effective_url,
error = %e,
"auth_channel_relay: initiate_oauth call to channel-relay failed"
);
Err(ExtensionError::AuthFailed(e.to_string()))
}
} }
} }
@@ -4413,112 +4233,46 @@ impl ExtensionManager {
name: &str, name: &str,
user_id: &str, user_id: &str,
) -> Result<ActivateResult, ExtensionError> { ) -> Result<ActivateResult, ExtensionError> {
tracing::trace!( let token_key = format!("relay:{}:stream_token", name);
extension = %name,
user_id = %user_id,
"activate_channel_relay: starting"
);
let team_id_key = format!("relay:{}:team_id", name); let team_id_key = format!("relay:{}:team_id", name);
// Get team_id from settings (stored by the OAuth callback) // Check if we have a stream token
let team_id = if let Some(ref store) = self.store { // Verify auth: stream token must exist (even though we don't use it in this constructor path)
match store.get_setting(user_id, &team_id_key).await { let _stream_token = match self.secrets.get_decrypted(user_id, &token_key).await {
Ok(Some(v)) => { Ok(secret) => secret.expose().to_string(),
let id = v.as_str().map(|s| s.to_string()).unwrap_or_default(); Err(_) => {
tracing::trace!( return Err(ExtensionError::AuthRequired);
extension = %name,
team_id_empty = id.is_empty(),
"activate_channel_relay: loaded team_id from store"
);
id
}
Ok(None) => {
tracing::trace!(
extension = %name,
setting_key = %team_id_key,
"activate_channel_relay: no team_id in settings store"
);
String::new()
}
Err(e) => {
tracing::warn!(
extension = %name,
error = %e,
"activate_channel_relay: failed to read team_id from settings store"
);
String::new()
}
} }
};
// Get team_id from settings
let team_id = if let Some(ref store) = self.store {
store
.get_setting(user_id, &team_id_key)
.await
.ok()
.flatten()
.and_then(|v| v.as_str().map(|s| s.to_string()))
.unwrap_or_default()
} else { } else {
tracing::trace!(
extension = %name,
"activate_channel_relay: no settings store available"
);
String::new() String::new()
}; };
if team_id.is_empty() {
tracing::trace!(
extension = %name,
"activate_channel_relay: team_id is empty, returning AuthRequired"
);
return Err(ExtensionError::AuthRequired);
}
// Use relay config captured at startup // Use relay config captured at startup
let relay_config = self.relay_config().map_err(|e| { let relay_config = self.relay_config()?;
tracing::warn!(
extension = %name,
error = %e,
"activate_channel_relay: relay config not available"
);
e
})?;
// Allow per-extension URL override from settings
let effective_url = self
.effective_relay_url(name)
.await
.unwrap_or_else(|| relay_config.url.clone());
tracing::trace!(
extension = %name,
relay_url = %effective_url,
"activate_channel_relay: relay config loaded"
);
let instance_id = self.relay_instance_id(relay_config, user_id); let instance_id = self.relay_instance_id(relay_config, user_id);
let client = crate::channels::relay::RelayClient::new( let client = crate::channels::relay::RelayClient::new(
effective_url.clone(), relay_config.url.clone(),
relay_config.api_key.clone(), relay_config.api_key.clone(),
relay_config.request_timeout_secs, relay_config.request_timeout_secs,
) )
.map_err(|e| { .map_err(|e| ExtensionError::ActivationFailed(e.to_string()))?;
tracing::warn!(
extension = %name,
relay_url = %effective_url,
error = %e,
"activate_channel_relay: failed to create relay HTTP client"
);
ExtensionError::ActivationFailed(e.to_string())
})?;
// Fetch the per-instance signing secret from channel-relay. // Fetch the per-instance signing secret from channel-relay.
// This must succeed — there is no fallback. // This must succeed — there is no fallback.
tracing::trace!(
extension = %name,
relay_url = %effective_url,
"activate_channel_relay: fetching signing secret from channel-relay"
);
let signing_secret = client.get_signing_secret(&team_id).await.map_err(|e| { let signing_secret = client.get_signing_secret(&team_id).await.map_err(|e| {
tracing::warn!(
extension = %name,
relay_url = %effective_url,
error = %e,
"activate_channel_relay: failed to fetch signing secret from channel-relay"
);
ExtensionError::Config(format!("Failed to fetch relay signing secret: {e}")) ExtensionError::Config(format!("Failed to fetch relay signing secret: {e}"))
})?; })?;
@@ -4537,29 +4291,16 @@ impl ExtensionManager {
// Hot-add to channel manager // Hot-add to channel manager
let cm_guard = self.relay_channel_manager.read().await; let cm_guard = self.relay_channel_manager.read().await;
let channel_mgr = cm_guard.as_ref().ok_or_else(|| { let channel_mgr = cm_guard.as_ref().ok_or_else(|| {
tracing::warn!(
extension = %name,
"activate_channel_relay: channel manager not initialized"
);
ExtensionError::ActivationFailed("Channel manager not initialized".to_string()) ExtensionError::ActivationFailed("Channel manager not initialized".to_string())
})?; })?;
channel_mgr.hot_add(Box::new(channel)).await.map_err(|e| { channel_mgr
tracing::warn!( .hot_add(Box::new(channel))
extension = %name, .await
error = %e, .map_err(|e| ExtensionError::ActivationFailed(e.to_string()))?;
"activate_channel_relay: hot_add to channel manager failed"
);
ExtensionError::ActivationFailed(e.to_string())
})?;
if let Ok(mut cache) = self.relay_signing_secret_cache.lock() { if let Ok(mut cache) = self.relay_signing_secret_cache.lock() {
*cache = Some(signing_secret); *cache = Some(signing_secret);
} else {
tracing::warn!(
extension = %name,
"activate_channel_relay: failed to cache signing secret (mutex poisoned)"
);
} }
// Store the event sender so the web gateway's relay webhook endpoint can push events // Store the event sender so the web gateway's relay webhook endpoint can push events
@@ -4577,12 +4318,6 @@ impl ExtensionManager {
self.broadcast_extension_status(name, "active", Some(&status_msg)) self.broadcast_extension_status(name, "active", Some(&status_msg))
.await; .await;
tracing::info!(
extension = %name,
instance_id = %instance_id,
"activate_channel_relay: relay channel activated successfully"
);
Ok(ActivateResult { Ok(ActivateResult {
name: name.to_string(), name: name.to_string(),
kind: ExtensionKind::ChannelRelay, kind: ExtensionKind::ChannelRelay,
@@ -4632,11 +4367,11 @@ impl ExtensionManager {
return Ok(ExtensionKind::WasmChannel); return Ok(ExtensionKind::WasmChannel);
} }
// Check channel-relay extensions (installed in memory or has stored team_id) // Check channel-relay extensions (installed in memory or has stored token)
if self.installed_relay_extensions.read().await.contains(name) { if self.installed_relay_extensions.read().await.contains(name) {
return Ok(ExtensionKind::ChannelRelay); return Ok(ExtensionKind::ChannelRelay);
} }
// Also check if there's a stored team_id setting (persisted across restarts) // Also check if there's a stored stream token (persisted across restarts)
if self.is_relay_channel(name, user_id).await { if self.is_relay_channel(name, user_id).await {
return Ok(ExtensionKind::ChannelRelay); return Ok(ExtensionKind::ChannelRelay);
} }
@@ -4862,41 +4597,6 @@ impl ExtensionManager {
} }
Ok(ExtensionSetupSchema { secrets, fields }) Ok(ExtensionSetupSchema { secrets, fields })
} }
ExtensionKind::ChannelRelay => {
let relay_url_key = format!("extensions.{name}.relay_url");
let current_url = if let Some(ref store) = self.store {
match store.get_setting(&self.user_id, &relay_url_key).await {
Ok(value_opt) => value_opt
.and_then(|v| v.as_str().map(|s| s.to_string()))
.filter(|s| !s.is_empty()),
Err(e) => {
tracing::warn!(
extension = %name,
setting_key = %relay_url_key,
error = %e,
"get_setup_schema: failed to read relay_url from settings"
);
None
}
}
} else {
None
};
let env_url = self.relay_config.as_ref().map(|c| c.url.as_str());
Ok(ExtensionSetupSchema {
secrets: Vec::new(),
fields: vec![crate::channels::web::types::SetupFieldInfo {
name: "relay_url".to_string(),
prompt: format!(
"Channel-relay service URL (leave empty to use env default{})",
env_url.map(|u| format!(": {u}")).unwrap_or_default()
),
optional: true,
provided: current_url.is_some(),
input_type: crate::tools::wasm::ToolSetupFieldInputType::Text,
}],
})
}
_ => Ok(ExtensionSetupSchema { _ => Ok(ExtensionSetupSchema {
secrets: Vec::new(), secrets: Vec::new(),
fields: Vec::new(), fields: Vec::new(),
@@ -5300,15 +5000,9 @@ impl ExtensionManager {
(names, Vec::new()) (names, Vec::new())
} }
ExtensionKind::ChannelRelay => { ExtensionKind::ChannelRelay => {
let relay_fields = vec![crate::tools::wasm::ToolFieldSetupSchema { let mut names = std::collections::HashSet::new();
name: "relay_url".to_string(), names.insert(format!("relay:{}:stream_token", name));
prompt: "Channel-relay service URL override".to_string(), (names, Vec::new())
optional: true,
setting_path: Some(format!("extensions.{name}.relay_url")),
input_type: crate::tools::wasm::ToolSetupFieldInputType::Text,
restart_required: false,
}];
(std::collections::HashSet::new(), relay_fields)
} }
}; };
@@ -5400,28 +5094,13 @@ impl ExtensionManager {
))); )));
} }
let trimmed = field_value.trim(); let trimmed = field_value.trim();
let field_def = setup_field_defs.get(field_name);
// Empty value on an optional field with a setting_path: clear the
// stored override so the system reverts to the env/default value.
if trimmed.is_empty() { if trimmed.is_empty() {
if let Some(def) = field_def
&& def.optional
{
stored_fields.remove(field_name);
if let Some(setting_path) = &def.setting_path {
Self::validate_setup_setting_path(name, setting_path)?;
if let Some(store) = self.store.as_ref() {
let _ = store.delete_setting(&self.user_id, setting_path).await;
}
}
}
continue; continue;
} }
stored_fields.insert(field_name.clone(), trimmed.to_string()); stored_fields.insert(field_name.clone(), trimmed.to_string());
if let Some(field_def) = field_def { if let Some(field_def) = setup_field_defs.get(field_name) {
if field_def.restart_required { if field_def.restart_required {
restart_required = true; restart_required = true;
} }
@@ -5755,9 +5434,7 @@ impl ExtensionManager {
.map_err(|e| ExtensionError::NotInstalled(e.to_string()))?; .map_err(|e| ExtensionError::NotInstalled(e.to_string()))?;
server.token_secret_name() server.token_secret_name()
} }
ExtensionKind::ChannelRelay => { ExtensionKind::ChannelRelay => format!("relay:{}:stream_token", name),
return Err(ExtensionError::AuthRequired);
}
}; };
let mut secrets = std::collections::HashMap::new(); let mut secrets = std::collections::HashMap::new();
@@ -6025,7 +5702,7 @@ mod tests {
use crate::extensions::manager::{ use crate::extensions::manager::{
ChannelRuntimeState, FallbackDecision, TelegramBindingData, TelegramBindingResult, ChannelRuntimeState, FallbackDecision, TelegramBindingData, TelegramBindingResult,
TelegramOwnerBindingState, build_wasm_channel_runtime_config_updates, TelegramOwnerBindingState, build_wasm_channel_runtime_config_updates,
combine_install_errors, fallback_decision, infer_kind_from_url, combine_install_errors, fallback_decision, hosted_proxy_client_secret, infer_kind_from_url,
normalize_hosted_callback_url, send_telegram_text_message, normalize_hosted_callback_url, send_telegram_text_message,
telegram_message_matches_verification_code, telegram_message_matches_verification_code,
}; };
@@ -7366,7 +7043,7 @@ mod tests {
let dir = tempfile::tempdir().expect("temp dir"); let dir = tempfile::tempdir().expect("temp dir");
let mgr = make_test_manager(None, dir.path().to_path_buf()); let mgr = make_test_manager(None, dir.path().to_path_buf());
// No store configured, no team_id → not a relay channel // No token stored → not a relay channel
assert!(!mgr.is_relay_channel("slack-relay", "test").await); assert!(!mgr.is_relay_channel("slack-relay", "test").await);
} }
@@ -7385,39 +7062,6 @@ mod tests {
); );
} }
/// Regression: installed-but-not-authenticated relay must NOT short-circuit
/// `auth_channel_relay()` to "authenticated". Previously, `auth_channel_relay`
/// called `is_relay_channel()` which checked the in-memory
/// `installed_relay_extensions` set; that returned `true` even when no team_id
/// existed in the store, so the OAuth URL was never offered.
#[tokio::test]
async fn test_auth_channel_relay_installed_without_team_id_is_not_authenticated() {
let dir = tempfile::tempdir().expect("temp dir");
let mgr = make_test_manager(None, dir.path().to_path_buf());
// Mark as installed (simulates clicking Install in the UI)
mgr.installed_relay_extensions
.write()
.await
.insert("slack-relay".to_string());
// Without a stored team_id, auth should NOT return authenticated.
// It should fail because relay config is missing (no CHANNEL_RELAY_URL),
// but the key assertion is that it does NOT return Ok(authenticated).
let result = mgr.auth_channel_relay("slack-relay", "test").await;
match result {
Ok(ref auth_result) if auth_result.is_authenticated() => {
panic!(
"auth_channel_relay returned authenticated for installed-but-no-team-id relay; \
expected either an OAuth URL or a config error"
);
}
_ => {
// Config error (no relay URL) or awaiting_authorization — both are correct
}
}
}
#[tokio::test] #[tokio::test]
async fn test_remove_relay_shuts_down_via_relay_channel_manager() { async fn test_remove_relay_shuts_down_via_relay_channel_manager() {
// Regression: remove() only checked channel_runtime for shutdown, missing // Regression: remove() only checked channel_runtime for shutdown, missing
@@ -8218,13 +7862,19 @@ mod tests {
.await .await
.insert("test-relay".to_string()); .insert("test-relay".to_string());
// configure() with empty secrets should dispatch to // configure() should dispatch to activate_channel_relay(), not
// activate_channel_relay(), not activate_wasm_channel(). Relay auth // activate_wasm_channel(). Both will fail (no runtime configured),
// is OAuth-only so there are no manual secrets to pass. // but the error should be about relay config, not WASM channels.
let mut secrets = std::collections::HashMap::new();
secrets.insert(
"relay:test-relay:stream_token".to_string(),
"tok".to_string(),
);
let result = mgr let result = mgr
.configure( .configure(
"test-relay", "test-relay",
&std::collections::HashMap::new(), &secrets,
&std::collections::HashMap::new(), &std::collections::HashMap::new(),
"test", "test",
) )
@@ -8236,6 +7886,7 @@ mod tests {
); );
let result = result.unwrap(); let result = result.unwrap();
// Activation will fail (no relay config), but secrets should still be stored
assert!( assert!(
!result.activated, !result.activated,
"activation should fail without relay config" "activation should fail without relay config"
@@ -8245,6 +7896,15 @@ mod tests {
"error should not mention WASM — got: {}", "error should not mention WASM — got: {}",
result.message result.message
); );
// Verify the secret was stored
assert!(
mgr.secrets
.exists("test", "relay:test-relay:stream_token")
.await
.unwrap_or(false),
"configure should have stored the relay stream token"
);
} }
#[test] #[test]
fn test_validation_failed_is_distinct_error_variant() { fn test_validation_failed_is_distinct_error_variant() {
@@ -8310,8 +7970,7 @@ mod tests {
let builtin_ref = builtin.as_ref(); let builtin_ref = builtin.as_ref();
let secret = Some(builtin_ref.unwrap().client_secret.to_string()); let secret = Some(builtin_ref.unwrap().client_secret.to_string());
let result = let result = hosted_proxy_client_secret(&secret, builtin_ref, true);
crate::cli::oauth_defaults::hosted_proxy_client_secret(&secret, builtin_ref, true);
assert_eq!( assert_eq!(
result, None, result, None,
"built-in desktop secret must be suppressed when the exchange proxy is configured" "built-in desktop secret must be suppressed when the exchange proxy is configured"
@@ -8323,8 +7982,7 @@ mod tests {
let builtin = crate::cli::oauth_defaults::builtin_credentials("google_oauth_token"); let builtin = crate::cli::oauth_defaults::builtin_credentials("google_oauth_token");
let secret = Some("user-entered-custom-secret".to_string()); let secret = Some("user-entered-custom-secret".to_string());
let result = let result = hosted_proxy_client_secret(&secret, builtin.as_ref(), true);
crate::cli::oauth_defaults::hosted_proxy_client_secret(&secret, builtin.as_ref(), true);
assert_eq!( assert_eq!(
result, result,
Some("user-entered-custom-secret".to_string()), Some("user-entered-custom-secret".to_string()),
@@ -8338,8 +7996,7 @@ mod tests {
let builtin_ref = builtin.as_ref(); let builtin_ref = builtin.as_ref();
let secret = Some(builtin_ref.unwrap().client_secret.to_string()); let secret = Some(builtin_ref.unwrap().client_secret.to_string());
let result = let result = hosted_proxy_client_secret(&secret, builtin_ref, false);
crate::cli::oauth_defaults::hosted_proxy_client_secret(&secret, builtin_ref, false);
assert_eq!( assert_eq!(
result, secret, result, secret,
"built-in secret must be kept when the callback will exchange directly" "built-in secret must be kept when the callback will exchange directly"
@@ -8350,8 +8007,7 @@ mod tests {
fn test_proxy_client_secret_none_stays_none() { fn test_proxy_client_secret_none_stays_none() {
let builtin = crate::cli::oauth_defaults::builtin_credentials("google_oauth_token"); let builtin = crate::cli::oauth_defaults::builtin_credentials("google_oauth_token");
let result = let result = hosted_proxy_client_secret(&None, builtin.as_ref(), true);
crate::cli::oauth_defaults::hosted_proxy_client_secret(&None, builtin.as_ref(), true);
assert_eq!( assert_eq!(
result, None, result, None,
"None secret stays None even when the exchange proxy is configured" "None secret stays None even when the exchange proxy is configured"
@@ -8365,8 +8021,7 @@ mod tests {
assert!(builtin.is_none()); assert!(builtin.is_none());
let secret = Some("dcr-secret".to_string()); let secret = Some("dcr-secret".to_string());
let result = let result = hosted_proxy_client_secret(&secret, builtin.as_ref(), true);
crate::cli::oauth_defaults::hosted_proxy_client_secret(&secret, builtin.as_ref(), true);
assert_eq!( assert_eq!(
result, result,
Some("dcr-secret".to_string()), Some("dcr-secret".to_string()),
+3 -47
View File
@@ -1162,25 +1162,15 @@ impl Store {
pub async fn get_webhook_routine_by_path( pub 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.conn().await?; let conn = self.conn().await?;
let row = if let Some(uid) = user_id { let row = conn
conn.query_opt( .query_opt(
"SELECT * FROM routines WHERE enabled AND trigger_type = 'webhook' \
AND user_id = $2 \
AND (trigger_config->>'path' = $1 OR (trigger_config->>'path' IS NULL AND id::text = $1))",
&[&path, &uid],
)
.await?
} else {
conn.query_opt(
"SELECT * FROM routines WHERE enabled AND trigger_type = 'webhook' \ "SELECT * FROM routines WHERE enabled AND trigger_type = 'webhook' \
AND (trigger_config->>'path' = $1 OR (trigger_config->>'path' IS NULL AND id::text = $1))", AND (trigger_config->>'path' = $1 OR (trigger_config->>'path' IS NULL AND id::text = $1))",
&[&path], &[&path],
) )
.await? .await?;
};
row.as_ref().map(row_to_routine).transpose() row.as_ref().map(row_to_routine).transpose()
} }
@@ -1413,40 +1403,6 @@ impl Store {
Ok(counts) Ok(counts)
} }
/// Batch-load the most recent run status for multiple routines in a single query.
/// Uses a window function to pick only the latest run per routine.
#[cfg(feature = "postgres")]
pub async fn batch_get_last_run_status(
&self,
routine_ids: &[Uuid],
) -> Result<HashMap<Uuid, RunStatus>, DatabaseError> {
if routine_ids.is_empty() {
return Ok(HashMap::new());
}
let conn = self.conn().await?;
let rows = conn
.query(
"SELECT DISTINCT ON (routine_id) routine_id, status
FROM routine_runs
WHERE routine_id = ANY($1)
ORDER BY routine_id, started_at DESC",
&[&routine_ids],
)
.await?;
let mut statuses = HashMap::new();
for row in rows {
let id: Uuid = row.get("routine_id");
let status_str: String = row.get("status");
if let std::result::Result::Ok(status) = status_str.parse::<RunStatus>() {
statuses.insert(id, status);
}
}
Ok(statuses)
}
/// Link a routine run to a dispatched job. /// Link a routine run to a dispatched job.
pub async fn link_routine_run_to_job( pub async fn link_routine_run_to_job(
&self, &self,
-1
View File
@@ -69,7 +69,6 @@ pub mod service;
pub mod settings; pub mod settings;
pub mod setup; pub mod setup;
pub mod skills; pub mod skills;
pub mod tenant;
pub mod timezone; pub mod timezone;
pub mod tools; pub mod tools;
pub mod tracing_fmt; pub mod tracing_fmt;
-2
View File
@@ -575,7 +575,6 @@ fn extract_response_content(response: &AnthropicResponse) -> (Option<String>, Ve
id: id.clone(), id: id.clone(),
name: name.clone(), name: name.clone(),
arguments: input.clone(), arguments: input.clone(),
reasoning: None,
}); });
} }
} }
@@ -624,7 +623,6 @@ mod tests {
id: "call_1".to_string(), id: "call_1".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("Search for test"), ChatMessage::user("Search for test"),
-7
View File
@@ -522,7 +522,6 @@ fn extract_content_blocks(
id: tu.tool_use_id().to_string(), id: tu.tool_use_id().to_string(),
name: tu.name().to_string(), name: tu.name().to_string(),
arguments: document_to_json(tu.input()), arguments: document_to_json(tu.input()),
reasoning: None,
}); });
} }
// Ignore reasoning, citations, images, etc. // Ignore reasoning, citations, images, etc.
@@ -760,13 +759,11 @@ mod tests {
id: "call_1".to_string(), id: "call_1".to_string(),
name: "echo".to_string(), name: "echo".to_string(),
arguments: serde_json::json!({"text": "hi"}), arguments: serde_json::json!({"text": "hi"}),
reasoning: None,
}; };
let tc2 = crate::llm::provider::ToolCall { let tc2 = crate::llm::provider::ToolCall {
id: "call_2".to_string(), id: "call_2".to_string(),
name: "time".to_string(), name: "time".to_string(),
arguments: serde_json::json!({}), arguments: serde_json::json!({}),
reasoning: None,
}; };
let messages = vec![ let messages = vec![
@@ -805,7 +802,6 @@ mod tests {
id: "call_1".to_string(), id: "call_1".to_string(),
name: "search".to_string(), name: "search".to_string(),
arguments: serde_json::json!({"query": "test"}), arguments: serde_json::json!({"query": "test"}),
reasoning: None,
}; };
let messages = vec![ let messages = vec![
@@ -829,7 +825,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 messages = vec![ let messages = vec![
@@ -994,13 +989,11 @@ mod tests {
id: "call_abc".to_string(), id: "call_abc".to_string(),
name: "get_weather".to_string(), name: "get_weather".to_string(),
arguments: serde_json::json!({"city": "NYC"}), arguments: serde_json::json!({"city": "NYC"}),
reasoning: None,
}; };
let tc2 = crate::llm::provider::ToolCall { let tc2 = crate::llm::provider::ToolCall {
id: "call_def".to_string(), id: "call_def".to_string(),
name: "get_time".to_string(), name: "get_time".to_string(),
arguments: serde_json::json!({"tz": "EST"}), arguments: serde_json::json!({"tz": "EST"}),
reasoning: None,
}; };
let messages = vec![ let messages = vec![

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