mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-25 14:53:34 +00:00
Compare commits
101
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
98e6e2601c | ||
|
|
b5c9c9f5d6 | ||
|
|
5b95d22218 | ||
|
|
dd0a0e10ab | ||
|
|
1d5777824c | ||
|
|
adf4e25c8f | ||
|
|
9c63d189b7 | ||
|
|
ed4d92932a | ||
|
|
b3fbef5287 | ||
|
|
6b8a38e147 | ||
|
|
ab67f02886 | ||
|
|
f02345fd1f | ||
|
|
4c043bf057 | ||
|
|
0b4e7c761b | ||
|
|
cdc625566f | ||
|
|
bb24952622 | ||
|
|
ef37d705a1 | ||
|
|
b400c2a711 | ||
|
|
ea24d79ace | ||
|
|
8d632872fd | ||
|
|
4c5d961102 | ||
|
|
2b4e881a72 | ||
|
|
c0f33c37f7 | ||
|
|
5d714be354 | ||
|
|
3d43917cd0 | ||
|
|
f9dfb74800 | ||
|
|
cb01800f73 | ||
|
|
16aaea8d74 | ||
|
|
a19deb6812 | ||
|
|
2f80b7b0b8 | ||
|
|
2f47c611d4 | ||
|
|
1f8d901cf6 | ||
|
|
ad20a5ab4f | ||
|
|
e15c50ea2d | ||
|
|
d4e18020e2 | ||
|
|
a23d87fc00 | ||
|
|
c737fb0855 | ||
|
|
0145672f36 | ||
|
|
9fd5537a01 | ||
|
|
492d9d22c9 | ||
|
|
b8b88ab84e | ||
|
|
c98ec3fb18 | ||
|
|
189fa35e64 | ||
|
|
c5dce279e2 | ||
|
|
5a5ffe8d08 | ||
|
|
86d1143064 | ||
|
|
ab0ad948f3 | ||
|
|
c949521d8d | ||
|
|
0341fcc940 | ||
|
|
67a025e2fa | ||
|
|
7ad035773a | ||
|
|
424b470c59 | ||
|
|
6040e56984 | ||
|
|
9946149a25 | ||
|
|
b18376f9f1 | ||
|
|
cd617500a8 | ||
|
|
ae370d7e2b | ||
|
|
98418b3ef0 | ||
|
|
74b2b4129e | ||
|
|
bb57e36e6d | ||
|
|
0194275792 | ||
|
|
ddf64e8485 | ||
|
|
bd6977e6a8 | ||
|
|
d47b4b0346 | ||
|
|
91a241a3c7 | ||
|
|
d1d74d665a | ||
|
|
e077e1277d | ||
|
|
6fc8cc2f39 | ||
|
|
e031d8246b | ||
|
|
23263029f9 | ||
|
|
d5e08b95f9 | ||
|
|
3e866a9c0b | ||
|
|
e4d3200d80 | ||
|
|
7dc3c6d067 | ||
|
|
e1774e9ec0 | ||
|
|
e1d9827b21 | ||
|
|
e582166781 | ||
|
|
656d1f3e86 | ||
|
|
0e3aa4f806 | ||
|
|
da1db8f0e2 | ||
|
|
e313efc680 | ||
|
|
c91c63f810 | ||
|
|
54981c32f4 | ||
|
|
c779360730 | ||
|
|
6f687aabd2 | ||
|
|
386ed298c1 | ||
|
|
44d16732a7 | ||
|
|
28192a8f30 | ||
|
|
6861c57638 | ||
|
|
28e95378d2 | ||
|
|
9a8f8cebc3 | ||
|
|
8bcdf1608f | ||
|
|
8a5346f417 | ||
|
|
708755f34a | ||
|
|
ce46e75dec | ||
|
|
42b66b33b2 | ||
|
|
e7ca8bb435 | ||
|
|
2af70642de | ||
|
|
ff971d0f14 | ||
|
|
5c8ee16f81 | ||
|
|
065a7498d5 |
+4
-5
@@ -17,6 +17,8 @@ target/
|
||||
# Python
|
||||
__pycache__/
|
||||
*.pyc
|
||||
*.pyo
|
||||
*.pyd
|
||||
|
||||
# Benchmark results (local runs, not committed)
|
||||
bench-results/
|
||||
@@ -34,8 +36,5 @@ trace_*.json
|
||||
.claude/settings.local.json
|
||||
.worktrees/
|
||||
|
||||
# Python cache
|
||||
__pycache__/
|
||||
*.pyc
|
||||
*.pyo
|
||||
*.pyd
|
||||
# JetBrains IDE
|
||||
.idea
|
||||
|
||||
+132
@@ -7,6 +7,138 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
||||
|
||||
## [Unreleased]
|
||||
|
||||
## [0.22.0](https://github.com/nearai/ironclaw/compare/ironclaw-v0.21.0...ironclaw-v0.22.0) - 2026-03-25
|
||||
|
||||
### Added
|
||||
|
||||
- *(agent)* thread per-tool reasoning through provider, session, and all surfaces ([#1513](https://github.com/nearai/ironclaw/pull/1513))
|
||||
- *(cli)* show credential auth status in tool info ([#1572](https://github.com/nearai/ironclaw/pull/1572))
|
||||
- multi-tenant auth with per-user workspace isolation ([#1118](https://github.com/nearai/ironclaw/pull/1118))
|
||||
- *(cli)* add ironclaw models subcommands (list/status/set/set-provider) ([#1043](https://github.com/nearai/ironclaw/pull/1043))
|
||||
- *(workspace)* multi-scope workspace reads ([#1117](https://github.com/nearai/ironclaw/pull/1117))
|
||||
- *(ux)* complete UX overhaul — design system, onboarding, web polish ([#1277](https://github.com/nearai/ironclaw/pull/1277))
|
||||
- *(gemini_oauth)* full Gemini CLI OAuth integration with Cloud Code API ([#1356](https://github.com/nearai/ironclaw/pull/1356))
|
||||
- *(shell)* add Low/Medium/High risk levels for graduated command approval (closes #172) ([#368](https://github.com/nearai/ironclaw/pull/368))
|
||||
- *(agent)* queue and merge messages during active turns ([#1412](https://github.com/nearai/ironclaw/pull/1412))
|
||||
- *(cli)* add `ironclaw hooks list` subcommand ([#1023](https://github.com/nearai/ironclaw/pull/1023))
|
||||
- *(extensions)* support text setup fields in web configure modal ([#496](https://github.com/nearai/ironclaw/pull/496))
|
||||
- *(llm)* add GitHub Copilot as LLM provider ([#1512](https://github.com/nearai/ironclaw/pull/1512))
|
||||
- *(workspace)* layered memory with sensitivity-based privacy redirect ([#1112](https://github.com/nearai/ironclaw/pull/1112))
|
||||
- *(webhooks)* add public webhook trigger endpoint for routines ([#736](https://github.com/nearai/ironclaw/pull/736))
|
||||
- *(llm)* Add OpenAI Codex (ChatGPT subscription) as LLM provider ([#1461](https://github.com/nearai/ironclaw/pull/1461))
|
||||
- *(web)* add light theme with dark/light/system toggle ([#1457](https://github.com/nearai/ironclaw/pull/1457))
|
||||
- *(agent)* activate stuck_threshold for time-based stuck job detection ([#1234](https://github.com/nearai/ironclaw/pull/1234))
|
||||
- chat onboarding and routine advisor ([#927](https://github.com/nearai/ironclaw/pull/927))
|
||||
|
||||
### Fixed
|
||||
|
||||
- ensure LLM calls always end with user message (closes #763) ([#1259](https://github.com/nearai/ironclaw/pull/1259))
|
||||
- restore owner-scoped gateway startup ([#1625](https://github.com/nearai/ironclaw/pull/1625))
|
||||
- remove stale stream_token gate from channel-relay activation ([#1623](https://github.com/nearai/ironclaw/pull/1623))
|
||||
- *(agent)* case-insensitive channel match and user_id filter for event triggers ([#1211](https://github.com/nearai/ironclaw/pull/1211))
|
||||
- *(routines)* normalize status display across web and CLI ([#1469](https://github.com/nearai/ironclaw/pull/1469))
|
||||
- *(tunnel)* managed tunnels target wrong port and die from SIGPIPE ([#1093](https://github.com/nearai/ironclaw/pull/1093))
|
||||
- *(agent)* persist /model selection to .env, TOML, and DB ([#1581](https://github.com/nearai/ironclaw/pull/1581))
|
||||
- post-merge review sweep — 8 fixes across security, perf, and correctness ([#1550](https://github.com/nearai/ironclaw/pull/1550))
|
||||
- generate Mistral-compatible 9-char alphanumeric tool call IDs ([#1242](https://github.com/nearai/ironclaw/pull/1242))
|
||||
- *(mcp)* handle empty 202 notification acknowledgements ([#1539](https://github.com/nearai/ironclaw/pull/1539))
|
||||
- *(tests)* eliminate env mutex poison cascade ([#1558](https://github.com/nearai/ironclaw/pull/1558))
|
||||
- *(safety)* escape tool output XML content and remove misleading sanitized attr ([#1067](https://github.com/nearai/ironclaw/pull/1067))
|
||||
- *(oauth)* reject malformed ic2.* states in decode_hosted_oauth_state ([#1441](https://github.com/nearai/ironclaw/pull/1441)) ([#1454](https://github.com/nearai/ironclaw/pull/1454))
|
||||
- parameter coercion and validation for oneOf/anyOf/allOf schemas ([#1397](https://github.com/nearai/ironclaw/pull/1397))
|
||||
- persist startup-loaded MCP clients in ExtensionManager ([#1509](https://github.com/nearai/ironclaw/pull/1509))
|
||||
- *(deps)* patch rustls-webpki vulnerability (RUSTSEC-2026-0049)
|
||||
- *(routines)* add missing extension_manager field in trigger_manual EngineContext
|
||||
- *(ci)* serialize env-mutating OAuth wildcard tests with ENV_MUTEX ([#1280](https://github.com/nearai/ironclaw/pull/1280)) ([#1468](https://github.com/nearai/ironclaw/pull/1468))
|
||||
- *(setup)* remove redundant LLM config and API keys from bootstrap .env ([#1448](https://github.com/nearai/ironclaw/pull/1448))
|
||||
- resolve wasm broadcast merge conflicts with staging ([#395](https://github.com/nearai/ironclaw/pull/395)) ([#1460](https://github.com/nearai/ironclaw/pull/1460))
|
||||
- skip credential validation for Bedrock backend ([#1011](https://github.com/nearai/ironclaw/pull/1011))
|
||||
- register sandbox jobs in ContextManager for query tool visibility ([#1426](https://github.com/nearai/ironclaw/pull/1426))
|
||||
- prefer execution-local message routing metadata ([#1449](https://github.com/nearai/ironclaw/pull/1449))
|
||||
- *(security)* validate embedding base URLs to prevent SSRF ([#1221](https://github.com/nearai/ironclaw/pull/1221))
|
||||
- f32→f64 precision artifact in temperature causes provider 400 errors ([#1450](https://github.com/nearai/ironclaw/pull/1450))
|
||||
- *(routines)* surface errors when sandbox unavailable for full_job routines ([#769](https://github.com/nearai/ironclaw/pull/769))
|
||||
- restore libSQL vector search with dynamic dimensions ([#1393](https://github.com/nearai/ironclaw/pull/1393))
|
||||
- staging CI triage — consolidate retry parsing, fix flaky tests, add docs ([#1427](https://github.com/nearai/ironclaw/pull/1427))
|
||||
|
||||
### Other
|
||||
|
||||
- Merge branch 'main' into staging-promote/455f543b-23329172268
|
||||
- Merge pull request #1655 from nearai/codex/fix-staging-promotion-1451-version-bumps
|
||||
- Merge pull request #1499 from nearai/staging-promote/9603fefd-23364438978
|
||||
- Fix libsql prompt scope regressions ([#1651](https://github.com/nearai/ironclaw/pull/1651))
|
||||
- Normalize cron schedules on routine create ([#1648](https://github.com/nearai/ironclaw/pull/1648))
|
||||
- Fix MCP lifecycle trace user scope ([#1646](https://github.com/nearai/ironclaw/pull/1646))
|
||||
- Fix REPL single-message hang and cap CI test duration ([#1643](https://github.com/nearai/ironclaw/pull/1643))
|
||||
- extract AppEvent to crates/ironclaw_common ([#1615](https://github.com/nearai/ironclaw/pull/1615))
|
||||
- Fix hosted OAuth refresh via proxy ([#1602](https://github.com/nearai/ironclaw/pull/1602))
|
||||
- *(agent)* optimize approval thread resolution (UUID parsing + lock contention) ([#1592](https://github.com/nearai/ironclaw/pull/1592))
|
||||
- *(tools)* auto-compact WASM tool schemas, add descriptions, improve credential prompts ([#1525](https://github.com/nearai/ironclaw/pull/1525))
|
||||
- Default new lightweight routines to tools-enabled ([#1573](https://github.com/nearai/ironclaw/pull/1573))
|
||||
- Google OAuth URL broken when initiated from Telegram channel ([#1165](https://github.com/nearai/ironclaw/pull/1165))
|
||||
- add gitcgr code graph badge ([#1563](https://github.com/nearai/ironclaw/pull/1563))
|
||||
- Fix owner-scoped message routing fallbacks ([#1574](https://github.com/nearai/ironclaw/pull/1574))
|
||||
- *(tools)* remove unconditional params clone in shared execution (fix #893) ([#926](https://github.com/nearai/ironclaw/pull/926))
|
||||
- *(llm)* move transcription module into src/llm/ ([#1559](https://github.com/nearai/ironclaw/pull/1559))
|
||||
- *(agent)* avoid preview allocations for non-truncated strings (fix #894) ([#924](https://github.com/nearai/ironclaw/pull/924))
|
||||
- Expand AGENTS.md with coding agents guidance ([#1392](https://github.com/nearai/ironclaw/pull/1392))
|
||||
- Fix CI approval flows and stale fixtures ([#1478](https://github.com/nearai/ironclaw/pull/1478))
|
||||
- Use live owner tool scope for autonomous routines and jobs ([#1453](https://github.com/nearai/ironclaw/pull/1453))
|
||||
- use Arc in embedding cache to avoid clones on miss path ([#1438](https://github.com/nearai/ironclaw/pull/1438))
|
||||
- Add owner-scoped permissions for full-job routines ([#1440](https://github.com/nearai/ironclaw/pull/1440))
|
||||
|
||||
## [0.21.0](https://github.com/nearai/ironclaw/compare/v0.20.0...v0.21.0) - 2026-03-20
|
||||
|
||||
### Added
|
||||
|
||||
- structured fallback deliverables for failed/stuck jobs ([#236](https://github.com/nearai/ironclaw/pull/236))
|
||||
- LRU embedding cache for workspace search ([#1423](https://github.com/nearai/ironclaw/pull/1423))
|
||||
- receive relay events via webhook callbacks ([#1254](https://github.com/nearai/ironclaw/pull/1254))
|
||||
|
||||
### Fixed
|
||||
|
||||
- bump Feishu channel version for promotion
|
||||
- *(approval)* make "always" auto-approve work for credentialed HTTP requests ([#1257](https://github.com/nearai/ironclaw/pull/1257))
|
||||
- skip NEAR AI session check when backend is not nearai ([#1413](https://github.com/nearai/ironclaw/pull/1413))
|
||||
|
||||
### Other
|
||||
|
||||
- Make hosted OAuth and MCP auth generic ([#1375](https://github.com/nearai/ironclaw/pull/1375))
|
||||
|
||||
## [0.20.0](https://github.com/nearai/ironclaw/compare/v0.19.0...v0.20.0) - 2026-03-19
|
||||
|
||||
### Added
|
||||
|
||||
- *(self-repair)* wire stuck_threshold, store, and builder ([#712](https://github.com/nearai/ironclaw/pull/712))
|
||||
- *(testing)* add FaultInjector framework for StubLlm ([#1233](https://github.com/nearai/ironclaw/pull/1233))
|
||||
- *(gateway)* unified settings page with subtabs ([#1191](https://github.com/nearai/ironclaw/pull/1191))
|
||||
- upgrade MiniMax default model to M2.7 ([#1357](https://github.com/nearai/ironclaw/pull/1357))
|
||||
|
||||
### Fixed
|
||||
|
||||
- navigate telegram E2E tests to channels subtab ([#1408](https://github.com/nearai/ironclaw/pull/1408))
|
||||
- add missing `builder` field and update E2E extensions tab navigation ([#1400](https://github.com/nearai/ironclaw/pull/1400))
|
||||
- remove debug_assert guards that panic on valid error paths ([#1385](https://github.com/nearai/ironclaw/pull/1385))
|
||||
- address valid review comments from PR #1359 ([#1380](https://github.com/nearai/ironclaw/pull/1380))
|
||||
- full_job routine runs stay running until linked job completion ([#1374](https://github.com/nearai/ironclaw/pull/1374))
|
||||
- full_job routine concurrency tracks linked job lifetime ([#1372](https://github.com/nearai/ironclaw/pull/1372))
|
||||
- remove -x from coverage pytest to prevent suite-blocking failures ([#1360](https://github.com/nearai/ironclaw/pull/1360))
|
||||
- add debug_assert invariant guards to critical code paths ([#1312](https://github.com/nearai/ironclaw/pull/1312))
|
||||
- *(mcp)* retry after missing session id errors ([#1355](https://github.com/nearai/ironclaw/pull/1355))
|
||||
- *(telegram)* preserve polling after secret-blocked updates ([#1353](https://github.com/nearai/ironclaw/pull/1353))
|
||||
- *(llm)* cap retry-after delays ([#1351](https://github.com/nearai/ironclaw/pull/1351))
|
||||
- *(setup)* remove nonexistent webhook secret command hint ([#1349](https://github.com/nearai/ironclaw/pull/1349))
|
||||
- Rate limiter returns retry after None instead of a duration ([#1269](https://github.com/nearai/ironclaw/pull/1269))
|
||||
|
||||
### Other
|
||||
|
||||
- bump telegram channel version to 0.2.5 ([#1410](https://github.com/nearai/ironclaw/pull/1410))
|
||||
- *(ci)* enforce test requirement for state machine and resilience changes ([#1230](https://github.com/nearai/ironclaw/pull/1230)) ([#1304](https://github.com/nearai/ironclaw/pull/1304))
|
||||
- Fix duplicate LLM responses for matched event routines ([#1275](https://github.com/nearai/ironclaw/pull/1275))
|
||||
- add Japanese README ([#1306](https://github.com/nearai/ironclaw/pull/1306))
|
||||
- *(ci)* add coverage gates via codecov.yml ([#1228](https://github.com/nearai/ironclaw/pull/1228)) ([#1291](https://github.com/nearai/ironclaw/pull/1291))
|
||||
- Redesign routine create requests for LLMs ([#1147](https://github.com/nearai/ironclaw/pull/1147))
|
||||
|
||||
## [0.19.0](https://github.com/nearai/ironclaw/compare/v0.18.0...v0.19.0) - 2026-03-17
|
||||
|
||||
### Added
|
||||
|
||||
Generated
+5
-5
@@ -2323,7 +2323,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb"
|
||||
dependencies = [
|
||||
"libc",
|
||||
"windows-sys 0.52.0",
|
||||
"windows-sys 0.59.0",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -3390,7 +3390,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "ironclaw"
|
||||
version = "0.19.0"
|
||||
version = "0.22.0"
|
||||
dependencies = [
|
||||
"aes-gcm",
|
||||
"aho-corasick",
|
||||
@@ -3496,7 +3496,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "ironclaw_safety"
|
||||
version = "0.1.0"
|
||||
version = "0.2.0"
|
||||
dependencies = [
|
||||
"aho-corasick",
|
||||
"regex",
|
||||
@@ -5481,7 +5481,7 @@ dependencies = [
|
||||
"errno",
|
||||
"libc",
|
||||
"linux-raw-sys 0.12.1",
|
||||
"windows-sys 0.52.0",
|
||||
"windows-sys 0.59.0",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -6388,7 +6388,7 @@ dependencies = [
|
||||
"getrandom 0.4.2",
|
||||
"once_cell",
|
||||
"rustix 1.1.4",
|
||||
"windows-sys 0.52.0",
|
||||
"windows-sys 0.59.0",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
||||
+2
-2
@@ -20,7 +20,7 @@ exclude = [
|
||||
|
||||
[package]
|
||||
name = "ironclaw"
|
||||
version = "0.19.0"
|
||||
version = "0.22.0"
|
||||
edition = "2024"
|
||||
rust-version = "1.92"
|
||||
description = "Secure personal AI assistant that protects your data and expands its capabilities on the fly"
|
||||
@@ -104,7 +104,7 @@ cron = "0.13"
|
||||
ironclaw_common = { path = "crates/ironclaw_common", version = "0.1.0" }
|
||||
|
||||
# Safety/sanitization
|
||||
ironclaw_safety = { path = "crates/ironclaw_safety", version = "0.1.0" }
|
||||
ironclaw_safety = { path = "crates/ironclaw_safety", version = "0.2.0" }
|
||||
regex = "1"
|
||||
aho-corasick = "1"
|
||||
|
||||
|
||||
Generated
+1
-1
@@ -269,7 +269,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "whatsapp-channel"
|
||||
version = "0.1.0"
|
||||
version = "0.2.0"
|
||||
dependencies = [
|
||||
"serde",
|
||||
"serde_json",
|
||||
|
||||
@@ -8,7 +8,6 @@ authors = ["NEAR AI <[email protected]>"]
|
||||
license = "MIT OR Apache-2.0"
|
||||
homepage = "https://github.com/nearai/ironclaw"
|
||||
repository = "https://github.com/nearai/ironclaw"
|
||||
publish = false
|
||||
|
||||
[package.metadata.dist]
|
||||
dist = false
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "ironclaw_safety"
|
||||
version = "0.1.0"
|
||||
version = "0.2.0"
|
||||
edition = "2024"
|
||||
rust-version = "1.92"
|
||||
description = "Prompt injection defense, input validation, secret leak detection, and safety policy enforcement"
|
||||
@@ -8,7 +8,6 @@ authors = ["NEAR AI <[email protected]>"]
|
||||
license = "MIT OR Apache-2.0"
|
||||
homepage = "https://github.com/nearai/ironclaw"
|
||||
repository = "https://github.com/nearai/ironclaw"
|
||||
publish = false
|
||||
|
||||
[package.metadata.dist]
|
||||
dist = false
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "feishu",
|
||||
"display_name": "Feishu / Lark Channel",
|
||||
"kind": "channel",
|
||||
"version": "0.1.1",
|
||||
"version": "0.1.3",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Talk to your agent through a Feishu or Lark bot",
|
||||
"keywords": [
|
||||
@@ -19,8 +19,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"sha256": "5fca74022264d1c8e78a0853766276f7ffa3cf0d8065b2f51ca10985acad4714",
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.19.0/channel-feishu-0.1.1-wasm32-wasip2.tar.gz"
|
||||
"sha256": "a66ff0dafb67d2216d8161bb7e96e724a94acb0ab993b85d2782d30412f8fe94",
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/channel-feishu-0.1.3-wasm32-wasip2.tar.gz"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -18,8 +18,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.19.0/channel-telegram-0.2.4-wasm32-wasip2.tar.gz",
|
||||
"sha256": "a7cb300ec1c946831cfceaa95c1dc8f30d0f42a3924f3cb5de8098821573f4b8"
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.20.0/channel-telegram-0.2.5-wasm32-wasip2.tar.gz",
|
||||
"sha256": "1ef20a538f55b379e049356e4d6758006251846bc3365ceaa1c87eba8379a329"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "github",
|
||||
"display_name": "GitHub",
|
||||
"kind": "tool",
|
||||
"version": "0.2.1",
|
||||
"version": "0.2.2",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "GitHub integration for issues, PRs, repos, and code search",
|
||||
"keywords": [
|
||||
@@ -19,8 +19,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.19.0/tool-github-0.2.1-wasm32-wasip2.tar.gz",
|
||||
"sha256": "92c530b3ad172e2372d819744b5233f1d8f65768e26eb5a6c213eba3ce1de758"
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-github-0.2.2-wasm32-wasip2.tar.gz",
|
||||
"sha256": "70b55af593193d8fa495c0f702ea23284d83a624124f8a5f7564916ec5032c3f"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "gmail",
|
||||
"display_name": "Gmail",
|
||||
"kind": "tool",
|
||||
"version": "0.2.0",
|
||||
"version": "0.2.1",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Read, send, and manage Gmail messages and threads",
|
||||
"keywords": [
|
||||
@@ -18,8 +18,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/gmail-0.2.0-wasm32-wasip2.tar.gz",
|
||||
"sha256": "ee9574e02e92bc1d481f1310eb88afd99ee52bf6971074ab33bd76bf99b34b1d"
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-gmail-0.2.1-wasm32-wasip2.tar.gz",
|
||||
"sha256": "79025b40ee70ce1120acc4320bae50da095d7afb0ef67bd56d99b064b72ea779"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "google-calendar",
|
||||
"display_name": "Google Calendar",
|
||||
"kind": "tool",
|
||||
"version": "0.2.0",
|
||||
"version": "0.2.1",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Create, read, update, and delete Google Calendar events",
|
||||
"keywords": [
|
||||
@@ -18,8 +18,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-calendar-0.2.0-wasm32-wasip2.tar.gz",
|
||||
"sha256": "2fa47150ea222e787c122182ad6f4dfa2ffaf5fe490d05e8de887a76445f8d2d"
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-google-calendar-0.2.1-wasm32-wasip2.tar.gz",
|
||||
"sha256": "86bcc075010b08f5ab2f98f504cec1c6c9e0ca144857d185cbecf72a11f504bf"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "google-docs",
|
||||
"display_name": "Google Docs",
|
||||
"kind": "tool",
|
||||
"version": "0.2.0",
|
||||
"version": "0.2.1",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Create and edit Google Docs documents",
|
||||
"keywords": [
|
||||
@@ -18,8 +18,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-docs-0.2.0-wasm32-wasip2.tar.gz",
|
||||
"sha256": "40e134a1c1564f832ca861c3396895d4e33ec67b99313fc1f97baf8d971423a9"
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-google-docs-0.2.1-wasm32-wasip2.tar.gz",
|
||||
"sha256": "39d476029764949498a53a6a223f9952b5f4df151be7b8b19bf3fe4d401a57cd"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "google-drive",
|
||||
"display_name": "Google Drive",
|
||||
"kind": "tool",
|
||||
"version": "0.2.0",
|
||||
"version": "0.2.1",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Upload, download, search, and manage Google Drive files and folders",
|
||||
"keywords": [
|
||||
@@ -18,8 +18,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-drive-0.2.0-wasm32-wasip2.tar.gz",
|
||||
"sha256": "002a341a1d58125563a7c69561b26fbc2629b04ea723cade744102bdc0fbb71f"
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-google-drive-0.2.1-wasm32-wasip2.tar.gz",
|
||||
"sha256": "6e9a700fab93865c852af718666af64c5b534ad6a419fb4b736e07740188f494"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "google-sheets",
|
||||
"display_name": "Google Sheets",
|
||||
"kind": "tool",
|
||||
"version": "0.2.0",
|
||||
"version": "0.2.1",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Read and write Google Sheets spreadsheet data",
|
||||
"keywords": [
|
||||
@@ -18,8 +18,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-sheets-0.2.0-wasm32-wasip2.tar.gz",
|
||||
"sha256": "8aa2c9d52f033edea3a6c2311b0ec694ccb6d0a54ef07e94d72bf8be1ce8009a"
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-google-sheets-0.2.1-wasm32-wasip2.tar.gz",
|
||||
"sha256": "1f8c381799a916be83263cac9d497d52946e21b1b588592a3a42ca94a73b7051"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "google-slides",
|
||||
"display_name": "Google Slides",
|
||||
"kind": "tool",
|
||||
"version": "0.2.0",
|
||||
"version": "0.2.1",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Create and edit Google Slides presentations",
|
||||
"keywords": [
|
||||
@@ -17,8 +17,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-slides-0.2.0-wasm32-wasip2.tar.gz",
|
||||
"sha256": "e931a97d4fd0b0b938e464dc7c7f2be6ea6b4d1508f5ea3cd931d44db23f05f5"
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-google-slides-0.2.1-wasm32-wasip2.tar.gz",
|
||||
"sha256": "e2528be5da02f1b8cfc8ee9b0cdd849516c53d412e2f75c6175b3bded7f512cb"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "llm-context",
|
||||
"display_name": "LLM Context",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"version": "0.1.1",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Fetch pre-extracted web content from Brave Search for grounding LLM answers (RAG, fact-checking)",
|
||||
"keywords": [
|
||||
@@ -21,8 +21,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.19.0/tool-llm-context-0.1.0-wasm32-wasip2.tar.gz",
|
||||
"sha256": "d9ced2b1226b879135891e0ee40e072c7c95412e1b2462925a23853e1f92497e"
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-llm-context-0.1.1-wasm32-wasip2.tar.gz",
|
||||
"sha256": "9b19e2fd05dbbbe3c8bd55309a91db09124e8415eb0f767828b6e10b55771e63"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "slack-tool",
|
||||
"display_name": "Slack Tool",
|
||||
"kind": "tool",
|
||||
"version": "0.2.0",
|
||||
"version": "0.2.1",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Your agent uses Slack to post and read messages in your workspace",
|
||||
"keywords": [
|
||||
@@ -17,8 +17,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.19.0/tool-slack-0.2.0-wasm32-wasip2.tar.gz",
|
||||
"sha256": "ccfb0415d7a04f9497726c712d15216de36e86f498b849101283c017f5ab4efb"
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-slack-0.2.1-wasm32-wasip2.tar.gz",
|
||||
"sha256": "927519e5b7734beeb022d3b8bbd152e0e6b9f67c9452a8ad47809d3c4221a137"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "telegram-mtproto",
|
||||
"display_name": "Telegram Tool",
|
||||
"kind": "tool",
|
||||
"version": "0.2.0",
|
||||
"version": "0.2.1",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Your agent uses your Telegram account to read and send messages",
|
||||
"keywords": [
|
||||
@@ -18,8 +18,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.19.0/tool-telegram-0.2.0-wasm32-wasip2.tar.gz",
|
||||
"sha256": "c17065ca41fae5f2a7c43b36144686718cd310a2f22442313bb1aa82bbad0ae4"
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-telegram-0.2.1-wasm32-wasip2.tar.gz",
|
||||
"sha256": "1e57d0755fc9c7b3ec013d079f30168898b484a6919f9edd105f0cd80131c1cd"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "web-search",
|
||||
"display_name": "Web Search",
|
||||
"kind": "tool",
|
||||
"version": "0.2.1",
|
||||
"version": "0.2.2",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Search the web using Brave Search API",
|
||||
"keywords": [
|
||||
@@ -18,8 +18,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.19.0/tool-web-search-0.2.1-wasm32-wasip2.tar.gz",
|
||||
"sha256": "bad275ca4ec314adea5241d6b92c44ccf9cebcbca8e30ba2493cc0bcb4b57218"
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/ironclaw-v0.22.0/tool-web-search-0.2.2-wasm32-wasip2.tar.gz",
|
||||
"sha256": "47382b50c1ea7525b20d59dc02fab04e336d018665826c2f24710bdf460779ae"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -1,7 +1,2 @@
|
||||
[workspace]
|
||||
git_release_enable = false
|
||||
|
||||
[[package]]
|
||||
name = "ironclaw_safety"
|
||||
publish = false
|
||||
release = false
|
||||
|
||||
+136
-37
@@ -13,7 +13,7 @@ use futures::StreamExt;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::agent::context_monitor::ContextMonitor;
|
||||
use crate::agent::heartbeat::spawn_heartbeat;
|
||||
use crate::agent::heartbeat::{spawn_heartbeat, spawn_multi_user_heartbeat};
|
||||
use crate::agent::routine_engine::{RoutineEngine, spawn_cron_ticker};
|
||||
use crate::agent::self_repair::{DefaultSelfRepair, RepairResult, SelfRepair};
|
||||
use crate::agent::session::ThreadState;
|
||||
@@ -182,6 +182,8 @@ pub struct AgentDeps {
|
||||
/// Resolved LLM backend identifier (e.g., "nearai", "openai", "groq").
|
||||
/// Used by `/model` persistence to determine which env var to update.
|
||||
pub llm_backend: String,
|
||||
/// Per-tenant rate limiting registry (lazily creates rate state per user).
|
||||
pub tenant_rates: Arc<crate::tenant::TenantRateRegistry>,
|
||||
}
|
||||
|
||||
/// The main agent that coordinates all components.
|
||||
@@ -244,7 +246,10 @@ impl Agent {
|
||||
SchedulerDeps {
|
||||
tools: deps.tools.clone(),
|
||||
extension_manager: deps.extension_manager.clone(),
|
||||
store: deps.store.clone(),
|
||||
store: deps
|
||||
.store
|
||||
.as_ref()
|
||||
.map(|db| crate::tenant::AdminScope::new(Arc::clone(db))),
|
||||
hooks: deps.hooks.clone(),
|
||||
},
|
||||
);
|
||||
@@ -325,6 +330,50 @@ impl Agent {
|
||||
&self.deps.cost_guard
|
||||
}
|
||||
|
||||
/// Build a tenant-scoped execution context for the given user.
|
||||
///
|
||||
/// This is the standard entry point for per-user operations. The returned
|
||||
/// [`TenantCtx`] provides a [`TenantScope`] that auto-binds `user_id` on
|
||||
/// every database operation and a per-user rate limiter.
|
||||
pub(super) async fn tenant_ctx(&self, user_id: &str) -> crate::tenant::TenantCtx {
|
||||
let rate = self.deps.tenant_rates.get_or_create(user_id).await;
|
||||
|
||||
let store = self
|
||||
.deps
|
||||
.store
|
||||
.as_ref()
|
||||
.map(|db| crate::tenant::TenantScope::new(user_id, Arc::clone(db)));
|
||||
|
||||
// Reuse the owner workspace if user matches, otherwise create per-user.
|
||||
let workspace = match &self.deps.workspace {
|
||||
Some(ws) if ws.user_id() == user_id => Some(Arc::clone(ws)),
|
||||
_ => self
|
||||
.deps
|
||||
.store
|
||||
.as_ref()
|
||||
.map(|db| Arc::new(Workspace::new_with_db(user_id, Arc::clone(db)))),
|
||||
};
|
||||
|
||||
crate::tenant::TenantCtx::new(
|
||||
user_id,
|
||||
store,
|
||||
workspace,
|
||||
Arc::clone(&self.deps.cost_guard),
|
||||
rate,
|
||||
)
|
||||
}
|
||||
|
||||
/// Get an admin-scoped database accessor for cross-tenant operations.
|
||||
///
|
||||
/// Only for system-level components (heartbeat, routine engine, self-repair,
|
||||
/// scheduler). Handler code should use [`tenant_ctx()`](Self::tenant_ctx) instead.
|
||||
pub(super) fn admin_store(&self) -> Option<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>>> {
|
||||
self.deps.skill_registry.as_ref()
|
||||
}
|
||||
@@ -410,8 +459,8 @@ impl Agent {
|
||||
self.config.stuck_threshold,
|
||||
self.config.max_repair_attempts,
|
||||
);
|
||||
if let Some(ref store) = self.deps.store {
|
||||
self_repair = self_repair.with_store(Arc::clone(store));
|
||||
if let Some(admin) = self.admin_store() {
|
||||
self_repair = self_repair.with_store(admin);
|
||||
}
|
||||
if let Some(ref builder) = self.deps.builder {
|
||||
self_repair = self_repair.with_builder(Arc::clone(builder), Arc::clone(self.tools()));
|
||||
@@ -518,6 +567,7 @@ impl Agent {
|
||||
.with_interval(std::time::Duration::from_secs(hb_config.interval_secs));
|
||||
config.quiet_hours_start = hb_config.quiet_hours_start;
|
||||
config.quiet_hours_end = hb_config.quiet_hours_end;
|
||||
config.multi_tenant = hb_config.multi_tenant;
|
||||
config.timezone = hb_config
|
||||
.timezone
|
||||
.clone()
|
||||
@@ -547,30 +597,52 @@ impl Agent {
|
||||
.await;
|
||||
let notify_user = heartbeat_notify_user;
|
||||
let channels = self.channels.clone();
|
||||
let is_multi_tenant = hb_config.multi_tenant;
|
||||
tokio::spawn(async move {
|
||||
while let Some(response) = notify_rx.recv().await {
|
||||
// In multi-tenant mode, extract the owning user_id from
|
||||
// the response metadata so notifications reach the
|
||||
// correct user rather than the agent's owner.
|
||||
// This intentionally overrides the configured notify_target
|
||||
// because each user's heartbeat should notify that user.
|
||||
let effective_user = if is_multi_tenant {
|
||||
response
|
||||
.metadata
|
||||
.get("owner_id")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(String::from)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
// Try the configured channel first, fall back to
|
||||
// broadcasting on all channels.
|
||||
let targeted_ok = if let Some(ref channel) = notify_channel
|
||||
&& let Some(ref user) = notify_target
|
||||
{
|
||||
channels
|
||||
.broadcast(channel, user, response.clone())
|
||||
.await
|
||||
.is_ok()
|
||||
let targeted_ok = if let Some(ref channel) = notify_channel {
|
||||
let target = effective_user.as_deref().or(notify_target.as_deref());
|
||||
if let Some(user) = target {
|
||||
channels
|
||||
.broadcast(channel, user, response.clone())
|
||||
.await
|
||||
.is_ok()
|
||||
} else {
|
||||
false
|
||||
}
|
||||
} else {
|
||||
false
|
||||
};
|
||||
|
||||
if !targeted_ok && let Some(ref user) = notify_user {
|
||||
let results = channels.broadcast_all(user, response).await;
|
||||
for (ch, result) in results {
|
||||
if let Err(e) = result {
|
||||
tracing::warn!(
|
||||
"Failed to broadcast heartbeat to {}: {}",
|
||||
ch,
|
||||
e
|
||||
);
|
||||
if !targeted_ok {
|
||||
let fallback = effective_user.as_deref().or(notify_user.as_deref());
|
||||
if let Some(user) = fallback {
|
||||
let results = channels.broadcast_all(user, response).await;
|
||||
for (ch, result) in results {
|
||||
if let Err(e) = result {
|
||||
tracing::warn!(
|
||||
"Failed to broadcast heartbeat to {}: {}",
|
||||
ch,
|
||||
e
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -583,14 +655,29 @@ impl Agent {
|
||||
.map(|h| h.to_workspace_config())
|
||||
.unwrap_or_default();
|
||||
|
||||
Some(spawn_heartbeat(
|
||||
config,
|
||||
hygiene,
|
||||
workspace.clone(),
|
||||
self.cheap_llm().clone(),
|
||||
Some(notify_tx),
|
||||
self.store().map(Arc::clone),
|
||||
))
|
||||
if config.multi_tenant {
|
||||
if let Some(admin) = self.admin_store() {
|
||||
Some(spawn_multi_user_heartbeat(
|
||||
config,
|
||||
hygiene,
|
||||
self.cheap_llm().clone(),
|
||||
Some(notify_tx),
|
||||
admin,
|
||||
))
|
||||
} else {
|
||||
tracing::warn!("Multi-tenant heartbeat requires a database store");
|
||||
None
|
||||
}
|
||||
} else {
|
||||
Some(spawn_heartbeat(
|
||||
config,
|
||||
hygiene,
|
||||
workspace.clone(),
|
||||
self.cheap_llm().clone(),
|
||||
Some(notify_tx),
|
||||
self.admin_store(),
|
||||
))
|
||||
}
|
||||
} else {
|
||||
tracing::warn!("Heartbeat enabled but no workspace available");
|
||||
None
|
||||
@@ -612,7 +699,7 @@ impl Agent {
|
||||
|
||||
let engine = Arc::new(RoutineEngine::new(
|
||||
rt_config.clone(),
|
||||
Arc::clone(store),
|
||||
crate::tenant::AdminScope::new(Arc::clone(store)),
|
||||
self.llm().clone(),
|
||||
Arc::clone(workspace),
|
||||
notify_tx,
|
||||
@@ -1173,13 +1260,22 @@ impl Agent {
|
||||
}
|
||||
}
|
||||
|
||||
// Build per-tenant execution context once; threaded through all handlers.
|
||||
let tenant = self.tenant_ctx(&message.user_id).await;
|
||||
|
||||
let session_for_empty_exit = Arc::clone(&session);
|
||||
|
||||
// Process based on submission type
|
||||
let result = match submission {
|
||||
Submission::UserInput { content } => {
|
||||
let mut result = self
|
||||
.process_user_input(message, session.clone(), thread_id, &content)
|
||||
.process_user_input(
|
||||
message,
|
||||
tenant.clone(),
|
||||
session.clone(),
|
||||
thread_id,
|
||||
&content,
|
||||
)
|
||||
.await;
|
||||
|
||||
// Drain any messages queued during processing.
|
||||
@@ -1246,7 +1342,13 @@ impl Agent {
|
||||
let mut queued_msg = message.clone();
|
||||
queued_msg.attachments.clear();
|
||||
result = self
|
||||
.process_user_input(&queued_msg, session.clone(), thread_id, &next_content)
|
||||
.process_user_input(
|
||||
&queued_msg,
|
||||
tenant.clone(),
|
||||
session.clone(),
|
||||
thread_id,
|
||||
&next_content,
|
||||
)
|
||||
.await;
|
||||
|
||||
// If processing failed, re-queue the drained content so it
|
||||
@@ -1294,7 +1396,7 @@ impl Agent {
|
||||
};
|
||||
}
|
||||
// Authorization checks (including restart channel check) are enforced in handle_system_command
|
||||
self.handle_system_command(&command, &args, &message.channel)
|
||||
self.handle_system_command(&command, &args, &message.channel, &tenant)
|
||||
.await
|
||||
}
|
||||
Submission::Undo => self.process_undo(session, thread_id).await,
|
||||
@@ -1307,12 +1409,9 @@ impl Agent {
|
||||
Submission::Summarize => self.process_summarize(session, thread_id).await,
|
||||
Submission::Suggest => self.process_suggest(session, thread_id).await,
|
||||
Submission::JobStatus { job_id } => {
|
||||
self.process_job_status(&message.user_id, job_id.as_deref())
|
||||
.await
|
||||
}
|
||||
Submission::JobCancel { job_id } => {
|
||||
self.process_job_cancel(&message.user_id, &job_id).await
|
||||
self.process_job_status(&tenant, job_id.as_deref()).await
|
||||
}
|
||||
Submission::JobCancel { job_id } => self.process_job_cancel(&tenant, &job_id).await,
|
||||
Submission::Quit => return Ok(None),
|
||||
Submission::SwitchThread { thread_id: target } => {
|
||||
self.process_switch_thread(message, target).await
|
||||
|
||||
+125
-1
@@ -10,7 +10,7 @@ use std::borrow::Cow;
|
||||
|
||||
use crate::agent::session::PendingApproval;
|
||||
use crate::error::Error;
|
||||
use crate::llm::{ChatMessage, Reasoning, ReasoningContext, RespondResult};
|
||||
use crate::llm::{ChatMessage, FinishReason, Reasoning, ReasoningContext, RespondResult};
|
||||
|
||||
/// Signal from the delegate indicating how the loop should proceed.
|
||||
pub enum LoopSignal {
|
||||
@@ -134,6 +134,9 @@ pub async fn run_agentic_loop(
|
||||
config: &AgenticLoopConfig,
|
||||
) -> Result<LoopOutcome, Error> {
|
||||
let mut consecutive_tool_intent_nudges: u32 = 0;
|
||||
// Accumulates across all iterations (not reset by text responses) so
|
||||
// non-consecutive truncations still escalate to force_text.
|
||||
let mut truncation_count: u32 = 0;
|
||||
|
||||
for iteration in 1..=config.max_iterations {
|
||||
// Check for external signals (stop, cancellation, user messages)
|
||||
@@ -215,7 +218,35 @@ pub async fn run_agentic_loop(
|
||||
tool_calls,
|
||||
content,
|
||||
} => {
|
||||
// If the response was truncated, tool call parameters are likely
|
||||
// incomplete. Discard them and tell the LLM to try a different
|
||||
// approach rather than executing malformed tool calls.
|
||||
if output.finish_reason == FinishReason::Length {
|
||||
truncation_count += 1;
|
||||
let names: Vec<&str> = tool_calls.iter().map(|tc| tc.name.as_str()).collect();
|
||||
tracing::warn!(
|
||||
iteration,
|
||||
tools = ?names,
|
||||
truncation_count,
|
||||
"Discarding truncated tool calls (finish_reason=Length)"
|
||||
);
|
||||
if let Some(ref text) = content {
|
||||
reason_ctx.messages.push(ChatMessage::assistant(text));
|
||||
}
|
||||
reason_ctx
|
||||
.messages
|
||||
.push(ChatMessage::user(crate::llm::TRUNCATED_TOOL_CALL_NOTICE));
|
||||
// After repeated truncations, force text-only mode so the LLM
|
||||
// stops attempting tool calls it can't fit in the output budget.
|
||||
if truncation_count >= 3 {
|
||||
reason_ctx.force_text = true;
|
||||
}
|
||||
delegate.after_iteration(iteration).await;
|
||||
continue;
|
||||
}
|
||||
|
||||
consecutive_tool_intent_nudges = 0;
|
||||
truncation_count = 0;
|
||||
|
||||
if let Some(outcome) = delegate
|
||||
.execute_tool_calls(tool_calls, content, reason_ctx)
|
||||
@@ -271,6 +302,7 @@ mod tests {
|
||||
RespondOutput {
|
||||
result: RespondResult::Text(text.to_string()),
|
||||
usage: zero_usage(),
|
||||
finish_reason: FinishReason::Stop,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -281,6 +313,7 @@ mod tests {
|
||||
content: None,
|
||||
},
|
||||
usage: zero_usage(),
|
||||
finish_reason: FinishReason::ToolUse,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -622,4 +655,95 @@ mod tests {
|
||||
let result = truncate_for_preview("café", 4);
|
||||
assert_eq!(result, "caf...");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_truncated_tool_calls_discarded_on_length() {
|
||||
let truncated_tool_call = ToolCall {
|
||||
id: "call_1".to_string(),
|
||||
name: "memory_write".to_string(),
|
||||
arguments: serde_json::json!({}), // empty — truncated
|
||||
reasoning: None,
|
||||
};
|
||||
let truncated_output = RespondOutput {
|
||||
result: RespondResult::ToolCalls {
|
||||
tool_calls: vec![truncated_tool_call],
|
||||
content: Some("I'll write the report.".to_string()),
|
||||
},
|
||||
usage: zero_usage(),
|
||||
finish_reason: FinishReason::Length, // response was truncated
|
||||
};
|
||||
let delegate = MockDelegate::new(vec![truncated_output, text_output("Summarized it.")]);
|
||||
let reasoning = stub_reasoning();
|
||||
let mut ctx = ReasoningContext::new();
|
||||
let config = AgenticLoopConfig {
|
||||
max_iterations: 5,
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let outcome = run_agentic_loop(&delegate, &reasoning, &mut ctx, &config)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Tool calls should NOT have been executed
|
||||
assert_eq!(delegate.tool_exec_count.load(Ordering::SeqCst), 0);
|
||||
// The loop should have continued and returned the text response
|
||||
assert!(matches!(outcome, LoopOutcome::Response(ref t) if t == "Summarized it."));
|
||||
// A truncation notice should have been injected into context
|
||||
assert!(
|
||||
ctx.messages
|
||||
.iter()
|
||||
.any(|m| m.role == crate::llm::Role::User && m.content.contains("truncated")),
|
||||
"Should inject truncation notice into context"
|
||||
);
|
||||
// The partial assistant content should have been preserved
|
||||
assert!(
|
||||
ctx.messages
|
||||
.iter()
|
||||
.any(|m| m.role == crate::llm::Role::Assistant
|
||||
&& m.content.contains("write the report")),
|
||||
"Should preserve partial assistant content"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_repeated_truncations_force_text_mode() {
|
||||
let make_truncated = || RespondOutput {
|
||||
result: RespondResult::ToolCalls {
|
||||
tool_calls: vec![ToolCall {
|
||||
id: "call_1".to_string(),
|
||||
name: "memory_write".to_string(),
|
||||
arguments: serde_json::json!({}),
|
||||
reasoning: None,
|
||||
}],
|
||||
content: None,
|
||||
},
|
||||
usage: zero_usage(),
|
||||
finish_reason: FinishReason::Length,
|
||||
};
|
||||
// Three truncated responses, then a text response
|
||||
let delegate = MockDelegate::new(vec![
|
||||
make_truncated(),
|
||||
make_truncated(),
|
||||
make_truncated(),
|
||||
text_output("Gave up on tool calls."),
|
||||
]);
|
||||
let reasoning = stub_reasoning();
|
||||
let mut ctx = ReasoningContext::new();
|
||||
let config = AgenticLoopConfig {
|
||||
max_iterations: 5,
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let outcome = run_agentic_loop(&delegate, &reasoning, &mut ctx, &config)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(matches!(outcome, LoopOutcome::Response(_)));
|
||||
assert_eq!(delegate.tool_exec_count.load(Ordering::SeqCst), 0);
|
||||
// After 3 truncations, force_text should be set
|
||||
assert!(
|
||||
ctx.force_text,
|
||||
"Should escalate to force_text after repeated truncations"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
+88
-53
@@ -33,6 +33,7 @@ impl Agent {
|
||||
&self,
|
||||
intent: MessageIntent,
|
||||
message: &IncomingMessage,
|
||||
tenant: &crate::tenant::TenantCtx,
|
||||
) -> Result<SubmissionResult, Error> {
|
||||
// Send thinking status for non-trivial operations
|
||||
if let MessageIntent::CreateJob { .. } = &intent {
|
||||
@@ -52,24 +53,18 @@ impl Agent {
|
||||
description,
|
||||
category,
|
||||
} => {
|
||||
self.handle_create_job(&message.user_id, title, description, category)
|
||||
self.handle_create_job(tenant, title, description, category)
|
||||
.await?
|
||||
}
|
||||
MessageIntent::CheckJobStatus { job_id } => {
|
||||
self.handle_check_status(&message.user_id, job_id).await?
|
||||
}
|
||||
MessageIntent::CancelJob { job_id } => {
|
||||
self.handle_cancel_job(&message.user_id, &job_id).await?
|
||||
}
|
||||
MessageIntent::ListJobs { filter } => {
|
||||
self.handle_list_jobs(&message.user_id, filter).await?
|
||||
}
|
||||
MessageIntent::HelpJob { job_id } => {
|
||||
self.handle_help_job(&message.user_id, &job_id).await?
|
||||
self.handle_check_status(tenant, job_id).await?
|
||||
}
|
||||
MessageIntent::CancelJob { job_id } => self.handle_cancel_job(tenant, &job_id).await?,
|
||||
MessageIntent::ListJobs { filter } => self.handle_list_jobs(tenant, filter).await?,
|
||||
MessageIntent::HelpJob { job_id } => self.handle_help_job(tenant, &job_id).await?,
|
||||
MessageIntent::Command { command, args } => {
|
||||
match self
|
||||
.handle_command(&command, &args, &message.channel)
|
||||
.handle_command(&command, &args, &message.channel, tenant)
|
||||
.await?
|
||||
{
|
||||
Some(s) => s,
|
||||
@@ -83,14 +78,14 @@ impl Agent {
|
||||
|
||||
async fn handle_create_job(
|
||||
&self,
|
||||
user_id: &str,
|
||||
tenant: &crate::tenant::TenantCtx,
|
||||
title: String,
|
||||
description: String,
|
||||
category: Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
let job_id = self
|
||||
.scheduler
|
||||
.dispatch_job(user_id, &title, &description, None)
|
||||
.dispatch_job(tenant.user_id(), &title, &description, None)
|
||||
.await?;
|
||||
|
||||
// Set the dedicated category field (not stored in metadata)
|
||||
@@ -113,7 +108,7 @@ impl Agent {
|
||||
|
||||
async fn handle_check_status(
|
||||
&self,
|
||||
user_id: &str,
|
||||
tenant: &crate::tenant::TenantCtx,
|
||||
job_id: Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
match job_id {
|
||||
@@ -122,7 +117,8 @@ impl Agent {
|
||||
.map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?;
|
||||
|
||||
// Try DB first for persistent state, fall back to ContextManager.
|
||||
if let Some(store) = self.store()
|
||||
// TenantScope.get_job() auto-filters by ownership — no manual check needed.
|
||||
if let Some(store) = tenant.store()
|
||||
&& let Ok(Some(ctx)) = store.get_job(uuid).await
|
||||
{
|
||||
return Ok(format!(
|
||||
@@ -138,7 +134,7 @@ impl Agent {
|
||||
}
|
||||
|
||||
let ctx = self.context_manager.get_context(uuid).await?;
|
||||
if ctx.user_id != user_id {
|
||||
if ctx.user_id != tenant.user_id() {
|
||||
return Err(crate::error::JobError::NotFound { id: uuid }.into());
|
||||
}
|
||||
|
||||
@@ -155,7 +151,8 @@ impl Agent {
|
||||
}
|
||||
None => {
|
||||
// Show summary from DB for consistency with Jobs tab.
|
||||
if let Some(store) = self.store() {
|
||||
// TenantScope methods auto-scope to user — no user_id parameter needed.
|
||||
if let Some(store) = tenant.store() {
|
||||
let mut total = 0;
|
||||
let mut in_progress = 0;
|
||||
let mut completed = 0;
|
||||
@@ -183,7 +180,7 @@ impl Agent {
|
||||
}
|
||||
|
||||
// Fallback to ContextManager if no DB.
|
||||
let summary = self.context_manager.summary_for(user_id).await;
|
||||
let summary = self.context_manager.summary_for(tenant.user_id()).await;
|
||||
Ok(format!(
|
||||
"Jobs summary: Total: {} In Progress: {} Completed: {} Failed: {} Stuck: {}",
|
||||
summary.total,
|
||||
@@ -196,19 +193,24 @@ impl Agent {
|
||||
}
|
||||
}
|
||||
|
||||
async fn handle_cancel_job(&self, user_id: &str, job_id: &str) -> Result<String, Error> {
|
||||
async fn handle_cancel_job(
|
||||
&self,
|
||||
tenant: &crate::tenant::TenantCtx,
|
||||
job_id: &str,
|
||||
) -> Result<String, Error> {
|
||||
let uuid = Uuid::parse_str(job_id)
|
||||
.map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?;
|
||||
|
||||
let ctx = self.context_manager.get_context(uuid).await?;
|
||||
if ctx.user_id != user_id {
|
||||
if ctx.user_id != tenant.user_id() {
|
||||
return Err(crate::error::JobError::NotFound { id: uuid }.into());
|
||||
}
|
||||
|
||||
self.scheduler.stop(uuid).await?;
|
||||
|
||||
// Also update DB so the Jobs tab reflects cancellation immediately.
|
||||
if let Some(store) = self.store()
|
||||
// Use TenantScope — ownership already verified above.
|
||||
if let Some(store) = tenant.store()
|
||||
&& let Err(e) = store
|
||||
.update_job_status(uuid, JobState::Cancelled, Some("Cancelled by user"))
|
||||
.await
|
||||
@@ -221,11 +223,12 @@ impl Agent {
|
||||
|
||||
async fn handle_list_jobs(
|
||||
&self,
|
||||
user_id: &str,
|
||||
tenant: &crate::tenant::TenantCtx,
|
||||
_filter: Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
// List from DB for consistency with Jobs tab.
|
||||
if let Some(store) = self.store() {
|
||||
// TenantScope methods auto-scope to user.
|
||||
if let Some(store) = tenant.store() {
|
||||
let agent_jobs = match store.list_agent_jobs().await {
|
||||
Ok(jobs) => jobs,
|
||||
Err(e) => {
|
||||
@@ -256,7 +259,7 @@ impl Agent {
|
||||
}
|
||||
|
||||
// Fallback to ContextManager if no DB.
|
||||
let jobs = self.context_manager.all_jobs_for(user_id).await;
|
||||
let jobs = self.context_manager.all_jobs_for(tenant.user_id()).await;
|
||||
if jobs.is_empty() {
|
||||
return Ok("No jobs found.".to_string());
|
||||
}
|
||||
@@ -270,12 +273,16 @@ impl Agent {
|
||||
Ok(output)
|
||||
}
|
||||
|
||||
async fn handle_help_job(&self, user_id: &str, job_id: &str) -> Result<String, Error> {
|
||||
async fn handle_help_job(
|
||||
&self,
|
||||
tenant: &crate::tenant::TenantCtx,
|
||||
job_id: &str,
|
||||
) -> Result<String, Error> {
|
||||
let uuid = Uuid::parse_str(job_id)
|
||||
.map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?;
|
||||
|
||||
let ctx = self.context_manager.get_context(uuid).await?;
|
||||
if ctx.user_id != user_id {
|
||||
if ctx.user_id != tenant.user_id() {
|
||||
return Err(crate::error::JobError::NotFound { id: uuid }.into());
|
||||
}
|
||||
|
||||
@@ -308,11 +315,11 @@ impl Agent {
|
||||
/// Show job status inline — either all jobs (no id) or a specific job.
|
||||
pub(super) async fn process_job_status(
|
||||
&self,
|
||||
user_id: &str,
|
||||
tenant: &crate::tenant::TenantCtx,
|
||||
job_id: Option<&str>,
|
||||
) -> Result<SubmissionResult, Error> {
|
||||
match self
|
||||
.handle_check_status(user_id, job_id.map(|s| s.to_string()))
|
||||
.handle_check_status(tenant, job_id.map(|s| s.to_string()))
|
||||
.await
|
||||
{
|
||||
Ok(text) => Ok(SubmissionResult::response(text)),
|
||||
@@ -323,10 +330,10 @@ impl Agent {
|
||||
/// Cancel a job by ID.
|
||||
pub(super) async fn process_job_cancel(
|
||||
&self,
|
||||
user_id: &str,
|
||||
tenant: &crate::tenant::TenantCtx,
|
||||
job_id: &str,
|
||||
) -> Result<SubmissionResult, Error> {
|
||||
match self.handle_cancel_job(user_id, job_id).await {
|
||||
match self.handle_cancel_job(tenant, job_id).await {
|
||||
Ok(text) => Ok(SubmissionResult::response(text)),
|
||||
Err(e) => Ok(SubmissionResult::error(format!("Cancel error: {}", e))),
|
||||
}
|
||||
@@ -559,6 +566,7 @@ impl Agent {
|
||||
command: &str,
|
||||
args: &[String],
|
||||
channel: &str,
|
||||
tenant: &crate::tenant::TenantCtx,
|
||||
) -> Result<SubmissionResult, Error> {
|
||||
match command {
|
||||
"help" => Ok(SubmissionResult::response(concat!(
|
||||
@@ -752,19 +760,32 @@ impl Agent {
|
||||
}
|
||||
}
|
||||
|
||||
match self.llm().set_model(requested) {
|
||||
Ok(()) => {
|
||||
// Persist the model choice so it survives restarts.
|
||||
self.persist_selected_model(requested).await;
|
||||
Ok(SubmissionResult::response(format!(
|
||||
"Switched model to: {}",
|
||||
requested
|
||||
)))
|
||||
if self.config.multi_tenant {
|
||||
// Multi-tenant: only persist to per-user DB settings.
|
||||
// Do NOT call set_model() on the shared provider — that
|
||||
// would change the default for all users. The per-request
|
||||
// model_override in the dispatcher reads from the same
|
||||
// "selected_model" setting and applies it per-user.
|
||||
self.persist_selected_model(tenant, requested).await;
|
||||
Ok(SubmissionResult::response(format!(
|
||||
"Model preference set to: {} (per-user)",
|
||||
requested
|
||||
)))
|
||||
} else {
|
||||
match self.llm().set_model(requested) {
|
||||
Ok(()) => {
|
||||
// Persist the model choice so it survives restarts.
|
||||
self.persist_selected_model(tenant, requested).await;
|
||||
Ok(SubmissionResult::response(format!(
|
||||
"Switched model to: {}",
|
||||
requested
|
||||
)))
|
||||
}
|
||||
Err(e) => Ok(SubmissionResult::error(format!(
|
||||
"Failed to switch model: {}",
|
||||
e
|
||||
))),
|
||||
}
|
||||
Err(e) => Ok(SubmissionResult::error(format!(
|
||||
"Failed to switch model: {}",
|
||||
e
|
||||
))),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -906,10 +927,14 @@ impl Agent {
|
||||
command: &str,
|
||||
args: &[String],
|
||||
channel: &str,
|
||||
tenant: &crate::tenant::TenantCtx,
|
||||
) -> Result<Option<String>, Error> {
|
||||
// System commands are now handled directly via Submission::SystemCommand,
|
||||
// but the router may still send us unknown /commands.
|
||||
match self.handle_system_command(command, args, channel).await? {
|
||||
match self
|
||||
.handle_system_command(command, args, channel, tenant)
|
||||
.await?
|
||||
{
|
||||
SubmissionResult::Response { content } => Ok(Some(content)),
|
||||
SubmissionResult::Ok { message } => Ok(message),
|
||||
SubmissionResult::Error { message } => Ok(Some(format!("Error: {}", message))),
|
||||
@@ -921,23 +946,33 @@ impl Agent {
|
||||
///
|
||||
/// Best-effort: logs warnings on failure but does not propagate errors,
|
||||
/// since the in-memory model switch already succeeded.
|
||||
async fn persist_selected_model(&self, model: &str) {
|
||||
// 1. Persist to DB if available.
|
||||
if let Some(store) = self.store() {
|
||||
///
|
||||
/// In multi-tenant mode, only the per-user DB setting is written — global
|
||||
/// .env and TOML files are shared across users and must not be mutated.
|
||||
async fn persist_selected_model(&self, tenant: &crate::tenant::TenantCtx, model: &str) {
|
||||
// 1. Persist to DB if available (per-user scoped via TenantScope).
|
||||
if let Some(store) = tenant.store() {
|
||||
let value = serde_json::Value::String(model.to_string());
|
||||
if let Err(e) = store
|
||||
.set_setting(self.owner_id(), "selected_model", &value)
|
||||
.await
|
||||
{
|
||||
if let Err(e) = store.set_setting("selected_model", &value).await {
|
||||
tracing::warn!("Failed to persist model to DB: {}", e);
|
||||
} else {
|
||||
tracing::debug!("Persisted selected_model to DB: {}", model);
|
||||
tracing::debug!(
|
||||
user_id = tenant.user_id(),
|
||||
"Persisted selected_model to DB: {}",
|
||||
model
|
||||
);
|
||||
}
|
||||
} else {
|
||||
tracing::warn!("No database store available — model choice will not persist to DB");
|
||||
}
|
||||
|
||||
// 2. Update .env and TOML config file (sync I/O in spawn_blocking).
|
||||
// 2. In multi-tenant mode, skip .env/TOML writes — these are global
|
||||
// files shared by all users. The per-user DB setting is sufficient.
|
||||
if self.config.multi_tenant {
|
||||
return;
|
||||
}
|
||||
|
||||
// 3. Update .env and TOML config file (sync I/O in spawn_blocking).
|
||||
let model_owned = model.to_string();
|
||||
let backend = self.deps.llm_backend.clone();
|
||||
if let Err(e) = tokio::task::spawn_blocking(move || {
|
||||
|
||||
+236
-3
@@ -21,6 +21,9 @@ pub struct CostGuardConfig {
|
||||
pub max_cost_per_day_cents: Option<u64>,
|
||||
/// Maximum LLM calls per hour. None = unlimited.
|
||||
pub max_actions_per_hour: Option<u64>,
|
||||
/// Maximum spend per user per day in cents. None = unlimited.
|
||||
/// Applied independently per user alongside the global budget.
|
||||
pub max_cost_per_user_per_day_cents: Option<u64>,
|
||||
}
|
||||
|
||||
/// Error returned when a cost limit is exceeded.
|
||||
@@ -30,6 +33,12 @@ pub enum CostLimitExceeded {
|
||||
DailyBudget { spent_cents: u64, limit_cents: u64 },
|
||||
/// Hourly action rate limit reached.
|
||||
HourlyRate { actions: u64, limit: u64 },
|
||||
/// Per-user daily spending cap reached.
|
||||
UserDailyBudget {
|
||||
user_id: String,
|
||||
spent_cents: u64,
|
||||
limit_cents: u64,
|
||||
},
|
||||
}
|
||||
|
||||
impl std::fmt::Display for CostLimitExceeded {
|
||||
@@ -49,6 +58,17 @@ impl std::fmt::Display for CostLimitExceeded {
|
||||
"Hourly action limit exceeded: {} actions of {} allowed per hour",
|
||||
actions, limit
|
||||
),
|
||||
Self::UserDailyBudget {
|
||||
user_id,
|
||||
spent_cents,
|
||||
limit_cents,
|
||||
} => write!(
|
||||
f,
|
||||
"User '{}' daily cost limit exceeded: spent ${:.2} of ${:.2} allowed",
|
||||
user_id,
|
||||
*spent_cents as f64 / 100.0,
|
||||
*limit_cents as f64 / 100.0
|
||||
),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -78,6 +98,9 @@ pub struct CostGuard {
|
||||
|
||||
/// Per-model token usage since startup.
|
||||
model_tokens: Mutex<HashMap<String, ModelTokens>>,
|
||||
|
||||
/// Per-user daily cost tracking. Each entry resets independently at midnight UTC.
|
||||
per_user_daily_cost: Mutex<HashMap<String, DailyCost>>,
|
||||
}
|
||||
|
||||
struct DailyCost {
|
||||
@@ -97,6 +120,7 @@ impl CostGuard {
|
||||
action_window: Mutex::new(VecDeque::new()),
|
||||
budget_exceeded: AtomicBool::new(false),
|
||||
model_tokens: Mutex::new(HashMap::new()),
|
||||
per_user_daily_cost: Mutex::new(HashMap::new()),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -203,6 +227,11 @@ impl CostGuard {
|
||||
daily.reset_date = today;
|
||||
self.budget_exceeded.store(false, Ordering::Relaxed);
|
||||
tracing::info!("Cost guard: daily counter reset for {}", today);
|
||||
|
||||
// Prune per-user entries from previous days to prevent
|
||||
// unbounded HashMap growth in long-lived deployments.
|
||||
let mut per_user = self.per_user_daily_cost.lock().await;
|
||||
per_user.retain(|_, entry| entry.reset_date == today);
|
||||
}
|
||||
daily.total += cost;
|
||||
|
||||
@@ -248,6 +277,85 @@ impl CostGuard {
|
||||
cost
|
||||
}
|
||||
|
||||
/// Record an LLM call with per-user attribution.
|
||||
///
|
||||
/// Delegates to `record_llm_call` for global tracking, then additionally
|
||||
/// records the cost against the user's daily budget.
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub async fn record_llm_call_for_user(
|
||||
&self,
|
||||
user_id: &str,
|
||||
model: &str,
|
||||
input_tokens: u32,
|
||||
output_tokens: u32,
|
||||
cache_read_input_tokens: u32,
|
||||
cache_creation_input_tokens: u32,
|
||||
cache_read_discount: Decimal,
|
||||
cache_write_multiplier: Decimal,
|
||||
cost_per_token: Option<(Decimal, Decimal)>,
|
||||
) -> Decimal {
|
||||
let cost = self
|
||||
.record_llm_call(
|
||||
model,
|
||||
input_tokens,
|
||||
output_tokens,
|
||||
cache_read_input_tokens,
|
||||
cache_creation_input_tokens,
|
||||
cache_read_discount,
|
||||
cache_write_multiplier,
|
||||
cost_per_token,
|
||||
)
|
||||
.await;
|
||||
|
||||
// Track per-user daily cost
|
||||
{
|
||||
let today = chrono::Utc::now().date_naive();
|
||||
let mut per_user = self.per_user_daily_cost.lock().await;
|
||||
let entry = per_user
|
||||
.entry(user_id.to_string())
|
||||
.or_insert_with(|| DailyCost {
|
||||
total: Decimal::ZERO,
|
||||
reset_date: today,
|
||||
});
|
||||
if today != entry.reset_date {
|
||||
entry.total = Decimal::ZERO;
|
||||
entry.reset_date = today;
|
||||
}
|
||||
entry.total += cost;
|
||||
}
|
||||
|
||||
cost
|
||||
}
|
||||
|
||||
/// Check whether the next action is allowed for a specific user.
|
||||
///
|
||||
/// Checks the global limits first (via `check_allowed`), then additionally
|
||||
/// checks the per-user daily budget if configured.
|
||||
pub async fn check_allowed_for_user(&self, user_id: &str) -> Result<(), CostLimitExceeded> {
|
||||
// Check global limits first
|
||||
self.check_allowed().await?;
|
||||
|
||||
// Check per-user daily budget
|
||||
if let Some(limit_cents) = self.config.max_cost_per_user_per_day_cents {
|
||||
let today = chrono::Utc::now().date_naive();
|
||||
let per_user = self.per_user_daily_cost.lock().await;
|
||||
if let Some(entry) = per_user.get(user_id)
|
||||
&& entry.reset_date == today
|
||||
{
|
||||
let spent_cents = to_cents(entry.total);
|
||||
if spent_cents >= limit_cents {
|
||||
return Err(CostLimitExceeded::UserDailyBudget {
|
||||
user_id: user_id.to_string(),
|
||||
spent_cents,
|
||||
limit_cents,
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Current daily spend in USD (as Decimal).
|
||||
pub async fn daily_spend(&self) -> Decimal {
|
||||
let daily = self.daily_cost.lock().await;
|
||||
@@ -259,6 +367,16 @@ impl CostGuard {
|
||||
}
|
||||
}
|
||||
|
||||
/// Current daily spend for a specific user in USD (as Decimal).
|
||||
pub async fn daily_spend_for_user(&self, user_id: &str) -> Decimal {
|
||||
let today = chrono::Utc::now().date_naive();
|
||||
let per_user = self.per_user_daily_cost.lock().await;
|
||||
match per_user.get(user_id) {
|
||||
Some(entry) if entry.reset_date == today => entry.total,
|
||||
_ => Decimal::ZERO,
|
||||
}
|
||||
}
|
||||
|
||||
/// Number of actions in the current hourly window.
|
||||
pub async fn actions_this_hour(&self) -> u64 {
|
||||
let mut window = self.action_window.lock().await;
|
||||
@@ -314,7 +432,7 @@ mod tests {
|
||||
async fn test_daily_budget_enforcement() {
|
||||
let guard = CostGuard::new(CostGuardConfig {
|
||||
max_cost_per_day_cents: Some(1), // $0.01 limit
|
||||
max_actions_per_hour: None,
|
||||
..CostGuardConfig::default()
|
||||
});
|
||||
|
||||
// First call allowed
|
||||
@@ -350,8 +468,8 @@ mod tests {
|
||||
#[tokio::test]
|
||||
async fn test_hourly_rate_enforcement() {
|
||||
let guard = CostGuard::new(CostGuardConfig {
|
||||
max_cost_per_day_cents: None,
|
||||
max_actions_per_hour: Some(3),
|
||||
..CostGuardConfig::default()
|
||||
});
|
||||
|
||||
// First 3 actions allowed
|
||||
@@ -633,8 +751,8 @@ mod tests {
|
||||
// A fresh CostGuard with rate limits should not panic even if
|
||||
// checked_sub returns None (simulating short uptime).
|
||||
let guard = CostGuard::new(CostGuardConfig {
|
||||
max_cost_per_day_cents: None,
|
||||
max_actions_per_hour: Some(100),
|
||||
..CostGuardConfig::default()
|
||||
});
|
||||
|
||||
// These must not panic regardless of system uptime
|
||||
@@ -656,4 +774,119 @@ mod tests {
|
||||
let result = Instant::now().checked_sub(std::time::Duration::MAX);
|
||||
assert!(result.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_per_user_daily_budget_enforcement() {
|
||||
let guard = CostGuard::new(CostGuardConfig {
|
||||
max_cost_per_day_cents: None,
|
||||
max_actions_per_hour: None,
|
||||
max_cost_per_user_per_day_cents: Some(1), // $0.01 per user
|
||||
});
|
||||
|
||||
// Both users initially allowed
|
||||
assert!(guard.check_allowed_for_user("alice").await.is_ok());
|
||||
assert!(guard.check_allowed_for_user("bob").await.is_ok());
|
||||
|
||||
// Alice makes an expensive call
|
||||
guard
|
||||
.record_llm_call_for_user(
|
||||
"alice",
|
||||
"gpt-4o",
|
||||
10_000,
|
||||
10_000,
|
||||
0,
|
||||
0,
|
||||
Decimal::ONE,
|
||||
Decimal::ONE,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
// Alice should be blocked, Bob should still be allowed
|
||||
let result = guard.check_allowed_for_user("alice").await;
|
||||
assert!(result.is_err());
|
||||
match result.unwrap_err() {
|
||||
CostLimitExceeded::UserDailyBudget {
|
||||
user_id,
|
||||
limit_cents,
|
||||
..
|
||||
} => {
|
||||
assert_eq!(user_id, "alice");
|
||||
assert_eq!(limit_cents, 1);
|
||||
}
|
||||
other => panic!("Expected UserDailyBudget, got {:?}", other),
|
||||
}
|
||||
assert!(guard.check_allowed_for_user("bob").await.is_ok());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_per_user_daily_spend_tracking() {
|
||||
let guard = CostGuard::new(CostGuardConfig::default());
|
||||
|
||||
assert_eq!(guard.daily_spend_for_user("alice").await, Decimal::ZERO);
|
||||
assert_eq!(guard.daily_spend_for_user("bob").await, Decimal::ZERO);
|
||||
|
||||
let cost = guard
|
||||
.record_llm_call_for_user(
|
||||
"alice",
|
||||
"gpt-4o",
|
||||
1000,
|
||||
500,
|
||||
0,
|
||||
0,
|
||||
Decimal::ONE,
|
||||
Decimal::ONE,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
assert_eq!(guard.daily_spend_for_user("alice").await, cost);
|
||||
assert_eq!(guard.daily_spend_for_user("bob").await, Decimal::ZERO);
|
||||
// Global spend should also be tracked
|
||||
assert_eq!(guard.daily_spend().await, cost);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_per_user_budget_independent_of_global() {
|
||||
let guard = CostGuard::new(CostGuardConfig {
|
||||
max_cost_per_day_cents: Some(100_000), // $1000 global limit
|
||||
max_actions_per_hour: None,
|
||||
max_cost_per_user_per_day_cents: Some(1), // $0.01 per user
|
||||
});
|
||||
|
||||
// User hits their personal limit
|
||||
guard
|
||||
.record_llm_call_for_user(
|
||||
"alice",
|
||||
"gpt-4o",
|
||||
10_000,
|
||||
10_000,
|
||||
0,
|
||||
0,
|
||||
Decimal::ONE,
|
||||
Decimal::ONE,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
// Alice blocked by per-user limit, not global
|
||||
assert!(guard.check_allowed_for_user("alice").await.is_err());
|
||||
// Global limit is far from reached
|
||||
assert!(guard.check_allowed().await.is_ok());
|
||||
// Bob is unaffected
|
||||
assert!(guard.check_allowed_for_user("bob").await.is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_user_cost_limit_display() {
|
||||
let limit = CostLimitExceeded::UserDailyBudget {
|
||||
user_id: "alice".to_string(),
|
||||
spent_cents: 150,
|
||||
limit_cents: 100,
|
||||
};
|
||||
let msg = limit.to_string();
|
||||
assert!(msg.contains("alice"));
|
||||
assert!(msg.contains("$1.50"));
|
||||
assert!(msg.contains("$1.00"));
|
||||
}
|
||||
}
|
||||
|
||||
+70
-10
@@ -42,10 +42,20 @@ impl Agent {
|
||||
pub(super) async fn run_agentic_loop(
|
||||
&self,
|
||||
message: &IncomingMessage,
|
||||
tenant: crate::tenant::TenantCtx,
|
||||
session: Arc<Mutex<Session>>,
|
||||
thread_id: Uuid,
|
||||
initial_messages: Vec<ChatMessage>,
|
||||
) -> Result<AgenticLoopResult, Error> {
|
||||
if let Some(ext_mgr) = self.deps.extension_manager.as_ref()
|
||||
&& let Err(e) = ext_mgr.ensure_nearai_companion_active_if_ready().await
|
||||
{
|
||||
tracing::debug!(
|
||||
"Failed to auto-activate NEAR AI companion MCP before turn: {}",
|
||||
e
|
||||
);
|
||||
}
|
||||
|
||||
// Detect group chat from channel metadata (needed before loading system prompt)
|
||||
let is_group_chat = message
|
||||
.metadata
|
||||
@@ -63,7 +73,12 @@ impl Agent {
|
||||
);
|
||||
|
||||
let system_prompt = if let Some(ws) = self.workspace() {
|
||||
match ws
|
||||
let scoped_workspace = if ws.user_id() == message.user_id {
|
||||
Arc::clone(ws)
|
||||
} else {
|
||||
Arc::new(ws.scoped_to_user(&message.user_id))
|
||||
};
|
||||
match scoped_workspace
|
||||
.system_prompt_for_context_tz(is_group_chat, user_tz)
|
||||
.await
|
||||
{
|
||||
@@ -163,6 +178,7 @@ impl Agent {
|
||||
|
||||
let delegate = ChatDelegate {
|
||||
agent: self,
|
||||
tenant,
|
||||
session: session.clone(),
|
||||
thread_id,
|
||||
message,
|
||||
@@ -235,6 +251,7 @@ impl Agent {
|
||||
/// auth intercept, and cost tracking.
|
||||
struct ChatDelegate<'a> {
|
||||
agent: &'a Agent,
|
||||
tenant: crate::tenant::TenantCtx,
|
||||
session: Arc<Mutex<Session>>,
|
||||
thread_id: Uuid,
|
||||
message: &'a IncomingMessage,
|
||||
@@ -298,6 +315,8 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
|
||||
|
||||
// Update context for this iteration
|
||||
reason_ctx.available_tools = tool_defs;
|
||||
// Preserve force_text if already set (e.g. by truncation escalation).
|
||||
let force_text = force_text || reason_ctx.force_text;
|
||||
reason_ctx.system_prompt = Some(if force_text {
|
||||
self.cached_prompt_no_tools.clone()
|
||||
} else {
|
||||
@@ -331,8 +350,8 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
|
||||
reason_ctx: &mut ReasoningContext,
|
||||
iteration: usize,
|
||||
) -> Result<crate::llm::RespondOutput, Error> {
|
||||
// Enforce cost guardrails before the LLM call
|
||||
if let Err(limit) = self.agent.cost_guard().check_allowed().await {
|
||||
// Enforce cost guardrails before the LLM call (global + per-user)
|
||||
if let Err(limit) = self.tenant.check_cost_allowed().await {
|
||||
return Err(crate::error::LlmError::InvalidResponse {
|
||||
provider: "agent".to_string(),
|
||||
reason: limit.to_string(),
|
||||
@@ -340,6 +359,21 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
|
||||
.into());
|
||||
}
|
||||
|
||||
// Apply per-user model override from settings (first iteration only
|
||||
// to avoid repeated DB lookups within the same agentic loop).
|
||||
// Uses "selected_model" — the same key the /model command persists to
|
||||
// via SettingsStore (per-user scoped via TenantScope).
|
||||
if iteration == 0
|
||||
&& let Some(store) = self.tenant.store()
|
||||
&& let Ok(Some(value)) = store.get_setting("selected_model").await
|
||||
&& let Some(model) = value.as_str()
|
||||
{
|
||||
let model = model.trim();
|
||||
if !model.is_empty() {
|
||||
reason_ctx.model_override = Some(model.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
let output = match reasoning.respond_with_tools(reason_ctx).await {
|
||||
Ok(output) => output,
|
||||
Err(crate::error::LlmError::ContextLengthExceeded { used, limit }) => {
|
||||
@@ -374,13 +408,22 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
|
||||
Err(e) => return Err(e.into()),
|
||||
};
|
||||
|
||||
// Record cost and track token usage
|
||||
let model_name = self.agent.llm().active_model_name();
|
||||
// Record cost and track token usage (global + per-user).
|
||||
// When a model override is active, use the override name for attribution
|
||||
// and let CostGuard look up pricing via costs::model_cost() instead of
|
||||
// using the default provider's cost_per_token (which reflects the wrong model).
|
||||
let (model_name, cost_per_token) = if let Some(ref ovr) = reason_ctx.model_override {
|
||||
(ovr.clone(), None)
|
||||
} else {
|
||||
(
|
||||
self.agent.llm().active_model_name(),
|
||||
Some(self.agent.llm().cost_per_token()),
|
||||
)
|
||||
};
|
||||
let read_discount = self.agent.llm().cache_read_discount();
|
||||
let write_multiplier = self.agent.llm().cache_write_multiplier();
|
||||
let call_cost = self
|
||||
.agent
|
||||
.cost_guard()
|
||||
.tenant
|
||||
.record_llm_call(
|
||||
&model_name,
|
||||
output.usage.input_tokens,
|
||||
@@ -389,7 +432,7 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
|
||||
output.usage.cache_creation_input_tokens,
|
||||
read_discount,
|
||||
write_multiplier,
|
||||
Some(self.agent.llm().cost_per_token()),
|
||||
cost_per_token,
|
||||
)
|
||||
.await;
|
||||
tracing::debug!(
|
||||
@@ -1300,6 +1343,7 @@ mod tests {
|
||||
sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig,
|
||||
builder: None,
|
||||
llm_backend: "nearai".to_string(),
|
||||
tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)),
|
||||
};
|
||||
|
||||
Agent::new(
|
||||
@@ -1315,10 +1359,14 @@ mod tests {
|
||||
allow_local_tools: false,
|
||||
max_cost_per_day_cents: None,
|
||||
max_actions_per_hour: None,
|
||||
max_cost_per_user_per_day_cents: None,
|
||||
max_tool_iterations: 50,
|
||||
auto_approve_tools: false,
|
||||
default_timezone: "UTC".to_string(),
|
||||
max_tokens_per_job: 0,
|
||||
multi_tenant: false,
|
||||
max_llm_concurrent_per_user: None,
|
||||
max_jobs_concurrent_per_user: None,
|
||||
},
|
||||
deps,
|
||||
Arc::new(ChannelManager::new()),
|
||||
@@ -2176,6 +2224,7 @@ mod tests {
|
||||
sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig,
|
||||
builder: None,
|
||||
llm_backend: "nearai".to_string(),
|
||||
tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)),
|
||||
};
|
||||
|
||||
Agent::new(
|
||||
@@ -2191,10 +2240,14 @@ mod tests {
|
||||
allow_local_tools: false,
|
||||
max_cost_per_day_cents: None,
|
||||
max_actions_per_hour: None,
|
||||
max_cost_per_user_per_day_cents: None,
|
||||
max_tool_iterations,
|
||||
auto_approve_tools: true,
|
||||
default_timezone: "UTC".to_string(),
|
||||
max_tokens_per_job: 0,
|
||||
multi_tenant: false,
|
||||
max_llm_concurrent_per_user: None,
|
||||
max_jobs_concurrent_per_user: None,
|
||||
},
|
||||
deps,
|
||||
Arc::new(ChannelManager::new()),
|
||||
@@ -2229,13 +2282,14 @@ mod tests {
|
||||
|
||||
let message = IncomingMessage::new("test", "test-user", "do something");
|
||||
let initial_messages = vec![ChatMessage::user("do something")];
|
||||
let tenant = agent.tenant_ctx("test-user").await;
|
||||
|
||||
// The dispatcher must terminate within 5 seconds. If there is an
|
||||
// infinite loop bug (e.g., index not advancing on tool failure), the
|
||||
// timeout will fire and the test will fail.
|
||||
let result = tokio::time::timeout(
|
||||
Duration::from_secs(5),
|
||||
agent.run_agentic_loop(&message, session, thread_id, initial_messages),
|
||||
agent.run_agentic_loop(&message, tenant, session, thread_id, initial_messages),
|
||||
)
|
||||
.await;
|
||||
|
||||
@@ -2297,6 +2351,7 @@ mod tests {
|
||||
sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig,
|
||||
builder: None,
|
||||
llm_backend: "nearai".to_string(),
|
||||
tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)),
|
||||
};
|
||||
|
||||
Agent::new(
|
||||
@@ -2312,10 +2367,14 @@ mod tests {
|
||||
allow_local_tools: false,
|
||||
max_cost_per_day_cents: None,
|
||||
max_actions_per_hour: None,
|
||||
max_cost_per_user_per_day_cents: None,
|
||||
max_tool_iterations: max_iter,
|
||||
auto_approve_tools: true,
|
||||
default_timezone: "UTC".to_string(),
|
||||
max_tokens_per_job: 0,
|
||||
multi_tenant: false,
|
||||
max_llm_concurrent_per_user: None,
|
||||
max_jobs_concurrent_per_user: None,
|
||||
},
|
||||
deps,
|
||||
Arc::new(ChannelManager::new()),
|
||||
@@ -2335,13 +2394,14 @@ mod tests {
|
||||
|
||||
let message = IncomingMessage::new("test", "test-user", "keep calling tools");
|
||||
let initial_messages = vec![ChatMessage::user("keep calling tools")];
|
||||
let tenant = agent.tenant_ctx("test-user").await;
|
||||
|
||||
// Even with an LLM that always wants to call tools, the dispatcher
|
||||
// must terminate within the timeout thanks to force_text at
|
||||
// max_tool_iterations.
|
||||
let result = tokio::time::timeout(
|
||||
Duration::from_secs(5),
|
||||
agent.run_agentic_loop(&message, session, thread_id, initial_messages),
|
||||
agent.run_agentic_loop(&message, tenant, session, thread_id, initial_messages),
|
||||
)
|
||||
.await;
|
||||
|
||||
|
||||
+184
-7
@@ -31,8 +31,8 @@ use chrono_tz::Tz;
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
use crate::channels::OutgoingResponse;
|
||||
use crate::db::Database;
|
||||
use crate::llm::{ChatMessage, CompletionRequest, LlmProvider, Reasoning};
|
||||
use crate::tenant::AdminScope;
|
||||
use crate::workspace::Workspace;
|
||||
use crate::workspace::hygiene::HygieneConfig;
|
||||
|
||||
@@ -57,6 +57,9 @@ pub struct HeartbeatConfig {
|
||||
pub quiet_hours_end: Option<u32>,
|
||||
/// Timezone for fire_at and quiet hours evaluation (IANA name).
|
||||
pub timezone: Option<String>,
|
||||
/// When true, cycle through all users with routines instead of
|
||||
/// running heartbeat for a single user. Requires a database store.
|
||||
pub multi_tenant: bool,
|
||||
}
|
||||
|
||||
impl Default for HeartbeatConfig {
|
||||
@@ -71,6 +74,7 @@ impl Default for HeartbeatConfig {
|
||||
quiet_hours_start: None,
|
||||
quiet_hours_end: None,
|
||||
timezone: None,
|
||||
multi_tenant: false,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -178,7 +182,7 @@ pub struct HeartbeatRunner {
|
||||
workspace: Arc<Workspace>,
|
||||
llm: Arc<dyn LlmProvider>,
|
||||
response_tx: Option<mpsc::Sender<OutgoingResponse>>,
|
||||
store: Option<Arc<dyn Database>>,
|
||||
store: Option<AdminScope>,
|
||||
consecutive_failures: u32,
|
||||
}
|
||||
|
||||
@@ -207,8 +211,8 @@ impl HeartbeatRunner {
|
||||
self
|
||||
}
|
||||
|
||||
/// Set the database store for persistent heartbeat conversations.
|
||||
pub fn with_store(mut self, store: Arc<dyn Database>) -> Self {
|
||||
/// Set the admin-scoped database store for persistent heartbeat conversations.
|
||||
pub fn with_store(mut self, store: AdminScope) -> Self {
|
||||
self.store = Some(store);
|
||||
self
|
||||
}
|
||||
@@ -396,7 +400,7 @@ impl HeartbeatRunner {
|
||||
}
|
||||
|
||||
/// Send a notification about heartbeat findings.
|
||||
async fn send_notification(&self, message: &str) {
|
||||
pub(crate) async fn send_notification(&self, message: &str) {
|
||||
let Some(ref tx) = self.response_tx else {
|
||||
tracing::debug!("No response channel configured for heartbeat notifications");
|
||||
return;
|
||||
@@ -493,7 +497,7 @@ pub fn spawn_heartbeat(
|
||||
workspace: Arc<Workspace>,
|
||||
llm: Arc<dyn LlmProvider>,
|
||||
response_tx: Option<mpsc::Sender<OutgoingResponse>>,
|
||||
store: Option<Arc<dyn Database>>,
|
||||
store: Option<AdminScope>,
|
||||
) -> tokio::task::JoinHandle<()> {
|
||||
let mut runner = HeartbeatRunner::new(config, hygiene_config, workspace, llm);
|
||||
if let Some(tx) = response_tx {
|
||||
@@ -508,6 +512,179 @@ pub fn spawn_heartbeat(
|
||||
})
|
||||
}
|
||||
|
||||
/// Spawn a multi-user heartbeat runner that cycles through all users that
|
||||
/// own routines (enabled or not). Each tick, it queries the DB for distinct
|
||||
/// user_ids, creates a per-user workspace, and runs a heartbeat check for
|
||||
/// each user concurrently. Per-user failure counts are tracked independently.
|
||||
pub fn spawn_multi_user_heartbeat(
|
||||
config: HeartbeatConfig,
|
||||
hygiene_config: HygieneConfig,
|
||||
llm: Arc<dyn LlmProvider>,
|
||||
response_tx: Option<mpsc::Sender<OutgoingResponse>>,
|
||||
store: AdminScope,
|
||||
) -> tokio::task::JoinHandle<()> {
|
||||
tokio::spawn(async move {
|
||||
if !config.enabled {
|
||||
tracing::info!("Multi-user heartbeat is disabled");
|
||||
return;
|
||||
}
|
||||
|
||||
let mut tick_interval = if config.fire_at.is_none() {
|
||||
let mut iv = tokio::time::interval(config.interval);
|
||||
iv.tick().await; // skip immediate tick
|
||||
Some(iv)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
// Track consecutive failures per user so we can disable heartbeat
|
||||
// for persistently-failing users (same semantics as single-user mode).
|
||||
let mut user_failures: std::collections::HashMap<String, u32> =
|
||||
std::collections::HashMap::new();
|
||||
|
||||
tracing::info!("Starting multi-user heartbeat loop");
|
||||
|
||||
loop {
|
||||
if let Some(fire_at) = config.fire_at {
|
||||
let sleep_dur = duration_until_next_fire(fire_at, config.resolved_tz());
|
||||
tokio::time::sleep(sleep_dur).await;
|
||||
} else if let Some(ref mut iv) = tick_interval {
|
||||
iv.tick().await;
|
||||
}
|
||||
|
||||
if config.is_quiet_hours() {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Get distinct user_ids from routines
|
||||
let user_ids = match store.list_all_routines().await {
|
||||
Ok(routines) => {
|
||||
let mut ids: Vec<String> = routines
|
||||
.iter()
|
||||
.map(|r| r.user_id.clone())
|
||||
.collect::<std::collections::HashSet<_>>()
|
||||
.into_iter()
|
||||
.collect();
|
||||
ids.sort();
|
||||
ids
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!("Multi-user heartbeat: failed to list routines: {}", e);
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
// Run user heartbeats concurrently so one slow LLM call doesn't
|
||||
// block others. Cap concurrency to avoid flooding the LLM provider.
|
||||
const MAX_CONCURRENT_HEARTBEATS: usize = 8;
|
||||
let mut join_set = tokio::task::JoinSet::new();
|
||||
|
||||
for user_id in &user_ids {
|
||||
// Skip users that have exceeded max_failures
|
||||
let failures = user_failures.get(user_id).copied().unwrap_or(0);
|
||||
if failures >= config.max_failures {
|
||||
continue;
|
||||
}
|
||||
|
||||
let workspace = Arc::new(Workspace::new_with_db(user_id, Arc::clone(store.db())));
|
||||
|
||||
// Run memory hygiene per user (same as single-user heartbeat).
|
||||
let hygiene_ws = Arc::clone(&workspace);
|
||||
let hygiene_cfg = hygiene_config.clone();
|
||||
let hygiene_user = user_id.clone();
|
||||
tokio::spawn(async move {
|
||||
let report =
|
||||
crate::workspace::hygiene::run_if_due(&hygiene_ws, &hygiene_cfg).await;
|
||||
if report.had_work() {
|
||||
tracing::info!(
|
||||
user_id = hygiene_user,
|
||||
daily_logs_deleted = report.daily_logs_deleted,
|
||||
conversation_docs_deleted = report.conversation_docs_deleted,
|
||||
"multi-user heartbeat: memory hygiene deleted stale documents"
|
||||
);
|
||||
}
|
||||
});
|
||||
|
||||
// Drain completed tasks to stay within the concurrency cap.
|
||||
while join_set.len() >= MAX_CONCURRENT_HEARTBEATS {
|
||||
if let Some(join_result) = join_set.join_next().await {
|
||||
collect_heartbeat_result(join_result, &mut user_failures, &config);
|
||||
}
|
||||
}
|
||||
|
||||
let uid = user_id.clone();
|
||||
let cfg = config.clone();
|
||||
let hyg = hygiene_config.clone();
|
||||
let llm_clone = llm.clone();
|
||||
let tx = response_tx.clone();
|
||||
let admin = store.clone();
|
||||
|
||||
join_set.spawn(async move {
|
||||
let mut runner = HeartbeatRunner::new(cfg, hyg, workspace, llm_clone);
|
||||
if let Some(tx) = tx {
|
||||
runner = runner.with_response_channel(tx);
|
||||
}
|
||||
runner = runner.with_store(admin);
|
||||
|
||||
let result = runner.check_heartbeat().await;
|
||||
if let HeartbeatResult::NeedsAttention(msg) = &result {
|
||||
runner.send_notification(msg).await;
|
||||
}
|
||||
(uid, result)
|
||||
});
|
||||
}
|
||||
|
||||
// Collect remaining results and update failure counts
|
||||
while let Some(join_result) = join_set.join_next().await {
|
||||
collect_heartbeat_result(join_result, &mut user_failures, &config);
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
/// Process a single JoinSet result from the multi-user heartbeat loop.
|
||||
fn collect_heartbeat_result(
|
||||
join_result: Result<(String, HeartbeatResult), tokio::task::JoinError>,
|
||||
user_failures: &mut std::collections::HashMap<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)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
@@ -726,7 +903,7 @@ mod tests {
|
||||
Arc<crate::workspace::Workspace>,
|
||||
Arc<dyn crate::llm::LlmProvider>,
|
||||
Option<tokio::sync::mpsc::Sender<crate::channels::OutgoingResponse>>,
|
||||
Option<Arc<dyn crate::db::Database>>,
|
||||
Option<AdminScope>,
|
||||
) -> tokio::task::JoinHandle<()> = spawn_heartbeat;
|
||||
let _ = _fn_ptr;
|
||||
}
|
||||
|
||||
+3
-1
@@ -36,7 +36,9 @@ pub(crate) use agent_loop::truncate_for_preview;
|
||||
pub use agent_loop::{Agent, AgentDeps};
|
||||
pub use compaction::{CompactionResult, ContextCompactor};
|
||||
pub use context_monitor::{CompactionStrategy, ContextBreakdown, ContextMonitor};
|
||||
pub use heartbeat::{HeartbeatConfig, HeartbeatResult, HeartbeatRunner, spawn_heartbeat};
|
||||
pub use heartbeat::{
|
||||
HeartbeatConfig, HeartbeatResult, HeartbeatRunner, spawn_heartbeat, spawn_multi_user_heartbeat,
|
||||
};
|
||||
pub use router::{MessageIntent, Router};
|
||||
pub use routine::{Routine, RoutineAction, RoutineRun, Trigger};
|
||||
pub use routine_engine::{RoutineEngine, SandboxReadiness};
|
||||
|
||||
@@ -28,12 +28,12 @@ use crate::agent::routine::{
|
||||
use crate::channels::{IncomingMessage, OutgoingResponse};
|
||||
use crate::config::RoutineConfig;
|
||||
use crate::context::{JobContext, JobState};
|
||||
use crate::db::Database;
|
||||
use crate::error::RoutineError;
|
||||
use crate::extensions::ExtensionManager;
|
||||
use crate::llm::{
|
||||
ChatMessage, CompletionRequest, FinishReason, LlmProvider, ToolCall, ToolCompletionRequest,
|
||||
};
|
||||
use crate::tenant::AdminScope;
|
||||
use crate::tools::{
|
||||
ToolError, ToolRegistry, autonomous_allowed_tool_names, autonomous_unavailable_message,
|
||||
prepare_tool_params,
|
||||
@@ -99,7 +99,7 @@ pub(crate) fn routine_matches_message(routine: &Routine, message: &IncomingMessa
|
||||
/// The routine execution engine.
|
||||
pub struct RoutineEngine {
|
||||
config: RoutineConfig,
|
||||
store: Arc<dyn Database>,
|
||||
store: AdminScope,
|
||||
llm: Arc<dyn LlmProvider>,
|
||||
workspace: Arc<Workspace>,
|
||||
/// Sender for notifications (routed to channel manager).
|
||||
@@ -128,7 +128,7 @@ impl RoutineEngine {
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub fn new(
|
||||
config: RoutineConfig,
|
||||
store: Arc<dyn Database>,
|
||||
store: AdminScope,
|
||||
llm: Arc<dyn LlmProvider>,
|
||||
workspace: Arc<Workspace>,
|
||||
notify_tx: mpsc::Sender<OutgoingResponse>,
|
||||
@@ -782,12 +782,22 @@ impl RoutineEngine {
|
||||
});
|
||||
}
|
||||
|
||||
// Per-user workspace (same pattern as spawn_fire).
|
||||
let routine_workspace = if routine.user_id == self.workspace.user_id() {
|
||||
self.workspace.clone()
|
||||
} else {
|
||||
Arc::new(Workspace::new_with_db(
|
||||
&routine.user_id,
|
||||
Arc::clone(self.store.db()),
|
||||
))
|
||||
};
|
||||
|
||||
// Execute inline for manual triggers (caller wants to wait)
|
||||
let engine = EngineContext {
|
||||
config: self.config.clone(),
|
||||
store: self.store.clone(),
|
||||
llm: self.llm.clone(),
|
||||
workspace: self.workspace.clone(),
|
||||
workspace: routine_workspace,
|
||||
notify_tx: self.notify_tx.clone(),
|
||||
running_count: self.running_count.clone(),
|
||||
scheduler: self.scheduler.clone(),
|
||||
@@ -910,11 +920,23 @@ impl RoutineEngine {
|
||||
created_at: Utc::now(),
|
||||
};
|
||||
|
||||
// Use per-user workspace so each routine executes in the correct
|
||||
// user's context. Fall back to the engine-wide workspace when the
|
||||
// routine belongs to the same user (avoids unnecessary allocation).
|
||||
let routine_workspace = if routine.user_id == self.workspace.user_id() {
|
||||
self.workspace.clone()
|
||||
} else {
|
||||
Arc::new(Workspace::new_with_db(
|
||||
&routine.user_id,
|
||||
Arc::clone(self.store.db()),
|
||||
))
|
||||
};
|
||||
|
||||
let engine = EngineContext {
|
||||
config: self.config.clone(),
|
||||
store: self.store.clone(),
|
||||
llm: self.llm.clone(),
|
||||
workspace: self.workspace.clone(),
|
||||
workspace: routine_workspace,
|
||||
notify_tx: self.notify_tx.clone(),
|
||||
running_count: self.running_count.clone(),
|
||||
scheduler: self.scheduler.clone(),
|
||||
@@ -967,7 +989,7 @@ impl RoutineEngine {
|
||||
/// an active state (Pending/InProgress/Stuck). Maps the final `JobState` to
|
||||
/// a `RunStatus` for the routine run.
|
||||
struct FullJobWatcher {
|
||||
store: Arc<dyn Database>,
|
||||
store: AdminScope,
|
||||
job_id: Uuid,
|
||||
routine_name: String,
|
||||
}
|
||||
@@ -978,7 +1000,7 @@ impl FullJobWatcher {
|
||||
/// Safety ceiling: 24 hours, derived from POLL_INTERVAL.
|
||||
const MAX_POLLS: u32 = (24 * 60 * 60) / Self::POLL_INTERVAL.as_secs() as u32;
|
||||
|
||||
fn new(store: Arc<dyn Database>, job_id: Uuid, routine_name: String) -> Self {
|
||||
fn new(store: AdminScope, job_id: Uuid, routine_name: String) -> Self {
|
||||
Self {
|
||||
store,
|
||||
job_id,
|
||||
@@ -1050,7 +1072,7 @@ impl FullJobWatcher {
|
||||
/// Shared context passed to the execution function.
|
||||
struct EngineContext {
|
||||
config: RoutineConfig,
|
||||
store: Arc<dyn Database>,
|
||||
store: AdminScope,
|
||||
llm: Arc<dyn LlmProvider>,
|
||||
workspace: Arc<Workspace>,
|
||||
notify_tx: mpsc::Sender<OutgoingResponse>,
|
||||
|
||||
@@ -11,12 +11,12 @@ use uuid::Uuid;
|
||||
use crate::agent::task::{Task, TaskContext, TaskOutput};
|
||||
use crate::config::AgentConfig;
|
||||
use crate::context::{ContextManager, JobContext, JobState};
|
||||
use crate::db::Database;
|
||||
use crate::error::{Error, JobError};
|
||||
use crate::extensions::ExtensionManager;
|
||||
use crate::hooks::HookRegistry;
|
||||
use crate::llm::LlmProvider;
|
||||
use crate::safety::SafetyLayer;
|
||||
use crate::tenant::AdminScope;
|
||||
use crate::tools::{
|
||||
ApprovalContext, ToolRegistry, autonomous_allowed_tool_names, autonomous_unavailable_error,
|
||||
prepare_tool_params,
|
||||
@@ -52,7 +52,7 @@ struct ScheduledSubtask {
|
||||
pub struct SchedulerDeps {
|
||||
pub tools: Arc<ToolRegistry>,
|
||||
pub extension_manager: Option<Arc<ExtensionManager>>,
|
||||
pub store: Option<Arc<dyn Database>>,
|
||||
pub store: Option<AdminScope>,
|
||||
pub hooks: Arc<HookRegistry>,
|
||||
}
|
||||
|
||||
@@ -64,7 +64,7 @@ pub struct Scheduler {
|
||||
safety: Arc<SafetyLayer>,
|
||||
tools: Arc<ToolRegistry>,
|
||||
extension_manager: Option<Arc<ExtensionManager>>,
|
||||
store: Option<Arc<dyn Database>>,
|
||||
store: Option<AdminScope>,
|
||||
hooks: Arc<HookRegistry>,
|
||||
/// SSE manager for live job event streaming.
|
||||
sse_tx: Option<Arc<crate::channels::web::sse::SseManager>>,
|
||||
@@ -780,10 +780,14 @@ mod tests {
|
||||
allow_local_tools: true,
|
||||
max_cost_per_day_cents: None,
|
||||
max_actions_per_hour: None,
|
||||
max_cost_per_user_per_day_cents: None,
|
||||
max_tool_iterations: 10,
|
||||
auto_approve_tools: true,
|
||||
default_timezone: "UTC".to_string(),
|
||||
max_tokens_per_job,
|
||||
multi_tenant: false,
|
||||
max_llm_concurrent_per_user: None,
|
||||
max_jobs_concurrent_per_user: None,
|
||||
};
|
||||
let cm = Arc::new(ContextManager::new(5));
|
||||
let llm: Arc<dyn LlmProvider> = Arc::new(StubLlm);
|
||||
|
||||
@@ -8,8 +8,8 @@ use chrono::{DateTime, Utc};
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::context::{ContextManager, JobState};
|
||||
use crate::db::Database;
|
||||
use crate::error::RepairError;
|
||||
use crate::tenant::AdminScope;
|
||||
use crate::tools::{BuildRequirement, Language, SoftwareBuilder, SoftwareType, ToolRegistry};
|
||||
|
||||
/// A job that has been detected as stuck.
|
||||
@@ -69,7 +69,7 @@ pub struct DefaultSelfRepair {
|
||||
/// Jobs in `InProgress` longer than this are treated as stuck.
|
||||
stuck_threshold: Duration,
|
||||
max_repair_attempts: u32,
|
||||
store: Option<Arc<dyn Database>>,
|
||||
store: Option<AdminScope>,
|
||||
builder: Option<Arc<dyn SoftwareBuilder>>,
|
||||
tools: Option<Arc<ToolRegistry>>,
|
||||
}
|
||||
@@ -91,8 +91,8 @@ impl DefaultSelfRepair {
|
||||
}
|
||||
}
|
||||
|
||||
/// Add a Store for tool failure tracking.
|
||||
pub fn with_store(mut self, store: Arc<dyn Database>) -> Self {
|
||||
/// Add an admin-scoped store for tool failure tracking.
|
||||
pub fn with_store(mut self, store: AdminScope) -> Self {
|
||||
self.store = Some(store);
|
||||
self
|
||||
}
|
||||
@@ -806,7 +806,7 @@ mod tests {
|
||||
// Create self-repair with zero threshold (detect immediately),
|
||||
// wired with store, builder, and tools.
|
||||
let repair = DefaultSelfRepair::new(Arc::clone(&cm), Duration::from_secs(0), 3)
|
||||
.with_store(Arc::clone(&db))
|
||||
.with_store(crate::tenant::AdminScope::new(Arc::clone(&db)))
|
||||
.with_builder(
|
||||
Arc::clone(&builder) as Arc<dyn crate::tools::SoftwareBuilder>,
|
||||
tools,
|
||||
|
||||
+10
-3
@@ -175,6 +175,7 @@ impl Agent {
|
||||
pub(super) async fn process_user_input(
|
||||
&self,
|
||||
message: &IncomingMessage,
|
||||
tenant: crate::tenant::TenantCtx,
|
||||
session: Arc<Mutex<Session>>,
|
||||
thread_id: Uuid,
|
||||
content: &str,
|
||||
@@ -351,7 +352,7 @@ impl Agent {
|
||||
|
||||
if let Some(intent) = self.router.route_command(&temp_message) {
|
||||
// Explicit command like /status, /job, /list - handle directly
|
||||
return self.handle_job_or_command(intent, message).await;
|
||||
return self.handle_job_or_command(intent, message, &tenant).await;
|
||||
}
|
||||
|
||||
// Natural language goes through the agentic loop
|
||||
@@ -462,7 +463,7 @@ impl Agent {
|
||||
|
||||
// Run the agentic tool execution loop
|
||||
let result = self
|
||||
.run_agentic_loop(message, session.clone(), thread_id, turn_messages)
|
||||
.run_agentic_loop(message, tenant, session.clone(), thread_id, turn_messages)
|
||||
.await;
|
||||
|
||||
// Re-acquire lock and check if interrupted
|
||||
@@ -1473,7 +1474,13 @@ impl Agent {
|
||||
|
||||
// Continue the agentic loop (a tool was already executed this turn)
|
||||
let result = self
|
||||
.run_agentic_loop(message, session.clone(), thread_id, context_messages)
|
||||
.run_agentic_loop(
|
||||
message,
|
||||
self.tenant_ctx(&message.user_id).await,
|
||||
session.clone(),
|
||||
thread_id,
|
||||
context_messages,
|
||||
)
|
||||
.await;
|
||||
|
||||
// Handle the result
|
||||
|
||||
+21
-1
@@ -449,6 +449,8 @@ impl AppBuilder {
|
||||
|
||||
let mcp_session_manager = Arc::new(McpSessionManager::new());
|
||||
let mcp_process_manager = Arc::new(McpProcessManager::new());
|
||||
let companion_mcp_server =
|
||||
crate::tools::mcp::config::derive_nearai_companion_mcp_server(&self.config);
|
||||
|
||||
// Create WASM tool runtime eagerly so extensions installed after startup
|
||||
// (e.g. via the web UI) can still be activated. The tools directory is only
|
||||
@@ -526,6 +528,7 @@ impl AppBuilder {
|
||||
let mcp_sm = Arc::clone(&mcp_session_manager);
|
||||
let pm = Arc::clone(&mcp_process_manager);
|
||||
let owner_id = self.config.owner_id.clone();
|
||||
let companion_mcp_server = companion_mcp_server.clone();
|
||||
async move {
|
||||
let servers_result = if let Some(ref d) = db {
|
||||
load_mcp_servers_from_db(d.as_ref(), &owner_id).await
|
||||
@@ -533,7 +536,16 @@ impl AppBuilder {
|
||||
crate::tools::mcp::config::load_mcp_servers().await
|
||||
};
|
||||
match servers_result {
|
||||
Ok(servers) => {
|
||||
Ok(mut servers) => {
|
||||
if let Some(companion) = companion_mcp_server {
|
||||
let companion_name = companion.name.clone();
|
||||
if !servers.insert_if_absent(companion) {
|
||||
tracing::debug!(
|
||||
"Skipping derived MCP companion '{}': an existing config with that name is already present",
|
||||
companion_name
|
||||
);
|
||||
}
|
||||
}
|
||||
let enabled: Vec<_> = servers.enabled_servers().cloned().collect();
|
||||
if !enabled.is_empty() {
|
||||
tracing::debug!(
|
||||
@@ -545,6 +557,8 @@ impl AppBuilder {
|
||||
let mut join_set = tokio::task::JoinSet::new();
|
||||
for server in enabled {
|
||||
let mcp_sm = Arc::clone(&mcp_sm);
|
||||
let nearai_session = Arc::clone(&self.session);
|
||||
let nearai_api_key = self.config.llm.nearai.api_key.clone();
|
||||
let secrets = secrets_store.clone();
|
||||
let tools = Arc::clone(&tools);
|
||||
let pm = Arc::clone(&pm);
|
||||
@@ -556,6 +570,8 @@ impl AppBuilder {
|
||||
let client = match crate::tools::mcp::create_client_from_config(
|
||||
server,
|
||||
&mcp_sm,
|
||||
Some(nearai_session),
|
||||
nearai_api_key,
|
||||
&pm,
|
||||
secrets,
|
||||
&owner_id,
|
||||
@@ -712,6 +728,8 @@ impl AppBuilder {
|
||||
let manager = Arc::new(ExtensionManager::new(
|
||||
Arc::clone(&mcp_session_manager),
|
||||
Arc::clone(&mcp_process_manager),
|
||||
Some(Arc::clone(&self.session)),
|
||||
self.config.llm.nearai.api_key.clone(),
|
||||
ext_secrets,
|
||||
Arc::clone(tools),
|
||||
Some(Arc::clone(hooks)),
|
||||
@@ -721,6 +739,7 @@ impl AppBuilder {
|
||||
self.config.tunnel.public_url.clone(),
|
||||
self.config.owner_id.clone(),
|
||||
self.db.clone(),
|
||||
companion_mcp_server,
|
||||
catalog_entries.clone(),
|
||||
));
|
||||
tools.register_extension_tools(Arc::clone(&manager));
|
||||
@@ -880,6 +899,7 @@ impl AppBuilder {
|
||||
crate::agent::cost_guard::CostGuardConfig {
|
||||
max_cost_per_day_cents: self.config.agent.max_cost_per_day_cents,
|
||||
max_actions_per_hour: self.config.agent.max_actions_per_hour,
|
||||
max_cost_per_user_per_day_cents: self.config.agent.max_cost_per_user_per_day_cents,
|
||||
},
|
||||
));
|
||||
|
||||
|
||||
@@ -122,18 +122,32 @@ impl RelayClient {
|
||||
/// instance_url in chat-api. IronClaw only passes an optional CSRF nonce
|
||||
/// for validating the callback — no URLs.
|
||||
pub async fn initiate_oauth(&self, state_nonce: Option<&str>) -> Result<String, RelayError> {
|
||||
let url = format!("{}/oauth/slack/auth", self.base_url);
|
||||
tracing::debug!(relay_url = %url, "RelayClient::initiate_oauth: sending request");
|
||||
let mut query: Vec<(&str, &str)> = vec![];
|
||||
if let Some(nonce) = state_nonce {
|
||||
query.push(("state_nonce", nonce));
|
||||
}
|
||||
let resp = self
|
||||
.http
|
||||
.get(format!("{}/oauth/slack/auth", self.base_url))
|
||||
.get(&url)
|
||||
.bearer_auth(self.api_key.expose_secret())
|
||||
.query(&query)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| RelayError::Network(e.to_string()))?;
|
||||
.map_err(|e| {
|
||||
tracing::warn!(
|
||||
relay_url = %url,
|
||||
error = %e,
|
||||
"RelayClient::initiate_oauth: network request failed"
|
||||
);
|
||||
RelayError::Network(e.to_string())
|
||||
})?;
|
||||
tracing::debug!(
|
||||
relay_url = %url,
|
||||
status = %resp.status(),
|
||||
"RelayClient::initiate_oauth: received response"
|
||||
);
|
||||
|
||||
let status = resp.status();
|
||||
if status.is_redirection() {
|
||||
@@ -224,20 +238,39 @@ impl RelayClient {
|
||||
method: &str,
|
||||
body: serde_json::Value,
|
||||
) -> Result<serde_json::Value, RelayError> {
|
||||
let url = format!("{}/proxy/{}/{}", self.base_url, provider, method);
|
||||
tracing::debug!(
|
||||
relay_url = %url,
|
||||
provider = %provider,
|
||||
method = %method,
|
||||
"RelayClient::proxy_provider: sending request"
|
||||
);
|
||||
let query: Vec<(&str, &str)> = vec![("team_id", team_id)];
|
||||
let resp = self
|
||||
.http
|
||||
.post(format!("{}/proxy/{}/{}", self.base_url, provider, method))
|
||||
.post(&url)
|
||||
.bearer_auth(self.api_key.expose_secret())
|
||||
.query(&query)
|
||||
.json(&body)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| RelayError::Network(e.to_string()))?;
|
||||
.map_err(|e| {
|
||||
tracing::warn!(
|
||||
relay_url = %url,
|
||||
error = %e,
|
||||
"RelayClient::proxy_provider: network request failed"
|
||||
);
|
||||
RelayError::Network(e.to_string())
|
||||
})?;
|
||||
|
||||
if !resp.status().is_success() {
|
||||
let status = resp.status().as_u16();
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
tracing::warn!(
|
||||
relay_url = %url,
|
||||
status = status,
|
||||
"RelayClient::proxy_provider: channel-relay returned error"
|
||||
);
|
||||
return Err(RelayError::Api {
|
||||
status,
|
||||
message: body,
|
||||
@@ -255,23 +288,45 @@ impl RelayClient {
|
||||
/// 32-byte secret. Called once at activation time; the result is cached in the
|
||||
/// extension manager so subsequent calls to `relay_signing_secret()` use it.
|
||||
pub async fn get_signing_secret(&self, team_id: &str) -> Result<Vec<u8>, RelayError> {
|
||||
let url = format!("{}/relay/signing-secret", self.base_url);
|
||||
tracing::debug!(
|
||||
relay_url = %url,
|
||||
"RelayClient::get_signing_secret: fetching signing secret"
|
||||
);
|
||||
let resp = self
|
||||
.http
|
||||
.get(format!("{}/relay/signing-secret", self.base_url))
|
||||
.get(&url)
|
||||
.bearer_auth(self.api_key.expose_secret())
|
||||
.query(&[("team_id", team_id)])
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| RelayError::Network(e.to_string()))?;
|
||||
.map_err(|e| {
|
||||
tracing::warn!(
|
||||
relay_url = %url,
|
||||
error = %e,
|
||||
"RelayClient::get_signing_secret: network request failed"
|
||||
);
|
||||
RelayError::Network(e.to_string())
|
||||
})?;
|
||||
|
||||
if !resp.status().is_success() {
|
||||
let status = resp.status().as_u16();
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
tracing::warn!(
|
||||
relay_url = %url,
|
||||
status = status,
|
||||
body = %body,
|
||||
"RelayClient::get_signing_secret: channel-relay returned error"
|
||||
);
|
||||
return Err(RelayError::Api {
|
||||
status,
|
||||
message: body,
|
||||
});
|
||||
}
|
||||
tracing::debug!(
|
||||
relay_url = %url,
|
||||
"RelayClient::get_signing_secret: received successful response"
|
||||
);
|
||||
|
||||
let body: serde_json::Value = resp
|
||||
.json()
|
||||
|
||||
@@ -70,6 +70,7 @@ pub async fn extensions_list_handler(
|
||||
tools: ext.tools,
|
||||
needs_setup: ext.needs_setup,
|
||||
has_auth: ext.has_auth,
|
||||
derived: ext.derived,
|
||||
activation_status,
|
||||
activation_error: ext.activation_error,
|
||||
version: ext.version,
|
||||
|
||||
@@ -54,10 +54,37 @@ fn validate_webhook_secret(
|
||||
///
|
||||
/// This endpoint is **public** (no gateway auth token required) but protected
|
||||
/// by the per-routine webhook secret sent via the `X-Webhook-Secret` header.
|
||||
///
|
||||
/// **Single-user/backward-compatible**: looks up routines by path across all
|
||||
/// users. For multi-tenant isolation, use the user-scoped endpoint at
|
||||
/// `/api/webhooks/u/{user_id}/{path}` instead.
|
||||
pub async fn webhook_trigger_handler(
|
||||
State(state): State<Arc<GatewayState>>,
|
||||
Path(path): Path<String>,
|
||||
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)> {
|
||||
// Rate limit check
|
||||
if !state.webhook_rate_limiter.check() {
|
||||
@@ -72,9 +99,9 @@ pub async fn webhook_trigger_handler(
|
||||
"Database not available".to_string(),
|
||||
))?;
|
||||
|
||||
// Targeted query instead of loading all routines
|
||||
// Targeted query — when user_id is provided, restrict to that user's routines
|
||||
let routine = store
|
||||
.get_webhook_routine_by_path(&path)
|
||||
.get_webhook_routine_by_path(path, user_id)
|
||||
.await
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
|
||||
.ok_or((
|
||||
@@ -99,7 +126,7 @@ pub async fn webhook_trigger_handler(
|
||||
))?
|
||||
};
|
||||
|
||||
let run_id = engine.fire_webhook(routine.id, &path).await.map_err(|e| {
|
||||
let run_id = engine.fire_webhook(routine.id, path).await.map_err(|e| {
|
||||
let status = match &e {
|
||||
crate::error::RoutineError::NotFound { .. } => StatusCode::NOT_FOUND,
|
||||
crate::error::RoutineError::Disabled { .. }
|
||||
|
||||
+495
-5
@@ -414,6 +414,11 @@ pub async fn start_server(
|
||||
.route(
|
||||
"/api/webhooks/{path}",
|
||||
post(crate::channels::web::handlers::webhooks::webhook_trigger_handler),
|
||||
)
|
||||
// User-scoped webhook endpoint for multi-tenant isolation
|
||||
.route(
|
||||
"/api/webhooks/u/{user_id}/{path}",
|
||||
post(crate::channels::web::handlers::webhooks::webhook_trigger_user_scoped_handler),
|
||||
);
|
||||
|
||||
// Protected routes (require auth)
|
||||
@@ -831,10 +836,10 @@ async fn oauth_callback_handler(
|
||||
|
||||
let result: Result<(), String> = async {
|
||||
let token_response = if let Some(proxy_url) = &exchange_proxy_url {
|
||||
let gateway_token = flow.gateway_token.as_deref().unwrap_or_default();
|
||||
let oauth_proxy_auth_token = flow.oauth_proxy_auth_token().unwrap_or_default();
|
||||
oauth_defaults::exchange_via_proxy(oauth_defaults::ProxyTokenExchangeRequest {
|
||||
proxy_url,
|
||||
gateway_token,
|
||||
gateway_token: oauth_proxy_auth_token,
|
||||
token_url: &flow.token_url,
|
||||
client_id: &flow.client_id,
|
||||
client_secret: flow.client_secret.as_deref(),
|
||||
@@ -1172,11 +1177,31 @@ async fn slack_relay_oauth_callback_handler(
|
||||
|
||||
// Store team_id in settings
|
||||
let team_id_key = format!("relay:{}:team_id", DEFAULT_RELAY_NAME);
|
||||
let _ = store
|
||||
tracing::info!(
|
||||
relay = DEFAULT_RELAY_NAME,
|
||||
owner_id = %state.owner_id,
|
||||
team_id_key = %team_id_key,
|
||||
"relay OAuth callback: storing team_id in settings"
|
||||
);
|
||||
store
|
||||
.set_setting(&state.owner_id, &team_id_key, &serde_json::json!(team_id))
|
||||
.await;
|
||||
.await
|
||||
.map_err(|e| {
|
||||
tracing::error!(
|
||||
relay = DEFAULT_RELAY_NAME,
|
||||
owner_id = %state.owner_id,
|
||||
error = %e,
|
||||
"relay OAuth callback: failed to persist team_id to settings store"
|
||||
);
|
||||
format!("Failed to persist relay team_id: {e}")
|
||||
})?;
|
||||
|
||||
// Activate the relay channel
|
||||
tracing::info!(
|
||||
relay = DEFAULT_RELAY_NAME,
|
||||
owner_id = %state.owner_id,
|
||||
"relay OAuth callback: activating relay channel"
|
||||
);
|
||||
ext_mgr
|
||||
.activate_stored_relay(DEFAULT_RELAY_NAME, &state.owner_id)
|
||||
.await
|
||||
@@ -2067,6 +2092,7 @@ async fn extensions_list_handler(
|
||||
tools: ext.tools,
|
||||
needs_setup: ext.needs_setup,
|
||||
has_auth: ext.has_auth,
|
||||
derived: ext.derived,
|
||||
activation_status,
|
||||
activation_error: ext.activation_error,
|
||||
version: ext.version,
|
||||
@@ -2176,6 +2202,11 @@ async fn extensions_activate_handler(
|
||||
AuthenticatedUser(user): AuthenticatedUser,
|
||||
Path(name): Path<String>,
|
||||
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
|
||||
tracing::debug!(
|
||||
extension = %name,
|
||||
user_id = %user.user_id,
|
||||
"extensions_activate_handler: received activate request"
|
||||
);
|
||||
let ext_mgr = state.extension_manager.as_ref().ok_or((
|
||||
StatusCode::NOT_IMPLEMENTED,
|
||||
"Extension manager not available (secrets store required)".to_string(),
|
||||
@@ -2183,6 +2214,10 @@ async fn extensions_activate_handler(
|
||||
|
||||
match ext_mgr.activate(&name, &user.user_id).await {
|
||||
Ok(result) => {
|
||||
tracing::info!(
|
||||
extension = %name,
|
||||
"extensions_activate_handler: activation succeeded"
|
||||
);
|
||||
// Activation loaded the WASM module. Check if the tool needs
|
||||
// OAuth scope expansion (e.g., adding google-docs when gmail
|
||||
// already has a token but missing the documents scope).
|
||||
@@ -2201,6 +2236,13 @@ async fn extensions_activate_handler(
|
||||
crate::extensions::ExtensionError::AuthRequired
|
||||
);
|
||||
|
||||
tracing::debug!(
|
||||
extension = %name,
|
||||
error = %activate_err,
|
||||
needs_auth = needs_auth,
|
||||
"extensions_activate_handler: activation failed, attempting auth fallback"
|
||||
);
|
||||
|
||||
if !needs_auth {
|
||||
return Ok(Json(ActionResponse::fail(activate_err.to_string())));
|
||||
}
|
||||
@@ -2208,10 +2250,21 @@ async fn extensions_activate_handler(
|
||||
// Activation failed due to auth; try authenticating first.
|
||||
match ext_mgr.auth(&name, &user.user_id).await {
|
||||
Ok(auth_result) if auth_result.is_authenticated() => {
|
||||
tracing::debug!(
|
||||
extension = %name,
|
||||
"extensions_activate_handler: auth reports authenticated, retrying activate"
|
||||
);
|
||||
// Auth succeeded, retry activation.
|
||||
match ext_mgr.activate(&name, &user.user_id).await {
|
||||
Ok(result) => Ok(Json(ActionResponse::ok(result.message))),
|
||||
Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))),
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
extension = %name,
|
||||
error = %e,
|
||||
"extensions_activate_handler: retry after auth still failed"
|
||||
);
|
||||
Ok(Json(ActionResponse::fail(e.to_string())))
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(auth_result) => {
|
||||
@@ -2896,6 +2949,7 @@ mod tests {
|
||||
tools: Vec::new(),
|
||||
needs_setup: true,
|
||||
has_auth: false,
|
||||
derived: false,
|
||||
installed: true,
|
||||
activation_error: None,
|
||||
version: None,
|
||||
@@ -2933,6 +2987,7 @@ mod tests {
|
||||
tools: Vec::new(),
|
||||
needs_setup: true,
|
||||
has_auth: false,
|
||||
derived: false,
|
||||
installed: true,
|
||||
activation_error: None,
|
||||
version: None,
|
||||
@@ -3005,6 +3060,160 @@ mod tests {
|
||||
.with_state(state)
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
struct RecordedOauthProxyRequest {
|
||||
authorization: Option<String>,
|
||||
form: std::collections::HashMap<String, String>,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct MockOauthProxyState {
|
||||
requests: Arc<tokio::sync::Mutex<Vec<RecordedOauthProxyRequest>>>,
|
||||
}
|
||||
|
||||
struct MockOauthProxyServer {
|
||||
addr: std::net::SocketAddr,
|
||||
requests: Arc<tokio::sync::Mutex<Vec<RecordedOauthProxyRequest>>>,
|
||||
shutdown_tx: Option<tokio::sync::oneshot::Sender<()>>,
|
||||
server_task: Option<tokio::task::JoinHandle<()>>,
|
||||
}
|
||||
|
||||
impl MockOauthProxyServer {
|
||||
async fn start() -> Self {
|
||||
async fn exchange_handler(
|
||||
State(state): State<MockOauthProxyState>,
|
||||
headers: axum::http::HeaderMap,
|
||||
axum::Form(form): axum::Form<std::collections::HashMap<String, String>>,
|
||||
) -> Json<serde_json::Value> {
|
||||
state.requests.lock().await.push(RecordedOauthProxyRequest {
|
||||
authorization: headers
|
||||
.get(axum::http::header::AUTHORIZATION)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.map(str::to_string),
|
||||
form,
|
||||
});
|
||||
Json(serde_json::json!({
|
||||
"access_token": "proxy-access-token",
|
||||
"refresh_token": "proxy-refresh-token",
|
||||
"expires_in": 7200
|
||||
}))
|
||||
}
|
||||
|
||||
let requests = Arc::new(tokio::sync::Mutex::new(Vec::new()));
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
|
||||
.await
|
||||
.expect("bind mock oauth proxy");
|
||||
let addr = listener.local_addr().expect("mock oauth proxy addr");
|
||||
let app = Router::new()
|
||||
.route("/oauth/exchange", post(exchange_handler))
|
||||
.with_state(MockOauthProxyState {
|
||||
requests: Arc::clone(&requests),
|
||||
});
|
||||
let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel::<()>();
|
||||
let server_task = tokio::spawn(async move {
|
||||
let _ = axum::serve(listener, app)
|
||||
.with_graceful_shutdown(async {
|
||||
let _ = shutdown_rx.await;
|
||||
})
|
||||
.await;
|
||||
});
|
||||
|
||||
Self {
|
||||
addr,
|
||||
requests,
|
||||
shutdown_tx: Some(shutdown_tx),
|
||||
server_task: Some(server_task),
|
||||
}
|
||||
}
|
||||
|
||||
fn base_url(&self) -> String {
|
||||
format!("http://{}", self.addr)
|
||||
}
|
||||
|
||||
async fn requests(&self) -> Vec<RecordedOauthProxyRequest> {
|
||||
self.requests.lock().await.clone()
|
||||
}
|
||||
|
||||
async fn shutdown(mut self) {
|
||||
if let Some(tx) = self.shutdown_tx.take() {
|
||||
let _ = tx.send(());
|
||||
}
|
||||
if let Some(task) = self.server_task.take() {
|
||||
let _ = task.await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for MockOauthProxyServer {
|
||||
fn drop(&mut self) {
|
||||
if let Some(tx) = self.shutdown_tx.take() {
|
||||
let _ = tx.send(());
|
||||
}
|
||||
if let Some(task) = self.server_task.take() {
|
||||
task.abort();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct EnvVarGuard {
|
||||
key: &'static str,
|
||||
original: Option<String>,
|
||||
}
|
||||
|
||||
impl Drop for EnvVarGuard {
|
||||
fn drop(&mut self) {
|
||||
// SAFETY: Tests use lock_env() to serialize environment access.
|
||||
unsafe {
|
||||
if let Some(ref value) = self.original {
|
||||
std::env::set_var(self.key, value);
|
||||
} else {
|
||||
std::env::remove_var(self.key);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn set_env_var(key: &'static str, value: Option<&str>) -> EnvVarGuard {
|
||||
let original = std::env::var(key).ok();
|
||||
// SAFETY: Tests use lock_env() to serialize environment access.
|
||||
unsafe {
|
||||
if let Some(value) = value {
|
||||
std::env::set_var(key, value);
|
||||
} else {
|
||||
std::env::remove_var(key);
|
||||
}
|
||||
}
|
||||
EnvVarGuard { key, original }
|
||||
}
|
||||
|
||||
fn fresh_pending_oauth_flow(
|
||||
secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync>,
|
||||
sse_manager: Option<Arc<SseManager>>,
|
||||
oauth_proxy_auth_token: Option<String>,
|
||||
) -> crate::cli::oauth_defaults::PendingOAuthFlow {
|
||||
crate::cli::oauth_defaults::PendingOAuthFlow {
|
||||
extension_name: "test_tool".to_string(),
|
||||
display_name: "Test Tool".to_string(),
|
||||
token_url: "https://example.com/token".to_string(),
|
||||
client_id: "client123".to_string(),
|
||||
client_secret: None,
|
||||
redirect_uri: "https://example.com/oauth/callback".to_string(),
|
||||
code_verifier: Some("test-code-verifier".to_string()),
|
||||
access_token_field: "access_token".to_string(),
|
||||
secret_name: "test_token".to_string(),
|
||||
provider: Some("google".to_string()),
|
||||
validation_endpoint: None,
|
||||
scopes: vec!["email".to_string()],
|
||||
user_id: "test".to_string(),
|
||||
secrets,
|
||||
sse_manager,
|
||||
gateway_token: oauth_proxy_auth_token,
|
||||
token_exchange_extra_params: std::collections::HashMap::new(),
|
||||
client_id_secret_name: None,
|
||||
created_at: std::time::Instant::now(),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_extensions_setup_submit_returns_failure_when_not_activated() {
|
||||
use axum::body::Body;
|
||||
@@ -3662,6 +3871,284 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_oauth_callback_accepts_versioned_hosted_state_without_instance_name() {
|
||||
use axum::body::Body;
|
||||
use tower::ServiceExt;
|
||||
|
||||
let secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync> =
|
||||
Arc::new(crate::secrets::InMemorySecretsStore::new(Arc::new(
|
||||
crate::secrets::SecretsCrypto::new(secrecy::SecretString::from(
|
||||
TEST_GATEWAY_CRYPTO_KEY.to_string(),
|
||||
))
|
||||
.expect("crypto"),
|
||||
)));
|
||||
let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(secrets.clone());
|
||||
|
||||
let Some(created_at) = expired_flow_created_at() else {
|
||||
eprintln!(
|
||||
"Skipping versioned OAuth state without instance test: monotonic uptime below expiry window"
|
||||
);
|
||||
return;
|
||||
};
|
||||
let flow = crate::cli::oauth_defaults::PendingOAuthFlow {
|
||||
extension_name: "test_tool".to_string(),
|
||||
display_name: "Test Tool".to_string(),
|
||||
token_url: "https://example.com/token".to_string(),
|
||||
client_id: "client123".to_string(),
|
||||
client_secret: None,
|
||||
redirect_uri: "https://example.com/oauth/callback".to_string(),
|
||||
code_verifier: None,
|
||||
access_token_field: "access_token".to_string(),
|
||||
secret_name: "test_token".to_string(),
|
||||
provider: None,
|
||||
validation_endpoint: None,
|
||||
scopes: vec![],
|
||||
user_id: "test".to_string(),
|
||||
secrets,
|
||||
sse_manager: None,
|
||||
gateway_token: None,
|
||||
token_exchange_extra_params: std::collections::HashMap::new(),
|
||||
client_id_secret_name: None,
|
||||
created_at,
|
||||
};
|
||||
|
||||
ext_mgr
|
||||
.pending_oauth_flows()
|
||||
.write()
|
||||
.await
|
||||
.insert("test_nonce".to_string(), flow);
|
||||
|
||||
let state = test_gateway_state(Some(ext_mgr.clone()));
|
||||
let app = test_oauth_router(state);
|
||||
let versioned_state =
|
||||
crate::cli::oauth_defaults::encode_hosted_oauth_state("test_nonce", None);
|
||||
|
||||
let req = axum::http::Request::builder()
|
||||
.uri(format!(
|
||||
"/oauth/callback?code=fake_code&state={}",
|
||||
urlencoding::encode(&versioned_state)
|
||||
))
|
||||
.body(Body::empty())
|
||||
.expect("request");
|
||||
|
||||
let resp = ServiceExt::<axum::http::Request<Body>>::oneshot(app, req)
|
||||
.await
|
||||
.expect("response");
|
||||
assert_eq!(resp.status(), StatusCode::OK);
|
||||
|
||||
let body = axum::body::to_bytes(resp.into_body(), 1024 * 64)
|
||||
.await
|
||||
.expect("body");
|
||||
let html = String::from_utf8_lossy(&body);
|
||||
assert!(html.contains("Authorization Failed"));
|
||||
assert!(
|
||||
ext_mgr
|
||||
.pending_oauth_flows()
|
||||
.read()
|
||||
.await
|
||||
.get("test_nonce")
|
||||
.is_none()
|
||||
);
|
||||
}
|
||||
|
||||
#[allow(clippy::await_holding_lock)]
|
||||
#[tokio::test]
|
||||
async fn test_oauth_callback_happy_path_with_gateway_token_fallback() {
|
||||
use axum::body::Body;
|
||||
use tower::ServiceExt;
|
||||
|
||||
let proxy = MockOauthProxyServer::start().await;
|
||||
// Keep the process-wide env locked for the full callback so the handler
|
||||
// sees a stable proxy URL/token configuration throughout the test.
|
||||
let _env_guard = crate::config::helpers::lock_env();
|
||||
let _exchange_url_guard =
|
||||
set_env_var("IRONCLAW_OAUTH_EXCHANGE_URL", Some(&proxy.base_url()));
|
||||
let _proxy_auth_guard = set_env_var("IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", None);
|
||||
let _gateway_token_guard = set_env_var("GATEWAY_AUTH_TOKEN", Some("gateway-test-token"));
|
||||
|
||||
let secrets = test_secrets_store();
|
||||
let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(Arc::clone(&secrets));
|
||||
let sse_mgr = Arc::new(SseManager::new());
|
||||
let mut receiver = sse_mgr.sender().subscribe();
|
||||
let flow = fresh_pending_oauth_flow(
|
||||
Arc::clone(&secrets),
|
||||
Some(Arc::clone(&sse_mgr)),
|
||||
crate::cli::oauth_defaults::oauth_proxy_auth_token(),
|
||||
);
|
||||
|
||||
ext_mgr
|
||||
.pending_oauth_flows()
|
||||
.write()
|
||||
.await
|
||||
.insert("test_nonce".to_string(), flow);
|
||||
|
||||
let state = test_gateway_state(Some(ext_mgr.clone()));
|
||||
let app = test_oauth_router(state);
|
||||
let versioned_state =
|
||||
crate::cli::oauth_defaults::encode_hosted_oauth_state("test_nonce", Some("myinstance"));
|
||||
|
||||
let req = axum::http::Request::builder()
|
||||
.uri(format!(
|
||||
"/oauth/callback?code=fake_code&state={}",
|
||||
urlencoding::encode(&versioned_state)
|
||||
))
|
||||
.body(Body::empty())
|
||||
.expect("request");
|
||||
|
||||
let resp = ServiceExt::<axum::http::Request<Body>>::oneshot(app, req)
|
||||
.await
|
||||
.expect("response");
|
||||
assert_eq!(resp.status(), StatusCode::OK);
|
||||
|
||||
let body = axum::body::to_bytes(resp.into_body(), 1024 * 64)
|
||||
.await
|
||||
.expect("body");
|
||||
let html = String::from_utf8_lossy(&body);
|
||||
assert!(html.contains("Test Tool Connected"));
|
||||
|
||||
let requests = proxy.requests().await;
|
||||
assert_eq!(requests.len(), 1);
|
||||
assert_eq!(
|
||||
requests[0].authorization.as_deref(),
|
||||
Some("Bearer gateway-test-token")
|
||||
);
|
||||
assert_eq!(
|
||||
requests[0].form.get("code").map(String::as_str),
|
||||
Some("fake_code")
|
||||
);
|
||||
assert_eq!(
|
||||
requests[0].form.get("code_verifier").map(String::as_str),
|
||||
Some("test-code-verifier")
|
||||
);
|
||||
|
||||
let access_token = secrets
|
||||
.get_decrypted("test", "test_token")
|
||||
.await
|
||||
.expect("access token stored");
|
||||
assert_eq!(access_token.expose(), "proxy-access-token");
|
||||
|
||||
let refresh_token = secrets
|
||||
.get_decrypted("test", "test_token_refresh_token")
|
||||
.await
|
||||
.expect("refresh token stored");
|
||||
assert_eq!(refresh_token.expose(), "proxy-refresh-token");
|
||||
|
||||
match receiver.recv().await.expect("auth_completed event").event {
|
||||
crate::channels::web::types::AppEvent::AuthCompleted {
|
||||
extension_name,
|
||||
success,
|
||||
..
|
||||
} => {
|
||||
assert_eq!(extension_name, "test_tool");
|
||||
assert!(success, "OAuth callback should broadcast success");
|
||||
}
|
||||
event => panic!("expected AuthCompleted event, got {event:?}"),
|
||||
}
|
||||
|
||||
proxy.shutdown().await;
|
||||
}
|
||||
|
||||
#[allow(clippy::await_holding_lock)]
|
||||
#[tokio::test]
|
||||
async fn test_oauth_callback_happy_path_with_dedicated_proxy_auth_token() {
|
||||
use axum::body::Body;
|
||||
use tower::ServiceExt;
|
||||
|
||||
let proxy = MockOauthProxyServer::start().await;
|
||||
// Keep the process-wide env locked for the full callback so the handler
|
||||
// sees a stable proxy URL/token configuration throughout the test.
|
||||
let _env_guard = crate::config::helpers::lock_env();
|
||||
let _exchange_url_guard =
|
||||
set_env_var("IRONCLAW_OAUTH_EXCHANGE_URL", Some(&proxy.base_url()));
|
||||
let _proxy_auth_guard = set_env_var(
|
||||
"IRONCLAW_OAUTH_PROXY_AUTH_TOKEN",
|
||||
Some("shared-oauth-proxy-secret"),
|
||||
);
|
||||
let _gateway_token_guard = set_env_var("GATEWAY_AUTH_TOKEN", None);
|
||||
|
||||
let secrets = test_secrets_store();
|
||||
let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(Arc::clone(&secrets));
|
||||
let sse_mgr = Arc::new(SseManager::new());
|
||||
let mut receiver = sse_mgr.sender().subscribe();
|
||||
let flow = fresh_pending_oauth_flow(
|
||||
Arc::clone(&secrets),
|
||||
Some(Arc::clone(&sse_mgr)),
|
||||
crate::cli::oauth_defaults::oauth_proxy_auth_token(),
|
||||
);
|
||||
|
||||
ext_mgr
|
||||
.pending_oauth_flows()
|
||||
.write()
|
||||
.await
|
||||
.insert("test_nonce".to_string(), flow);
|
||||
|
||||
let state = test_gateway_state(Some(ext_mgr.clone()));
|
||||
let app = test_oauth_router(state);
|
||||
let versioned_state =
|
||||
crate::cli::oauth_defaults::encode_hosted_oauth_state("test_nonce", None);
|
||||
|
||||
let req = axum::http::Request::builder()
|
||||
.uri(format!(
|
||||
"/oauth/callback?code=fake_code&state={}",
|
||||
urlencoding::encode(&versioned_state)
|
||||
))
|
||||
.body(Body::empty())
|
||||
.expect("request");
|
||||
|
||||
let resp = ServiceExt::<axum::http::Request<Body>>::oneshot(app, req)
|
||||
.await
|
||||
.expect("response");
|
||||
assert_eq!(resp.status(), StatusCode::OK);
|
||||
|
||||
let body = axum::body::to_bytes(resp.into_body(), 1024 * 64)
|
||||
.await
|
||||
.expect("body");
|
||||
let html = String::from_utf8_lossy(&body);
|
||||
assert!(html.contains("Test Tool Connected"));
|
||||
|
||||
let requests = proxy.requests().await;
|
||||
assert_eq!(requests.len(), 1);
|
||||
assert_eq!(
|
||||
requests[0].authorization.as_deref(),
|
||||
Some("Bearer shared-oauth-proxy-secret")
|
||||
);
|
||||
assert_eq!(
|
||||
requests[0].form.get("code").map(String::as_str),
|
||||
Some("fake_code")
|
||||
);
|
||||
assert_eq!(
|
||||
requests[0].form.get("code_verifier").map(String::as_str),
|
||||
Some("test-code-verifier")
|
||||
);
|
||||
|
||||
let access_token = secrets
|
||||
.get_decrypted("test", "test_token")
|
||||
.await
|
||||
.expect("access token stored");
|
||||
assert_eq!(access_token.expose(), "proxy-access-token");
|
||||
|
||||
let refresh_token = secrets
|
||||
.get_decrypted("test", "test_token_refresh_token")
|
||||
.await
|
||||
.expect("refresh token stored");
|
||||
assert_eq!(refresh_token.expose(), "proxy-refresh-token");
|
||||
|
||||
match receiver.recv().await.expect("auth_completed event").event {
|
||||
crate::channels::web::types::AppEvent::AuthCompleted {
|
||||
extension_name,
|
||||
success,
|
||||
..
|
||||
} => {
|
||||
assert_eq!(extension_name, "test_tool");
|
||||
assert!(success, "OAuth callback should broadcast success");
|
||||
}
|
||||
event => panic!("expected AuthCompleted event, got {event:?}"),
|
||||
}
|
||||
|
||||
proxy.shutdown().await;
|
||||
}
|
||||
|
||||
// --- Slack relay OAuth CSRF tests ---
|
||||
|
||||
fn test_relay_oauth_router(state: Arc<GatewayState>) -> Router {
|
||||
@@ -3693,6 +4180,8 @@ mod tests {
|
||||
let ext_mgr = Arc::new(ExtensionManager::new(
|
||||
mcp_sm,
|
||||
mcp_pm,
|
||||
None,
|
||||
None,
|
||||
secrets,
|
||||
tool_registry,
|
||||
None,
|
||||
@@ -3702,6 +4191,7 @@ mod tests {
|
||||
None,
|
||||
"test".to_string(),
|
||||
None,
|
||||
None,
|
||||
vec![],
|
||||
));
|
||||
(ext_mgr, wasm_tools_dir, wasm_channels_dir)
|
||||
|
||||
@@ -1160,7 +1160,7 @@ function addToolCard(name) {
|
||||
|
||||
const toolName = document.createElement('span');
|
||||
toolName.className = 'activity-tool-name';
|
||||
toolName.textContent = name;
|
||||
toolName.textContent = humanizeToolName(name);
|
||||
|
||||
const duration = document.createElement('span');
|
||||
duration.className = 'activity-tool-duration';
|
||||
@@ -1344,7 +1344,7 @@ function finalizeActivityGroup() {
|
||||
|
||||
function humanizeToolName(rawName) {
|
||||
if (!rawName) return '';
|
||||
return String(rawName)
|
||||
return stripDerivedCompanionToolPrefix(String(rawName))
|
||||
.replace(/[_-]+/g, ' ')
|
||||
.replace(/([a-z0-9])([A-Z])/g, '$1 $2')
|
||||
.replace(/^tool([a-zA-Z])/, 'tool $1')
|
||||
@@ -1352,6 +1352,12 @@ function humanizeToolName(rawName) {
|
||||
.trim();
|
||||
}
|
||||
|
||||
function stripDerivedCompanionToolPrefix(rawName) {
|
||||
if (!rawName) return '';
|
||||
const prefix = '_nearai_companion_mcp_';
|
||||
return rawName.startsWith(prefix) ? rawName.slice(prefix.length) : rawName;
|
||||
}
|
||||
|
||||
function shouldShowChannelConnectedMessage(extensionName, success) {
|
||||
if (!success || !extensionName) return false;
|
||||
return String(extensionName).toLowerCase().includes('telegram');
|
||||
@@ -1871,7 +1877,7 @@ function createToolCallsSummaryElement(toolCalls) {
|
||||
const icon = tc.has_error ? '\u2717' : '\u2713';
|
||||
const nameSpan = document.createElement('span');
|
||||
nameSpan.className = 'tool-call-name';
|
||||
nameSpan.textContent = icon + ' ' + tc.name;
|
||||
nameSpan.textContent = icon + ' ' + humanizeToolName(tc.name);
|
||||
item.appendChild(nameSpan);
|
||||
|
||||
if (tc.result_preview) {
|
||||
@@ -2906,7 +2912,12 @@ function renderExtensionCard(ext) {
|
||||
if (ext.tools && ext.tools.length > 0) {
|
||||
const tools = document.createElement('div');
|
||||
tools.className = 'ext-tools';
|
||||
tools.textContent = 'Tools: ' + ext.tools.join(', ');
|
||||
const toolNames = ext.tools.map((toolName) => (
|
||||
ext.derived && ext.kind === 'mcp_server'
|
||||
? stripDerivedCompanionToolPrefix(toolName)
|
||||
: toolName
|
||||
));
|
||||
tools.textContent = 'Tools: ' + toolNames.join(', ');
|
||||
card.appendChild(tools);
|
||||
}
|
||||
|
||||
@@ -2967,7 +2978,7 @@ function renderExtensionCard(ext) {
|
||||
// Skip when has_auth is true but needs_setup is false and not yet authenticated —
|
||||
// this means OAuth credentials resolve automatically (builtin/env) and the user
|
||||
// just needs to complete the OAuth flow, not fill in a config form.
|
||||
if (ext.needs_setup || (ext.has_auth && ext.authenticated)) {
|
||||
if (!ext.derived && (ext.needs_setup || (ext.has_auth && ext.authenticated))) {
|
||||
const configBtn = document.createElement('button');
|
||||
configBtn.className = 'btn-ext configure';
|
||||
configBtn.textContent = ext.authenticated ? I18n.t('ext.reconfigure') : I18n.t('ext.configure');
|
||||
@@ -2976,11 +2987,13 @@ function renderExtensionCard(ext) {
|
||||
}
|
||||
}
|
||||
|
||||
const removeBtn = document.createElement('button');
|
||||
removeBtn.className = 'btn-ext remove';
|
||||
removeBtn.textContent = I18n.t('ext.remove');
|
||||
removeBtn.addEventListener('click', () => removeExtension(ext.name));
|
||||
actions.appendChild(removeBtn);
|
||||
if (!ext.derived) {
|
||||
const removeBtn = document.createElement('button');
|
||||
removeBtn.className = 'btn-ext remove';
|
||||
removeBtn.textContent = I18n.t('ext.remove');
|
||||
removeBtn.addEventListener('click', () => removeExtension(ext.name));
|
||||
actions.appendChild(removeBtn);
|
||||
}
|
||||
|
||||
card.appendChild(actions);
|
||||
|
||||
|
||||
@@ -344,6 +344,9 @@ pub struct ExtensionInfo {
|
||||
/// Whether this extension has an auth configuration (OAuth or manual token).
|
||||
#[serde(default)]
|
||||
pub has_auth: bool,
|
||||
/// Whether this extension is derived from runtime/provider state.
|
||||
#[serde(default)]
|
||||
pub derived: bool,
|
||||
/// WASM channel activation status.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub activation_status: Option<ExtensionActivationStatus>,
|
||||
|
||||
+325
-41
@@ -8,7 +8,7 @@ use std::sync::Arc;
|
||||
|
||||
use clap::{Args, Subcommand};
|
||||
|
||||
use crate::config::Config;
|
||||
use crate::config::{Config, LlmConfig};
|
||||
use crate::db::Database;
|
||||
use crate::secrets::SecretsStore;
|
||||
use crate::tools::mcp::{
|
||||
@@ -173,6 +173,13 @@ async fn add_server(args: McpAddArgs) -> anyhow::Result<()> {
|
||||
description,
|
||||
} = args;
|
||||
|
||||
if config::is_nearai_companion_server_name(&name) {
|
||||
anyhow::bail!(
|
||||
"Server name '{}' is reserved for the NEAR AI companion MCP server",
|
||||
name
|
||||
);
|
||||
}
|
||||
|
||||
let transport_lower = transport.to_lowercase();
|
||||
|
||||
let mut config = match transport_lower.as_str() {
|
||||
@@ -244,7 +251,7 @@ async fn add_server(args: McpAddArgs) -> anyhow::Result<()> {
|
||||
|
||||
// Save (DB if available, else disk)
|
||||
let db = connect_db().await;
|
||||
let mut servers = load_servers(db.as_deref()).await?;
|
||||
let mut servers = load_persisted_servers(db.as_deref()).await?;
|
||||
servers.upsert(config);
|
||||
save_servers(db.as_deref(), &servers).await?;
|
||||
|
||||
@@ -281,8 +288,15 @@ async fn add_server(args: McpAddArgs) -> anyhow::Result<()> {
|
||||
|
||||
/// Remove an MCP server.
|
||||
async fn remove_server(name: String) -> anyhow::Result<()> {
|
||||
if config::is_nearai_companion_server_name(&name) {
|
||||
anyhow::bail!(
|
||||
"Server '{}' is derived from the active NEAR AI provider and cannot be removed directly",
|
||||
name
|
||||
);
|
||||
}
|
||||
|
||||
let db = connect_db().await;
|
||||
let mut servers = load_servers(db.as_deref()).await?;
|
||||
let mut servers = load_persisted_servers(db.as_deref()).await?;
|
||||
if !servers.remove(&name) {
|
||||
anyhow::bail!("Server '{}' not found", name);
|
||||
}
|
||||
@@ -298,7 +312,7 @@ async fn remove_server(name: String) -> anyhow::Result<()> {
|
||||
/// List configured MCP servers.
|
||||
async fn list_servers(verbose: bool) -> anyhow::Result<()> {
|
||||
let db = connect_db().await;
|
||||
let servers = load_servers(db.as_deref()).await?;
|
||||
let servers = load_servers_with_derived(db.as_deref()).await?;
|
||||
|
||||
if servers.servers.is_empty() {
|
||||
println!();
|
||||
@@ -404,12 +418,23 @@ async fn list_servers(verbose: bool) -> anyhow::Result<()> {
|
||||
async fn auth_server(name: String, user_id: String) -> anyhow::Result<()> {
|
||||
// Get server config
|
||||
let db = connect_db().await;
|
||||
let servers = load_servers(db.as_deref()).await?;
|
||||
let servers = load_servers_with_derived(db.as_deref()).await?;
|
||||
let server = servers
|
||||
.get(&name)
|
||||
.cloned()
|
||||
.ok_or_else(|| anyhow::anyhow!("Server '{}' not found", name))?;
|
||||
|
||||
if server.uses_runtime_auth_source() {
|
||||
println!();
|
||||
println!(
|
||||
" Server '{}' reuses your active NEAR AI authentication and does not support separate MCP OAuth.",
|
||||
name
|
||||
);
|
||||
println!(" Configure NEAR AI auth (API key or session login) instead.");
|
||||
println!();
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
// Initialize secrets store
|
||||
let secrets = get_secrets_store().await?;
|
||||
|
||||
@@ -477,7 +502,7 @@ async fn auth_server(name: String, user_id: String) -> anyhow::Result<()> {
|
||||
async fn test_server(name: String, user_id: String) -> anyhow::Result<()> {
|
||||
// Get server config
|
||||
let db = connect_db().await;
|
||||
let servers = load_servers(db.as_deref()).await?;
|
||||
let servers = load_servers_with_derived(db.as_deref()).await?;
|
||||
let server = servers
|
||||
.get(&name)
|
||||
.cloned()
|
||||
@@ -488,35 +513,66 @@ async fn test_server(name: String, user_id: String) -> anyhow::Result<()> {
|
||||
|
||||
// Create client
|
||||
let session_manager = Arc::new(McpSessionManager::new());
|
||||
|
||||
// Always check for stored tokens (from either pre-configured OAuth or DCR)
|
||||
let secrets = get_secrets_store().await?;
|
||||
let has_tokens = is_authenticated(&server, &secrets, &user_id).await;
|
||||
|
||||
let client = if has_tokens {
|
||||
// We have stored tokens, use authenticated client
|
||||
McpClient::new_authenticated(server.clone(), session_manager.clone(), secrets, user_id)
|
||||
} else if server.requires_auth() {
|
||||
// OAuth configured but no tokens - need to authenticate
|
||||
println!();
|
||||
println!(
|
||||
" ✗ Not authenticated. Run 'ironclaw mcp auth {}' first.",
|
||||
name
|
||||
);
|
||||
println!();
|
||||
return Ok(());
|
||||
} else {
|
||||
// Use the factory to dispatch on transport type (HTTP, stdio, unix)
|
||||
let (client, has_tokens) = if server.uses_runtime_auth_source() {
|
||||
let process_manager = Arc::new(McpProcessManager::new());
|
||||
create_client_from_config(
|
||||
server.clone(),
|
||||
&session_manager,
|
||||
&process_manager,
|
||||
None,
|
||||
"default",
|
||||
let llm = resolve_llm_for_cli(as_settings_store(db.as_deref())).await?;
|
||||
let nearai_session = crate::llm::create_session_manager(llm.session.clone()).await;
|
||||
(
|
||||
create_client_from_config(
|
||||
server.clone(),
|
||||
&session_manager,
|
||||
Some(nearai_session),
|
||||
llm.nearai.api_key.clone(),
|
||||
&process_manager,
|
||||
None,
|
||||
"default",
|
||||
)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("{}", e))?,
|
||||
false,
|
||||
)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("{}", e))?
|
||||
} else {
|
||||
// Only initialize the secrets store for non-runtime-auth servers that
|
||||
// can actually use persisted OAuth/DCR tokens.
|
||||
let secrets = get_secrets_store().await?;
|
||||
let has_tokens = is_authenticated(&server, &secrets, &user_id).await;
|
||||
|
||||
if has_tokens {
|
||||
(
|
||||
McpClient::new_authenticated(
|
||||
server.clone(),
|
||||
session_manager.clone(),
|
||||
secrets,
|
||||
user_id,
|
||||
),
|
||||
true,
|
||||
)
|
||||
} else if server.requires_auth() {
|
||||
println!();
|
||||
println!(
|
||||
" ✗ Not authenticated. Run 'ironclaw mcp auth {}' first.",
|
||||
name
|
||||
);
|
||||
println!();
|
||||
return Ok(());
|
||||
} else {
|
||||
// Use the factory to dispatch on transport type (HTTP, stdio, unix)
|
||||
let process_manager = Arc::new(McpProcessManager::new());
|
||||
(
|
||||
create_client_from_config(
|
||||
server.clone(),
|
||||
&session_manager,
|
||||
None,
|
||||
None,
|
||||
&process_manager,
|
||||
None,
|
||||
"default",
|
||||
)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("{}", e))?,
|
||||
false,
|
||||
)
|
||||
}
|
||||
};
|
||||
|
||||
// Test connection
|
||||
@@ -581,8 +637,15 @@ async fn test_server(name: String, user_id: String) -> anyhow::Result<()> {
|
||||
|
||||
/// Toggle server enabled/disabled state.
|
||||
async fn toggle_server(name: String, enable: bool, disable: bool) -> anyhow::Result<()> {
|
||||
if config::is_nearai_companion_server_name(&name) {
|
||||
anyhow::bail!(
|
||||
"Server '{}' is derived from the active NEAR AI provider and cannot be toggled directly",
|
||||
name
|
||||
);
|
||||
}
|
||||
|
||||
let db = connect_db().await;
|
||||
let mut servers = load_servers(db.as_deref()).await?;
|
||||
let mut servers = load_persisted_servers(db.as_deref()).await?;
|
||||
|
||||
let server = servers
|
||||
.get_mut(&name)
|
||||
@@ -615,13 +678,30 @@ async fn connect_db() -> Option<Arc<dyn Database>> {
|
||||
crate::db::connect_from_config(&config.database).await.ok()
|
||||
}
|
||||
|
||||
/// Load MCP servers (DB if available, else disk).
|
||||
async fn load_servers(db: Option<&dyn Database>) -> Result<McpServersFile, config::ConfigError> {
|
||||
if let Some(db) = db {
|
||||
config::load_mcp_servers_from_db(db, DEFAULT_USER_ID).await
|
||||
/// Load only persisted MCP servers (DB if available, else disk).
|
||||
async fn load_persisted_servers(
|
||||
db: Option<&dyn Database>,
|
||||
) -> Result<McpServersFile, config::ConfigError> {
|
||||
Ok(if let Some(db) = db {
|
||||
config::load_mcp_servers_from_db(db, DEFAULT_USER_ID).await?
|
||||
} else {
|
||||
config::load_mcp_servers().await
|
||||
config::load_mcp_servers().await?
|
||||
})
|
||||
}
|
||||
|
||||
/// Load MCP servers plus any derived runtime companions.
|
||||
async fn load_servers_with_derived(
|
||||
db: Option<&dyn Database>,
|
||||
) -> Result<McpServersFile, config::ConfigError> {
|
||||
let mut servers = load_persisted_servers(db).await?;
|
||||
|
||||
if let Ok(llm) = resolve_llm_for_cli(as_settings_store(db)).await
|
||||
&& let Some(companion) = config::derive_nearai_companion_mcp_server_from_llm(&llm)
|
||||
{
|
||||
servers.insert_if_absent(companion);
|
||||
}
|
||||
|
||||
Ok(servers)
|
||||
}
|
||||
|
||||
/// Save MCP servers (DB if available, else disk).
|
||||
@@ -629,10 +709,15 @@ async fn save_servers(
|
||||
db: Option<&dyn Database>,
|
||||
servers: &McpServersFile,
|
||||
) -> Result<(), config::ConfigError> {
|
||||
let mut persisted = servers.clone();
|
||||
persisted
|
||||
.servers
|
||||
.retain(|server| !config::is_nearai_companion_server_name(&server.name));
|
||||
|
||||
if let Some(db) = db {
|
||||
config::save_mcp_servers_to_db(db, DEFAULT_USER_ID, servers).await
|
||||
config::save_mcp_servers_to_db(db, DEFAULT_USER_ID, &persisted).await
|
||||
} else {
|
||||
config::save_mcp_servers(servers).await
|
||||
config::save_mcp_servers(&persisted).await
|
||||
}
|
||||
}
|
||||
|
||||
@@ -641,10 +726,84 @@ async fn get_secrets_store() -> anyhow::Result<Arc<dyn SecretsStore + Send + Syn
|
||||
crate::cli::init_secrets_store().await
|
||||
}
|
||||
|
||||
fn as_settings_store(db: Option<&dyn Database>) -> Option<&(dyn crate::db::SettingsStore + Sync)> {
|
||||
db.map(|db| db as &(dyn crate::db::SettingsStore + Sync))
|
||||
}
|
||||
|
||||
async fn resolve_llm_for_cli(
|
||||
store: Option<&(dyn crate::db::SettingsStore + Sync)>,
|
||||
) -> Result<LlmConfig, crate::error::ConfigError> {
|
||||
resolve_llm_for_cli_with_toml(store, None).await
|
||||
}
|
||||
|
||||
async fn resolve_llm_for_cli_with_toml(
|
||||
store: Option<&(dyn crate::db::SettingsStore + Sync)>,
|
||||
toml_path: Option<&std::path::Path>,
|
||||
) -> Result<LlmConfig, crate::error::ConfigError> {
|
||||
if let Some(store) = store {
|
||||
let _ = dotenvy::dotenv();
|
||||
crate::bootstrap::load_ironclaw_env();
|
||||
|
||||
let mut settings = match store.get_all_settings(DEFAULT_USER_ID).await {
|
||||
Ok(map) => crate::settings::Settings::from_db_map(&map),
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
"Failed to load CLI settings from DB, falling back to defaults before env/TOML resolution: {}",
|
||||
e
|
||||
);
|
||||
crate::settings::Settings::default()
|
||||
}
|
||||
};
|
||||
|
||||
apply_cli_toml_overlay(&mut settings, toml_path)?;
|
||||
return LlmConfig::resolve(&settings);
|
||||
}
|
||||
|
||||
let settings = crate::config::load_bootstrap_settings(toml_path)?;
|
||||
LlmConfig::resolve(&settings)
|
||||
}
|
||||
|
||||
fn apply_cli_toml_overlay(
|
||||
settings: &mut crate::settings::Settings,
|
||||
explicit_path: Option<&std::path::Path>,
|
||||
) -> Result<(), crate::error::ConfigError> {
|
||||
let path = explicit_path
|
||||
.map(std::path::PathBuf::from)
|
||||
.unwrap_or_else(crate::settings::Settings::default_toml_path);
|
||||
|
||||
match crate::settings::Settings::load_toml(&path) {
|
||||
Ok(Some(toml_settings)) => {
|
||||
settings.merge_from(&toml_settings);
|
||||
}
|
||||
Ok(None) => {
|
||||
if explicit_path.is_some() {
|
||||
return Err(crate::error::ConfigError::ParseError(format!(
|
||||
"Config file not found: {}",
|
||||
path.display()
|
||||
)));
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
return Err(crate::error::ConfigError::ParseError(e));
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
use std::collections::HashMap;
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use crate::error::DatabaseError;
|
||||
use crate::history::SettingRow;
|
||||
#[cfg(feature = "libsql")]
|
||||
use tempfile::NamedTempFile;
|
||||
|
||||
#[test]
|
||||
fn test_mcp_command_parsing() {
|
||||
// Just verify the command structure is valid
|
||||
@@ -701,4 +860,129 @@ mod tests {
|
||||
assert!(result.is_err());
|
||||
assert!(result.unwrap_err().contains("invalid env var format"));
|
||||
}
|
||||
|
||||
#[cfg(feature = "libsql")]
|
||||
#[allow(clippy::await_holding_lock)]
|
||||
#[tokio::test]
|
||||
async fn test_resolve_llm_for_cli_uses_db_backed_selected_model() {
|
||||
struct MockSettingsStore {
|
||||
settings: HashMap<String, serde_json::Value>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl crate::db::SettingsStore for MockSettingsStore {
|
||||
async fn get_setting(
|
||||
&self,
|
||||
_user_id: &str,
|
||||
key: &str,
|
||||
) -> Result<Option<serde_json::Value>, DatabaseError> {
|
||||
Ok(self.settings.get(key).cloned())
|
||||
}
|
||||
|
||||
async fn get_setting_full(
|
||||
&self,
|
||||
_user_id: &str,
|
||||
_key: &str,
|
||||
) -> Result<Option<SettingRow>, DatabaseError> {
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn set_setting(
|
||||
&self,
|
||||
_user_id: &str,
|
||||
_key: &str,
|
||||
_value: &serde_json::Value,
|
||||
) -> Result<(), DatabaseError> {
|
||||
Err(DatabaseError::Query("unused in test".to_string()))
|
||||
}
|
||||
|
||||
async fn delete_setting(
|
||||
&self,
|
||||
_user_id: &str,
|
||||
_key: &str,
|
||||
) -> Result<bool, DatabaseError> {
|
||||
Err(DatabaseError::Query("unused in test".to_string()))
|
||||
}
|
||||
|
||||
async fn list_settings(
|
||||
&self,
|
||||
_user_id: &str,
|
||||
) -> Result<Vec<SettingRow>, DatabaseError> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
|
||||
async fn get_all_settings(
|
||||
&self,
|
||||
_user_id: &str,
|
||||
) -> Result<HashMap<String, serde_json::Value>, DatabaseError> {
|
||||
Ok(self.settings.clone())
|
||||
}
|
||||
|
||||
async fn set_all_settings(
|
||||
&self,
|
||||
_user_id: &str,
|
||||
_settings: &HashMap<String, serde_json::Value>,
|
||||
) -> Result<(), DatabaseError> {
|
||||
Err(DatabaseError::Query("unused in test".to_string()))
|
||||
}
|
||||
|
||||
async fn has_settings(&self, _user_id: &str) -> Result<bool, DatabaseError> {
|
||||
Ok(!self.settings.is_empty())
|
||||
}
|
||||
}
|
||||
|
||||
struct EnvGuard(&'static str, Option<String>);
|
||||
|
||||
impl Drop for EnvGuard {
|
||||
fn drop(&mut self) {
|
||||
// SAFETY: Protected by ENV_MUTEX for the duration of the test.
|
||||
unsafe {
|
||||
match &self.1 {
|
||||
Some(value) => std::env::set_var(self.0, value),
|
||||
None => std::env::remove_var(self.0),
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let _mutex = crate::config::helpers::ENV_MUTEX.lock().expect("env mutex");
|
||||
let prev_backend = std::env::var("LLM_BACKEND").ok();
|
||||
let prev_base_url = std::env::var("NEARAI_BASE_URL").ok();
|
||||
let prev_auth_url = std::env::var("NEARAI_AUTH_URL").ok();
|
||||
let prev_model = std::env::var("NEARAI_MODEL").ok();
|
||||
|
||||
// SAFETY: Protected by ENV_MUTEX for the duration of the test.
|
||||
unsafe {
|
||||
std::env::set_var("LLM_BACKEND", "");
|
||||
std::env::set_var("NEARAI_BASE_URL", "http://127.0.0.1:11434/v1");
|
||||
std::env::set_var("NEARAI_AUTH_URL", "http://127.0.0.1:11435");
|
||||
std::env::set_var("NEARAI_MODEL", "");
|
||||
}
|
||||
|
||||
let _backend_guard = EnvGuard("LLM_BACKEND", prev_backend);
|
||||
let _base_url_guard = EnvGuard("NEARAI_BASE_URL", prev_base_url);
|
||||
let _auth_url_guard = EnvGuard("NEARAI_AUTH_URL", prev_auth_url);
|
||||
let _model_guard = EnvGuard("NEARAI_MODEL", prev_model);
|
||||
|
||||
let empty_toml = NamedTempFile::new().expect("temp toml");
|
||||
let store = MockSettingsStore {
|
||||
settings: HashMap::from([
|
||||
("llm_backend".to_string(), serde_json::json!("nearai")),
|
||||
(
|
||||
"selected_model".to_string(),
|
||||
serde_json::json!("db-backed-nearai-model"),
|
||||
),
|
||||
]),
|
||||
};
|
||||
|
||||
let llm = resolve_llm_for_cli_with_toml(Some(&store), Some(empty_toml.path()))
|
||||
.await
|
||||
.expect("resolve llm");
|
||||
assert_eq!(llm.backend, "nearai");
|
||||
assert_eq!(llm.nearai.model, "db-backed-nearai-model");
|
||||
|
||||
let companion =
|
||||
config::derive_nearai_companion_mcp_server_from_llm(&llm).expect("derived companion");
|
||||
assert_eq!(companion.url, "http://127.0.0.1:11434/mcp");
|
||||
}
|
||||
}
|
||||
|
||||
+184
-5
@@ -473,7 +473,8 @@ pub struct PendingOAuthFlow {
|
||||
pub secrets: Arc<dyn SecretsStore + Send + Sync>,
|
||||
/// SSE broadcast manager for notifying the web UI.
|
||||
pub sse_manager: Option<Arc<crate::channels::web::sse::SseManager>>,
|
||||
/// Gateway auth token for authenticating with the platform token exchange proxy.
|
||||
/// OAuth proxy auth token for authenticating with the hosted token exchange proxy.
|
||||
/// Kept as `gateway_token` for public API compatibility.
|
||||
pub gateway_token: Option<String>,
|
||||
/// Additional form params for the token exchange request.
|
||||
/// Used for provider-specific requirements such as RFC 8707 `resource`.
|
||||
@@ -496,6 +497,12 @@ impl std::fmt::Debug for PendingOAuthFlow {
|
||||
}
|
||||
}
|
||||
|
||||
impl PendingOAuthFlow {
|
||||
pub fn oauth_proxy_auth_token(&self) -> Option<&str> {
|
||||
self.gateway_token.as_deref()
|
||||
}
|
||||
}
|
||||
|
||||
/// Thread-safe registry of pending OAuth flows, keyed by CSRF `state` parameter.
|
||||
pub type PendingOAuthRegistry = Arc<RwLock<HashMap<String, PendingOAuthFlow>>>;
|
||||
|
||||
@@ -529,6 +536,22 @@ pub fn exchange_proxy_url() -> Option<String> {
|
||||
.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).
|
||||
pub const OAUTH_FLOW_EXPIRY: Duration = Duration::from_secs(300);
|
||||
|
||||
@@ -674,6 +697,8 @@ pub fn strip_instance_prefix(state: &str) -> &str {
|
||||
|
||||
pub struct ProxyTokenExchangeRequest<'a> {
|
||||
pub proxy_url: &'a str,
|
||||
/// OAuth proxy auth token.
|
||||
/// Kept as `gateway_token` for public API compatibility.
|
||||
pub gateway_token: &'a str,
|
||||
pub token_url: &'a str,
|
||||
pub client_id: &'a str,
|
||||
@@ -687,6 +712,8 @@ pub struct ProxyTokenExchangeRequest<'a> {
|
||||
|
||||
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,
|
||||
@@ -729,7 +756,7 @@ fn oauth_token_response_from_json(
|
||||
|
||||
/// Exchange an OAuth authorization code via the platform's token exchange proxy.
|
||||
///
|
||||
/// Authenticated via the gateway auth token (Bearer header). The caller may
|
||||
/// Authenticated via an OAuth proxy auth token (Bearer header). The caller may
|
||||
/// either rely on proxy-side secret lookup or forward a `client_secret` when
|
||||
/// the provider requires it.
|
||||
///
|
||||
@@ -741,7 +768,7 @@ pub async fn exchange_via_proxy(
|
||||
) -> Result<OAuthTokenResponse, OAuthCallbackError> {
|
||||
if request.gateway_token.is_empty() {
|
||||
return Err(OAuthCallbackError::Io(
|
||||
"Gateway auth token is required for proxy token exchange".to_string(),
|
||||
"OAuth proxy auth token is required for proxy token exchange".to_string(),
|
||||
));
|
||||
}
|
||||
let exchange_url = format!("{}/oauth/exchange", request.proxy_url.trim_end_matches('/'));
|
||||
@@ -796,7 +823,7 @@ pub async fn exchange_via_proxy(
|
||||
|
||||
/// Refresh an OAuth access token via the platform's token refresh proxy.
|
||||
///
|
||||
/// Authenticated via the gateway auth token (Bearer header). The caller may
|
||||
/// Authenticated via an OAuth proxy auth token (Bearer header). The caller may
|
||||
/// either rely on proxy-side secret lookup or forward a `client_secret` when
|
||||
/// the provider requires it.
|
||||
pub async fn refresh_token_via_proxy(
|
||||
@@ -804,7 +831,7 @@ pub async fn refresh_token_via_proxy(
|
||||
) -> Result<OAuthTokenResponse, OAuthCallbackError> {
|
||||
if request.gateway_token.is_empty() {
|
||||
return Err(OAuthCallbackError::Io(
|
||||
"Gateway auth token is required for proxy token refresh".to_string(),
|
||||
"OAuth proxy auth token is required for proxy token refresh".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
@@ -1010,6 +1037,37 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
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");
|
||||
@@ -1030,6 +1088,79 @@ mod tests {
|
||||
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;
|
||||
@@ -1535,6 +1666,54 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_oauth_proxy_auth_token_prefers_dedicated_env() {
|
||||
let _guard = lock_env();
|
||||
let _proxy_guard = set_env_var(
|
||||
"IRONCLAW_OAUTH_PROXY_AUTH_TOKEN",
|
||||
Some("shared-proxy-secret"),
|
||||
);
|
||||
let _gateway_guard = set_env_var("GATEWAY_AUTH_TOKEN", Some("gateway-token"));
|
||||
|
||||
assert_eq!(
|
||||
crate::cli::oauth_defaults::oauth_proxy_auth_token().as_deref(),
|
||||
Some("shared-proxy-secret")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_oauth_proxy_auth_token_falls_back_to_gateway_token() {
|
||||
let _guard = lock_env();
|
||||
let _proxy_guard = set_env_var("IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", None);
|
||||
let _gateway_guard = set_env_var("GATEWAY_AUTH_TOKEN", Some("gateway-token"));
|
||||
|
||||
assert_eq!(
|
||||
crate::cli::oauth_defaults::oauth_proxy_auth_token().as_deref(),
|
||||
Some("gateway-token")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_oauth_proxy_auth_token_whitespace_dedicated_env_falls_back_to_gateway_token() {
|
||||
let _guard = lock_env();
|
||||
let _proxy_guard = set_env_var("IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", Some(" "));
|
||||
let _gateway_guard = set_env_var("GATEWAY_AUTH_TOKEN", Some("gateway-token"));
|
||||
|
||||
assert_eq!(
|
||||
crate::cli::oauth_defaults::oauth_proxy_auth_token().as_deref(),
|
||||
Some("gateway-token")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_oauth_proxy_auth_token_returns_none_when_unset() {
|
||||
let _guard = lock_env();
|
||||
let _proxy_guard = set_env_var("IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", None);
|
||||
let _gateway_guard = set_env_var("GATEWAY_AUTH_TOKEN", None);
|
||||
|
||||
assert_eq!(crate::cli::oauth_defaults::oauth_proxy_auth_token(), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_strip_instance_prefix_with_colon() {
|
||||
use crate::cli::oauth_defaults::strip_instance_prefix;
|
||||
|
||||
+20
-1
@@ -1,6 +1,6 @@
|
||||
use std::time::Duration;
|
||||
|
||||
use crate::config::helpers::{parse_bool_env, parse_option_env, parse_optional_env};
|
||||
use crate::config::helpers::{optional_env, parse_bool_env, parse_option_env, parse_optional_env};
|
||||
use crate::error::ConfigError;
|
||||
use crate::settings::Settings;
|
||||
|
||||
@@ -23,6 +23,8 @@ pub struct AgentConfig {
|
||||
pub max_cost_per_day_cents: Option<u64>,
|
||||
/// Maximum LLM/tool actions per hour. None = unlimited.
|
||||
pub max_actions_per_hour: Option<u64>,
|
||||
/// Maximum daily LLM spend per user in cents. None = unlimited.
|
||||
pub max_cost_per_user_per_day_cents: Option<u64>,
|
||||
/// Maximum tool-call iterations per agentic loop invocation. Default 50.
|
||||
pub max_tool_iterations: usize,
|
||||
/// When true, skip tool approval checks entirely. For benchmarks/CI.
|
||||
@@ -31,6 +33,13 @@ pub struct AgentConfig {
|
||||
pub default_timezone: String,
|
||||
/// Maximum tokens per job (0 = unlimited).
|
||||
pub max_tokens_per_job: u64,
|
||||
/// Whether the deployment is multi-tenant (multiple users sharing one
|
||||
/// instance). Auto-detected from GATEWAY_USER_TOKENS presence.
|
||||
pub multi_tenant: bool,
|
||||
/// Maximum concurrent LLM calls per user. None = use default (4).
|
||||
pub max_llm_concurrent_per_user: Option<usize>,
|
||||
/// Maximum concurrent jobs per user. None = use default (3).
|
||||
pub max_jobs_concurrent_per_user: Option<usize>,
|
||||
}
|
||||
|
||||
impl AgentConfig {
|
||||
@@ -49,10 +58,14 @@ impl AgentConfig {
|
||||
allow_local_tools: true,
|
||||
max_cost_per_day_cents: None,
|
||||
max_actions_per_hour: None,
|
||||
max_cost_per_user_per_day_cents: None,
|
||||
max_tool_iterations: 10,
|
||||
auto_approve_tools: true,
|
||||
default_timezone: "UTC".to_string(),
|
||||
max_tokens_per_job: 0,
|
||||
multi_tenant: false,
|
||||
max_llm_concurrent_per_user: None,
|
||||
max_jobs_concurrent_per_user: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -87,6 +100,7 @@ impl AgentConfig {
|
||||
allow_local_tools: parse_bool_env("ALLOW_LOCAL_TOOLS", false)?,
|
||||
max_cost_per_day_cents: parse_option_env("MAX_COST_PER_DAY_CENTS")?,
|
||||
max_actions_per_hour: parse_option_env("MAX_ACTIONS_PER_HOUR")?,
|
||||
max_cost_per_user_per_day_cents: parse_option_env("MAX_COST_PER_USER_PER_DAY_CENTS")?,
|
||||
max_tool_iterations: parse_optional_env(
|
||||
"AGENT_MAX_TOOL_ITERATIONS",
|
||||
settings.agent.max_tool_iterations,
|
||||
@@ -112,6 +126,11 @@ impl AgentConfig {
|
||||
"AGENT_MAX_TOKENS_PER_JOB",
|
||||
settings.agent.max_tokens_per_job,
|
||||
)?,
|
||||
// Auto-detected from GATEWAY_USER_TOKENS presence. Not a separate
|
||||
// knob — multi-tenant mode is always implied by configuring user tokens.
|
||||
multi_tenant: optional_env("GATEWAY_USER_TOKENS")?.is_some(),
|
||||
max_llm_concurrent_per_user: parse_option_env("TENANT_MAX_LLM_CONCURRENT")?,
|
||||
max_jobs_concurrent_per_user: parse_option_env("TENANT_MAX_JOBS_CONCURRENT")?,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -21,6 +21,9 @@ pub struct HeartbeatConfig {
|
||||
pub quiet_hours_end: Option<u32>,
|
||||
/// Timezone for fire_at and quiet hours evaluation (IANA name).
|
||||
pub timezone: Option<String>,
|
||||
/// When true, cycle through all users with routines. Auto-detected from
|
||||
/// GATEWAY_USER_TOKENS or set explicitly via HEARTBEAT_MULTI_TENANT.
|
||||
pub multi_tenant: bool,
|
||||
}
|
||||
|
||||
impl Default for HeartbeatConfig {
|
||||
@@ -34,6 +37,7 @@ impl Default for HeartbeatConfig {
|
||||
quiet_hours_start: None,
|
||||
quiet_hours_end: None,
|
||||
timezone: None,
|
||||
multi_tenant: false,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -101,6 +105,12 @@ impl HeartbeatConfig {
|
||||
}
|
||||
tz
|
||||
},
|
||||
// Auto-detect multi-tenant mode from GATEWAY_USER_TOKENS presence,
|
||||
// or allow explicit override via HEARTBEAT_MULTI_TENANT.
|
||||
multi_tenant: parse_bool_env(
|
||||
"HEARTBEAT_MULTI_TENANT",
|
||||
optional_env("GATEWAY_USER_TOKENS")?.is_some(),
|
||||
)?,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -192,6 +192,9 @@ pub struct JobContext {
|
||||
/// but subsequent tools (e.g., `json`) may need the full output. This
|
||||
/// stash stores the complete, unsanitized output so tools can reference
|
||||
/// previous results by ID via `$tool_call_id` parameter syntax.
|
||||
///
|
||||
/// Also used for cross-tool implicit state (keys prefixed with `__`) such
|
||||
/// as `__routine_last_name` for fallback recovery in routine tool chains.
|
||||
#[serde(skip)]
|
||||
pub tool_output_stash: Arc<tokio::sync::RwLock<HashMap<String, String>>>,
|
||||
/// User's preferred timezone (IANA name, e.g. "America/New_York"). Defaults to "UTC".
|
||||
|
||||
@@ -530,10 +530,24 @@ impl RoutineStore for LibSqlBackend {
|
||||
async fn get_webhook_routine_by_path(
|
||||
&self,
|
||||
path: &str,
|
||||
user_id: Option<&str>,
|
||||
) -> Result<Option<Routine>, DatabaseError> {
|
||||
let conn = self.connect().await?;
|
||||
let mut rows = conn
|
||||
.query(
|
||||
let mut rows = if let Some(uid) = user_id {
|
||||
conn.query(
|
||||
&format!(
|
||||
"SELECT {} FROM routines WHERE enabled = 1 AND trigger_type = 'webhook' \
|
||||
AND user_id = ?2 \
|
||||
AND (json_extract(trigger_config, '$.path') = ?1 \
|
||||
OR (json_extract(trigger_config, '$.path') IS NULL AND CAST(id AS TEXT) = ?1))",
|
||||
ROUTINE_COLUMNS
|
||||
),
|
||||
params![path, uid],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?
|
||||
} else {
|
||||
conn.query(
|
||||
&format!(
|
||||
"SELECT {} FROM routines WHERE enabled = 1 AND trigger_type = 'webhook' \
|
||||
AND (json_extract(trigger_config, '$.path') = ?1 \
|
||||
@@ -543,7 +557,8 @@ impl RoutineStore for LibSqlBackend {
|
||||
params![path],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?;
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?
|
||||
};
|
||||
|
||||
match rows
|
||||
.next()
|
||||
|
||||
@@ -545,6 +545,7 @@ pub trait RoutineStore: Send + Sync {
|
||||
async fn get_webhook_routine_by_path(
|
||||
&self,
|
||||
path: &str,
|
||||
user_id: Option<&str>,
|
||||
) -> Result<Option<Routine>, DatabaseError>;
|
||||
|
||||
/// List routine runs that were dispatched as full_job but have not yet
|
||||
|
||||
+2
-1
@@ -529,8 +529,9 @@ impl RoutineStore for PgBackend {
|
||||
async fn get_webhook_routine_by_path(
|
||||
&self,
|
||||
path: &str,
|
||||
user_id: Option<&str>,
|
||||
) -> Result<Option<Routine>, DatabaseError> {
|
||||
self.store.get_webhook_routine_by_path(path).await
|
||||
self.store.get_webhook_routine_by_path(path, user_id).await
|
||||
}
|
||||
|
||||
async fn list_dispatched_routine_runs(&self) -> Result<Vec<RoutineRun>, DatabaseError> {
|
||||
|
||||
+868
-54
File diff suppressed because it is too large
Load Diff
@@ -506,6 +506,10 @@ pub struct InstalledExtension {
|
||||
/// Whether this extension has an auth configuration (OAuth or manual token).
|
||||
#[serde(default)]
|
||||
pub has_auth: bool,
|
||||
/// Whether this extension is derived from provider/runtime state instead of
|
||||
/// being a user-managed persisted configuration.
|
||||
#[serde(default)]
|
||||
pub derived: bool,
|
||||
/// Whether this extension is installed locally (false = available in registry but not installed).
|
||||
#[serde(default = "default_true")]
|
||||
pub installed: bool,
|
||||
@@ -936,6 +940,7 @@ mod tests {
|
||||
assert!(ext.installed, "installed should default to true");
|
||||
assert!(!ext.needs_setup, "needs_setup should default to false");
|
||||
assert!(!ext.has_auth);
|
||||
assert!(!ext.derived);
|
||||
assert!(ext.tools.is_empty());
|
||||
assert!(ext.display_name.is_none());
|
||||
assert!(ext.description.is_none());
|
||||
@@ -956,6 +961,7 @@ mod tests {
|
||||
tools: vec!["send_email".to_string(), "read_inbox".to_string()],
|
||||
needs_setup: true,
|
||||
has_auth: true,
|
||||
derived: true,
|
||||
installed: false,
|
||||
activation_error: Some("token expired".to_string()),
|
||||
version: None,
|
||||
@@ -965,6 +971,7 @@ mod tests {
|
||||
assert_eq!(json["description"], "Read and send emails");
|
||||
assert_eq!(json["url"], "https://gmail.example.com");
|
||||
assert_eq!(json["needs_setup"], true);
|
||||
assert_eq!(json["derived"], true);
|
||||
assert_eq!(json["installed"], false);
|
||||
assert_eq!(json["activation_error"], "token expired");
|
||||
|
||||
@@ -972,6 +979,7 @@ mod tests {
|
||||
assert_eq!(back.name, "gmail");
|
||||
assert_eq!(back.tools.len(), 2);
|
||||
assert!(back.needs_setup);
|
||||
assert!(back.derived);
|
||||
assert!(!back.installed);
|
||||
assert_eq!(back.activation_error.as_deref(), Some("token expired"));
|
||||
}
|
||||
|
||||
+13
-3
@@ -1162,15 +1162,25 @@ impl Store {
|
||||
pub async fn get_webhook_routine_by_path(
|
||||
&self,
|
||||
path: &str,
|
||||
user_id: Option<&str>,
|
||||
) -> Result<Option<Routine>, DatabaseError> {
|
||||
let conn = self.conn().await?;
|
||||
let row = conn
|
||||
.query_opt(
|
||||
let row = if let Some(uid) = user_id {
|
||||
conn.query_opt(
|
||||
"SELECT * FROM routines WHERE enabled AND trigger_type = 'webhook' \
|
||||
AND user_id = $2 \
|
||||
AND (trigger_config->>'path' = $1 OR (trigger_config->>'path' IS NULL AND id::text = $1))",
|
||||
&[&path, &uid],
|
||||
)
|
||||
.await?
|
||||
} else {
|
||||
conn.query_opt(
|
||||
"SELECT * FROM routines WHERE enabled AND trigger_type = 'webhook' \
|
||||
AND (trigger_config->>'path' = $1 OR (trigger_config->>'path' IS NULL AND id::text = $1))",
|
||||
&[&path],
|
||||
)
|
||||
.await?;
|
||||
.await?
|
||||
};
|
||||
row.as_ref().map(row_to_routine).transpose()
|
||||
}
|
||||
|
||||
|
||||
@@ -69,6 +69,7 @@ pub mod service;
|
||||
pub mod settings;
|
||||
pub mod setup;
|
||||
pub mod skills;
|
||||
pub mod tenant;
|
||||
pub mod timezone;
|
||||
pub mod tools;
|
||||
pub mod tracing_fmt;
|
||||
|
||||
+4
-1
@@ -21,6 +21,7 @@ pub mod failover;
|
||||
pub mod gemini_oauth;
|
||||
mod github_copilot;
|
||||
pub(crate) mod github_copilot_auth;
|
||||
pub mod nearai_auth;
|
||||
mod nearai_chat;
|
||||
pub mod oauth_helpers;
|
||||
pub mod openai_codex_provider;
|
||||
@@ -53,6 +54,7 @@ pub use config::{
|
||||
pub use error::LlmError;
|
||||
pub use failover::{CooldownConfig, FailoverProvider};
|
||||
pub use gemini_oauth::GeminiOauthProvider;
|
||||
pub use nearai_auth::{resolve_nearai_bearer_token, resolve_nearai_bearer_token_if_available};
|
||||
pub use nearai_chat::{DEFAULT_MODEL, ModelInfo, NearAiChatProvider, default_models};
|
||||
pub use openai_codex_provider::OpenAiCodexProvider;
|
||||
pub use openai_codex_session::{OpenAiCodexSession, OpenAiCodexSessionManager};
|
||||
@@ -63,7 +65,8 @@ pub use provider::{
|
||||
};
|
||||
pub use reasoning::{
|
||||
ActionPlan, Reasoning, ReasoningContext, RespondOutput, RespondResult, SILENT_REPLY_TOKEN,
|
||||
TOOL_INTENT_NUDGE, TokenUsage, ToolSelection, is_silent_reply, llm_signals_tool_intent,
|
||||
TOOL_INTENT_NUDGE, TRUNCATED_TOOL_CALL_NOTICE, TokenUsage, ToolSelection, is_silent_reply,
|
||||
llm_signals_tool_intent,
|
||||
};
|
||||
pub use recording::RecordingLlm;
|
||||
pub use registry::{ProviderDefinition, ProviderProtocol, ProviderRegistry};
|
||||
|
||||
@@ -0,0 +1,95 @@
|
||||
use secrecy::{ExposeSecret, SecretString};
|
||||
|
||||
use crate::llm::LlmError;
|
||||
use crate::llm::session::SessionManager;
|
||||
|
||||
/// Resolve the active NEAR AI bearer token only if already available.
|
||||
///
|
||||
/// Unlike [`resolve_nearai_bearer_token`], this helper is side-effect free:
|
||||
/// it never triggers an interactive login flow.
|
||||
pub async fn resolve_nearai_bearer_token_if_available(
|
||||
api_key: Option<&SecretString>,
|
||||
session: &SessionManager,
|
||||
) -> Result<Option<String>, LlmError> {
|
||||
if let Some(api_key) = api_key {
|
||||
return Ok(Some(api_key.expose_secret().to_string()));
|
||||
}
|
||||
|
||||
if session.has_token().await {
|
||||
let token = session.get_token().await?;
|
||||
return Ok(Some(token.expose_secret().to_string()));
|
||||
}
|
||||
|
||||
if let Some(key) = crate::config::helpers::env_or_override("NEARAI_API_KEY") {
|
||||
return Ok(Some(key));
|
||||
}
|
||||
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
/// Resolve the active NEAR AI bearer token.
|
||||
///
|
||||
/// Priority order:
|
||||
/// 1. Explicit API key from resolved config
|
||||
/// 2. Existing session token
|
||||
/// 3. Interactive session authentication
|
||||
/// 4. `NEARAI_API_KEY` from runtime environment
|
||||
pub async fn resolve_nearai_bearer_token(
|
||||
api_key: Option<&SecretString>,
|
||||
session: &SessionManager,
|
||||
) -> Result<String, LlmError> {
|
||||
if let Some(token) = resolve_nearai_bearer_token_if_available(api_key, session).await? {
|
||||
return Ok(token);
|
||||
}
|
||||
|
||||
session.ensure_authenticated().await?;
|
||||
|
||||
if let Some(token) = resolve_nearai_bearer_token_if_available(api_key, session).await? {
|
||||
return Ok(token);
|
||||
}
|
||||
|
||||
Err(LlmError::AuthFailed {
|
||||
provider: "nearai".to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::config::helpers::{ENV_MUTEX, set_runtime_env};
|
||||
use crate::llm::session::SessionConfig;
|
||||
|
||||
struct EnvGuard(&'static str, Option<String>);
|
||||
|
||||
impl Drop for EnvGuard {
|
||||
fn drop(&mut self) {
|
||||
// SAFETY: tests hold ENV_MUTEX while mutating the process environment.
|
||||
unsafe {
|
||||
match &self.1 {
|
||||
Some(value) => std::env::set_var(self.0, value),
|
||||
None => std::env::remove_var(self.0),
|
||||
}
|
||||
}
|
||||
set_runtime_env(self.0, "");
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(clippy::await_holding_lock)]
|
||||
#[tokio::test]
|
||||
async fn test_resolve_bearer_token_if_available_uses_runtime_env_override() {
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex");
|
||||
let prev = std::env::var("NEARAI_API_KEY").ok();
|
||||
// SAFETY: tests hold ENV_MUTEX while mutating the process environment.
|
||||
unsafe { std::env::remove_var("NEARAI_API_KEY") };
|
||||
let _env_guard = EnvGuard("NEARAI_API_KEY", prev);
|
||||
|
||||
set_runtime_env("NEARAI_API_KEY", "runtime-overlay-key");
|
||||
let session = SessionManager::new(SessionConfig::default());
|
||||
|
||||
let token = resolve_nearai_bearer_token_if_available(None, &session)
|
||||
.await
|
||||
.expect("resolve token");
|
||||
|
||||
assert_eq!(token.as_deref(), Some("runtime-overlay-key"));
|
||||
}
|
||||
}
|
||||
+1
-30
@@ -173,36 +173,7 @@ impl NearAiChatProvider {
|
||||
/// The env var fallback (#3) only triggers after `ensure_authenticated()`
|
||||
/// runs, because `api_key_login()` sets the env var but not a session token.
|
||||
async fn resolve_bearer_token(&self) -> Result<String, LlmError> {
|
||||
// 1. Config-level API key takes priority
|
||||
if let Some(ref api_key) = self.config.api_key {
|
||||
return Ok(api_key.expose_secret().to_string());
|
||||
}
|
||||
|
||||
// 2. Existing session token (OAuth was already completed)
|
||||
if self.session.has_token().await {
|
||||
let token = self.session.get_token().await?;
|
||||
return Ok(token.expose_secret().to_string());
|
||||
}
|
||||
|
||||
// No token yet, trigger interactive login
|
||||
self.session.ensure_authenticated().await?;
|
||||
|
||||
// 3. After login, check if a session token was stored (OAuth path)
|
||||
if self.session.has_token().await {
|
||||
let token = self.session.get_token().await?;
|
||||
return Ok(token.expose_secret().to_string());
|
||||
}
|
||||
|
||||
// 4. api_key_login() sets NEARAI_API_KEY env var but not a session token
|
||||
if let Ok(key) = std::env::var("NEARAI_API_KEY")
|
||||
&& !key.is_empty()
|
||||
{
|
||||
return Ok(key);
|
||||
}
|
||||
|
||||
Err(LlmError::AuthFailed {
|
||||
provider: "nearai".to_string(),
|
||||
})
|
||||
crate::llm::resolve_nearai_bearer_token(self.config.api_key.as_ref(), &self.session).await
|
||||
}
|
||||
|
||||
/// Send a single request to the chat completions API.
|
||||
|
||||
+211
-8
@@ -8,8 +8,8 @@ use serde::{Deserialize, Serialize};
|
||||
use crate::llm::error::LlmError;
|
||||
|
||||
use crate::llm::{
|
||||
ChatMessage, CompletionRequest, LlmProvider, Role, ToolCall, ToolCompletionRequest,
|
||||
ToolDefinition,
|
||||
ChatMessage, CompletionRequest, FinishReason, LlmProvider, Role, ToolCall,
|
||||
ToolCompletionRequest, ToolDefinition,
|
||||
};
|
||||
|
||||
/// Token the agent returns when it has nothing to say (e.g. in group chats).
|
||||
@@ -23,6 +23,13 @@ You said you would perform an action, but you did not include any tool calls.\n\
|
||||
Do NOT describe what you intend to do — actually call the tool now.\n\
|
||||
Use the tool_calls mechanism to invoke the appropriate tool.";
|
||||
|
||||
/// Notice injected when the LLM's response was truncated mid-tool-call,
|
||||
/// causing incomplete parameters. Tells the LLM to try a different approach.
|
||||
pub const TRUNCATED_TOOL_CALL_NOTICE: &str = "\
|
||||
Your previous response was truncated while generating tool call parameters. \
|
||||
The tool calls were discarded. Please try a different approach — \
|
||||
summarize or transform the data instead of echoing it verbatim in a tool call.";
|
||||
|
||||
/// Seed value used as the second argument to `generate_tool_call_id` when
|
||||
/// recovering tool calls from malformed LLM text responses. This must differ
|
||||
/// from the `0` seed used in `rig_adapter::normalized_tool_call_id` to avoid
|
||||
@@ -194,11 +201,17 @@ pub struct ReasoningContext {
|
||||
pub metadata: std::collections::HashMap<String, String>,
|
||||
/// When true, force a text-only response (ignore available tools).
|
||||
/// Used by the agentic loop to guarantee termination near the iteration limit.
|
||||
/// Sticky: once set, never cleared within a loop invocation. Callers must
|
||||
/// create a fresh `ReasoningContext` per `run_agentic_loop()` call.
|
||||
pub force_text: bool,
|
||||
/// Pre-built system prompt. When set, `respond_with_tools` uses this directly
|
||||
/// instead of calling `build_system_prompt_with_tools`. Allows callers to build
|
||||
/// the prompt once and reuse it across iterations.
|
||||
pub system_prompt: Option<String>,
|
||||
/// Per-user model override. When set, completion requests use this model
|
||||
/// instead of the provider's default. Only effective with providers that
|
||||
/// support per-request model overrides (e.g. NearAI).
|
||||
pub model_override: Option<String>,
|
||||
}
|
||||
|
||||
impl ReasoningContext {
|
||||
@@ -212,6 +225,7 @@ impl ReasoningContext {
|
||||
metadata: std::collections::HashMap::new(),
|
||||
force_text: false,
|
||||
system_prompt: None,
|
||||
model_override: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -344,6 +358,7 @@ pub enum RespondResult {
|
||||
pub struct RespondOutput {
|
||||
pub result: RespondResult,
|
||||
pub usage: TokenUsage,
|
||||
pub finish_reason: FinishReason,
|
||||
}
|
||||
|
||||
/// Reasoning engine for the agent.
|
||||
@@ -525,6 +540,17 @@ impl Reasoning {
|
||||
|
||||
let response = self.llm.complete_with_tools(request).await?;
|
||||
|
||||
// If the response was truncated, tool call parameters are likely incomplete.
|
||||
// Return empty so the caller can fall through to respond_with_tools() which
|
||||
// has a larger output token budget.
|
||||
if response.finish_reason == FinishReason::Length {
|
||||
tracing::warn!(
|
||||
"select_tools response truncated (finish_reason=Length), \
|
||||
discarding potentially incomplete tool selections"
|
||||
);
|
||||
return Ok(vec![]);
|
||||
}
|
||||
|
||||
let shared_reasoning = response
|
||||
.content
|
||||
.map(|c| {
|
||||
@@ -671,6 +697,9 @@ Respond in JSON format:
|
||||
.with_temperature(0.7)
|
||||
.with_tool_choice("auto");
|
||||
request.metadata = context.metadata.clone();
|
||||
if let Some(ref model) = context.model_override {
|
||||
request.model = Some(model.clone());
|
||||
}
|
||||
|
||||
let response = self.llm.complete_with_tools(request).await?;
|
||||
let usage = TokenUsage {
|
||||
@@ -714,6 +743,7 @@ Respond in JSON format:
|
||||
content: narrative,
|
||||
},
|
||||
usage,
|
||||
finish_reason: response.finish_reason,
|
||||
});
|
||||
}
|
||||
|
||||
@@ -741,6 +771,7 @@ Respond in JSON format:
|
||||
},
|
||||
},
|
||||
usage,
|
||||
finish_reason: response.finish_reason,
|
||||
});
|
||||
}
|
||||
|
||||
@@ -766,6 +797,7 @@ Respond in JSON format:
|
||||
Ok(RespondOutput {
|
||||
result: RespondResult::Text(final_text),
|
||||
usage,
|
||||
finish_reason: response.finish_reason,
|
||||
})
|
||||
} else {
|
||||
// No tools, use simple completion
|
||||
@@ -773,6 +805,9 @@ Respond in JSON format:
|
||||
.with_max_tokens(4096)
|
||||
.with_temperature(0.7);
|
||||
request.metadata = context.metadata.clone();
|
||||
if let Some(ref model) = context.model_override {
|
||||
request.model = Some(model.clone());
|
||||
}
|
||||
|
||||
let response = self.llm.complete(request).await?;
|
||||
let pre_truncated = truncate_at_tool_tags(&response.content);
|
||||
@@ -794,6 +829,7 @@ Respond in JSON format:
|
||||
cache_read_input_tokens: response.cache_read_input_tokens,
|
||||
cache_creation_input_tokens: response.cache_creation_input_tokens,
|
||||
},
|
||||
finish_reason: response.finish_reason,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1334,6 +1370,49 @@ fn is_inside_code(pos: usize, regions: &[CodeRegion]) -> bool {
|
||||
regions.iter().any(|r| pos >= r.start && pos < r.end)
|
||||
}
|
||||
|
||||
/// Check whether a byte range overlaps any code region.
|
||||
fn overlaps_code_region(start: usize, end: usize, regions: &[CodeRegion]) -> bool {
|
||||
regions.iter().any(|r| start < r.end && end > r.start)
|
||||
}
|
||||
|
||||
/// Return the byte bounds of the line containing `pos`, excluding the trailing newline.
|
||||
fn line_bounds(text: &str, pos: usize) -> (usize, usize) {
|
||||
let start = text[..pos].rfind('\n').map_or(0, |idx| idx + 1);
|
||||
let end = text[pos..].find('\n').map_or(text.len(), |idx| pos + idx);
|
||||
(start, end)
|
||||
}
|
||||
|
||||
/// Only recover XML-style tool calls when they are isolated content outside
|
||||
/// markdown code and quote contexts. This avoids converting code examples or
|
||||
/// quoted snippets into executable tool calls.
|
||||
fn is_recoverable_tool_call_segment(
|
||||
text: &str,
|
||||
start: usize,
|
||||
end: usize,
|
||||
code_regions: &[CodeRegion],
|
||||
) -> bool {
|
||||
if overlaps_code_region(start, end, code_regions) {
|
||||
return false;
|
||||
}
|
||||
|
||||
let (first_line_start, first_line_end) = line_bounds(text, start);
|
||||
let first_line = &text[first_line_start..first_line_end];
|
||||
|
||||
if first_line.trim_start().starts_with('>') {
|
||||
return false;
|
||||
}
|
||||
|
||||
let (_, last_line_end) = line_bounds(text, end.saturating_sub(1));
|
||||
let first_line_prefix = &text[first_line_start..start];
|
||||
let last_line_suffix = &text[end..last_line_end];
|
||||
|
||||
if !first_line_prefix.trim().is_empty() || !last_line_suffix.trim().is_empty() {
|
||||
return false;
|
||||
}
|
||||
|
||||
true
|
||||
}
|
||||
|
||||
/// Clean up LLM response by stripping model-internal tags and reasoning patterns.
|
||||
///
|
||||
/// Some models (GLM-4.7, etc.) emit XML-tagged internal state like
|
||||
@@ -1353,6 +1432,7 @@ fn recover_tool_calls_from_content(
|
||||
) -> Vec<ToolCall> {
|
||||
let tool_names: std::collections::HashSet<&str> =
|
||||
available_tools.iter().map(|t| t.name.as_str()).collect();
|
||||
let code_regions = find_code_regions(content);
|
||||
let mut calls = Vec::new();
|
||||
|
||||
for (open, close) in &[
|
||||
@@ -1361,15 +1441,23 @@ fn recover_tool_calls_from_content(
|
||||
("<function_call>", "</function_call>"),
|
||||
("<|function_call|>", "<|/function_call|>"),
|
||||
] {
|
||||
let mut remaining = content;
|
||||
while let Some(start) = remaining.find(open) {
|
||||
let mut search_from = 0;
|
||||
while let Some(offset) = content[search_from..].find(open) {
|
||||
let start = search_from + offset;
|
||||
let inner_start = start + open.len();
|
||||
let after = &remaining[inner_start..];
|
||||
let Some(end) = after.find(close) else {
|
||||
let after = &content[inner_start..];
|
||||
let Some(end_offset) = after.find(close) else {
|
||||
break;
|
||||
};
|
||||
let inner = after[..end].trim();
|
||||
remaining = &after[end + close.len()..];
|
||||
let end = inner_start + end_offset;
|
||||
let segment_end = end + close.len();
|
||||
search_from = segment_end;
|
||||
|
||||
if !is_recoverable_tool_call_segment(content, start, segment_end, &code_regions) {
|
||||
continue;
|
||||
}
|
||||
|
||||
let inner = content[inner_start..end].trim();
|
||||
|
||||
if inner.is_empty() {
|
||||
continue;
|
||||
@@ -2302,6 +2390,40 @@ That's my plan."#;
|
||||
assert_eq!(calls[0].name, "tool_list");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_recover_tool_call_in_fenced_code_block_ignored() {
|
||||
let tools = make_tools(&["tool_list"]);
|
||||
let content = "Here is the XML format:\n\n```xml\n<tool_call>tool_list</tool_call>\n```";
|
||||
let calls = recover_tool_calls_from_content(content, &tools);
|
||||
assert!(calls.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_recover_tool_call_in_inline_code_ignored() {
|
||||
let tools = make_tools(&["tool_list"]);
|
||||
let content = "Use `<tool_call>tool_list</tool_call>` to illustrate the syntax.";
|
||||
let calls = recover_tool_calls_from_content(content, &tools);
|
||||
assert!(calls.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_recover_tool_call_in_blockquote_ignored() {
|
||||
let tools = make_tools(&["tool_list"]);
|
||||
let content = "The page replied:\n> <tool_call>tool_list</tool_call>";
|
||||
let calls = recover_tool_calls_from_content(content, &tools);
|
||||
assert!(calls.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_recover_multiline_json_tool_call_on_own_line() {
|
||||
let tools = make_tools(&["memory_search"]);
|
||||
let content = "Let me check.\n\n<tool_call>\n{\"name\": \"memory_search\", \"arguments\": {\"query\": \"test\"}}\n</tool_call>\n\nDone.";
|
||||
let calls = recover_tool_calls_from_content(content, &tools);
|
||||
assert_eq!(calls.len(), 1);
|
||||
assert_eq!(calls[0].name, "memory_search");
|
||||
assert_eq!(calls[0].arguments, serde_json::json!({"query": "test"}));
|
||||
}
|
||||
|
||||
// ---- System prompt building tests (issue #565) ----
|
||||
|
||||
fn make_test_reasoning() -> Reasoning {
|
||||
@@ -3218,4 +3340,85 @@ That's my plan."#;
|
||||
let cleaned = clean_response(&pre_truncated);
|
||||
assert!(cleaned.trim().is_empty());
|
||||
}
|
||||
|
||||
// ---- select_tools truncation guard ----
|
||||
|
||||
/// Mock provider that returns tool calls with a configurable finish_reason.
|
||||
struct TruncatingLlm {
|
||||
finish_reason: crate::llm::FinishReason,
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl crate::llm::LlmProvider for TruncatingLlm {
|
||||
fn model_name(&self) -> &str {
|
||||
"truncating-stub"
|
||||
}
|
||||
fn cost_per_token(&self) -> (rust_decimal::Decimal, rust_decimal::Decimal) {
|
||||
(rust_decimal::Decimal::ZERO, rust_decimal::Decimal::ZERO)
|
||||
}
|
||||
async fn complete(
|
||||
&self,
|
||||
_request: crate::llm::CompletionRequest,
|
||||
) -> Result<crate::llm::CompletionResponse, crate::llm::error::LlmError> {
|
||||
unimplemented!()
|
||||
}
|
||||
async fn complete_with_tools(
|
||||
&self,
|
||||
_request: crate::llm::ToolCompletionRequest,
|
||||
) -> Result<crate::llm::ToolCompletionResponse, crate::llm::error::LlmError> {
|
||||
Ok(crate::llm::ToolCompletionResponse {
|
||||
content: Some("I'll write the report.".to_string()),
|
||||
tool_calls: vec![ToolCall {
|
||||
id: "call_1".to_string(),
|
||||
name: "memory_write".to_string(),
|
||||
arguments: serde_json::json!({}),
|
||||
reasoning: None,
|
||||
}],
|
||||
input_tokens: 5000,
|
||||
output_tokens: 1024,
|
||||
finish_reason: self.finish_reason,
|
||||
cache_read_input_tokens: 0,
|
||||
cache_creation_input_tokens: 0,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_select_tools_returns_empty_on_truncation() {
|
||||
let llm = Arc::new(TruncatingLlm {
|
||||
finish_reason: FinishReason::Length,
|
||||
});
|
||||
let reasoning = Reasoning::new(llm);
|
||||
let mut ctx = ReasoningContext::new().with_message(ChatMessage::user("Write a report"));
|
||||
ctx.available_tools.push(ToolDefinition {
|
||||
name: "memory_write".to_string(),
|
||||
description: "Write to memory".to_string(),
|
||||
parameters: serde_json::json!({"type": "object"}),
|
||||
});
|
||||
|
||||
let selections = reasoning.select_tools(&ctx).await.unwrap();
|
||||
assert!(
|
||||
selections.is_empty(),
|
||||
"Truncated tool selections should be discarded (got {} selections)",
|
||||
selections.len()
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_select_tools_returns_selections_when_not_truncated() {
|
||||
let llm = Arc::new(TruncatingLlm {
|
||||
finish_reason: FinishReason::ToolUse,
|
||||
});
|
||||
let reasoning = Reasoning::new(llm);
|
||||
let mut ctx = ReasoningContext::new().with_message(ChatMessage::user("Write a report"));
|
||||
ctx.available_tools.push(ToolDefinition {
|
||||
name: "memory_write".to_string(),
|
||||
description: "Write to memory".to_string(),
|
||||
parameters: serde_json::json!({"type": "object"}),
|
||||
});
|
||||
|
||||
let selections = reasoning.select_tools(&ctx).await.unwrap();
|
||||
assert_eq!(selections.len(), 1);
|
||||
assert_eq!(selections[0].tool_name, "memory_write");
|
||||
}
|
||||
}
|
||||
|
||||
+32
-20
@@ -598,6 +598,30 @@ fn build_rig_request(
|
||||
})
|
||||
}
|
||||
|
||||
/// Inject a per-request model override into the rig request's `additional_params`.
|
||||
///
|
||||
/// Rig-core bakes the model name at construction time inside each provider's
|
||||
/// `CompletionModel` implementation. The actual HTTP request body includes a
|
||||
/// `model` field set by the provider. Rig-core's `#[serde(flatten)]` on
|
||||
/// `additional_params` emits these fields AFTER the provider's own fields.
|
||||
/// Most API servers (Python, Go) use last-key-wins when deserializing
|
||||
/// duplicate JSON keys, so the injected `model` value takes effect.
|
||||
fn inject_model_override(rig_req: &mut RigRequest, model_override: Option<&str>) {
|
||||
let Some(model) = model_override else {
|
||||
return;
|
||||
};
|
||||
match rig_req.additional_params {
|
||||
Some(ref mut params) => {
|
||||
if let Some(obj) = params.as_object_mut() {
|
||||
obj.insert("model".to_string(), serde_json::json!(model));
|
||||
}
|
||||
}
|
||||
None => {
|
||||
rig_req.additional_params = Some(serde_json::json!({ "model": model }));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl<M> LlmProvider for RigAdapter<M>
|
||||
where
|
||||
@@ -632,15 +656,7 @@ where
|
||||
&self,
|
||||
mut request: CompletionRequest,
|
||||
) -> Result<CompletionResponse, LlmError> {
|
||||
if let Some(requested_model) = request.model.as_deref()
|
||||
&& requested_model != self.model_name.as_str()
|
||||
{
|
||||
tracing::warn!(
|
||||
requested_model = requested_model,
|
||||
active_model = %self.model_name,
|
||||
"Per-request model override is not supported for this provider; using configured model"
|
||||
);
|
||||
}
|
||||
let model_override = request.model.take();
|
||||
|
||||
self.strip_unsupported_completion_params(&mut request);
|
||||
|
||||
@@ -648,7 +664,7 @@ where
|
||||
crate::llm::provider::sanitize_tool_messages(&mut messages);
|
||||
let (preamble, history) = convert_messages(&messages);
|
||||
|
||||
let rig_req = build_rig_request(
|
||||
let mut rig_req = build_rig_request(
|
||||
preamble,
|
||||
history,
|
||||
Vec::new(),
|
||||
@@ -658,6 +674,8 @@ where
|
||||
self.cache_retention,
|
||||
)?;
|
||||
|
||||
inject_model_override(&mut rig_req, model_override.as_deref());
|
||||
|
||||
let response =
|
||||
self.model
|
||||
.completion(rig_req)
|
||||
@@ -695,15 +713,7 @@ where
|
||||
&self,
|
||||
mut request: ToolCompletionRequest,
|
||||
) -> Result<ToolCompletionResponse, LlmError> {
|
||||
if let Some(requested_model) = request.model.as_deref()
|
||||
&& requested_model != self.model_name.as_str()
|
||||
{
|
||||
tracing::warn!(
|
||||
requested_model = requested_model,
|
||||
active_model = %self.model_name,
|
||||
"Per-request model override is not supported for this provider; using configured model"
|
||||
);
|
||||
}
|
||||
let model_override = request.model.take();
|
||||
|
||||
self.strip_unsupported_tool_params(&mut request);
|
||||
|
||||
@@ -716,7 +726,7 @@ where
|
||||
let tools = convert_tools(&request.tools);
|
||||
let tool_choice = convert_tool_choice(request.tool_choice.as_deref());
|
||||
|
||||
let rig_req = build_rig_request(
|
||||
let mut rig_req = build_rig_request(
|
||||
preamble,
|
||||
history,
|
||||
tools,
|
||||
@@ -726,6 +736,8 @@ where
|
||||
self.cache_retention,
|
||||
)?;
|
||||
|
||||
inject_model_override(&mut rig_req, model_override.as_deref());
|
||||
|
||||
let response =
|
||||
self.model
|
||||
.completion(rig_req)
|
||||
|
||||
@@ -914,6 +914,10 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
},
|
||||
builder: components.builder,
|
||||
llm_backend: config.llm.backend.clone(),
|
||||
tenant_rates: Arc::new(ironclaw::tenant::TenantRateRegistry::new(
|
||||
config.agent.max_llm_concurrent_per_user.unwrap_or(4),
|
||||
config.agent.max_jobs_concurrent_per_user.unwrap_or(3),
|
||||
)),
|
||||
};
|
||||
|
||||
let channels_for_warnings = Arc::clone(&channels);
|
||||
|
||||
+906
@@ -0,0 +1,906 @@
|
||||
//! Compile-time tenant isolation.
|
||||
//!
|
||||
//! Provides two database access tiers:
|
||||
//!
|
||||
//! - **[`TenantScope`]** (default): All operations are bound to a single user.
|
||||
//! ID-based lookups return `None` if the resource doesn't belong to this user.
|
||||
//! This is the only way handler code should access the database.
|
||||
//!
|
||||
//! - **[`AdminScope`]**: Cross-tenant access for system-level operations
|
||||
//! (heartbeat, routine engine, self-repair). Must be obtained explicitly via
|
||||
//! [`AgentDeps::admin_store()`](crate::agent::AgentDeps::admin_store).
|
||||
//!
|
||||
//! [`TenantCtx`] bundles a `TenantScope` with workspace, cost guard, and
|
||||
//! per-tenant rate limiting. Constructed once per request at the entry point
|
||||
//! where a `user_id` becomes known.
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
use chrono::{DateTime, Utc};
|
||||
use rust_decimal::Decimal;
|
||||
use tokio::sync::{Semaphore, SemaphorePermit};
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::agent::BrokenTool;
|
||||
use crate::agent::cost_guard::{CostGuard, CostLimitExceeded};
|
||||
use crate::agent::routine::{Routine, RoutineRun, RunStatus};
|
||||
use crate::context::{ActionRecord, JobContext, JobState};
|
||||
use crate::db::Database;
|
||||
use crate::error::DatabaseError;
|
||||
use crate::history::{
|
||||
AgentJobRecord, AgentJobSummary, ConversationMessage, ConversationSummary, LlmCallRecord,
|
||||
SandboxJobRecord, SandboxJobSummary, SettingRow,
|
||||
};
|
||||
use crate::workspace::Workspace;
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// TenantScope — scoped database access (default tier)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Scoped database view. All operations are bound to a single user.
|
||||
///
|
||||
/// This is the **only** way handler code should access the database.
|
||||
/// ID-based lookups (jobs, routines, sandbox jobs) automatically filter
|
||||
/// by ownership — returning `None` when the resource belongs to a
|
||||
/// different user.
|
||||
#[derive(Clone)]
|
||||
pub struct TenantScope {
|
||||
user_id: String,
|
||||
inner: Arc<dyn Database>,
|
||||
}
|
||||
|
||||
impl TenantScope {
|
||||
pub fn new(user_id: impl Into<String>, db: Arc<dyn Database>) -> Self {
|
||||
Self {
|
||||
user_id: user_id.into(),
|
||||
inner: db,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn user_id(&self) -> &str {
|
||||
&self.user_id
|
||||
}
|
||||
|
||||
// === Jobs ===
|
||||
|
||||
pub async fn list_agent_jobs(&self) -> Result<Vec<AgentJobRecord>, DatabaseError> {
|
||||
self.inner.list_agent_jobs_for_user(&self.user_id).await
|
||||
}
|
||||
|
||||
pub async fn agent_job_summary(&self) -> Result<AgentJobSummary, DatabaseError> {
|
||||
self.inner.agent_job_summary_for_user(&self.user_id).await
|
||||
}
|
||||
|
||||
/// Fetch a job by ID, returning `None` if it doesn't belong to this user.
|
||||
pub async fn get_job(&self, id: Uuid) -> Result<Option<JobContext>, DatabaseError> {
|
||||
match self.inner.get_job(id).await? {
|
||||
Some(ctx) if ctx.user_id == self.user_id => Ok(Some(ctx)),
|
||||
_ => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn get_agent_job_failure_reason(
|
||||
&self,
|
||||
id: Uuid,
|
||||
) -> Result<Option<String>, DatabaseError> {
|
||||
// Verify ownership first
|
||||
if self.get_job(id).await?.is_none() {
|
||||
return Ok(None);
|
||||
}
|
||||
self.inner.get_agent_job_failure_reason(id).await
|
||||
}
|
||||
|
||||
pub async fn update_job_status(
|
||||
&self,
|
||||
id: Uuid,
|
||||
status: JobState,
|
||||
failure_reason: Option<&str>,
|
||||
) -> Result<(), DatabaseError> {
|
||||
// Verify ownership before mutating
|
||||
if self.get_job(id).await?.is_none() {
|
||||
return Err(DatabaseError::NotFound {
|
||||
entity: "job".to_string(),
|
||||
id: id.to_string(),
|
||||
});
|
||||
}
|
||||
self.inner
|
||||
.update_job_status(id, status, failure_reason)
|
||||
.await
|
||||
}
|
||||
|
||||
// === Sandbox jobs ===
|
||||
|
||||
pub async fn list_sandbox_jobs(&self) -> Result<Vec<SandboxJobRecord>, DatabaseError> {
|
||||
self.inner.list_sandbox_jobs_for_user(&self.user_id).await
|
||||
}
|
||||
|
||||
pub async fn sandbox_job_summary(&self) -> Result<SandboxJobSummary, DatabaseError> {
|
||||
self.inner.sandbox_job_summary_for_user(&self.user_id).await
|
||||
}
|
||||
|
||||
/// Fetch a sandbox job by ID, returning `None` if it doesn't belong to this user.
|
||||
pub async fn get_sandbox_job(
|
||||
&self,
|
||||
id: Uuid,
|
||||
) -> Result<Option<SandboxJobRecord>, DatabaseError> {
|
||||
match self.inner.get_sandbox_job(id).await? {
|
||||
Some(job) if job.user_id == self.user_id => Ok(Some(job)),
|
||||
_ => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn sandbox_job_belongs_to_user(&self, job_id: Uuid) -> Result<bool, DatabaseError> {
|
||||
self.inner
|
||||
.sandbox_job_belongs_to_user(job_id, &self.user_id)
|
||||
.await
|
||||
}
|
||||
|
||||
// === Routines ===
|
||||
|
||||
pub async fn list_routines(&self) -> Result<Vec<Routine>, DatabaseError> {
|
||||
self.inner.list_routines(&self.user_id).await
|
||||
}
|
||||
|
||||
pub async fn get_routine_by_name(&self, name: &str) -> Result<Option<Routine>, DatabaseError> {
|
||||
self.inner.get_routine_by_name(&self.user_id, name).await
|
||||
}
|
||||
|
||||
/// Fetch a routine by ID, returning `None` if it doesn't belong to this user.
|
||||
pub async fn get_routine(&self, id: Uuid) -> Result<Option<Routine>, DatabaseError> {
|
||||
match self.inner.get_routine(id).await? {
|
||||
Some(r) if r.user_id == self.user_id => Ok(Some(r)),
|
||||
_ => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn create_routine(&self, routine: &Routine) -> Result<(), DatabaseError> {
|
||||
debug_assert_eq!(
|
||||
routine.user_id, self.user_id,
|
||||
"routine.user_id must match TenantScope user"
|
||||
);
|
||||
self.inner.create_routine(routine).await
|
||||
}
|
||||
|
||||
pub async fn update_routine(&self, routine: &Routine) -> Result<(), DatabaseError> {
|
||||
// Verify ownership
|
||||
if self.get_routine(routine.id).await?.is_none() {
|
||||
return Err(DatabaseError::NotFound {
|
||||
entity: "routine".to_string(),
|
||||
id: routine.id.to_string(),
|
||||
});
|
||||
}
|
||||
self.inner.update_routine(routine).await
|
||||
}
|
||||
|
||||
pub async fn delete_routine(&self, id: Uuid) -> Result<bool, DatabaseError> {
|
||||
// Verify ownership
|
||||
if self.get_routine(id).await?.is_none() {
|
||||
return Err(DatabaseError::NotFound {
|
||||
entity: "routine".to_string(),
|
||||
id: id.to_string(),
|
||||
});
|
||||
}
|
||||
self.inner.delete_routine(id).await
|
||||
}
|
||||
|
||||
/// List routine runs, verifying the routine belongs to this user.
|
||||
pub async fn list_routine_runs(
|
||||
&self,
|
||||
routine_id: Uuid,
|
||||
limit: i64,
|
||||
) -> Result<Vec<RoutineRun>, DatabaseError> {
|
||||
// Verify routine ownership first
|
||||
if self.get_routine(routine_id).await?.is_none() {
|
||||
return Err(DatabaseError::NotFound {
|
||||
entity: "routine".to_string(),
|
||||
id: routine_id.to_string(),
|
||||
});
|
||||
}
|
||||
self.inner.list_routine_runs(routine_id, limit).await
|
||||
}
|
||||
|
||||
pub async fn get_webhook_routine_by_path(
|
||||
&self,
|
||||
path: &str,
|
||||
) -> Result<Option<Routine>, DatabaseError> {
|
||||
self.inner
|
||||
.get_webhook_routine_by_path(path, Some(&self.user_id))
|
||||
.await
|
||||
}
|
||||
|
||||
// === Settings ===
|
||||
|
||||
pub async fn get_setting(&self, key: &str) -> Result<Option<serde_json::Value>, DatabaseError> {
|
||||
self.inner.get_setting(&self.user_id, key).await
|
||||
}
|
||||
|
||||
pub async fn get_setting_full(&self, key: &str) -> Result<Option<SettingRow>, DatabaseError> {
|
||||
self.inner.get_setting_full(&self.user_id, key).await
|
||||
}
|
||||
|
||||
pub async fn set_setting(
|
||||
&self,
|
||||
key: &str,
|
||||
value: &serde_json::Value,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner.set_setting(&self.user_id, key, value).await
|
||||
}
|
||||
|
||||
pub async fn delete_setting(&self, key: &str) -> Result<bool, DatabaseError> {
|
||||
self.inner.delete_setting(&self.user_id, key).await
|
||||
}
|
||||
|
||||
pub async fn list_settings(&self) -> Result<Vec<SettingRow>, DatabaseError> {
|
||||
self.inner.list_settings(&self.user_id).await
|
||||
}
|
||||
|
||||
pub async fn get_all_settings(
|
||||
&self,
|
||||
) -> Result<HashMap<String, serde_json::Value>, DatabaseError> {
|
||||
self.inner.get_all_settings(&self.user_id).await
|
||||
}
|
||||
|
||||
pub async fn set_all_settings(
|
||||
&self,
|
||||
settings: &HashMap<String, serde_json::Value>,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner.set_all_settings(&self.user_id, settings).await
|
||||
}
|
||||
|
||||
pub async fn has_settings(&self) -> Result<bool, DatabaseError> {
|
||||
self.inner.has_settings(&self.user_id).await
|
||||
}
|
||||
|
||||
// === Conversations ===
|
||||
|
||||
pub async fn create_conversation(
|
||||
&self,
|
||||
channel: &str,
|
||||
thread_id: Option<&str>,
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
self.inner
|
||||
.create_conversation(channel, &self.user_id, thread_id)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn ensure_conversation(
|
||||
&self,
|
||||
id: Uuid,
|
||||
channel: &str,
|
||||
thread_id: Option<&str>,
|
||||
) -> Result<bool, DatabaseError> {
|
||||
self.inner
|
||||
.ensure_conversation(id, channel, &self.user_id, thread_id)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn list_conversations_with_preview(
|
||||
&self,
|
||||
channel: &str,
|
||||
limit: i64,
|
||||
) -> Result<Vec<ConversationSummary>, DatabaseError> {
|
||||
self.inner
|
||||
.list_conversations_with_preview(&self.user_id, channel, limit)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn list_conversations_all_channels(
|
||||
&self,
|
||||
limit: i64,
|
||||
) -> Result<Vec<ConversationSummary>, DatabaseError> {
|
||||
self.inner
|
||||
.list_conversations_all_channels(&self.user_id, limit)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn get_or_create_routine_conversation(
|
||||
&self,
|
||||
routine_id: Uuid,
|
||||
routine_name: &str,
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
self.inner
|
||||
.get_or_create_routine_conversation(routine_id, routine_name, &self.user_id)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn get_or_create_heartbeat_conversation(&self) -> Result<Uuid, DatabaseError> {
|
||||
self.inner
|
||||
.get_or_create_heartbeat_conversation(&self.user_id)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn get_or_create_assistant_conversation(
|
||||
&self,
|
||||
channel: &str,
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
self.inner
|
||||
.get_or_create_assistant_conversation(&self.user_id, channel)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn conversation_belongs_to_user(
|
||||
&self,
|
||||
conversation_id: Uuid,
|
||||
) -> Result<bool, DatabaseError> {
|
||||
self.inner
|
||||
.conversation_belongs_to_user(conversation_id, &self.user_id)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Add a message to a conversation owned by this tenant.
|
||||
///
|
||||
/// Verifies the conversation belongs to this user before adding.
|
||||
pub async fn add_conversation_message(
|
||||
&self,
|
||||
conversation_id: Uuid,
|
||||
role: &str,
|
||||
content: &str,
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
self.inner
|
||||
.add_conversation_message(conversation_id, role, content)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn touch_conversation(&self, id: Uuid) -> Result<(), DatabaseError> {
|
||||
self.inner.touch_conversation(id).await
|
||||
}
|
||||
|
||||
pub async fn list_conversation_messages(
|
||||
&self,
|
||||
conversation_id: Uuid,
|
||||
) -> Result<Vec<ConversationMessage>, DatabaseError> {
|
||||
self.inner.list_conversation_messages(conversation_id).await
|
||||
}
|
||||
|
||||
pub async fn list_conversation_messages_paginated(
|
||||
&self,
|
||||
conversation_id: Uuid,
|
||||
before: Option<DateTime<Utc>>,
|
||||
limit: i64,
|
||||
) -> Result<(Vec<ConversationMessage>, bool), DatabaseError> {
|
||||
self.inner
|
||||
.list_conversation_messages_paginated(conversation_id, before, limit)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn create_conversation_with_metadata(
|
||||
&self,
|
||||
channel: &str,
|
||||
metadata: &serde_json::Value,
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
self.inner
|
||||
.create_conversation_with_metadata(channel, &self.user_id, metadata)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn update_conversation_metadata_field(
|
||||
&self,
|
||||
id: Uuid,
|
||||
key: &str,
|
||||
value: &serde_json::Value,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner
|
||||
.update_conversation_metadata_field(id, key, value)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn get_conversation_metadata(
|
||||
&self,
|
||||
id: Uuid,
|
||||
) -> Result<Option<serde_json::Value>, DatabaseError> {
|
||||
self.inner.get_conversation_metadata(id).await
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// AdminScope — explicit cross-tenant access
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Cross-tenant database access for system-level operations.
|
||||
///
|
||||
/// **Not** available through [`TenantCtx`] — must be obtained explicitly via
|
||||
/// [`AgentDeps::admin_store()`](crate::agent::AgentDeps::admin_store).
|
||||
///
|
||||
/// Used by: heartbeat enumeration, routine engine scheduling, self-repair,
|
||||
/// scheduler job persistence, worker status updates.
|
||||
#[derive(Clone)]
|
||||
pub struct AdminScope {
|
||||
inner: Arc<dyn Database>,
|
||||
}
|
||||
|
||||
impl AdminScope {
|
||||
pub fn new(db: Arc<dyn Database>) -> Self {
|
||||
Self { inner: db }
|
||||
}
|
||||
|
||||
/// Access the raw Database trait object.
|
||||
///
|
||||
/// Prefer using the typed methods on AdminScope instead. This is provided
|
||||
/// for call sites that need sub-trait access not yet wrapped here.
|
||||
pub fn db(&self) -> &Arc<dyn Database> {
|
||||
&self.inner
|
||||
}
|
||||
|
||||
// === Routine engine ===
|
||||
|
||||
pub async fn list_all_routines(&self) -> Result<Vec<Routine>, DatabaseError> {
|
||||
self.inner.list_all_routines().await
|
||||
}
|
||||
|
||||
pub async fn list_event_routines(&self) -> Result<Vec<Routine>, DatabaseError> {
|
||||
self.inner.list_event_routines().await
|
||||
}
|
||||
|
||||
pub async fn list_due_cron_routines(&self) -> Result<Vec<Routine>, DatabaseError> {
|
||||
self.inner.list_due_cron_routines().await
|
||||
}
|
||||
|
||||
pub async fn list_dispatched_routine_runs(&self) -> Result<Vec<RoutineRun>, DatabaseError> {
|
||||
self.inner.list_dispatched_routine_runs().await
|
||||
}
|
||||
|
||||
pub async fn count_running_routine_runs_batch(
|
||||
&self,
|
||||
routine_ids: &[Uuid],
|
||||
) -> Result<HashMap<Uuid, i64>, DatabaseError> {
|
||||
self.inner
|
||||
.count_running_routine_runs_batch(routine_ids)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn batch_get_last_run_status(
|
||||
&self,
|
||||
routine_ids: &[Uuid],
|
||||
) -> Result<HashMap<Uuid, RunStatus>, DatabaseError> {
|
||||
self.inner.batch_get_last_run_status(routine_ids).await
|
||||
}
|
||||
|
||||
pub async fn count_running_routine_runs(&self, routine_id: Uuid) -> Result<i64, DatabaseError> {
|
||||
self.inner.count_running_routine_runs(routine_id).await
|
||||
}
|
||||
|
||||
pub async fn update_routine_runtime(
|
||||
&self,
|
||||
id: Uuid,
|
||||
last_run_at: DateTime<Utc>,
|
||||
next_fire_at: Option<DateTime<Utc>>,
|
||||
run_count: u64,
|
||||
consecutive_failures: u32,
|
||||
state: &serde_json::Value,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner
|
||||
.update_routine_runtime(
|
||||
id,
|
||||
last_run_at,
|
||||
next_fire_at,
|
||||
run_count,
|
||||
consecutive_failures,
|
||||
state,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn create_routine_run(&self, run: &RoutineRun) -> Result<(), DatabaseError> {
|
||||
self.inner.create_routine_run(run).await
|
||||
}
|
||||
|
||||
pub async fn complete_routine_run(
|
||||
&self,
|
||||
id: Uuid,
|
||||
status: RunStatus,
|
||||
result_summary: Option<&str>,
|
||||
tokens_used: Option<i32>,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner
|
||||
.complete_routine_run(id, status, result_summary, tokens_used)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn link_routine_run_to_job(
|
||||
&self,
|
||||
run_id: Uuid,
|
||||
job_id: Uuid,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner.link_routine_run_to_job(run_id, job_id).await
|
||||
}
|
||||
|
||||
pub async fn get_routine(&self, id: Uuid) -> Result<Option<Routine>, DatabaseError> {
|
||||
self.inner.get_routine(id).await
|
||||
}
|
||||
|
||||
pub async fn update_routine(&self, routine: &Routine) -> Result<(), DatabaseError> {
|
||||
self.inner.update_routine(routine).await
|
||||
}
|
||||
|
||||
// === Self-repair ===
|
||||
|
||||
pub async fn get_stuck_jobs(&self) -> Result<Vec<Uuid>, DatabaseError> {
|
||||
self.inner.get_stuck_jobs().await
|
||||
}
|
||||
|
||||
pub async fn get_broken_tools(&self, threshold: i32) -> Result<Vec<BrokenTool>, DatabaseError> {
|
||||
self.inner.get_broken_tools(threshold).await
|
||||
}
|
||||
|
||||
pub async fn record_tool_failure(
|
||||
&self,
|
||||
tool_name: &str,
|
||||
error_message: &str,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner
|
||||
.record_tool_failure(tool_name, error_message)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn mark_tool_repaired(&self, tool_name: &str) -> Result<(), DatabaseError> {
|
||||
self.inner.mark_tool_repaired(tool_name).await
|
||||
}
|
||||
|
||||
pub async fn increment_repair_attempts(&self, tool_name: &str) -> Result<(), DatabaseError> {
|
||||
self.inner.increment_repair_attempts(tool_name).await
|
||||
}
|
||||
|
||||
// === Sandbox housekeeping ===
|
||||
|
||||
pub async fn cleanup_stale_sandbox_jobs(&self) -> Result<u64, DatabaseError> {
|
||||
self.inner.cleanup_stale_sandbox_jobs().await
|
||||
}
|
||||
|
||||
pub async fn get_sandbox_job(
|
||||
&self,
|
||||
id: Uuid,
|
||||
) -> Result<Option<SandboxJobRecord>, DatabaseError> {
|
||||
self.inner.get_sandbox_job(id).await
|
||||
}
|
||||
|
||||
pub async fn save_sandbox_job(&self, job: &SandboxJobRecord) -> Result<(), DatabaseError> {
|
||||
self.inner.save_sandbox_job(job).await
|
||||
}
|
||||
|
||||
pub async fn update_sandbox_job_status(
|
||||
&self,
|
||||
id: Uuid,
|
||||
status: &str,
|
||||
success: Option<bool>,
|
||||
message: Option<&str>,
|
||||
started_at: Option<DateTime<Utc>>,
|
||||
completed_at: Option<DateTime<Utc>>,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner
|
||||
.update_sandbox_job_status(id, status, success, message, started_at, completed_at)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn update_sandbox_job_mode(&self, id: Uuid, mode: &str) -> Result<(), DatabaseError> {
|
||||
self.inner.update_sandbox_job_mode(id, mode).await
|
||||
}
|
||||
|
||||
pub async fn get_sandbox_job_mode(&self, id: Uuid) -> Result<Option<String>, DatabaseError> {
|
||||
self.inner.get_sandbox_job_mode(id).await
|
||||
}
|
||||
|
||||
pub async fn save_job_event(
|
||||
&self,
|
||||
job_id: Uuid,
|
||||
event_type: &str,
|
||||
data: &serde_json::Value,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner.save_job_event(job_id, event_type, data).await
|
||||
}
|
||||
|
||||
pub async fn list_job_events(
|
||||
&self,
|
||||
job_id: Uuid,
|
||||
limit: Option<i64>,
|
||||
) -> Result<Vec<crate::history::JobEventRecord>, DatabaseError> {
|
||||
self.inner.list_job_events(job_id, limit).await
|
||||
}
|
||||
|
||||
// === Job persistence (scheduler, worker) ===
|
||||
|
||||
pub async fn get_job(&self, id: Uuid) -> Result<Option<JobContext>, DatabaseError> {
|
||||
self.inner.get_job(id).await
|
||||
}
|
||||
|
||||
pub async fn save_job(&self, ctx: &JobContext) -> Result<(), DatabaseError> {
|
||||
self.inner.save_job(ctx).await
|
||||
}
|
||||
|
||||
pub async fn update_job_status(
|
||||
&self,
|
||||
id: Uuid,
|
||||
status: JobState,
|
||||
failure_reason: Option<&str>,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner
|
||||
.update_job_status(id, status, failure_reason)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn mark_job_stuck(&self, id: Uuid) -> Result<(), DatabaseError> {
|
||||
self.inner.mark_job_stuck(id).await
|
||||
}
|
||||
|
||||
pub async fn list_agent_jobs(&self) -> Result<Vec<AgentJobRecord>, DatabaseError> {
|
||||
self.inner.list_agent_jobs().await
|
||||
}
|
||||
|
||||
pub async fn get_agent_job_failure_reason(
|
||||
&self,
|
||||
id: Uuid,
|
||||
) -> Result<Option<String>, DatabaseError> {
|
||||
self.inner.get_agent_job_failure_reason(id).await
|
||||
}
|
||||
|
||||
// === LLM call recording ===
|
||||
|
||||
pub async fn record_llm_call(&self, record: &LlmCallRecord<'_>) -> Result<Uuid, DatabaseError> {
|
||||
self.inner.record_llm_call(record).await
|
||||
}
|
||||
|
||||
pub async fn save_action(
|
||||
&self,
|
||||
job_id: Uuid,
|
||||
action: &ActionRecord,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner.save_action(job_id, action).await
|
||||
}
|
||||
|
||||
pub async fn get_job_actions(&self, job_id: Uuid) -> Result<Vec<ActionRecord>, DatabaseError> {
|
||||
self.inner.get_job_actions(job_id).await
|
||||
}
|
||||
|
||||
// === Estimation ===
|
||||
|
||||
pub async fn save_estimation_snapshot(
|
||||
&self,
|
||||
job_id: Uuid,
|
||||
category: &str,
|
||||
tool_names: &[String],
|
||||
estimated_cost: Decimal,
|
||||
estimated_time_secs: i32,
|
||||
estimated_value: Decimal,
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
self.inner
|
||||
.save_estimation_snapshot(
|
||||
job_id,
|
||||
category,
|
||||
tool_names,
|
||||
estimated_cost,
|
||||
estimated_time_secs,
|
||||
estimated_value,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn update_estimation_actuals(
|
||||
&self,
|
||||
id: Uuid,
|
||||
actual_cost: Decimal,
|
||||
actual_time_secs: i32,
|
||||
actual_value: Option<Decimal>,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner
|
||||
.update_estimation_actuals(id, actual_cost, actual_time_secs, actual_value)
|
||||
.await
|
||||
}
|
||||
|
||||
// === Conversations (admin context) ===
|
||||
|
||||
pub async fn add_conversation_message(
|
||||
&self,
|
||||
conversation_id: Uuid,
|
||||
role: &str,
|
||||
content: &str,
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
self.inner
|
||||
.add_conversation_message(conversation_id, role, content)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn get_or_create_routine_conversation(
|
||||
&self,
|
||||
routine_id: Uuid,
|
||||
routine_name: &str,
|
||||
user_id: &str,
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
self.inner
|
||||
.get_or_create_routine_conversation(routine_id, routine_name, user_id)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn get_or_create_heartbeat_conversation(
|
||||
&self,
|
||||
user_id: &str,
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
self.inner
|
||||
.get_or_create_heartbeat_conversation(user_id)
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// TenantRateState / TenantRateRegistry — per-user concurrency
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Per-tenant concurrency limits.
|
||||
pub struct TenantRateState {
|
||||
/// Limits concurrent LLM calls for this user.
|
||||
pub llm_semaphore: Arc<Semaphore>,
|
||||
/// Limits concurrent jobs for this user.
|
||||
pub job_semaphore: Arc<Semaphore>,
|
||||
}
|
||||
|
||||
impl TenantRateState {
|
||||
pub fn new(max_llm_concurrent: usize, max_job_concurrent: usize) -> Self {
|
||||
Self {
|
||||
llm_semaphore: Arc::new(Semaphore::new(max_llm_concurrent)),
|
||||
job_semaphore: Arc::new(Semaphore::new(max_job_concurrent)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Registry that lazily creates per-tenant rate state.
|
||||
///
|
||||
/// Uses `tokio::sync::RwLock<HashMap>` (consistent with the rest of the
|
||||
/// codebase — no DashMap dependency).
|
||||
pub struct TenantRateRegistry {
|
||||
state: tokio::sync::RwLock<HashMap<String, Arc<TenantRateState>>>,
|
||||
max_llm_concurrent: usize,
|
||||
max_job_concurrent: usize,
|
||||
}
|
||||
|
||||
impl TenantRateRegistry {
|
||||
pub fn new(max_llm_concurrent: usize, max_job_concurrent: usize) -> Self {
|
||||
Self {
|
||||
state: tokio::sync::RwLock::new(HashMap::new()),
|
||||
max_llm_concurrent,
|
||||
max_job_concurrent,
|
||||
}
|
||||
}
|
||||
|
||||
/// Get or lazily create rate state for a user.
|
||||
pub async fn get_or_create(&self, user_id: &str) -> Arc<TenantRateState> {
|
||||
// Fast path: read lock
|
||||
{
|
||||
let map = self.state.read().await;
|
||||
if let Some(s) = map.get(user_id) {
|
||||
return Arc::clone(s);
|
||||
}
|
||||
}
|
||||
// Slow path: write lock with double-check
|
||||
let mut map = self.state.write().await;
|
||||
if let Some(s) = map.get(user_id) {
|
||||
return Arc::clone(s);
|
||||
}
|
||||
let s = Arc::new(TenantRateState::new(
|
||||
self.max_llm_concurrent,
|
||||
self.max_job_concurrent,
|
||||
));
|
||||
map.insert(user_id.to_string(), Arc::clone(&s));
|
||||
s
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// TenantCtx — per-request tenant execution context
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Per-request tenant execution context.
|
||||
///
|
||||
/// Bundles a [`TenantScope`] (scoped DB access), workspace, cost guard,
|
||||
/// and per-tenant rate limiting. Constructed once per request via
|
||||
/// [`AgentDeps::tenant_ctx()`](crate::agent::AgentDeps::tenant_ctx).
|
||||
///
|
||||
/// `Clone + Send + Sync` — safe to store on `ChatDelegate` without lifetime issues.
|
||||
#[derive(Clone)]
|
||||
pub struct TenantCtx {
|
||||
user_id: String,
|
||||
store: Option<TenantScope>,
|
||||
workspace: Option<Arc<Workspace>>,
|
||||
cost_guard: Arc<CostGuard>,
|
||||
rate: Arc<TenantRateState>,
|
||||
}
|
||||
|
||||
impl TenantCtx {
|
||||
pub fn new(
|
||||
user_id: impl Into<String>,
|
||||
store: Option<TenantScope>,
|
||||
workspace: Option<Arc<Workspace>>,
|
||||
cost_guard: Arc<CostGuard>,
|
||||
rate: Arc<TenantRateState>,
|
||||
) -> Self {
|
||||
Self {
|
||||
user_id: user_id.into(),
|
||||
store,
|
||||
workspace,
|
||||
cost_guard,
|
||||
rate,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn user_id(&self) -> &str {
|
||||
&self.user_id
|
||||
}
|
||||
|
||||
pub fn store(&self) -> Option<&TenantScope> {
|
||||
self.store.as_ref()
|
||||
}
|
||||
|
||||
pub fn workspace(&self) -> Option<&Arc<Workspace>> {
|
||||
self.workspace.as_ref()
|
||||
}
|
||||
|
||||
pub fn cost_guard(&self) -> &CostGuard {
|
||||
&self.cost_guard
|
||||
}
|
||||
|
||||
/// Check cost limits for this tenant (global + per-user).
|
||||
pub async fn check_cost_allowed(&self) -> Result<(), CostLimitExceeded> {
|
||||
self.cost_guard.check_allowed_for_user(&self.user_id).await
|
||||
}
|
||||
|
||||
/// Record an LLM call for this tenant.
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub async fn record_llm_call(
|
||||
&self,
|
||||
model: &str,
|
||||
input_tokens: u32,
|
||||
output_tokens: u32,
|
||||
cache_read_input_tokens: u32,
|
||||
cache_creation_input_tokens: u32,
|
||||
cache_read_discount: Decimal,
|
||||
cache_write_multiplier: Decimal,
|
||||
cost_per_token: Option<(Decimal, Decimal)>,
|
||||
) -> Decimal {
|
||||
self.cost_guard
|
||||
.record_llm_call_for_user(
|
||||
&self.user_id,
|
||||
model,
|
||||
input_tokens,
|
||||
output_tokens,
|
||||
cache_read_input_tokens,
|
||||
cache_creation_input_tokens,
|
||||
cache_read_discount,
|
||||
cache_write_multiplier,
|
||||
cost_per_token,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Acquire an LLM concurrency permit for this tenant.
|
||||
pub async fn acquire_llm_permit(&self) -> Result<SemaphorePermit<'_>, crate::error::Error> {
|
||||
self.rate.llm_semaphore.acquire().await.map_err(|_| {
|
||||
crate::error::Error::Config(crate::error::ConfigError::InvalidValue {
|
||||
key: "llm_semaphore".to_string(),
|
||||
message: "semaphore closed".to_string(),
|
||||
})
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_rate_registry_returns_same_state_for_same_user() {
|
||||
let registry = TenantRateRegistry::new(4, 3);
|
||||
let a1 = registry.get_or_create("alice").await;
|
||||
let a2 = registry.get_or_create("alice").await;
|
||||
assert!(Arc::ptr_eq(&a1, &a2));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_rate_registry_different_users_get_different_state() {
|
||||
let registry = TenantRateRegistry::new(4, 3);
|
||||
let alice = registry.get_or_create("alice").await;
|
||||
let bob = registry.get_or_create("bob").await;
|
||||
assert!(!Arc::ptr_eq(&alice, &bob));
|
||||
}
|
||||
}
|
||||
@@ -532,6 +532,7 @@ impl TestHarnessBuilder {
|
||||
let cost_guard = Arc::new(CostGuard::new(CostGuardConfig {
|
||||
max_cost_per_day_cents: None,
|
||||
max_actions_per_hour: None,
|
||||
max_cost_per_user_per_day_cents: None,
|
||||
}));
|
||||
|
||||
let channel = if self.stub_channel {
|
||||
@@ -564,6 +565,7 @@ impl TestHarnessBuilder {
|
||||
sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig,
|
||||
builder: None,
|
||||
llm_backend: "nearai".to_string(),
|
||||
tenant_rates: std::sync::Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)),
|
||||
};
|
||||
|
||||
TestHarness {
|
||||
|
||||
@@ -139,6 +139,8 @@ mod tests {
|
||||
Arc::new(ExtensionManager::new(
|
||||
Arc::new(McpSessionManager::new()),
|
||||
Arc::new(McpProcessManager::new()),
|
||||
None,
|
||||
None,
|
||||
secrets,
|
||||
tools,
|
||||
Some(Arc::new(HookRegistry::default())),
|
||||
@@ -148,6 +150,7 @@ mod tests {
|
||||
None,
|
||||
owner_id.to_string(),
|
||||
None,
|
||||
None,
|
||||
Vec::new(),
|
||||
))
|
||||
}
|
||||
|
||||
@@ -800,6 +800,8 @@ mod tests {
|
||||
Arc::new(ExtensionManager::new(
|
||||
Arc::new(McpSessionManager::new()),
|
||||
Arc::new(crate::tools::mcp::process::McpProcessManager::new()),
|
||||
None,
|
||||
None,
|
||||
Arc::new(InMemorySecretsStore::new(crypto)),
|
||||
Arc::new(ToolRegistry::new()),
|
||||
None,
|
||||
@@ -809,6 +811,7 @@ mod tests {
|
||||
None,
|
||||
"test".to_string(),
|
||||
None,
|
||||
None,
|
||||
Vec::new(),
|
||||
))
|
||||
}
|
||||
|
||||
@@ -650,6 +650,23 @@ pub(crate) fn routine_update_parameters_schema() -> Value {
|
||||
})
|
||||
}
|
||||
|
||||
const ROUTINE_LAST_NAME_STASH_KEY: &str = "__routine_last_name";
|
||||
|
||||
async fn stash_last_routine_name(ctx: &JobContext, name: &str) {
|
||||
ctx.tool_output_stash
|
||||
.write()
|
||||
.await
|
||||
.insert(ROUTINE_LAST_NAME_STASH_KEY.to_string(), name.to_string());
|
||||
}
|
||||
|
||||
async fn restore_last_routine_name(ctx: &JobContext) -> Option<String> {
|
||||
ctx.tool_output_stash
|
||||
.read()
|
||||
.await
|
||||
.get(ROUTINE_LAST_NAME_STASH_KEY)
|
||||
.cloned()
|
||||
}
|
||||
|
||||
fn nested_object<'a>(params: &'a Value, field: &str) -> Option<&'a Map<String, Value>> {
|
||||
params.get(field).and_then(Value::as_object)
|
||||
}
|
||||
@@ -915,7 +932,7 @@ fn parse_routine_create_request(
|
||||
fn build_routine_trigger(trigger: &NormalizedTriggerRequest) -> Trigger {
|
||||
match trigger {
|
||||
NormalizedTriggerRequest::Cron { schedule, timezone } => Trigger::Cron {
|
||||
schedule: schedule.clone(),
|
||||
schedule: normalize_cron_expression(schedule),
|
||||
timezone: timezone.clone(),
|
||||
},
|
||||
NormalizedTriggerRequest::Manual => Trigger::Manual,
|
||||
@@ -1093,6 +1110,7 @@ impl Tool for RoutineCreateTool {
|
||||
) -> Result<ToolOutput, ToolError> {
|
||||
let start = std::time::Instant::now();
|
||||
let normalized = parse_routine_create_request(¶ms)?;
|
||||
stash_last_routine_name(ctx, &normalized.name).await;
|
||||
let trigger = build_routine_trigger(&normalized.trigger);
|
||||
let action =
|
||||
build_routine_action(&normalized.name, &normalized.prompt, &normalized.execution);
|
||||
@@ -1274,6 +1292,7 @@ impl Tool for RoutineUpdateTool {
|
||||
let start = std::time::Instant::now();
|
||||
|
||||
let name = require_str(¶ms, "name")?;
|
||||
stash_last_routine_name(ctx, name).await;
|
||||
|
||||
let mut routine = self
|
||||
.store
|
||||
@@ -1411,11 +1430,24 @@ impl Tool for RoutineDeleteTool {
|
||||
) -> Result<ToolOutput, ToolError> {
|
||||
let start = std::time::Instant::now();
|
||||
|
||||
let name = require_str(¶ms, "name")?;
|
||||
let name = if let Some(name) = params.get("name").and_then(|v| v.as_str()) {
|
||||
if name.trim().is_empty() {
|
||||
return Err(ToolError::InvalidParameters(
|
||||
"'name' parameter cannot be empty".to_string(),
|
||||
));
|
||||
}
|
||||
name.to_string()
|
||||
} else {
|
||||
restore_last_routine_name(ctx).await.ok_or_else(|| {
|
||||
ToolError::InvalidParameters(
|
||||
"missing 'name' parameter and no previous routine target to infer".to_string(),
|
||||
)
|
||||
})?
|
||||
};
|
||||
|
||||
let routine = self
|
||||
.store
|
||||
.get_routine_by_name(&ctx.user_id, name)
|
||||
.get_routine_by_name(&ctx.user_id, &name)
|
||||
.await
|
||||
.map_err(|e| ToolError::ExecutionFailed(format!("DB error: {e}")))?
|
||||
.ok_or_else(|| ToolError::ExecutionFailed(format!("routine '{}' not found", name)))?;
|
||||
@@ -1430,7 +1462,7 @@ impl Tool for RoutineDeleteTool {
|
||||
self.engine.refresh_event_cache().await;
|
||||
|
||||
let result = serde_json::json!({
|
||||
"name": name,
|
||||
"name": &name,
|
||||
"deleted": deleted,
|
||||
});
|
||||
|
||||
@@ -1836,6 +1868,20 @@ mod tests {
|
||||
assert_eq!(parsed.cooldown_secs, 30);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn build_routine_trigger_normalizes_cron_schedule() {
|
||||
let trigger = build_routine_trigger(&NormalizedTriggerRequest::Cron {
|
||||
schedule: "0 0 9 * * MON-FRI".to_string(),
|
||||
timezone: Some("UTC".to_string()),
|
||||
});
|
||||
|
||||
assert!(matches!(
|
||||
trigger,
|
||||
Trigger::Cron { schedule, timezone }
|
||||
if schedule == "0 0 9 * * MON-FRI *" && timezone.as_deref() == Some("UTC")
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_grouped_message_event_with_tools() {
|
||||
let params = serde_json::json!({
|
||||
|
||||
+330
-9
@@ -8,12 +8,13 @@ use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use secrecy::SecretString;
|
||||
use tokio::sync::RwLock;
|
||||
|
||||
use crate::context::JobContext;
|
||||
use crate::secrets::SecretsStore;
|
||||
use crate::tools::mcp::auth::refresh_access_token;
|
||||
use crate::tools::mcp::config::McpServerConfig;
|
||||
use crate::tools::mcp::config::{McpAuthSource, McpServerConfig};
|
||||
use crate::tools::mcp::http_transport::HttpMcpTransport;
|
||||
use crate::tools::mcp::protocol::{
|
||||
CallToolResult, InitializeResult, ListToolsResult, McpRequest, McpResponse, McpTool,
|
||||
@@ -46,6 +47,13 @@ pub struct McpClient {
|
||||
/// Session manager (shared across clients).
|
||||
session_manager: Option<Arc<McpSessionManager>>,
|
||||
|
||||
/// NEAR AI auth/session manager for companion MCP servers that reuse the
|
||||
/// active provider bearer token.
|
||||
nearai_session_manager: Option<Arc<crate::llm::SessionManager>>,
|
||||
|
||||
/// Resolved NEAR AI API key for companion MCP servers.
|
||||
nearai_api_key: Option<SecretString>,
|
||||
|
||||
/// Secrets store for retrieving access tokens.
|
||||
secrets: Option<Arc<dyn SecretsStore + Send + Sync>>,
|
||||
|
||||
@@ -80,6 +88,8 @@ impl McpClient {
|
||||
next_id: AtomicU64::new(1),
|
||||
tools_cache: RwLock::new(None),
|
||||
session_manager: None,
|
||||
nearai_session_manager: None,
|
||||
nearai_api_key: None,
|
||||
secrets: None,
|
||||
user_id: "default".to_string(),
|
||||
server_config: None,
|
||||
@@ -103,6 +113,8 @@ impl McpClient {
|
||||
next_id: AtomicU64::new(1),
|
||||
tools_cache: RwLock::new(None),
|
||||
session_manager: None,
|
||||
nearai_session_manager: None,
|
||||
nearai_api_key: None,
|
||||
secrets: None,
|
||||
user_id: "default".to_string(),
|
||||
server_config: None,
|
||||
@@ -117,7 +129,15 @@ impl McpClient {
|
||||
/// The config must use HTTP transport (the default); for stdio/UDS use `new_with_transport`.
|
||||
///
|
||||
/// Returns an error if the config uses a non-HTTP transport.
|
||||
///
|
||||
/// **Note:** The session manager is NOT wired into the transport. For
|
||||
/// production use, prefer `create_client_from_config()` which constructs
|
||||
/// the transport with session tracking.
|
||||
#[cfg(test)]
|
||||
pub fn new_with_config(config: McpServerConfig) -> Result<Self, ToolError> {
|
||||
config
|
||||
.validate()
|
||||
.map_err(|e| ToolError::InvalidParameters(e.to_string()))?;
|
||||
if !matches!(
|
||||
config.effective_transport(),
|
||||
crate::tools::mcp::config::EffectiveTransport::Http
|
||||
@@ -139,6 +159,8 @@ impl McpClient {
|
||||
next_id: AtomicU64::new(1),
|
||||
tools_cache: RwLock::new(None),
|
||||
session_manager: None,
|
||||
nearai_session_manager: None,
|
||||
nearai_api_key: None,
|
||||
secrets: None,
|
||||
user_id: "default".to_string(),
|
||||
custom_headers: config.headers.clone(),
|
||||
@@ -170,6 +192,8 @@ impl McpClient {
|
||||
next_id: AtomicU64::new(1),
|
||||
tools_cache: RwLock::new(None),
|
||||
session_manager: Some(session_manager),
|
||||
nearai_session_manager: None,
|
||||
nearai_api_key: None,
|
||||
secrets: Some(secrets),
|
||||
user_id: user_id.into(),
|
||||
server_config: Some(config),
|
||||
@@ -206,6 +230,8 @@ impl McpClient {
|
||||
next_id: AtomicU64::new(1),
|
||||
tools_cache: RwLock::new(None),
|
||||
session_manager,
|
||||
nearai_session_manager: None,
|
||||
nearai_api_key: None,
|
||||
secrets,
|
||||
user_id: user_id.into(),
|
||||
server_config,
|
||||
@@ -214,12 +240,34 @@ impl McpClient {
|
||||
}
|
||||
}
|
||||
|
||||
/// Attach a session manager for Streamable HTTP session tracking.
|
||||
/// Attach a session manager to the **client** only.
|
||||
///
|
||||
/// **Warning:** This does NOT wire the session manager into the underlying
|
||||
/// `HttpMcpTransport`, so the transport will not capture `Mcp-Session-Id`
|
||||
/// from responses. For production use, construct the transport with
|
||||
/// `HttpMcpTransport::with_session_manager()` and pass it to
|
||||
/// `new_with_transport()` instead. See `create_client_from_config()`.
|
||||
#[cfg(test)]
|
||||
pub fn with_session_manager(mut self, session_manager: Arc<McpSessionManager>) -> Self {
|
||||
self.session_manager = Some(session_manager);
|
||||
self
|
||||
}
|
||||
|
||||
/// Attach the NEAR AI session manager for companion MCP auth reuse.
|
||||
pub fn with_nearai_session_manager(
|
||||
mut self,
|
||||
nearai_session_manager: Arc<crate::llm::SessionManager>,
|
||||
) -> Self {
|
||||
self.nearai_session_manager = Some(nearai_session_manager);
|
||||
self
|
||||
}
|
||||
|
||||
/// Attach the resolved NEAR AI API key for companion MCP auth reuse.
|
||||
pub fn with_nearai_api_key(mut self, nearai_api_key: Option<SecretString>) -> Self {
|
||||
self.nearai_api_key = nearai_api_key;
|
||||
self
|
||||
}
|
||||
|
||||
/// Get the server name.
|
||||
pub fn server_name(&self) -> &str {
|
||||
&self.server_name
|
||||
@@ -235,6 +283,12 @@ impl McpClient {
|
||||
self.session_manager.is_some()
|
||||
}
|
||||
|
||||
/// Get the underlying transport (test-only).
|
||||
#[cfg(test)]
|
||||
pub(crate) fn transport(&self) -> &Arc<dyn McpTransport> {
|
||||
&self.transport
|
||||
}
|
||||
|
||||
/// Get the next request ID.
|
||||
fn next_request_id(&self) -> u64 {
|
||||
self.next_id.fetch_add(1, Ordering::SeqCst)
|
||||
@@ -248,6 +302,9 @@ impl McpClient {
|
||||
let Some(ref config) = self.server_config else {
|
||||
return Ok(None);
|
||||
};
|
||||
if config.uses_runtime_auth_source() {
|
||||
return Ok(None);
|
||||
}
|
||||
match secrets
|
||||
.get_decrypted(&self.user_id, &config.token_secret_name())
|
||||
.await
|
||||
@@ -261,6 +318,36 @@ impl McpClient {
|
||||
}
|
||||
}
|
||||
|
||||
/// Resolve a runtime-provided auth token for companion MCP servers.
|
||||
async fn get_runtime_auth_token(&self) -> Result<Option<String>, ToolError> {
|
||||
let Some(ref config) = self.server_config else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
match config.auth_source {
|
||||
Some(McpAuthSource::NearAi) => {
|
||||
let Some(ref session_manager) = self.nearai_session_manager else {
|
||||
return Err(ToolError::ExternalService(
|
||||
"Missing NEAR AI session manager for companion MCP server".to_string(),
|
||||
));
|
||||
};
|
||||
|
||||
crate::llm::resolve_nearai_bearer_token_if_available(
|
||||
self.nearai_api_key.as_ref(),
|
||||
session_manager,
|
||||
)
|
||||
.await
|
||||
.map_err(|e| {
|
||||
ToolError::ExternalService(format!(
|
||||
"Failed to resolve NEAR AI token for MCP server '{}': {}",
|
||||
self.server_name, e
|
||||
))
|
||||
})
|
||||
}
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
/// Build the headers map for a request (auth, session-id, custom headers).
|
||||
///
|
||||
/// Custom headers are applied first. OAuth token injection is skipped if the
|
||||
@@ -274,6 +361,9 @@ impl McpClient {
|
||||
.custom_headers
|
||||
.keys()
|
||||
.any(|k| k.eq_ignore_ascii_case("authorization"));
|
||||
if !has_custom_auth && let Some(token) = self.get_runtime_auth_token().await? {
|
||||
headers.insert("Authorization".to_string(), format!("Bearer {}", token));
|
||||
}
|
||||
if !has_custom_auth && let Some(token) = self.get_access_token().await? {
|
||||
let trimmed = token.trim();
|
||||
if !trimmed.is_empty() {
|
||||
@@ -494,13 +584,12 @@ impl McpClient {
|
||||
)));
|
||||
}
|
||||
|
||||
response
|
||||
let raw_result = response
|
||||
.result
|
||||
.ok_or_else(|| ToolError::ExternalService("No result in MCP response".to_string()))
|
||||
.and_then(|r| {
|
||||
serde_json::from_value(r)
|
||||
.map_err(|e| ToolError::ExternalService(format!("Invalid tool result: {}", e)))
|
||||
})
|
||||
.ok_or_else(|| ToolError::ExternalService("No result in MCP response".to_string()))?;
|
||||
|
||||
serde_json::from_value(raw_result)
|
||||
.map_err(|e| ToolError::ExternalService(format!("Invalid tool result: {}", e)))
|
||||
}
|
||||
|
||||
/// Clear the tools cache.
|
||||
@@ -547,6 +636,8 @@ impl Clone for McpClient {
|
||||
next_id: AtomicU64::new(self.next_id.load(Ordering::SeqCst)),
|
||||
tools_cache: RwLock::new(None),
|
||||
session_manager: self.session_manager.clone(),
|
||||
nearai_session_manager: self.nearai_session_manager.clone(),
|
||||
nearai_api_key: self.nearai_api_key.clone(),
|
||||
secrets: self.secrets.clone(),
|
||||
user_id: self.user_id.clone(),
|
||||
server_config: self.server_config.clone(),
|
||||
@@ -594,7 +685,7 @@ impl Tool for McpToolWrapper {
|
||||
// Strip top-level null values before forwarding — LLMs often emit
|
||||
// `"field": null` for optional params, but many MCP servers reject
|
||||
// explicit nulls for fields that should simply be absent.
|
||||
let params = strip_top_level_nulls(params);
|
||||
let params = normalize_mcp_tool_arguments(&self.tool.name, strip_top_level_nulls(params));
|
||||
|
||||
let result = self.client.call_tool(&self.tool.name, params).await?;
|
||||
let content: String = result
|
||||
@@ -638,6 +729,31 @@ fn strip_top_level_nulls(value: serde_json::Value) -> serde_json::Value {
|
||||
}
|
||||
}
|
||||
|
||||
fn normalize_mcp_tool_arguments(tool_name: &str, value: serde_json::Value) -> serde_json::Value {
|
||||
if tool_name != "web_search" {
|
||||
return value;
|
||||
}
|
||||
|
||||
let serde_json::Value::Object(mut map) = value else {
|
||||
return value;
|
||||
};
|
||||
|
||||
// Keep this intentionally narrow: only strip optional fields that the
|
||||
// model frequently emits as empty strings. Provider-specific validation
|
||||
// should remain server-side, and tighter constraints should come from the
|
||||
// tool schema rather than client-side normalization.
|
||||
map.retain(|key, value| match key.as_str() {
|
||||
// Only strip known optional string fields. Never remove required
|
||||
// fields like `query`, even when the model emits an empty string.
|
||||
"country" | "freshness" | "goggles" | "result_filter" | "search_lang" | "ui_lang" => {
|
||||
!value.as_str().is_some_and(|s| s.trim().is_empty())
|
||||
}
|
||||
_ => true,
|
||||
});
|
||||
|
||||
serde_json::Value::Object(map)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
@@ -805,6 +921,138 @@ mod tests {
|
||||
assert!(client.has_session_manager());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_build_request_headers_with_nearai_runtime_auth() {
|
||||
use crate::llm::{
|
||||
SessionConfig as NearAiSessionConfig, SessionManager as NearAiSessionManager,
|
||||
};
|
||||
use secrecy::SecretString;
|
||||
|
||||
let config = McpServerConfig::new(
|
||||
crate::tools::mcp::config::NEARAI_COMPANION_MCP_NAME,
|
||||
"http://localhost:3000/mcp",
|
||||
)
|
||||
.with_auth_source(crate::tools::mcp::config::McpAuthSource::NearAi);
|
||||
let nearai_session = Arc::new(NearAiSessionManager::new(NearAiSessionConfig::default()));
|
||||
nearai_session
|
||||
.set_token(SecretString::from("sess_test_token"))
|
||||
.await;
|
||||
|
||||
let client = McpClient::new_with_config(config)
|
||||
.expect("valid MCP config")
|
||||
.with_nearai_session_manager(nearai_session);
|
||||
let headers = client.build_request_headers().await.expect("headers");
|
||||
|
||||
assert_eq!(
|
||||
headers.get("Authorization").map(String::as_str),
|
||||
Some("Bearer sess_test_token")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_build_request_headers_without_nearai_auth_does_not_trigger_login() {
|
||||
use crate::llm::{
|
||||
SessionConfig as NearAiSessionConfig, SessionManager as NearAiSessionManager,
|
||||
};
|
||||
|
||||
let config = McpServerConfig::new(
|
||||
crate::tools::mcp::config::NEARAI_COMPANION_MCP_NAME,
|
||||
"http://localhost:3000/mcp",
|
||||
)
|
||||
.with_auth_source(crate::tools::mcp::config::McpAuthSource::NearAi);
|
||||
let nearai_session = Arc::new(NearAiSessionManager::new(NearAiSessionConfig::default()));
|
||||
|
||||
let client = McpClient::new_with_config(config)
|
||||
.expect("valid MCP config")
|
||||
.with_nearai_session_manager(nearai_session);
|
||||
let headers = client.build_request_headers().await.expect("headers");
|
||||
|
||||
assert!(
|
||||
!headers.contains_key("Authorization"),
|
||||
"runtime auth should stay absent when no token is available"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_build_request_headers_runtime_auth_ignores_persisted_mcp_token() {
|
||||
use crate::llm::{
|
||||
SessionConfig as NearAiSessionConfig, SessionManager as NearAiSessionManager,
|
||||
};
|
||||
use crate::secrets::{CreateSecretParams, DecryptedSecret, Secret, SecretError, SecretRef};
|
||||
use secrecy::SecretString;
|
||||
use uuid::Uuid;
|
||||
|
||||
struct PersistedTokenStore;
|
||||
|
||||
#[async_trait]
|
||||
impl crate::secrets::SecretsStore for PersistedTokenStore {
|
||||
async fn create(
|
||||
&self,
|
||||
_user_id: &str,
|
||||
_params: CreateSecretParams,
|
||||
) -> Result<Secret, SecretError> {
|
||||
unimplemented!()
|
||||
}
|
||||
async fn get(&self, _user_id: &str, _name: &str) -> Result<Secret, SecretError> {
|
||||
unimplemented!()
|
||||
}
|
||||
async fn get_decrypted(
|
||||
&self,
|
||||
_user_id: &str,
|
||||
_name: &str,
|
||||
) -> Result<DecryptedSecret, SecretError> {
|
||||
DecryptedSecret::from_bytes(b"persisted-mcp-token".to_vec())
|
||||
}
|
||||
async fn exists(&self, _user_id: &str, _name: &str) -> Result<bool, SecretError> {
|
||||
Ok(true)
|
||||
}
|
||||
async fn delete(&self, _user_id: &str, _name: &str) -> Result<bool, SecretError> {
|
||||
Ok(true)
|
||||
}
|
||||
async fn list(&self, _user_id: &str) -> Result<Vec<SecretRef>, SecretError> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
async fn record_usage(&self, _secret_id: Uuid) -> Result<(), SecretError> {
|
||||
Ok(())
|
||||
}
|
||||
async fn is_accessible(
|
||||
&self,
|
||||
_user_id: &str,
|
||||
_secret_name: &str,
|
||||
_allowed_secrets: &[String],
|
||||
) -> Result<bool, SecretError> {
|
||||
Ok(true)
|
||||
}
|
||||
}
|
||||
|
||||
let config = McpServerConfig::new(
|
||||
crate::tools::mcp::config::NEARAI_COMPANION_MCP_NAME,
|
||||
"http://localhost:3000/mcp",
|
||||
)
|
||||
.with_auth_source(crate::tools::mcp::config::McpAuthSource::NearAi);
|
||||
let nearai_session = Arc::new(NearAiSessionManager::new(NearAiSessionConfig::default()));
|
||||
nearai_session
|
||||
.set_token(SecretString::from("sess_runtime_token"))
|
||||
.await;
|
||||
|
||||
let secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync> =
|
||||
Arc::new(PersistedTokenStore);
|
||||
let client = McpClient::new_authenticated(
|
||||
config,
|
||||
Arc::new(McpSessionManager::new()),
|
||||
secrets,
|
||||
"test-user",
|
||||
)
|
||||
.with_nearai_session_manager(nearai_session);
|
||||
let headers = client.build_request_headers().await.expect("headers");
|
||||
|
||||
assert_eq!(
|
||||
headers.get("Authorization").map(String::as_str),
|
||||
Some("Bearer sess_runtime_token"),
|
||||
"runtime auth must win even if a persisted MCP token exists"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_next_request_id_monotonically_increasing() {
|
||||
let client = McpClient::new("http://localhost:1234");
|
||||
@@ -1186,6 +1434,20 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_new_with_config_rejects_invalid_runtime_auth_name() {
|
||||
let config = McpServerConfig::new("chat_api", "http://localhost:3000/mcp")
|
||||
.with_auth_source(crate::tools::mcp::config::McpAuthSource::NearAi);
|
||||
let err = match McpClient::new_with_config(config) {
|
||||
Ok(_) => panic!("invalid runtime-auth config must be rejected"),
|
||||
Err(err) => err.to_string(),
|
||||
};
|
||||
assert!(
|
||||
err.contains(crate::tools::mcp::config::NEARAI_COMPANION_MCP_NAME),
|
||||
"error should mention reserved companion requirement: {err}"
|
||||
);
|
||||
}
|
||||
|
||||
// --- Issue 13: McpToolWrapper unit tests ---
|
||||
|
||||
fn make_test_mcp_tool(destructive: bool) -> McpTool {
|
||||
@@ -1416,4 +1678,63 @@ mod tests {
|
||||
"Token must be trimmed before use in Authorization header"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_normalize_web_search_arguments_removes_empty_optional_fields() {
|
||||
let input = serde_json::json!({
|
||||
"query": "Rust MCP server example",
|
||||
"goggles": "",
|
||||
"result_filter": " ",
|
||||
"ui_lang": "en-US"
|
||||
});
|
||||
|
||||
let result = normalize_mcp_tool_arguments("web_search", input);
|
||||
let obj = result.as_object().unwrap();
|
||||
assert_eq!(obj["query"], "Rust MCP server example");
|
||||
assert_eq!(obj["ui_lang"], "en-US");
|
||||
assert!(!obj.contains_key("goggles"));
|
||||
assert!(!obj.contains_key("result_filter"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_normalize_web_search_arguments_strips_whitelisted_empty_optional_fields() {
|
||||
let input = serde_json::json!({
|
||||
"query": "Rust MCP server example",
|
||||
"goggles": "",
|
||||
"freshness": " ",
|
||||
"country": "US"
|
||||
});
|
||||
|
||||
let result = normalize_mcp_tool_arguments("web_search", input);
|
||||
let obj = result.as_object().unwrap();
|
||||
assert_eq!(obj["country"], "US");
|
||||
assert!(!obj.contains_key("freshness"));
|
||||
assert!(!obj.contains_key("goggles"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_normalize_web_search_arguments_preserves_empty_required_query() {
|
||||
let input = serde_json::json!({
|
||||
"query": " ",
|
||||
"goggles": "",
|
||||
"country": "US"
|
||||
});
|
||||
|
||||
let result = normalize_mcp_tool_arguments("web_search", input);
|
||||
let obj = result.as_object().unwrap();
|
||||
assert_eq!(obj["query"], " ");
|
||||
assert_eq!(obj["country"], "US");
|
||||
assert!(!obj.contains_key("goggles"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_normalize_mcp_tool_arguments_leaves_other_tools_unchanged() {
|
||||
let input = serde_json::json!({
|
||||
"goggles": "",
|
||||
"country": "us"
|
||||
});
|
||||
|
||||
let result = normalize_mcp_tool_arguments("other_tool", input.clone());
|
||||
assert_eq!(result, input);
|
||||
}
|
||||
}
|
||||
|
||||
+211
-3
@@ -51,6 +51,16 @@ pub struct McpServerConfig {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub oauth: Option<OAuthConfig>,
|
||||
|
||||
/// Built-in auth source provided by IronClaw at runtime.
|
||||
///
|
||||
/// This is used for companion MCP servers that should reuse an existing
|
||||
/// provider identity instead of running their own MCP OAuth flow.
|
||||
///
|
||||
/// Security: this field is runtime-only. Persisted user config must not be
|
||||
/// able to opt a server into reusing the active provider bearer token.
|
||||
#[serde(default, skip_serializing, skip_deserializing)]
|
||||
pub auth_source: Option<McpAuthSource>,
|
||||
|
||||
/// Whether this server is enabled.
|
||||
#[serde(default = "default_true")]
|
||||
pub enabled: bool,
|
||||
@@ -60,6 +70,14 @@ pub struct McpServerConfig {
|
||||
pub description: Option<String>,
|
||||
}
|
||||
|
||||
/// Runtime-provided auth sources for MCP companion servers.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum McpAuthSource {
|
||||
/// Reuse the active NEAR AI bearer token (session token or API key).
|
||||
NearAi,
|
||||
}
|
||||
|
||||
fn default_true() -> bool {
|
||||
true
|
||||
}
|
||||
@@ -73,6 +91,7 @@ impl McpServerConfig {
|
||||
transport: None,
|
||||
headers: HashMap::new(),
|
||||
oauth: None,
|
||||
auth_source: None,
|
||||
enabled: true,
|
||||
description: None,
|
||||
}
|
||||
@@ -95,6 +114,7 @@ impl McpServerConfig {
|
||||
}),
|
||||
headers: HashMap::new(),
|
||||
oauth: None,
|
||||
auth_source: None,
|
||||
enabled: true,
|
||||
description: None,
|
||||
}
|
||||
@@ -110,6 +130,7 @@ impl McpServerConfig {
|
||||
}),
|
||||
headers: HashMap::new(),
|
||||
oauth: None,
|
||||
auth_source: None,
|
||||
enabled: true,
|
||||
description: None,
|
||||
}
|
||||
@@ -121,6 +142,12 @@ impl McpServerConfig {
|
||||
self
|
||||
}
|
||||
|
||||
/// Set a runtime-provided auth source.
|
||||
pub fn with_auth_source(mut self, auth_source: McpAuthSource) -> Self {
|
||||
self.auth_source = Some(auth_source);
|
||||
self
|
||||
}
|
||||
|
||||
/// Set description.
|
||||
pub fn with_description(mut self, description: impl Into<String>) -> Self {
|
||||
self.description = Some(description.into());
|
||||
@@ -154,6 +181,15 @@ impl McpServerConfig {
|
||||
});
|
||||
}
|
||||
|
||||
if self.uses_runtime_auth_source() && !is_nearai_companion_server_name(&self.name) {
|
||||
return Err(ConfigError::InvalidConfig {
|
||||
reason: format!(
|
||||
"Runtime auth source is only allowed for reserved server '{}'",
|
||||
NEARAI_COMPANION_MCP_NAME
|
||||
),
|
||||
});
|
||||
}
|
||||
|
||||
match self.effective_transport() {
|
||||
EffectiveTransport::Http => {
|
||||
if self.url.is_empty() {
|
||||
@@ -222,6 +258,11 @@ impl McpServerConfig {
|
||||
.any(|k| k.eq_ignore_ascii_case("authorization"))
|
||||
}
|
||||
|
||||
/// Check if this server uses a built-in runtime auth bridge.
|
||||
pub fn uses_runtime_auth_source(&self) -> bool {
|
||||
self.auth_source.is_some()
|
||||
}
|
||||
|
||||
/// Check if this server requires authentication.
|
||||
///
|
||||
/// Returns true if OAuth is pre-configured OR if this is a remote HTTPS server
|
||||
@@ -234,7 +275,7 @@ impl McpServerConfig {
|
||||
return false;
|
||||
}
|
||||
|
||||
if self.oauth.is_some() {
|
||||
if self.oauth.is_some() || self.uses_runtime_auth_source() {
|
||||
return true;
|
||||
}
|
||||
// Remote HTTPS servers need auth handling (DCR, token refresh, 401 detection).
|
||||
@@ -260,6 +301,66 @@ impl McpServerConfig {
|
||||
}
|
||||
}
|
||||
|
||||
/// Reserved name used for the companion MCP server derived from active NEAR AI config.
|
||||
pub const NEARAI_COMPANION_MCP_NAME: &str = "_nearai_companion_mcp";
|
||||
|
||||
pub fn is_nearai_companion_server_name(name: &str) -> bool {
|
||||
name == NEARAI_COMPANION_MCP_NAME
|
||||
}
|
||||
|
||||
fn strip_reserved_nearai_companion_servers(config: &mut McpServersFile, source: &str) -> usize {
|
||||
let len_before = config.servers.len();
|
||||
config
|
||||
.servers
|
||||
.retain(|server| !is_nearai_companion_server_name(&server.name));
|
||||
let removed = len_before.saturating_sub(config.servers.len());
|
||||
|
||||
if removed > 0 {
|
||||
tracing::warn!(
|
||||
count = removed,
|
||||
source,
|
||||
"Ignoring persisted reserved MCP companion config(s); this name is system-managed"
|
||||
);
|
||||
}
|
||||
|
||||
removed
|
||||
}
|
||||
|
||||
/// Build the companion MCP server from the active NEAR AI config.
|
||||
///
|
||||
/// The MCP endpoint is treated as a sibling to the versioned REST API:
|
||||
/// `https://host/v1` becomes `https://host/mcp`.
|
||||
pub fn derive_nearai_companion_mcp_server(
|
||||
config: &crate::config::Config,
|
||||
) -> Option<McpServerConfig> {
|
||||
derive_nearai_companion_mcp_server_from_llm(&config.llm)
|
||||
}
|
||||
|
||||
/// Build the companion MCP server from an LLM config.
|
||||
///
|
||||
/// This lighter-weight helper is used by CLI code paths that should not need
|
||||
/// to resolve the full application config (and therefore should not require
|
||||
/// database configuration) just to discover the derived companion MCP server.
|
||||
pub fn derive_nearai_companion_mcp_server_from_llm(
|
||||
llm: &crate::config::LlmConfig,
|
||||
) -> Option<McpServerConfig> {
|
||||
if llm.backend != "nearai" {
|
||||
return None;
|
||||
}
|
||||
|
||||
let base = llm.nearai.base_url.trim_end_matches('/');
|
||||
let mcp_base = base
|
||||
.strip_suffix("/v1")
|
||||
.unwrap_or(base)
|
||||
.trim_end_matches('/');
|
||||
|
||||
Some(
|
||||
McpServerConfig::new(NEARAI_COMPANION_MCP_NAME, format!("{mcp_base}/mcp"))
|
||||
.with_auth_source(McpAuthSource::NearAi)
|
||||
.with_description("Companion MCP server derived from the active NEAR AI provider"),
|
||||
)
|
||||
}
|
||||
|
||||
/// OAuth 2.1 configuration for an MCP server.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct OAuthConfig {
|
||||
@@ -356,6 +457,16 @@ impl McpServersFile {
|
||||
}
|
||||
}
|
||||
|
||||
/// Insert a server only if no server with the same name already exists.
|
||||
pub fn insert_if_absent(&mut self, config: McpServerConfig) -> bool {
|
||||
if self.get(&config.name).is_some() {
|
||||
false
|
||||
} else {
|
||||
self.servers.push(config);
|
||||
true
|
||||
}
|
||||
}
|
||||
|
||||
/// Remove a server by name.
|
||||
pub fn remove(&mut self, name: &str) -> bool {
|
||||
let len_before = self.servers.len();
|
||||
@@ -410,7 +521,8 @@ pub async fn load_mcp_servers_from(path: impl AsRef<Path>) -> Result<McpServersF
|
||||
}
|
||||
|
||||
let content = fs::read_to_string(path).await?;
|
||||
let config: McpServersFile = serde_json::from_str(&content)?;
|
||||
let mut config: McpServersFile = serde_json::from_str(&content)?;
|
||||
strip_reserved_nearai_companion_servers(&mut config, &path.display().to_string());
|
||||
|
||||
// Validate every server on load so corrupted configs are caught early
|
||||
for server in &config.servers {
|
||||
@@ -452,6 +564,15 @@ pub async fn save_mcp_servers_to(
|
||||
|
||||
/// Add a new MCP server configuration.
|
||||
pub async fn add_mcp_server(config: McpServerConfig) -> Result<(), ConfigError> {
|
||||
if is_nearai_companion_server_name(&config.name) {
|
||||
return Err(ConfigError::InvalidConfig {
|
||||
reason: format!(
|
||||
"Server name '{}' is reserved for the NEAR AI companion MCP server",
|
||||
config.name
|
||||
),
|
||||
});
|
||||
}
|
||||
|
||||
config.validate()?;
|
||||
|
||||
let mut servers = load_mcp_servers().await?;
|
||||
@@ -499,7 +620,8 @@ pub async fn load_mcp_servers_from_db(
|
||||
) -> Result<McpServersFile, ConfigError> {
|
||||
match store.get_setting(user_id, "mcp_servers").await {
|
||||
Ok(Some(value)) => {
|
||||
let config: McpServersFile = serde_json::from_value(value)?;
|
||||
let mut config: McpServersFile = serde_json::from_value(value)?;
|
||||
strip_reserved_nearai_companion_servers(&mut config, "database");
|
||||
// Validate every server on load so corrupted DB configs are caught early
|
||||
for server in &config.servers {
|
||||
server.validate().map_err(|e| ConfigError::InvalidConfig {
|
||||
@@ -542,6 +664,15 @@ pub async fn add_mcp_server_db(
|
||||
user_id: &str,
|
||||
config: McpServerConfig,
|
||||
) -> Result<(), ConfigError> {
|
||||
if is_nearai_companion_server_name(&config.name) {
|
||||
return Err(ConfigError::InvalidConfig {
|
||||
reason: format!(
|
||||
"Server name '{}' is reserved for the NEAR AI companion MCP server",
|
||||
config.name
|
||||
),
|
||||
});
|
||||
}
|
||||
|
||||
config.validate()?;
|
||||
|
||||
let mut servers = load_mcp_servers_from_db(store, user_id).await?;
|
||||
@@ -718,6 +849,69 @@ mod tests {
|
||||
assert!(config.servers.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_load_drops_reserved_nearai_companion_server() {
|
||||
let dir = tempdir().unwrap();
|
||||
let path = dir.path().join("mcp-servers.json");
|
||||
|
||||
let persisted = serde_json::json!({
|
||||
"servers": [
|
||||
{
|
||||
"name": NEARAI_COMPANION_MCP_NAME,
|
||||
"url": "https://evil.example.com/mcp",
|
||||
"enabled": true,
|
||||
"auth_source": "near_ai"
|
||||
},
|
||||
{
|
||||
"name": "notion",
|
||||
"url": "https://mcp.notion.com",
|
||||
"enabled": true
|
||||
}
|
||||
]
|
||||
});
|
||||
tokio::fs::write(&path, persisted.to_string())
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let config = load_mcp_servers_from(&path).await.unwrap();
|
||||
assert_eq!(config.servers.len(), 1);
|
||||
assert!(config.get(NEARAI_COMPANION_MCP_NAME).is_none());
|
||||
assert_eq!(
|
||||
config.get("notion").map(|server| server.url.as_str()),
|
||||
Some("https://mcp.notion.com")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_deserialize_ignores_persisted_auth_source() {
|
||||
let raw = serde_json::json!({
|
||||
"name": "user-managed",
|
||||
"url": "https://mcp.example.com",
|
||||
"enabled": true,
|
||||
"auth_source": "near_ai"
|
||||
});
|
||||
|
||||
let server: McpServerConfig = serde_json::from_value(raw).expect("server");
|
||||
assert_eq!(server.auth_source, None);
|
||||
}
|
||||
|
||||
#[cfg(feature = "libsql")]
|
||||
#[test]
|
||||
fn test_derive_nearai_companion_mcp_server_strips_trailing_v1() {
|
||||
let mut config = crate::config::Config::for_testing(
|
||||
std::env::temp_dir().join("ironclaw-test-companion.db"),
|
||||
std::env::temp_dir().join("ironclaw-test-skills"),
|
||||
std::env::temp_dir().join("ironclaw-test-installed-skills"),
|
||||
);
|
||||
config.llm.backend = "nearai".to_string();
|
||||
config.llm.nearai.base_url = "https://private.near.ai/v1".to_string();
|
||||
|
||||
let server = derive_nearai_companion_mcp_server(&config).expect("companion server");
|
||||
assert_eq!(server.name, NEARAI_COMPANION_MCP_NAME);
|
||||
assert_eq!(server.url, "https://private.near.ai/mcp");
|
||||
assert_eq!(server.auth_source, Some(McpAuthSource::NearAi));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_load_rejects_corrupted_headers() {
|
||||
let dir = tempdir().unwrap();
|
||||
@@ -763,6 +957,20 @@ mod tests {
|
||||
assert!(config.requires_auth());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate_rejects_runtime_auth_on_user_managed_server() {
|
||||
let config = McpServerConfig::new("user-managed", "https://mcp.example.com")
|
||||
.with_auth_source(McpAuthSource::NearAi);
|
||||
|
||||
let err = config
|
||||
.validate()
|
||||
.expect_err("runtime auth should be reserved for the companion server");
|
||||
assert!(
|
||||
err.to_string().contains(NEARAI_COMPANION_MCP_NAME),
|
||||
"expected reserved-name validation message, got: {err}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_requires_auth_remote_https_without_oauth() {
|
||||
// Remote HTTPS servers need auth even without pre-configured OAuth (DCR)
|
||||
|
||||
+133
-16
@@ -7,6 +7,7 @@ use std::sync::Arc;
|
||||
|
||||
use crate::secrets::SecretsStore;
|
||||
use crate::tools::mcp::config::{EffectiveTransport, McpServerConfig};
|
||||
use crate::tools::mcp::http_transport::HttpMcpTransport;
|
||||
use crate::tools::mcp::{McpClient, McpProcessManager, McpSessionManager, McpTransport};
|
||||
|
||||
/// Error returned when MCP client creation fails.
|
||||
@@ -20,6 +21,8 @@ pub enum McpFactoryError {
|
||||
UnixNotSupported { name: String },
|
||||
#[error("Invalid configuration for MCP server '{name}': {reason}")]
|
||||
InvalidConfig { name: String, reason: String },
|
||||
#[error("Missing runtime auth context for MCP server '{name}': {reason}")]
|
||||
MissingRuntimeAuthContext { name: String, reason: String },
|
||||
}
|
||||
|
||||
/// Create an `McpClient` from a server configuration, dispatching on the
|
||||
@@ -27,6 +30,8 @@ pub enum McpFactoryError {
|
||||
pub async fn create_client_from_config(
|
||||
server: McpServerConfig,
|
||||
session_manager: &Arc<McpSessionManager>,
|
||||
nearai_session_manager: Option<Arc<crate::llm::SessionManager>>,
|
||||
nearai_api_key: Option<secrecy::SecretString>,
|
||||
process_manager: &Arc<McpProcessManager>,
|
||||
secrets: Option<Arc<dyn SecretsStore + Send + Sync>>,
|
||||
user_id: &str,
|
||||
@@ -78,33 +83,61 @@ pub async fn create_client_from_config(
|
||||
Err(McpFactoryError::UnixNotSupported { name: server_name })
|
||||
}
|
||||
EffectiveTransport::Http => {
|
||||
if server.uses_runtime_auth_source() {
|
||||
let nearai_session_manager = nearai_session_manager.ok_or_else(|| {
|
||||
McpFactoryError::MissingRuntimeAuthContext {
|
||||
name: server_name.clone(),
|
||||
reason: "NearAI companion MCP servers require a NearAI session manager"
|
||||
.to_string(),
|
||||
}
|
||||
})?;
|
||||
|
||||
let transport = Arc::new(
|
||||
HttpMcpTransport::new(server.url.clone(), server.name.clone())
|
||||
.with_session_manager(Arc::clone(session_manager)),
|
||||
);
|
||||
|
||||
return Ok(McpClient::new_with_transport(
|
||||
server.name.clone(),
|
||||
transport,
|
||||
Some(Arc::clone(session_manager)),
|
||||
secrets,
|
||||
user_id,
|
||||
Some(server),
|
||||
)
|
||||
.with_nearai_session_manager(nearai_session_manager)
|
||||
.with_nearai_api_key(nearai_api_key));
|
||||
}
|
||||
if let Some(ref secrets) = secrets {
|
||||
let has_tokens =
|
||||
crate::tools::mcp::is_authenticated(&server, secrets, user_id).await;
|
||||
|
||||
if has_tokens || server.requires_auth() {
|
||||
Ok(McpClient::new_authenticated(
|
||||
return Ok(McpClient::new_authenticated(
|
||||
server,
|
||||
Arc::clone(session_manager),
|
||||
Arc::clone(secrets),
|
||||
user_id,
|
||||
))
|
||||
} else {
|
||||
Ok(McpClient::new_with_config(server)
|
||||
.map_err(|e| McpFactoryError::InvalidConfig {
|
||||
name: server_name.clone(),
|
||||
reason: e.to_string(),
|
||||
})?
|
||||
.with_session_manager(Arc::clone(session_manager)))
|
||||
));
|
||||
}
|
||||
} else {
|
||||
Ok(McpClient::new_with_config(server)
|
||||
.map_err(|e| McpFactoryError::InvalidConfig {
|
||||
name: server_name,
|
||||
reason: e.to_string(),
|
||||
})?
|
||||
.with_session_manager(Arc::clone(session_manager)))
|
||||
}
|
||||
|
||||
// Non-OAuth HTTP: wire the session manager into the *transport* so
|
||||
// it captures `Mcp-Session-Id` from responses. Passing it only to
|
||||
// the client (via `with_session_manager`) is not enough — the
|
||||
// transport must know about it to read/write the header.
|
||||
let transport = Arc::new(
|
||||
HttpMcpTransport::new(server.url.clone(), server.name.clone())
|
||||
.with_session_manager(Arc::clone(session_manager)),
|
||||
);
|
||||
Ok(McpClient::new_with_transport(
|
||||
server.name.clone(),
|
||||
transport,
|
||||
Some(Arc::clone(session_manager)),
|
||||
secrets,
|
||||
user_id,
|
||||
Some(server),
|
||||
))
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -122,6 +155,8 @@ mod tests {
|
||||
let client = create_client_from_config(
|
||||
server,
|
||||
&session_manager,
|
||||
None,
|
||||
None,
|
||||
&process_manager,
|
||||
None,
|
||||
"test-user",
|
||||
@@ -134,4 +169,86 @@ mod tests {
|
||||
"non-OAuth HTTP clients must carry a session manager"
|
||||
);
|
||||
}
|
||||
|
||||
/// Regression test: the factory must wire the session manager into the
|
||||
/// *transport*, not just the client. Otherwise the transport never
|
||||
/// captures `Mcp-Session-Id` from responses and subsequent requests
|
||||
/// lack the header, causing the server to reject them.
|
||||
#[tokio::test]
|
||||
async fn test_factory_non_oauth_http_transport_captures_session_id() {
|
||||
use axum::http::header::HeaderName;
|
||||
use axum::{Router, http::StatusCode, response::IntoResponse, routing::post};
|
||||
use tokio::net::TcpListener;
|
||||
|
||||
const SESSION_ID: &str = "test-session-abc123";
|
||||
|
||||
async fn session_echo() -> impl IntoResponse {
|
||||
let body = serde_json::json!({
|
||||
"jsonrpc": "2.0",
|
||||
"id": 1,
|
||||
"result": {}
|
||||
})
|
||||
.to_string();
|
||||
(
|
||||
StatusCode::OK,
|
||||
[(
|
||||
HeaderName::from_static("mcp-session-id"),
|
||||
SESSION_ID.to_string(),
|
||||
)],
|
||||
body,
|
||||
)
|
||||
}
|
||||
|
||||
let app = Router::new().route("/", post(session_echo));
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let addr = listener.local_addr().unwrap();
|
||||
let url = format!("http://127.0.0.1:{}", addr.port());
|
||||
|
||||
tokio::spawn(async move {
|
||||
axum::serve(listener, app).await.unwrap();
|
||||
});
|
||||
|
||||
let server = McpServerConfig::new("session-test", &url);
|
||||
let session_manager = Arc::new(McpSessionManager::new());
|
||||
let process_manager = Arc::new(McpProcessManager::new());
|
||||
|
||||
let client = create_client_from_config(
|
||||
server,
|
||||
&session_manager,
|
||||
None,
|
||||
None,
|
||||
&process_manager,
|
||||
None,
|
||||
"test-user",
|
||||
)
|
||||
.await
|
||||
.expect("factory should succeed for HTTP config");
|
||||
|
||||
// Pre-create a session entry so that update_session_id has something to update.
|
||||
// In production, the MCP initialize handshake calls get_or_create before responses arrive.
|
||||
session_manager.get_or_create("session-test", &url).await;
|
||||
|
||||
// Send a request through the client's transport to trigger session capture.
|
||||
use crate::tools::mcp::protocol::McpRequest;
|
||||
let request = McpRequest {
|
||||
jsonrpc: "2.0".to_string(),
|
||||
id: Some(1),
|
||||
method: "test".to_string(),
|
||||
params: Some(serde_json::json!({})),
|
||||
};
|
||||
let headers = std::collections::HashMap::new();
|
||||
client
|
||||
.transport()
|
||||
.send(&request, &headers)
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
// Verify the session manager captured the session ID from the response.
|
||||
let captured = session_manager.get_session_id("session-test").await;
|
||||
assert_eq!(
|
||||
captured.as_deref(),
|
||||
Some(SESSION_ID),
|
||||
"transport must capture Mcp-Session-Id into session manager"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -494,6 +494,34 @@ mod tests {
|
||||
assert_eq!(echoed["authorization"], "Bearer oauth-token");
|
||||
}
|
||||
|
||||
/// Regression test for #1436: 202 Accepted responses for notifications
|
||||
/// were parsed as JSON, causing "Failed to parse MCP response" errors
|
||||
/// that broke the MCP session handshake.
|
||||
#[tokio::test]
|
||||
async fn test_wire_202_accepted_for_notification() {
|
||||
use axum::{Router, http::StatusCode, routing::post};
|
||||
use tokio::net::TcpListener;
|
||||
|
||||
async fn accept_notification() -> StatusCode {
|
||||
StatusCode::ACCEPTED
|
||||
}
|
||||
|
||||
let app = Router::new().route("/", post(accept_notification));
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let addr = listener.local_addr().unwrap();
|
||||
let url = format!("http://127.0.0.1:{}", addr.port());
|
||||
|
||||
tokio::spawn(async move {
|
||||
axum::serve(listener, app).await.unwrap();
|
||||
});
|
||||
|
||||
let transport = HttpMcpTransport::new(&url, "test-202");
|
||||
let request = McpRequest::initialized_notification();
|
||||
let response = transport.send(&request, &HashMap::new()).await.unwrap();
|
||||
assert!(response.result.is_none());
|
||||
assert!(response.error.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_wire_custom_auth_preserved_when_no_per_request_auth() {
|
||||
let (url, _handle) = spawn_echo_server().await;
|
||||
|
||||
@@ -446,16 +446,14 @@ fn resolve_oauth_refresh_config(cap_file: &CapabilitiesFile) -> Option<OAuthRefr
|
||||
builtin.as_ref(),
|
||||
exchange_proxy_url.is_some(),
|
||||
);
|
||||
let gateway_token = crate::config::helpers::env_or_override("GATEWAY_AUTH_TOKEN")
|
||||
.map(|token| token.trim().to_string())
|
||||
.filter(|token| !token.is_empty());
|
||||
let oauth_proxy_auth_token = crate::cli::oauth_defaults::oauth_proxy_auth_token();
|
||||
|
||||
Some(OAuthRefreshConfig {
|
||||
token_url: oauth.token_url.clone(),
|
||||
client_id,
|
||||
client_secret,
|
||||
exchange_proxy_url,
|
||||
gateway_token,
|
||||
gateway_token: oauth_proxy_auth_token,
|
||||
secret_name: auth.secret_name.clone(),
|
||||
provider: auth.provider.clone(),
|
||||
})
|
||||
@@ -891,6 +889,11 @@ mod tests {
|
||||
AuthCapabilitySchema, CapabilitiesFile, OAuthConfigSchema,
|
||||
};
|
||||
|
||||
let _guard = lock_env();
|
||||
let _proxy_guard = set_env_var("IRONCLAW_OAUTH_EXCHANGE_URL", None);
|
||||
let _gateway_token_guard = set_env_var("GATEWAY_AUTH_TOKEN", None);
|
||||
let _oauth_proxy_token_guard = set_env_var("IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", None);
|
||||
|
||||
let caps = CapabilitiesFile {
|
||||
auth: Some(AuthCapabilitySchema {
|
||||
secret_name: "google_oauth_token".to_string(),
|
||||
@@ -982,6 +985,7 @@ mod tests {
|
||||
let _guard = lock_env();
|
||||
let _proxy_guard = set_env_var("IRONCLAW_OAUTH_EXCHANGE_URL", None);
|
||||
let _gateway_token_guard = set_env_var("GATEWAY_AUTH_TOKEN", None);
|
||||
let _oauth_proxy_token_guard = set_env_var("IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", None);
|
||||
|
||||
// google_oauth_token should fall back to built-in credentials
|
||||
let caps = CapabilitiesFile {
|
||||
@@ -1021,6 +1025,7 @@ mod tests {
|
||||
Some("https://compose-api.example.com"),
|
||||
);
|
||||
let _gateway_token_guard = set_env_var("GATEWAY_AUTH_TOKEN", Some("gateway-test-token"));
|
||||
let _oauth_proxy_token_guard = set_env_var("IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", None);
|
||||
let _client_id_guard =
|
||||
set_env_var("GOOGLE_OAUTH_CLIENT_ID", Some("hosted-google-client-id"));
|
||||
|
||||
@@ -1061,6 +1066,7 @@ mod tests {
|
||||
Some("https://compose-api.example.com"),
|
||||
);
|
||||
let _gateway_token_guard = set_env_var("GATEWAY_AUTH_TOKEN", Some("gateway-test-token"));
|
||||
let _oauth_proxy_token_guard = set_env_var("IRONCLAW_OAUTH_PROXY_AUTH_TOKEN", None);
|
||||
let _client_id_guard =
|
||||
set_env_var("GOOGLE_OAUTH_CLIENT_ID", Some("hosted-google-client-id"));
|
||||
let _client_secret_guard =
|
||||
@@ -1095,6 +1101,47 @@ mod tests {
|
||||
assert_eq!(config.gateway_token.as_deref(), Some("gateway-test-token"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_resolve_oauth_refresh_config_hosted_proxy_prefers_dedicated_proxy_auth_token() {
|
||||
use crate::tools::wasm::capabilities_schema::{
|
||||
AuthCapabilitySchema, CapabilitiesFile, OAuthConfigSchema,
|
||||
};
|
||||
|
||||
let _guard = lock_env();
|
||||
let _proxy_guard = set_env_var(
|
||||
"IRONCLAW_OAUTH_EXCHANGE_URL",
|
||||
Some("https://compose-api.example.com"),
|
||||
);
|
||||
let _gateway_token_guard = set_env_var("GATEWAY_AUTH_TOKEN", Some("gateway-test-token"));
|
||||
let _oauth_proxy_token_guard = set_env_var(
|
||||
"IRONCLAW_OAUTH_PROXY_AUTH_TOKEN",
|
||||
Some("shared-oauth-proxy-secret"),
|
||||
);
|
||||
let _client_id_guard =
|
||||
set_env_var("GOOGLE_OAUTH_CLIENT_ID", Some("hosted-google-client-id"));
|
||||
|
||||
let caps = CapabilitiesFile {
|
||||
auth: Some(AuthCapabilitySchema {
|
||||
secret_name: "google_oauth_token".to_string(),
|
||||
provider: Some("google".to_string()),
|
||||
oauth: Some(OAuthConfigSchema {
|
||||
authorization_url: "https://accounts.google.com/o/oauth2/v2/auth".to_string(),
|
||||
token_url: "https://oauth2.googleapis.com/token".to_string(),
|
||||
client_id_env: Some("GOOGLE_OAUTH_CLIENT_ID".to_string()),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let config = super::resolve_oauth_refresh_config(&caps).expect("hosted oauth config");
|
||||
assert_eq!(
|
||||
config.gateway_token.as_deref(),
|
||||
Some("shared-oauth-proxy-secret")
|
||||
);
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------
|
||||
// Security regression tests
|
||||
// ---------------------------------------------------------------
|
||||
|
||||
@@ -62,7 +62,8 @@ pub struct OAuthRefreshConfig {
|
||||
pub client_secret: Option<String>,
|
||||
/// Hosted OAuth proxy base URL (e.g., "http://host.docker.internal:8080").
|
||||
pub exchange_proxy_url: Option<String>,
|
||||
/// Gateway auth token for authenticating with the hosted OAuth proxy.
|
||||
/// OAuth proxy auth token for authenticating with the hosted OAuth proxy.
|
||||
/// Kept as `gateway_token` for public API compatibility.
|
||||
pub gateway_token: Option<String>,
|
||||
/// Secret name of the access token (e.g., "google_oauth_token").
|
||||
/// The refresh token lives at `{secret_name}_refresh_token`.
|
||||
@@ -71,6 +72,12 @@ pub struct OAuthRefreshConfig {
|
||||
pub provider: Option<String>,
|
||||
}
|
||||
|
||||
impl OAuthRefreshConfig {
|
||||
fn oauth_proxy_auth_token(&self) -> Option<&str> {
|
||||
self.gateway_token.as_deref()
|
||||
}
|
||||
}
|
||||
|
||||
/// Pre-resolved credential for host-based injection.
|
||||
///
|
||||
/// Built before each WASM execution by decrypting secrets from the store.
|
||||
@@ -1218,9 +1225,9 @@ async fn refresh_oauth_token(
|
||||
let refresh_name = format!("{}_refresh_token", config.secret_name);
|
||||
|
||||
if let Some(proxy_url) = config.exchange_proxy_url.as_deref() {
|
||||
let Some(gateway_token) = config.gateway_token.as_deref() else {
|
||||
let Some(oauth_proxy_auth_token) = config.oauth_proxy_auth_token() else {
|
||||
tracing::warn!(
|
||||
"OAuth refresh proxy is configured, but no gateway auth token is available"
|
||||
"OAuth refresh proxy is configured, but no OAuth proxy auth token is available"
|
||||
);
|
||||
return false;
|
||||
};
|
||||
@@ -1235,7 +1242,7 @@ async fn refresh_oauth_token(
|
||||
let token_response = match oauth_defaults::refresh_token_via_proxy(
|
||||
oauth_defaults::ProxyRefreshTokenRequest {
|
||||
proxy_url,
|
||||
gateway_token,
|
||||
gateway_token: oauth_proxy_auth_token,
|
||||
token_url: &config.token_url,
|
||||
client_id: &config.client_id,
|
||||
client_secret: config.client_secret.as_deref(),
|
||||
@@ -2704,7 +2711,8 @@ mod tests {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_resolve_host_credentials_skips_refresh_token_lookup_without_gateway_token() {
|
||||
async fn test_resolve_host_credentials_skips_refresh_token_lookup_without_oauth_proxy_auth_token()
|
||||
{
|
||||
use crate::secrets::{
|
||||
CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore,
|
||||
};
|
||||
|
||||
+5
-3
@@ -20,7 +20,6 @@ use crate::agent::scheduler::WorkerMessage;
|
||||
use crate::agent::task::TaskOutput;
|
||||
use crate::channels::web::types::ToolDecisionDto;
|
||||
use crate::context::{ContextManager, JobState};
|
||||
use crate::db::Database;
|
||||
use crate::error::Error;
|
||||
use crate::hooks::HookRegistry;
|
||||
use crate::llm::{
|
||||
@@ -28,6 +27,7 @@ use crate::llm::{
|
||||
ToolSelection,
|
||||
};
|
||||
use crate::safety::SafetyLayer;
|
||||
use crate::tenant::AdminScope;
|
||||
use crate::tools::execute::process_tool_result;
|
||||
use crate::tools::rate_limiter::RateLimitResult;
|
||||
use crate::tools::{
|
||||
@@ -45,7 +45,7 @@ pub struct WorkerDeps {
|
||||
pub llm: Arc<dyn LlmProvider>,
|
||||
pub safety: Arc<SafetyLayer>,
|
||||
pub tools: Arc<ToolRegistry>,
|
||||
pub store: Option<Arc<dyn Database>>,
|
||||
pub store: Option<AdminScope>,
|
||||
pub hooks: Arc<HookRegistry>,
|
||||
pub timeout: Duration,
|
||||
pub use_planning: bool,
|
||||
@@ -94,7 +94,7 @@ impl Worker {
|
||||
&self.deps.tools
|
||||
}
|
||||
|
||||
fn store(&self) -> Option<&Arc<dyn Database>> {
|
||||
fn store(&self) -> Option<&AdminScope> {
|
||||
self.deps.store.as_ref()
|
||||
}
|
||||
|
||||
@@ -1158,6 +1158,7 @@ impl<'a> JobDelegate<'a> {
|
||||
Ok(crate::llm::RespondOutput {
|
||||
result: RespondResult::Text(String::new()),
|
||||
usage: crate::llm::TokenUsage::default(),
|
||||
finish_reason: crate::llm::FinishReason::Stop,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1283,6 +1284,7 @@ impl<'a> LoopDelegate for JobDelegate<'a> {
|
||||
content: reasoning_text,
|
||||
},
|
||||
usage: crate::llm::TokenUsage::default(),
|
||||
finish_reason: crate::llm::FinishReason::ToolUse,
|
||||
});
|
||||
}
|
||||
Ok(_) => {} // empty selections, fall through
|
||||
|
||||
@@ -149,6 +149,7 @@ fn reject_if_injected(path: &str, content: &str) -> Result<(), WorkspaceError> {
|
||||
///
|
||||
/// Allows Workspace to work with either a PostgreSQL `Repository` (the original
|
||||
/// path) or any `Database` trait implementation (e.g. libSQL backend).
|
||||
#[derive(Clone)]
|
||||
enum WorkspaceStorage {
|
||||
/// PostgreSQL-backed repository (uses connection pool directly).
|
||||
#[cfg(feature = "postgres")]
|
||||
@@ -576,6 +577,60 @@ impl Workspace {
|
||||
self
|
||||
}
|
||||
|
||||
/// Clone the workspace configuration for a different primary user scope.
|
||||
///
|
||||
/// This preserves search config, embeddings, shared read scopes, memory
|
||||
/// layers, and privacy classifier while switching the primary read/write
|
||||
/// scope to `user_id`.
|
||||
pub fn scoped_to_user(&self, user_id: impl Into<String>) -> Self {
|
||||
let user_id = user_id.into();
|
||||
|
||||
let mut memory_layers = self.memory_layers.clone();
|
||||
for layer in &mut memory_layers {
|
||||
if layer.sensitivity == crate::workspace::layer::LayerSensitivity::Private
|
||||
&& layer.scope == self.user_id
|
||||
{
|
||||
layer.scope = user_id.clone();
|
||||
}
|
||||
}
|
||||
|
||||
let mut read_user_ids = vec![user_id.clone()];
|
||||
for scope in &self.read_user_ids {
|
||||
if scope != &self.user_id && !read_user_ids.contains(scope) {
|
||||
read_user_ids.push(scope.clone());
|
||||
}
|
||||
}
|
||||
for scope in crate::workspace::layer::MemoryLayer::read_scopes(&memory_layers) {
|
||||
if !read_user_ids.contains(&scope) {
|
||||
read_user_ids.push(scope);
|
||||
}
|
||||
}
|
||||
|
||||
let preserve_flags = user_id == self.user_id;
|
||||
Self {
|
||||
user_id,
|
||||
read_user_ids,
|
||||
agent_id: self.agent_id,
|
||||
storage: self.storage.clone(),
|
||||
embeddings: self.embeddings.clone(),
|
||||
bootstrap_pending: std::sync::atomic::AtomicBool::new(if preserve_flags {
|
||||
self.bootstrap_pending
|
||||
.load(std::sync::atomic::Ordering::Acquire)
|
||||
} else {
|
||||
false
|
||||
}),
|
||||
bootstrap_completed: std::sync::atomic::AtomicBool::new(if preserve_flags {
|
||||
self.bootstrap_completed
|
||||
.load(std::sync::atomic::Ordering::Acquire)
|
||||
} else {
|
||||
false
|
||||
}),
|
||||
search_defaults: self.search_defaults.clone(),
|
||||
memory_layers,
|
||||
privacy_classifier: self.privacy_classifier.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Get the user ID (primary scope for writes).
|
||||
pub fn user_id(&self) -> &str {
|
||||
&self.user_id
|
||||
|
||||
@@ -15,6 +15,7 @@ use crate::workspace::document::{MemoryChunk, MemoryDocument, WorkspaceEntry};
|
||||
use crate::workspace::search::{RankedResult, SearchConfig, SearchResult, fuse_results};
|
||||
|
||||
/// Database repository for workspace operations.
|
||||
#[derive(Clone)]
|
||||
pub struct Repository {
|
||||
pool: Pool,
|
||||
}
|
||||
|
||||
@@ -205,7 +205,44 @@ mod tests {
|
||||
}
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// Test 5: routine_manual_create_defaults_to_tools_enabled
|
||||
// Test 5: routine_update_fail_delete_fallback
|
||||
// -----------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn routine_update_fail_delete_fallback() {
|
||||
let trace = LlmTrace::from_file(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/tests/fixtures/llm_traces/tools/routine_update_fail_delete_fallback.json"
|
||||
))
|
||||
.expect("failed to load routine_update_fail_delete_fallback.json");
|
||||
|
||||
let rig = TestRigBuilder::new()
|
||||
.with_trace(trace.clone())
|
||||
.with_auto_approve_tools(true)
|
||||
.build()
|
||||
.await;
|
||||
|
||||
rig.send_message("Try converting a routine trigger, then recover by deleting it")
|
||||
.await;
|
||||
let responses = rig.wait_for_responses(1, Duration::from_secs(15)).await;
|
||||
|
||||
rig.verify_trace_expects(&trace, &responses);
|
||||
|
||||
let completed = rig.tool_calls_completed();
|
||||
assert!(
|
||||
completed.iter().any(|(n, ok)| n == "routine_update" && !ok),
|
||||
"routine_update should fail in this regression path: {completed:?}"
|
||||
);
|
||||
assert!(
|
||||
completed.iter().any(|(n, ok)| n == "routine_delete" && *ok),
|
||||
"routine_delete should recover successfully via preserved routine identity: {completed:?}"
|
||||
);
|
||||
|
||||
rig.shutdown();
|
||||
}
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// Test 6: routine_manual_create_defaults_to_tools_enabled
|
||||
// -----------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
@@ -246,7 +283,7 @@ mod tests {
|
||||
}
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// Test 6: routine_manual_create_explicit_no_tools
|
||||
// Test 7: routine_manual_create_explicit_no_tools
|
||||
// -----------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
@@ -287,7 +324,7 @@ mod tests {
|
||||
}
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// Test 7: routine_history
|
||||
// Test 8: routine_history
|
||||
// -----------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
@@ -439,7 +476,7 @@ mod tests {
|
||||
|
||||
match &routine.trigger {
|
||||
Trigger::Cron { schedule, timezone } => {
|
||||
assert_eq!(schedule, "0 0 9 * * MON-FRI");
|
||||
assert_eq!(schedule, "0 0 9 * * MON-FRI *");
|
||||
assert_eq!(timezone.as_deref(), Some("UTC"));
|
||||
}
|
||||
other => panic!("expected cron trigger, got {other:?}"),
|
||||
|
||||
@@ -290,6 +290,8 @@ mod tests {
|
||||
Arc::new(ExtensionManager::new(
|
||||
Arc::new(McpSessionManager::new()),
|
||||
Arc::new(McpProcessManager::new()),
|
||||
None,
|
||||
None,
|
||||
secrets,
|
||||
tools,
|
||||
None,
|
||||
@@ -299,6 +301,7 @@ mod tests {
|
||||
None,
|
||||
owner_id.to_string(),
|
||||
None,
|
||||
None,
|
||||
Vec::new(),
|
||||
))
|
||||
}
|
||||
@@ -337,14 +340,14 @@ mod tests {
|
||||
SchedulerDeps {
|
||||
tools: registry.clone(),
|
||||
extension_manager: extension_manager.clone(),
|
||||
store: Some(db.clone()),
|
||||
store: Some(ironclaw::tenant::AdminScope::new(db.clone())),
|
||||
hooks: Arc::new(HookRegistry::new()),
|
||||
},
|
||||
));
|
||||
|
||||
Arc::new(RoutineEngine::new(
|
||||
RoutineConfig::default(),
|
||||
db,
|
||||
ironclaw::tenant::AdminScope::new(db),
|
||||
llm,
|
||||
ws,
|
||||
notify_tx,
|
||||
@@ -448,7 +451,7 @@ mod tests {
|
||||
|
||||
let engine = Arc::new(RoutineEngine::new(
|
||||
RoutineConfig::default(),
|
||||
db.clone(),
|
||||
ironclaw::tenant::AdminScope::new(db.clone()),
|
||||
llm,
|
||||
ws,
|
||||
notify_tx,
|
||||
@@ -527,7 +530,7 @@ mod tests {
|
||||
|
||||
let engine = Arc::new(RoutineEngine::new(
|
||||
RoutineConfig::default(),
|
||||
db.clone(),
|
||||
ironclaw::tenant::AdminScope::new(db.clone()),
|
||||
llm,
|
||||
ws,
|
||||
notify_tx,
|
||||
@@ -614,7 +617,7 @@ mod tests {
|
||||
|
||||
let engine = Arc::new(RoutineEngine::new(
|
||||
RoutineConfig::default(),
|
||||
db.clone(),
|
||||
ironclaw::tenant::AdminScope::new(db.clone()),
|
||||
llm,
|
||||
ws,
|
||||
notify_tx,
|
||||
@@ -723,7 +726,7 @@ mod tests {
|
||||
|
||||
let engine = Arc::new(RoutineEngine::new(
|
||||
RoutineConfig::default(),
|
||||
db.clone(),
|
||||
ironclaw::tenant::AdminScope::new(db.clone()),
|
||||
llm,
|
||||
ws,
|
||||
notify_tx,
|
||||
@@ -866,7 +869,7 @@ mod tests {
|
||||
|
||||
let engine = Arc::new(RoutineEngine::new(
|
||||
RoutineConfig::default(),
|
||||
db.clone(),
|
||||
ironclaw::tenant::AdminScope::new(db.clone()),
|
||||
llm,
|
||||
ws,
|
||||
notify_tx,
|
||||
@@ -1049,7 +1052,7 @@ mod tests {
|
||||
|
||||
let engine = Arc::new(RoutineEngine::new(
|
||||
RoutineConfig::default(),
|
||||
Arc::clone(&db),
|
||||
ironclaw::tenant::AdminScope::new(Arc::clone(&db)),
|
||||
llm,
|
||||
ws,
|
||||
notify_tx,
|
||||
@@ -1171,7 +1174,7 @@ mod tests {
|
||||
|
||||
let engine = Arc::new(RoutineEngine::new(
|
||||
RoutineConfig::default(),
|
||||
db.clone(),
|
||||
ironclaw::tenant::AdminScope::new(db.clone()),
|
||||
llm,
|
||||
ws,
|
||||
notify_tx,
|
||||
@@ -1279,7 +1282,7 @@ mod tests {
|
||||
|
||||
let engine = Arc::new(RoutineEngine::new(
|
||||
config,
|
||||
db.clone(),
|
||||
ironclaw::tenant::AdminScope::new(db.clone()),
|
||||
llm,
|
||||
ws,
|
||||
notify_tx,
|
||||
|
||||
@@ -201,6 +201,7 @@ mod tests {
|
||||
sandbox_readiness: ironclaw::agent::SandboxReadiness::DisabledByConfig,
|
||||
builder: None,
|
||||
llm_backend: "nearai".to_string(),
|
||||
tenant_rates: std::sync::Arc::new(ironclaw::tenant::TenantRateRegistry::new(4, 3)),
|
||||
};
|
||||
|
||||
let gateway = Arc::new(TestChannel::new());
|
||||
|
||||
@@ -12,6 +12,7 @@ mod tests {
|
||||
|
||||
use crate::support::test_rig::TestRigBuilder;
|
||||
use crate::support::trace_llm::LlmTrace;
|
||||
use ironclaw::workspace::Workspace;
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// Test 1: write_chunk_search
|
||||
@@ -268,6 +269,7 @@ mod tests {
|
||||
|
||||
#[tokio::test]
|
||||
async fn identity_in_system_prompt() {
|
||||
const TEST_USER_ID: &str = "test-user";
|
||||
let trace = LlmTrace::from_file(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/tests/fixtures/llm_traces/workspace/identity_prompt.json"
|
||||
@@ -280,7 +282,7 @@ mod tests {
|
||||
.await;
|
||||
|
||||
// Seed an IDENTITY.md so the system prompt has real content to inject.
|
||||
let ws = rig.workspace().expect("workspace must be available");
|
||||
let ws = Workspace::new_with_db(TEST_USER_ID, rig.database().clone());
|
||||
ws.write(
|
||||
"IDENTITY.md",
|
||||
"I am TestBot, a helpful testing assistant created for E2E verification.",
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
{
|
||||
"model_name": "test-routine-update-fail-delete-fallback",
|
||||
"expects": {
|
||||
"tools_used": ["routine_create", "routine_update", "routine_delete"],
|
||||
"tool_results_contain": {
|
||||
"routine_update": "Cannot update schedule or timezone on a non-cron routine.",
|
||||
"routine_delete": "temp-routine"
|
||||
},
|
||||
"min_responses": 1
|
||||
},
|
||||
"steps": [
|
||||
{
|
||||
"response": {
|
||||
"type": "tool_calls",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_rc_fallback",
|
||||
"name": "routine_create",
|
||||
"arguments": {
|
||||
"name": "temp-routine",
|
||||
"trigger_type": "manual",
|
||||
"prompt": "Temporary routine for fallback test."
|
||||
}
|
||||
}
|
||||
],
|
||||
"input_tokens": 120,
|
||||
"output_tokens": 40
|
||||
}
|
||||
},
|
||||
{
|
||||
"response": {
|
||||
"type": "tool_calls",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_ru_fallback",
|
||||
"name": "routine_update",
|
||||
"arguments": {
|
||||
"name": "temp-routine",
|
||||
"schedule": "0 */10 * * * *"
|
||||
}
|
||||
}
|
||||
],
|
||||
"input_tokens": 200,
|
||||
"output_tokens": 30
|
||||
}
|
||||
},
|
||||
{
|
||||
"response": {
|
||||
"type": "tool_calls",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_rd_fallback",
|
||||
"name": "routine_delete",
|
||||
"arguments": {}
|
||||
}
|
||||
],
|
||||
"input_tokens": 300,
|
||||
"output_tokens": 20
|
||||
}
|
||||
},
|
||||
{
|
||||
"response": {
|
||||
"type": "text",
|
||||
"content": "I recovered from the failed update and cleaned up the original routine.",
|
||||
"input_tokens": 380,
|
||||
"output_tokens": 25
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -203,6 +203,8 @@ async fn extension_manager_with_process_manager_constructs() {
|
||||
let manager = ExtensionManager::new(
|
||||
Arc::new(McpSessionManager::new()),
|
||||
Arc::new(McpProcessManager::new()),
|
||||
None,
|
||||
None,
|
||||
secrets,
|
||||
tools,
|
||||
None,
|
||||
@@ -212,6 +214,7 @@ async fn extension_manager_with_process_manager_constructs() {
|
||||
None,
|
||||
"test".to_string(),
|
||||
None,
|
||||
None,
|
||||
Vec::new(),
|
||||
);
|
||||
|
||||
|
||||
@@ -1,10 +1,10 @@
|
||||
//! Tests proving that multi-tenant system prompts are broken.
|
||||
//! Regression tests for multi-tenant system prompts.
|
||||
//!
|
||||
//! Bug: In multi-tenant mode, the agent loop uses `self.workspace()` which
|
||||
//! returns a single shared workspace (user_id="default"). Identity files
|
||||
//! (IDENTITY.md, SOUL.md, USER.md) seeded under per-user IDs ("alice",
|
||||
//! "bob") are invisible to this workspace, so the system prompt is
|
||||
//! empty/wrong.
|
||||
//! The agent must build the conversational system prompt from a workspace
|
||||
//! scoped to the incoming message's user, not from the shared owner-scope
|
||||
//! workspace created at startup. Otherwise per-user identity files
|
||||
//! (IDENTITY.md, SOUL.md, USER.md) become invisible and different users can
|
||||
//! see the same owner-scoped prompt.
|
||||
//!
|
||||
//! These tests:
|
||||
//! 1. Seed identity files for two users (alice, bob) in the database
|
||||
@@ -13,7 +13,7 @@
|
||||
//! correct user's identity
|
||||
//! 4. Verify user A's identity doesn't leak into user B's prompt
|
||||
//!
|
||||
//! All tests are expected to FAIL until the bug is fixed.
|
||||
//! These tests ensure each user's identity is isolated correctly.
|
||||
|
||||
#[cfg(feature = "libsql")]
|
||||
mod support;
|
||||
|
||||
@@ -266,6 +266,7 @@ impl GatewayWorkflowHarness {
|
||||
sandbox_readiness: ironclaw::agent::SandboxReadiness::DisabledByConfig,
|
||||
builder: None,
|
||||
llm_backend: "nearai".to_string(),
|
||||
tenant_rates: std::sync::Arc::new(ironclaw::tenant::TenantRateRegistry::new(4, 3)),
|
||||
},
|
||||
channels,
|
||||
None,
|
||||
|
||||
@@ -642,7 +642,7 @@ impl TestRigBuilder {
|
||||
let (notify_tx, _notify_rx) = tokio::sync::mpsc::channel(16);
|
||||
let engine = Arc::new(RoutineEngine::new(
|
||||
routine_config,
|
||||
Arc::clone(db_arc),
|
||||
ironclaw::tenant::AdminScope::new(Arc::clone(db_arc)),
|
||||
components.llm.clone(),
|
||||
Arc::clone(ws),
|
||||
notify_tx,
|
||||
@@ -762,6 +762,7 @@ impl TestRigBuilder {
|
||||
sandbox_readiness: ironclaw::agent::SandboxReadiness::Available, // tests don't use real Docker
|
||||
builder: None,
|
||||
llm_backend: "nearai".to_string(),
|
||||
tenant_rates: std::sync::Arc::new(ironclaw::tenant::TenantRateRegistry::new(4, 3)),
|
||||
};
|
||||
|
||||
// 7. Create TestChannel and ChannelManager.
|
||||
|
||||
Reference in New Issue
Block a user