Compare commits

..
Author SHA1 Message Date
ZakiandClaude Opus 4.6 aa289997e3 style: fix import ordering for routines module
Co-Authored-By: Claude Opus 4.6 <[email protected]>
2026-03-12 12:50:53 -07:00
ZakiandClaude Opus 4.6 4aad0cfbaa refactor(cli): rename cron subcommand to routines
The system manages all routine types (cron, webhook, event, manual),
not just cron schedules. Rename the CLI subcommand to reflect this:
- `ironclaw cron` -> `ironclaw routines` (with `cron` as hidden alias)
- List shows all routines by default, add --trigger filter
- Remove cron-trigger-only validation
- Simplify require_routine helper (no trigger type check)

Co-Authored-By: Claude Opus 4.6 <[email protected]>
2026-03-12 12:06:30 -07:00
Zaki 112a4087e7 fix(cli): reject invalid cron timezones 2026-03-12 11:00:16 -07:00
reidliu41andZaki 403f6f504f feat(cli): add cron subcommand for managing scheduled routines
Rebase onto staging branch and address collaborator review:
  - Fix .unwrap_or(None) → proper error propagation in set_enabled()
  - Add --yes/-y flag for non-interactive deletion with confirmation prompt
  - Add --json flag for machine-readable output in list and history
  - Preserve error context chain with {e:#} in run_cron_cli()

  Note: GATEWAY_USER_ID is trusted from the environment; future work may
  add authentication for multi-tenant deployments.
2026-03-12 11:00:16 -07:00
83 changed files with 469 additions and 6407 deletions
-6
View File
@@ -70,12 +70,6 @@ NEARAI_AUTH_URL=https://private.near.ai
# LLM_BASE_URL=https://api.fireworks.ai/inference/v1 # LLM_BASE_URL=https://api.fireworks.ai/inference/v1
# LLM_API_KEY=fw_... # LLM_API_KEY=fw_...
# === MiniMax ===
# LLM_BACKEND=minimax
# MINIMAX_API_KEY=...
# MINIMAX_MODEL=MiniMax-M2.5
# MINIMAX_BASE_URL=https://api.minimax.io/v1 # default (global); use https://api.minimaxi.com/v1 for China
# === Anthropic Direct === # === Anthropic Direct ===
# LLM_BACKEND=anthropic # LLM_BACKEND=anthropic
# ANTHROPIC_MODEL=claude-sonnet-4-6 # ANTHROPIC_MODEL=claude-sonnet-4-6
-23
View File
@@ -1,23 +0,0 @@
#!/usr/bin/env bash
set -euo pipefail
# Pre-push hook: run clippy and tests before pushing.
# Install: git config core.hooksPath .githooks
echo "pre-push: running clippy..."
if ! cargo clippy --all --benches --tests --examples --all-features -- -D warnings; then
echo ""
echo "Push blocked: clippy warnings found."
echo "To bypass: git push --no-verify"
exit 1
fi
echo "pre-push: running tests..."
if ! cargo test; then
echo ""
echo "Push blocked: tests failed."
echo "To bypass: git push --no-verify"
exit 1
fi
echo "pre-push: all checks passed."
+1 -1
View File
@@ -48,7 +48,7 @@ jobs:
matrix: matrix:
include: include:
- group: core - group: core
files: "tests/e2e/scenarios/test_connection.py tests/e2e/scenarios/test_chat.py tests/e2e/scenarios/test_sse_reconnect.py tests/e2e/scenarios/test_html_injection.py tests/e2e/scenarios/test_csp.py" files: "tests/e2e/scenarios/test_connection.py tests/e2e/scenarios/test_chat.py tests/e2e/scenarios/test_sse_reconnect.py tests/e2e/scenarios/test_html_injection.py"
- group: features - group: features
files: "tests/e2e/scenarios/test_skills.py tests/e2e/scenarios/test_tool_approval.py" files: "tests/e2e/scenarios/test_skills.py tests/e2e/scenarios/test_tool_approval.py"
- group: extensions - group: extensions
+5 -12
View File
@@ -13,11 +13,6 @@ jobs:
with: with:
fetch-depth: 0 fetch-depth: 0
- name: Fetch PR head and base
run: |
git fetch origin ${{ github.event.pull_request.base.ref }}
git fetch origin pull/${{ github.event.pull_request.number }}/head:pr-head
- name: Check for regression tests - name: Check for regression tests
env: env:
PR_TITLE: ${{ github.event.pull_request.title }} PR_TITLE: ${{ github.event.pull_request.title }}
@@ -26,8 +21,6 @@ jobs:
set -euo pipefail set -euo pipefail
BASE_REF="origin/${{ github.event.pull_request.base.ref }}" BASE_REF="origin/${{ github.event.pull_request.base.ref }}"
# Use the actual PR head, not the merge commit that actions/checkout checks out
HEAD_REF="pr-head"
# --- 1. Is this a fix PR? Check title first, then commit messages --- # --- 1. Is this a fix PR? Check title first, then commit messages ---
IS_FIX=false IS_FIX=false
@@ -37,7 +30,7 @@ jobs:
fi fi
if [ "$IS_FIX" = false ]; then if [ "$IS_FIX" = false ]; then
COMMITS=$(git log --format='%s' "${BASE_REF}..${HEAD_REF}") COMMITS=$(git log --format='%s' "${BASE_REF}..HEAD")
if grep -qiE '^(fix(\(.*\))?|hotfix|bugfix):' <<< "$COMMITS"; then if grep -qiE '^(fix(\(.*\))?|hotfix|bugfix):' <<< "$COMMITS"; then
IS_FIX=true IS_FIX=true
fi fi
@@ -56,14 +49,14 @@ jobs:
exit 0 exit 0
fi fi
COMMIT_BODIES=$(git log --format='%B' "${BASE_REF}..${HEAD_REF}") COMMIT_BODIES=$(git log --format='%B' "${BASE_REF}..HEAD")
if grep -qF '[skip-regression-check]' <<< "$COMMIT_BODIES"; then if grep -qF '[skip-regression-check]' <<< "$COMMIT_BODIES"; then
echo "[skip-regression-check] found in commit message — skipping." echo "[skip-regression-check] found in commit message — skipping."
exit 0 exit 0
fi fi
# --- 3. Exempt static-only / docs-only changes --- # --- 3. Exempt static-only / docs-only changes ---
CHANGED_FILES=$(git diff --name-only "${BASE_REF}...${HEAD_REF}") CHANGED_FILES=$(git diff --name-only "${BASE_REF}...HEAD")
if [ -z "$CHANGED_FILES" ]; then if [ -z "$CHANGED_FILES" ]; then
echo "No changed files — skipping." echo "No changed files — skipping."
@@ -87,13 +80,13 @@ jobs:
# --- 4. Look for test changes --- # --- 4. Look for test changes ---
# Fast path: new test attributes or test modules in added lines. # Fast path: new test attributes or test modules in added lines.
if git diff "${BASE_REF}...${HEAD_REF}" -U0 -- '*.rs' | grep -qE '^\+.*(#\[test\]|#\[tokio::test\]|#\[cfg\(test\)\]|mod tests)'; then if git diff "${BASE_REF}...HEAD" -U0 -- '*.rs' | grep -qE '^\+.*(#\[test\]|#\[tokio::test\]|#\[cfg\(test\)\]|mod tests)'; then
echo "Test changes found in .rs files." echo "Test changes found in .rs files."
exit 0 exit 0
fi fi
# Whole-function context: detect edits inside existing test functions. # Whole-function context: detect edits inside existing test functions.
if git diff "${BASE_REF}...${HEAD_REF}" -W -- '*.rs' | awk ' if git diff "${BASE_REF}...HEAD" -W -- '*.rs' | awk '
/^@@/ { if (has_test && has_add) { found=1; exit } has_test=0; has_add=0 } /^@@/ { if (has_test && has_add) { found=1; exit } has_test=0; has_add=0 }
/^ .*#\[test\]/ || /^ .*#\[tokio::test\]/ || /^ .*#\[cfg\(test\)\]/ || /^ .*mod tests/ { has_test=1 } /^ .*#\[test\]/ || /^ .*#\[tokio::test\]/ || /^ .*#\[cfg\(test\)\]/ || /^ .*mod tests/ { has_test=1 }
/^\+.*#\[test\]/ || /^\+.*#\[tokio::test\]/ || /^\+.*#\[cfg\(test\)\]/ || /^\+.*mod tests/ { has_test=1 } /^\+.*#\[test\]/ || /^\+.*#\[tokio::test\]/ || /^\+.*#\[cfg\(test\)\]/ || /^\+.*mod tests/ { has_test=1 }
+10 -13
View File
@@ -44,7 +44,6 @@ jobs:
id: check id: check
env: env:
FORCE_RUN: ${{ inputs.force }} FORCE_RUN: ${{ inputs.force }}
DEFAULT_BRANCH: ${{ github.event.repository.default_branch }}
run: | run: |
CURRENT_HEAD=$(git rev-parse HEAD) CURRENT_HEAD=$(git rev-parse HEAD)
echo "current_head=${CURRENT_HEAD}" >> "$GITHUB_OUTPUT" echo "current_head=${CURRENT_HEAD}" >> "$GITHUB_OUTPUT"
@@ -66,8 +65,8 @@ jobs:
echo "Found ${COMMIT_COUNT} new commit(s) since last tested" echo "Found ${COMMIT_COUNT} new commit(s) since last tested"
DIFF_RANGE="${LAST_TESTED}..${CURRENT_HEAD}" DIFF_RANGE="${LAST_TESTED}..${CURRENT_HEAD}"
else else
git fetch origin "${DEFAULT_BRANCH}" git fetch origin main
MERGE_BASE=$(git merge-base "origin/${DEFAULT_BRANCH}" HEAD) MERGE_BASE=$(git merge-base origin/main HEAD)
echo "First run -- reviewing from merge-base ${MERGE_BASE}" echo "First run -- reviewing from merge-base ${MERGE_BASE}"
DIFF_RANGE="${MERGE_BASE}..${CURRENT_HEAD}" DIFF_RANGE="${MERGE_BASE}..${CURRENT_HEAD}"
fi fi
@@ -130,19 +129,18 @@ jobs:
echo "token=${{ github.token }}" >> "$GITHUB_OUTPUT" echo "token=${{ github.token }}" >> "$GITHUB_OUTPUT"
fi fi
- name: Check if staging is ahead of target branch - name: Check if staging is ahead of main
id: ahead-check id: ahead-check
env: env:
GH_TOKEN: ${{ steps.token.outputs.token }} GH_TOKEN: ${{ steps.token.outputs.token }}
DEFAULT_BRANCH: ${{ github.event.repository.default_branch }}
run: | run: |
git fetch origin "${DEFAULT_BRANCH}" git fetch origin main
AHEAD=$(git rev-list --count "origin/${DEFAULT_BRANCH}..origin/staging") AHEAD=$(git rev-list --count origin/main..origin/staging)
echo "commits_ahead=${AHEAD}" >> "$GITHUB_OUTPUT" echo "commits_ahead=${AHEAD}" >> "$GITHUB_OUTPUT"
if [ "$AHEAD" -eq 0 ]; then if [ "$AHEAD" -eq 0 ]; then
echo "Staging is not ahead of ${DEFAULT_BRANCH}. Nothing to promote." echo "Staging is not ahead of main. Nothing to promote."
else else
echo "Staging is ${AHEAD} commits ahead of ${DEFAULT_BRANCH}." echo "Staging is ${AHEAD} commits ahead of main."
fi fi
- name: Create promotion branch - name: Create promotion branch
@@ -161,7 +159,6 @@ jobs:
if: steps.ahead-check.outputs.commits_ahead != '0' if: steps.ahead-check.outputs.commits_ahead != '0'
env: env:
GH_TOKEN: ${{ steps.token.outputs.token }} GH_TOKEN: ${{ steps.token.outputs.token }}
DEFAULT_BRANCH: ${{ github.event.repository.default_branch }}
run: | run: |
# Find the newest open promotion PR with a staging-promote/* head branch # Find the newest open promotion PR with a staging-promote/* head branch
LATEST=$(gh pr list --label staging-promotion --state open \ LATEST=$(gh pr list --label staging-promotion --state open \
@@ -171,8 +168,8 @@ jobs:
echo "base=${LATEST}" >> "$GITHUB_OUTPUT" echo "base=${LATEST}" >> "$GITHUB_OUTPUT"
echo "Chaining onto existing promotion branch: ${LATEST}" echo "Chaining onto existing promotion branch: ${LATEST}"
else else
echo "base=${DEFAULT_BRANCH}" >> "$GITHUB_OUTPUT" echo "base=main" >> "$GITHUB_OUTPUT"
echo "No existing promotion PR — targeting ${DEFAULT_BRANCH}" echo "No existing promotion PR — targeting main"
fi fi
- name: Create promotion PR - name: Create promotion PR
@@ -189,7 +186,7 @@ jobs:
PR_URL=$(gh pr create \ PR_URL=$(gh pr create \
--base "$BASE" \ --base "$BASE" \
--head "$BRANCH" \ --head "$BRANCH" \
--title "chore: promote staging to ${BASE} (${TIMESTAMP})" \ --title "chore: promote staging to main (${TIMESTAMP})" \
--body "## Auto-promotion from staging CI --body "## Auto-promotion from staging CI
**Batch range:** \`${RANGE}\` **Batch range:** \`${RANGE}\`
-9
View File
@@ -7,15 +7,6 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
## [Unreleased] ## [Unreleased]
## [0.18.0](https://github.com/nearai/ironclaw/compare/v0.17.0...v0.18.0) - 2026-03-11
### Other
- Merge pull request #907 from nearai/staging-promote/b0214fef-22930316561
- promote staging to main (2026-03-10 15:19 UTC) ([#865](https://github.com/nearai/ironclaw/pull/865))
- Merge pull request #830 from nearai/staging-promote/3a2989d0-22888378864
- update WASM artifact SHA256 checksums [skip ci] ([#876](https://github.com/nearai/ironclaw/pull/876))
## [0.17.0](https://github.com/nearai/ironclaw/compare/v0.16.1...v0.17.0) - 2026-03-10 ## [0.17.0](https://github.com/nearai/ironclaw/compare/v0.16.1...v0.17.0) - 2026-03-10
### Added ### Added
Generated
+1 -1
View File
@@ -3350,7 +3350,7 @@ dependencies = [
[[package]] [[package]]
name = "ironclaw" name = "ironclaw"
version = "0.18.0" version = "0.17.0"
dependencies = [ dependencies = [
"aes-gcm", "aes-gcm",
"aho-corasick", "aho-corasick",
+1 -1
View File
@@ -20,7 +20,7 @@ exclude = [
[package] [package]
name = "ironclaw" name = "ironclaw"
version = "0.18.0" version = "0.17.0"
edition = "2024" edition = "2024"
rust-version = "1.92" rust-version = "1.92"
description = "Secure personal AI assistant that protects your data and expands its capabilities on the fly" description = "Secure personal AI assistant that protects your data and expands its capabilities on the fly"
-205
View File
@@ -20,162 +20,33 @@ version = "1.0.102"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7f202df86484c868dbad7eaa557ef785d5c66295e41b460ef922eca0723b842c" checksum = "7f202df86484c868dbad7eaa557ef785d5c66295e41b460ef922eca0723b842c"
[[package]]
name = "base64ct"
version = "1.8.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2af50177e190e07a26ab74f8b1efbfe2ef87da2116221318cb1c2e82baf7de06"
[[package]] [[package]]
name = "bitflags" name = "bitflags"
version = "2.11.0" version = "2.11.0"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "843867be96c8daad0d758b57df9392b6d8d271134fce549de6ce169ff98a92af" checksum = "843867be96c8daad0d758b57df9392b6d8d271134fce549de6ce169ff98a92af"
[[package]]
name = "block-buffer"
version = "0.10.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3078c7629b62d3f0439517fa394996acacc5cbc91c5a20d8c658e77abd503a71"
dependencies = [
"generic-array",
]
[[package]] [[package]]
name = "cfg-if" name = "cfg-if"
version = "1.0.4" version = "1.0.4"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801" checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801"
[[package]]
name = "const-oid"
version = "0.9.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c2459377285ad874054d797f3ccebf984978aa39129f6eafde5cdc8315b612f8"
[[package]]
name = "cpufeatures"
version = "0.2.17"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "59ed5838eebb26a2bb2e58f6d5b5316989ae9d08bab10e0e6d103e656d1b0280"
dependencies = [
"libc",
]
[[package]]
name = "crypto-common"
version = "0.1.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a"
dependencies = [
"generic-array",
"typenum",
]
[[package]]
name = "curve25519-dalek"
version = "4.1.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "97fb8b7c4503de7d6ae7b42ab72a5a59857b4c937ec27a3d4539dba95b5ab2be"
dependencies = [
"cfg-if",
"cpufeatures",
"curve25519-dalek-derive",
"digest",
"fiat-crypto",
"rustc_version",
"subtle",
"zeroize",
]
[[package]]
name = "curve25519-dalek-derive"
version = "0.1.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f46882e17999c6cc590af592290432be3bce0428cb0d5f8b6715e4dc7b383eb3"
dependencies = [
"proc-macro2",
"quote",
"syn",
]
[[package]]
name = "der"
version = "0.7.10"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e7c1832837b905bbfb5101e07cc24c8deddf52f93225eee6ead5f4d63d53ddcb"
dependencies = [
"const-oid",
"zeroize",
]
[[package]]
name = "digest"
version = "0.10.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292"
dependencies = [
"block-buffer",
"crypto-common",
]
[[package]] [[package]]
name = "discord-channel" name = "discord-channel"
version = "0.1.0" version = "0.1.0"
dependencies = [ dependencies = [
"ed25519-dalek",
"hex",
"serde", "serde",
"serde_json", "serde_json",
"wit-bindgen", "wit-bindgen",
] ]
[[package]]
name = "ed25519"
version = "2.2.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "115531babc129696a58c64a4fef0a8bf9e9698629fb97e9e40767d235cfbcd53"
dependencies = [
"pkcs8",
"signature",
]
[[package]]
name = "ed25519-dalek"
version = "2.2.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "70e796c081cee67dc755e1a36a0a172b897fab85fc3f6bc48307991f64e4eca9"
dependencies = [
"curve25519-dalek",
"ed25519",
"serde",
"sha2",
"subtle",
"zeroize",
]
[[package]] [[package]]
name = "equivalent" name = "equivalent"
version = "1.0.2" version = "1.0.2"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f"
[[package]]
name = "fiat-crypto"
version = "0.2.9"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "28dea519a9695b9977216879a3ebfddf92f1c08c05d984f8996aecd6ecdc811d"
[[package]]
name = "generic-array"
version = "0.14.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a"
dependencies = [
"typenum",
"version_check",
]
[[package]] [[package]]
name = "hashbrown" name = "hashbrown"
version = "0.14.5" version = "0.14.5"
@@ -197,12 +68,6 @@ version = "0.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea"
[[package]]
name = "hex"
version = "0.4.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70"
[[package]] [[package]]
name = "id-arena" name = "id-arena"
version = "2.3.0" version = "2.3.0"
@@ -233,12 +98,6 @@ version = "0.2.5"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "884e2677b40cc8c339eaefcb701c32ef1fd2493d71118dc0ca4b6a736c93bd67" checksum = "884e2677b40cc8c339eaefcb701c32ef1fd2493d71118dc0ca4b6a736c93bd67"
[[package]]
name = "libc"
version = "0.2.182"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6800badb6cb2082ffd7b6a67e6125bb39f18782f793520caee8cb8846be06112"
[[package]] [[package]]
name = "log" name = "log"
version = "0.4.29" version = "0.4.29"
@@ -257,16 +116,6 @@ version = "1.21.3"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "42f5e15c9953c5e4ccceeb2e7382a716482c34515315f7b03532b8b4e8393d2d" checksum = "42f5e15c9953c5e4ccceeb2e7382a716482c34515315f7b03532b8b4e8393d2d"
[[package]]
name = "pkcs8"
version = "0.10.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f950b2377845cebe5cf8b5165cb3cc1a5e0fa5cfa3e1f7f55707d8fd82e0a7b7"
dependencies = [
"der",
"spki",
]
[[package]] [[package]]
name = "prettyplease" name = "prettyplease"
version = "0.2.37" version = "0.2.37"
@@ -295,15 +144,6 @@ dependencies = [
"proc-macro2", "proc-macro2",
] ]
[[package]]
name = "rustc_version"
version = "0.4.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "cfcb3a22ef46e85b45de6ee7e79d063319ebb6594faafcf1c225ea92ab6e9b92"
dependencies = [
"semver",
]
[[package]] [[package]]
name = "semver" name = "semver"
version = "1.0.27" version = "1.0.27"
@@ -353,23 +193,6 @@ dependencies = [
"zmij", "zmij",
] ]
[[package]]
name = "sha2"
version = "0.10.9"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283"
dependencies = [
"cfg-if",
"cpufeatures",
"digest",
]
[[package]]
name = "signature"
version = "2.2.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "77549399552de45a898a580c1b41d445bf730df867cc44e6c0233bbc4b8329de"
[[package]] [[package]]
name = "smallvec" name = "smallvec"
version = "1.15.1" version = "1.15.1"
@@ -385,22 +208,6 @@ dependencies = [
"smallvec", "smallvec",
] ]
[[package]]
name = "spki"
version = "0.7.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d91ed6c858b01f942cd56b37a94b3e0a1798290327d1236e4d9cf4eaca44d29d"
dependencies = [
"base64ct",
"der",
]
[[package]]
name = "subtle"
version = "2.6.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292"
[[package]] [[package]]
name = "syn" name = "syn"
version = "2.0.117" version = "2.0.117"
@@ -412,12 +219,6 @@ dependencies = [
"unicode-ident", "unicode-ident",
] ]
[[package]]
name = "typenum"
version = "1.19.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "562d481066bde0658276a35467c4af00bdc6ee726305698a55b86e61d7ad82bb"
[[package]] [[package]]
name = "unicode-ident" name = "unicode-ident"
version = "1.0.24" version = "1.0.24"
@@ -593,12 +394,6 @@ dependencies = [
"syn", "syn",
] ]
[[package]]
name = "zeroize"
version = "1.8.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b97154e67e32c85465826e8bcc1c59429aaaf107c1e4a9e53c8d8ccd5eff88d0"
[[package]] [[package]]
name = "zmij" name = "zmij"
version = "1.0.21" version = "1.0.21"
-2
View File
@@ -10,8 +10,6 @@ publish = false
serde = { version = "1.0", features = ["derive"] } serde = { version = "1.0", features = ["derive"] }
serde_json = "1.0" serde_json = "1.0"
wit-bindgen = "0.36" wit-bindgen = "0.36"
ed25519-dalek = { version = "2", default-features = false, features = ["alloc", "fast", "zeroize"] }
hex = "0.4"
[lib] [lib]
crate-type = ["cdylib"] crate-type = ["cdylib"]
+7 -33
View File
@@ -21,10 +21,11 @@ WASM channel for Discord integration - handle slash commands and button interact
ironclaw secret set discord_bot_token YOUR_BOT_TOKEN ironclaw secret set discord_bot_token YOUR_BOT_TOKEN
``` ```
**Note:** The `discord_bot_token` secret is used for Discord REST API calls. **Note:** The `discord_bot_token` secret is the only value read directly by this
Interaction signature verification is performed inside the Discord channel Discord channel WASM component. The `discord_app_id` and `discord_public_key`
module and uses the channel config field `webhook_secret` (set this to your secrets are used by the IronClaw host (for example, to verify Discord
Discord app public key hex). interaction signatures and manage slash command registration) and are not
accessed from the WASM module itself.
## Discord Configuration ## Discord Configuration
@@ -86,30 +87,6 @@ If an internal error occurs (e.g., metadata serialization failure), the tool att
Check the host logs for detailed error information. Check the host logs for detailed error information.
## Advanced Usage ## Advanced Usage
### Mention Polling
The Discord channel can also poll configured channels for `@bot` mentions.
Example channel config:
```json
{
"require_signature_verification": true,
"webhook_secret": "YOUR_DISCORD_PUBLIC_KEY_HEX",
"polling_enabled": true,
"poll_interval_ms": 30000,
"mention_channel_ids": ["123456789012345678"],
"owner_id": null,
"dm_policy": "pairing",
"allow_from": []
}
```
### Access Control
- `owner_id`: when set, only that Discord user can interact with the bot.
- `dm_policy`: `open` allows all DMs; `pairing` requires approval.
- `allow_from`: allowlist entries for DM pairing checks (`*`, user id, or username).
### Embeds ### Embeds
@@ -119,11 +96,8 @@ To send embeds, include an `embeds` array in the `metadata_json` field of the ag
### "Invalid Signature" ### "Invalid Signature"
- Check that `webhook_secret` is set to your Discord app public key hex in the - Check that `discord_public_key` is set correctly in IronClaw secrets.
Discord channel config. - This validation happens on the host before reaching the WASM.
- Validation happens inside the Discord WASM channel.
- If `require_signature_verification` is `true` and `webhook_secret` is empty,
the channel returns HTTP `500` with a configuration error.
### "401 Unauthorized" ### "401 Unauthorized"
@@ -3,7 +3,7 @@
"wit_version": "0.3.0", "wit_version": "0.3.0",
"type": "channel", "type": "channel",
"name": "discord", "name": "discord",
"description": "Discord webhook channel for slash commands, components, and optional mention polling", "description": "Discord Gateway/Webhook channel for handling slash commands, buttons, and messages",
"setup": { "setup": {
"required_secrets": [ "required_secrets": [
{ {
@@ -41,7 +41,7 @@
}, },
"channel": { "channel": {
"allowed_paths": ["/webhook/discord"], "allowed_paths": ["/webhook/discord"],
"allow_polling": true, "allow_polling": false,
"callback_timeout_secs": 45, "callback_timeout_secs": 45,
"workspace_prefix": "channels/discord/", "workspace_prefix": "channels/discord/",
"emit_rate_limit": { "emit_rate_limit": {
@@ -55,10 +55,6 @@
}, },
"config": { "config": {
"require_signature_verification": true, "require_signature_verification": true,
"webhook_secret": null,
"polling_enabled": false,
"poll_interval_ms": 30000,
"mention_channel_ids": [],
"owner_id": null, "owner_id": null,
"dm_policy": "pairing", "dm_policy": "pairing",
"allow_from": [] "allow_from": []
File diff suppressed because it is too large Load Diff
-5
View File
@@ -1,10 +1,5 @@
# WARNING: Replace all CHANGE_ME values before deploying. # WARNING: Replace all CHANGE_ME values before deploying.
# Do not use placeholder passwords in production. # Do not use placeholder passwords in production.
# Pin the Docker image version for deterministic deployments.
# Update this value when deploying a new release.
# IRONCLAW_VERSION=v1.0.0
DATABASE_URL=postgres://ironclaw:CHANGE_ME@localhost:5432/ironclaw DATABASE_URL=postgres://ironclaw:CHANGE_ME@localhost:5432/ironclaw
# NEAR AI Cloud (API key auth, Chat Completions API) # NEAR AI Cloud (API key auth, Chat Completions API)
+5 -9
View File
@@ -5,17 +5,13 @@ Requires=cloud-sql-proxy.service
[Service] [Service]
Type=simple Type=simple
EnvironmentFile=/opt/ironclaw/.env ExecStartPre=/usr/bin/docker pull us-central1-docker.pkg.dev/ironclaw-prod/ironclaw/agent:latest
# Pin to a specific version tag or digest instead of :latest to prevent ExecStart=/usr/bin/docker run --rm \
# uncontrolled deployments. Update IRONCLAW_VERSION in /opt/ironclaw/.env
# or replace the tag below when deploying a new release.
ExecStartPre=/bin/bash -c 'docker pull us-central1-docker.pkg.dev/ironclaw-prod/ironclaw/agent:${IRONCLAW_VERSION:-latest}'
ExecStart=/bin/bash -c 'docker run --rm \
--name ironclaw \ --name ironclaw \
--env-file /opt/ironclaw/.env \ --env-file /opt/ironclaw/.env \
-p 3000:3000 \ --network=host \
us-central1-docker.pkg.dev/ironclaw-prod/ironclaw/agent:${IRONCLAW_VERSION:-latest} \ us-central1-docker.pkg.dev/ironclaw-prod/ironclaw/agent:latest \
--no-onboard' --no-onboard
ExecStop=/usr/bin/docker stop ironclaw ExecStop=/usr/bin/docker stop ironclaw
Restart=always Restart=always
RestartSec=10 RestartSec=10
+1 -8
View File
@@ -24,15 +24,8 @@ systemctl enable docker
systemctl start docker systemctl start docker
echo "==> Installing Cloud SQL Auth Proxy" echo "==> Installing Cloud SQL Auth Proxy"
CLOUD_SQL_PROXY_VERSION="v2.14.3"
CLOUD_SQL_PROXY_SHA256="75e7cc1f158ab6f97b7810e9d8419c55735cff40bc56d4f19673adfdf2406a59"
curl -fsSL -o /usr/local/bin/cloud-sql-proxy \ curl -fsSL -o /usr/local/bin/cloud-sql-proxy \
"https://storage.googleapis.com/cloud-sql-connectors/cloud-sql-proxy/${CLOUD_SQL_PROXY_VERSION}/cloud-sql-proxy.linux.amd64" https://storage.googleapis.com/cloud-sql-connectors/cloud-sql-proxy/v2.14.3/cloud-sql-proxy.linux.amd64
echo "${CLOUD_SQL_PROXY_SHA256} /usr/local/bin/cloud-sql-proxy" | sha256sum -c - || {
echo "ERROR: Cloud SQL Auth Proxy checksum verification failed -- aborting"
rm -f /usr/local/bin/cloud-sql-proxy
exit 1
}
chmod +x /usr/local/bin/cloud-sql-proxy chmod +x /usr/local/bin/cloud-sql-proxy
echo "==> Installing systemd services" echo "==> Installing systemd services"
-20
View File
@@ -15,7 +15,6 @@ configurations.
| io.net | `ionet` | `IONET_API_KEY` | Intelligence API | | io.net | `ionet` | `IONET_API_KEY` | Intelligence API |
| Mistral | `mistral` | `MISTRAL_API_KEY` | Mistral models | | Mistral | `mistral` | `MISTRAL_API_KEY` | Mistral models |
| Yandex AI Studio | `yandex` | `YANDEX_API_KEY` | YandexGPT models | | Yandex AI Studio | `yandex` | `YANDEX_API_KEY` | YandexGPT models |
| MiniMax | `minimax` | `MINIMAX_API_KEY` | MiniMax-M2.5 models |
| Cloudflare Workers AI | `cloudflare` | `CLOUDFLARE_API_KEY` | Access to Workers AI | | Cloudflare Workers AI | `cloudflare` | `CLOUDFLARE_API_KEY` | Access to Workers AI |
| Ollama | `ollama` | No | Local inference | | Ollama | `ollama` | No | Local inference |
| AWS Bedrock | `bedrock` | AWS credentials | Native Converse API | | AWS Bedrock | `bedrock` | AWS credentials | Native Converse API |
@@ -75,25 +74,6 @@ Pull a model first: `ollama pull llama3.2`
--- ---
## MiniMax
[MiniMax](https://platform.minimax.io) provides high-performance language models with 204,800 token context windows.
```env
LLM_BACKEND=minimax
MINIMAX_API_KEY=...
```
Available models: `MiniMax-M2.5` (default), `MiniMax-M2.5-highspeed`
To use the China mainland endpoint, set:
```env
MINIMAX_BASE_URL=https://api.minimaxi.com/v1
```
---
## AWS Bedrock (requires `--features bedrock`) ## AWS Bedrock (requires `--features bedrock`)
Uses the native AWS Converse API via `aws-sdk-bedrockruntime`. Supports standard AWS Uses the native AWS Converse API via `aws-sdk-bedrockruntime`. Supports standard AWS
-21
View File
@@ -382,27 +382,6 @@
"can_list_models": true "can_list_models": true
} }
}, },
{
"id": "minimax",
"aliases": [
"mini_max"
],
"protocol": "open_ai_completions",
"default_base_url": "https://api.minimax.io/v1",
"api_key_env": "MINIMAX_API_KEY",
"api_key_required": true,
"base_url_env": "MINIMAX_BASE_URL",
"model_env": "MINIMAX_MODEL",
"default_model": "MiniMax-M2.5",
"description": "MiniMax API (MiniMax-M2.5 and MiniMax-M2.5-highspeed models)",
"setup": {
"kind": "api_key",
"secret_name": "llm_minimax_api_key",
"key_url": "https://platform.minimax.io",
"display_name": "MiniMax",
"can_list_models": false
}
},
{ {
"id": "cloudflare", "id": "cloudflare",
"aliases": [ "aliases": [
+2 -2
View File
@@ -18,8 +18,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/discord-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/discord-wasm32-wasip2.tar.gz",
"sha256": "efa1b9019fa33e243f8db1e1fcc732731d45836336bdd26ca19b6fe227ca8b69" "sha256": null
} }
}, },
"auth_summary": { "auth_summary": {
+2 -2
View File
@@ -18,8 +18,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/slack-0.2.1-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/slack-wasm32-wasip2.tar.gz",
"sha256": "d4667e35126986509d862bc3a0088777305d8f41c75de83c1e223b42312ede48" "sha256": null
} }
}, },
"auth_summary": { "auth_summary": {
+2 -2
View File
@@ -18,8 +18,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/telegram-0.2.2-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/telegram-wasm32-wasip2.tar.gz",
"sha256": "b9a83d5a2d1285ce0ec116b354336a1f245f893291ccb01dffbcaccf89d72aed" "sha256": null
} }
}, },
"auth_summary": { "auth_summary": {
+2 -2
View File
@@ -18,8 +18,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/whatsapp-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/whatsapp-wasm32-wasip2.tar.gz",
"sha256": "feb9194719d9bed796b070ab4dc30348dbfb5d3dec56f9f21e02d14137abab01" "sha256": null
} }
}, },
"auth_summary": { "auth_summary": {
+2 -2
View File
@@ -19,8 +19,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/github-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/github-wasm32-wasip2.tar.gz",
"sha256": "da9fac56b6f20197a415489bbaec9fefb085a5cf6324cab79ea48a47eb19c13b" "sha256": null
} }
}, },
"auth_summary": { "auth_summary": {
+2 -2
View File
@@ -18,8 +18,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/gmail-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/gmail-wasm32-wasip2.tar.gz",
"sha256": "ee9574e02e92bc1d481f1310eb88afd99ee52bf6971074ab33bd76bf99b34b1d" "sha256": null
} }
}, },
"auth_summary": { "auth_summary": {
+2 -2
View File
@@ -18,8 +18,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-calendar-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/google-calendar-wasm32-wasip2.tar.gz",
"sha256": "2fa47150ea222e787c122182ad6f4dfa2ffaf5fe490d05e8de887a76445f8d2d" "sha256": null
} }
}, },
"auth_summary": { "auth_summary": {
+2 -2
View File
@@ -18,8 +18,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-docs-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/google-docs-wasm32-wasip2.tar.gz",
"sha256": "40e134a1c1564f832ca861c3396895d4e33ec67b99313fc1f97baf8d971423a9" "sha256": null
} }
}, },
"auth_summary": { "auth_summary": {
+2 -2
View File
@@ -18,8 +18,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-drive-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/google-drive-wasm32-wasip2.tar.gz",
"sha256": "002a341a1d58125563a7c69561b26fbc2629b04ea723cade744102bdc0fbb71f" "sha256": null
} }
}, },
"auth_summary": { "auth_summary": {
+2 -2
View File
@@ -18,8 +18,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-sheets-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/google-sheets-wasm32-wasip2.tar.gz",
"sha256": "8aa2c9d52f033edea3a6c2311b0ec694ccb6d0a54ef07e94d72bf8be1ce8009a" "sha256": null
} }
}, },
"auth_summary": { "auth_summary": {
+2 -2
View File
@@ -17,8 +17,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-slides-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/google-slides-wasm32-wasip2.tar.gz",
"sha256": "e931a97d4fd0b0b938e464dc7c7f2be6ea6b4d1508f5ea3cd931d44db23f05f5" "sha256": null
} }
}, },
"auth_summary": { "auth_summary": {
-41
View File
@@ -1,41 +0,0 @@
{
"name": "llm-context",
"display_name": "LLM Context",
"kind": "tool",
"version": "0.1.0",
"wit_version": "0.3.0",
"description": "Fetch pre-extracted web content from Brave Search for grounding LLM answers (RAG, fact-checking)",
"keywords": [
"search",
"web",
"brave",
"rag",
"grounding",
"llm",
"context"
],
"source": {
"dir": "tools-src/llm-context",
"capabilities": "llm-context-tool.capabilities.json",
"crate_name": "llm-context-tool"
},
"artifacts": {
"wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/llm-context-wasm32-wasip2.tar.gz",
"sha256": "581cc5867ef3b75116b7ddc8161e63dd92befe2b53e6ad8213c007639aa243c3"
}
},
"auth_summary": {
"method": "manual",
"provider": "Brave",
"secrets": [
"brave_api_key"
],
"shared_auth": "Same API key as Web Search tool (brave_api_key)",
"setup_url": "https://brave.com/search/api/"
},
"tags": [
"default",
"search"
]
}
+2 -2
View File
@@ -17,8 +17,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/slack-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/slack-tool-wasm32-wasip2.tar.gz",
"sha256": "8af3f884240de8413d272845fad2164a347d7d2a502a0d148aa38425b93f62ed" "sha256": null
} }
}, },
"auth_summary": { "auth_summary": {
+2 -2
View File
@@ -18,8 +18,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/telegram-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/telegram-mtproto-wasm32-wasip2.tar.gz",
"sha256": "2c66245913854be4294021fc6bb479e43f7d65830c5cec25cf6c60a71d1af468" "sha256": null
} }
}, },
"auth_summary": { "auth_summary": {
+2 -2
View File
@@ -18,8 +18,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/web-search-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/web-search-wasm32-wasip2.tar.gz",
"sha256": "56834573c54ea2a33cea1eb0f04bbdf59f1ef8d8702995cf431b0921302eeccc" "sha256": null
} }
}, },
"auth_summary": { "auth_summary": {
+24 -5
View File
@@ -18,7 +18,7 @@ use crate::agent::self_repair::{DefaultSelfRepair, RepairResult, SelfRepair};
use crate::agent::session_manager::SessionManager; use crate::agent::session_manager::SessionManager;
use crate::agent::submission::{Submission, SubmissionParser, SubmissionResult}; use crate::agent::submission::{Submission, SubmissionParser, SubmissionResult};
use crate::agent::{HeartbeatConfig as AgentHeartbeatConfig, Router, Scheduler}; use crate::agent::{HeartbeatConfig as AgentHeartbeatConfig, Router, Scheduler};
use crate::channels::{ChannelManager, IncomingMessage, OutgoingResponse}; use crate::channels::{ChannelManager, IncomingMessage, OutgoingResponse, StatusUpdate};
use crate::config::{AgentConfig, HeartbeatConfig, RoutineConfig, SkillsConfig}; use crate::config::{AgentConfig, HeartbeatConfig, RoutineConfig, SkillsConfig};
use crate::context::ContextManager; use crate::context::ContextManager;
use crate::db::Database; use crate::db::Database;
@@ -936,10 +936,29 @@ impl Agent {
SubmissionResult::Ok { message } => Ok(message), SubmissionResult::Ok { message } => Ok(message),
SubmissionResult::Error { message } => Ok(Some(format!("Error: {}", message))), SubmissionResult::Error { message } => Ok(Some(format!("Error: {}", message))),
SubmissionResult::Interrupted => Ok(Some("Interrupted.".into())), SubmissionResult::Interrupted => Ok(Some("Interrupted.".into())),
SubmissionResult::NeedApproval { .. } => { SubmissionResult::NeedApproval {
// ApprovalNeeded status was already sent by thread_ops.rs before request_id,
// returning this result. Empty string signals the caller to skip tool_name,
// respond() (no duplicate text). description,
parameters,
} => {
// Each channel renders the approval prompt via send_status.
// Web gateway shows an inline card, REPL prints a formatted prompt, etc.
let _ = self
.channels
.send_status(
&message.channel,
StatusUpdate::ApprovalNeeded {
request_id: request_id.to_string(),
tool_name,
description,
parameters,
},
&message.metadata,
)
.await;
// Empty string signals the caller to skip respond() (no duplicate text)
Ok(Some(String::new())) Ok(Some(String::new()))
} }
} }
-72
View File
@@ -554,31 +554,6 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
}; };
if needs_approval { if needs_approval {
// In non-DM relay channels, auto-deny approval-
// requiring tools to prevent stuck AwaitingApproval
// state and prompt injection from other users.
let is_relay = self.message.channel.ends_with("-relay");
let is_dm = self
.message
.metadata
.get("event_type")
.and_then(|v| v.as_str())
== Some("direct_message");
if is_relay && !is_dm {
tracing::info!(
tool = %tc.name,
channel = %self.message.channel,
"Auto-denying approval-requiring tool in non-DM relay channel"
);
let reject_msg = format!(
"Tool '{}' requires approval and cannot run in shared channels. \
Ask the user to message me directly (DM) to use this tool.",
tc.name
);
preflight.push((tc, PreflightOutcome::Rejected(reject_msg)));
continue;
}
approval_needed = Some((idx, tc, tool)); approval_needed = Some((idx, tc, tool));
break; break;
} }
@@ -2260,51 +2235,4 @@ mod tests {
"Present 'data' field should produce non-empty string" "Present 'data' field should produce non-empty string"
); );
} }
/// Test the relay channel auto-deny decision logic:
/// approval-requiring tools in non-DM relay channels must be rejected.
#[test]
fn test_relay_non_dm_auto_deny_decision() {
use crate::channels::IncomingMessage;
// Case 1: relay channel + non-DM → should auto-deny
let msg = IncomingMessage::new("slack-relay", "u1", "hello")
.with_metadata(serde_json::json!({ "event_type": "message" }));
let is_relay = msg.channel.ends_with("-relay");
let is_dm =
msg.metadata.get("event_type").and_then(|v| v.as_str()) == Some("direct_message");
assert!(is_relay && !is_dm, "Should auto-deny in relay non-DM");
// Case 2: relay channel + DM → should NOT auto-deny
let msg_dm = IncomingMessage::new("slack-relay", "u1", "hello")
.with_metadata(serde_json::json!({ "event_type": "direct_message" }));
let is_dm_2 =
msg_dm.metadata.get("event_type").and_then(|v| v.as_str()) == Some("direct_message");
assert!(
!msg_dm.channel.ends_with("-relay") || is_dm_2,
"Should NOT auto-deny in relay DM"
);
// Case 3: non-relay channel → should NOT auto-deny
let msg_web = IncomingMessage::new("web", "u1", "hello")
.with_metadata(serde_json::json!({ "event_type": "message" }));
assert!(
!msg_web.channel.ends_with("-relay"),
"Non-relay channel should not trigger auto-deny"
);
}
/// Test that the auto-deny produces a PreflightOutcome::Rejected-style message.
#[test]
fn test_relay_auto_deny_message_format() {
let tool_name = "shell";
let result_msg = format!(
"Tool '{}' requires approval and cannot run in shared channels. \
Ask the user to message me directly (DM) to use this tool.",
tool_name
);
assert!(result_msg.contains("shell"));
assert!(result_msg.contains("approval"));
assert!(result_msg.contains("DM"));
}
} }
+1
View File
@@ -32,6 +32,7 @@ pub mod task;
mod thread_ops; mod thread_ops;
pub mod undo; pub mod undo;
pub use crate::worker::{Worker, WorkerDeps};
pub(crate) use agent_loop::truncate_for_preview; pub(crate) use agent_loop::truncate_for_preview;
pub use agent_loop::{Agent, AgentDeps}; pub use agent_loop::{Agent, AgentDeps};
pub use compaction::{CompactionResult, ContextCompactor}; pub use compaction::{CompactionResult, ContextCompactor};
+3 -116
View File
@@ -207,7 +207,7 @@ impl Trigger {
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")] #[serde(tag = "type", rename_all = "snake_case")]
pub enum RoutineAction { pub enum RoutineAction {
/// Single LLM call (optionally with tools). Cheap and fast. /// Single LLM call, no tools. Cheap and fast.
Lightweight { Lightweight {
/// The prompt sent to the LLM. /// The prompt sent to the LLM.
prompt: String, prompt: String,
@@ -217,14 +217,6 @@ pub enum RoutineAction {
/// Max output tokens (default: 4096). /// Max output tokens (default: 4096).
#[serde(default = "default_max_tokens")] #[serde(default = "default_max_tokens")]
max_tokens: u32, max_tokens: u32,
/// Enable tool access (default: false for backward compatibility).
/// When true, the LLM can call tools during execution.
/// Tools requiring approval are automatically filtered out.
#[serde(default)]
use_tools: bool,
/// Max tool call rounds (default: 3). Only used when use_tools is true.
#[serde(default = "default_max_tool_rounds")]
max_tool_rounds: u32,
}, },
/// Full multi-turn worker job with tool access. /// Full multi-turn worker job with tool access.
FullJob { FullJob {
@@ -251,19 +243,6 @@ fn default_max_iterations() -> u32 {
10 10
} }
fn default_max_tool_rounds() -> u32 {
3
}
/// Hard upper bound for max_tool_rounds to prevent runaway loops and cost explosion.
pub(crate) const MAX_TOOL_ROUNDS_LIMIT: u32 = 20;
/// Clamp max_tool_rounds to [1, MAX_TOOL_ROUNDS_LIMIT].
/// Accepts u64 to avoid truncation before clamping.
fn clamp_max_tool_rounds(value: u64) -> u32 {
value.clamp(1, MAX_TOOL_ROUNDS_LIMIT as u64) as u32
}
/// Parse a `tool_permissions` JSON array into a `Vec<String>`. /// Parse a `tool_permissions` JSON array into a `Vec<String>`.
pub fn parse_tool_permissions(value: &serde_json::Value) -> Vec<String> { pub fn parse_tool_permissions(value: &serde_json::Value) -> Vec<String> {
value value
@@ -311,22 +290,10 @@ impl RoutineAction {
.get("max_tokens") .get("max_tokens")
.and_then(|v| v.as_u64()) .and_then(|v| v.as_u64())
.unwrap_or(default_max_tokens() as u64) as u32; .unwrap_or(default_max_tokens() as u64) as u32;
let use_tools = config
.get("use_tools")
.and_then(|v| v.as_bool())
.unwrap_or(false);
let max_tool_rounds = clamp_max_tool_rounds(
config
.get("max_tool_rounds")
.and_then(|v| v.as_u64())
.unwrap_or(default_max_tool_rounds() as u64),
);
Ok(RoutineAction::Lightweight { Ok(RoutineAction::Lightweight {
prompt, prompt,
context_paths, context_paths,
max_tokens, max_tokens,
use_tools,
max_tool_rounds,
}) })
} }
"full_job" => { "full_job" => {
@@ -372,14 +339,10 @@ impl RoutineAction {
prompt, prompt,
context_paths, context_paths,
max_tokens, max_tokens,
use_tools,
max_tool_rounds,
} => serde_json::json!({ } => serde_json::json!({
"prompt": prompt, "prompt": prompt,
"context_paths": context_paths, "context_paths": context_paths,
"max_tokens": max_tokens, "max_tokens": max_tokens,
"use_tools": use_tools,
"max_tool_rounds": max_tool_rounds,
}), }),
RoutineAction::FullJob { RoutineAction::FullJob {
title, title,
@@ -541,8 +504,7 @@ pub fn next_cron_fire(
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use crate::agent::routine::{ use crate::agent::routine::{
MAX_TOOL_ROUNDS_LIMIT, RoutineAction, RoutineGuardrails, RunStatus, Trigger, content_hash, RoutineAction, RoutineGuardrails, RunStatus, Trigger, content_hash, next_cron_fire,
next_cron_fire,
}; };
#[test] #[test]
@@ -592,13 +554,11 @@ mod tests {
prompt: "Check PRs".to_string(), prompt: "Check PRs".to_string(),
context_paths: vec!["context/priorities.md".to_string()], context_paths: vec!["context/priorities.md".to_string()],
max_tokens: 2048, max_tokens: 2048,
use_tools: false,
max_tool_rounds: 3,
}; };
let json = action.to_config_json(); let json = action.to_config_json();
let parsed = RoutineAction::from_db("lightweight", json).expect("parse lightweight"); let parsed = RoutineAction::from_db("lightweight", json).expect("parse lightweight");
assert!( assert!(
matches!(parsed, RoutineAction::Lightweight { prompt, context_paths, max_tokens, .. } matches!(parsed, RoutineAction::Lightweight { prompt, context_paths, max_tokens }
if prompt == "Check PRs" && context_paths.len() == 1 && max_tokens == 2048) if prompt == "Check PRs" && context_paths.len() == 1 && max_tokens == 2048)
); );
} }
@@ -735,77 +695,4 @@ mod tests {
); );
assert_eq!(Trigger::Manual.type_tag(), "manual"); assert_eq!(Trigger::Manual.type_tag(), "manual");
} }
#[test]
fn test_action_lightweight_backward_compat_no_use_tools() {
// Simulate old DB record without use_tools field
let json = serde_json::json!({
"prompt": "old routine",
"context_paths": [],
"max_tokens": 4096
});
let parsed = RoutineAction::from_db("lightweight", json).expect("parse lightweight");
assert!(
matches!(parsed, RoutineAction::Lightweight { use_tools, max_tool_rounds, .. }
if !use_tools && max_tool_rounds == 3),
"missing use_tools should default to false, max_tool_rounds to 3"
);
}
#[test]
fn test_max_tool_rounds_clamped_to_upper_bound() {
let json = serde_json::json!({
"prompt": "test",
"use_tools": true,
"max_tool_rounds": 9999
});
let parsed = RoutineAction::from_db("lightweight", json).expect("parse");
match parsed {
RoutineAction::Lightweight {
max_tool_rounds, ..
} => {
assert_eq!(
max_tool_rounds, MAX_TOOL_ROUNDS_LIMIT,
"should clamp to MAX_TOOL_ROUNDS_LIMIT"
);
}
_ => panic!("expected Lightweight"),
}
}
#[test]
fn test_max_tool_rounds_clamped_to_lower_bound() {
let json = serde_json::json!({
"prompt": "test",
"use_tools": true,
"max_tool_rounds": 0
});
let parsed = RoutineAction::from_db("lightweight", json).expect("parse");
match parsed {
RoutineAction::Lightweight {
max_tool_rounds, ..
} => {
assert_eq!(max_tool_rounds, 1, "should clamp 0 to 1");
}
_ => panic!("expected Lightweight"),
}
}
#[test]
fn test_max_tool_rounds_normal_value_passes_through() {
let json = serde_json::json!({
"prompt": "test",
"use_tools": true,
"max_tool_rounds": 10
});
let parsed = RoutineAction::from_db("lightweight", json).expect("parse");
match parsed {
RoutineAction::Lightweight {
max_tool_rounds, ..
} => {
assert_eq!(max_tool_rounds, 10, "normal value should pass through");
}
_ => panic!("expected Lightweight"),
}
}
} }
+23 -84
View File
@@ -459,20 +459,7 @@ async fn execute_routine(ctx: EngineContext, routine: Routine, run: RoutineRun)
prompt, prompt,
context_paths, context_paths,
max_tokens, max_tokens,
use_tools, } => execute_lightweight(&ctx, &routine, prompt, context_paths, *max_tokens).await,
max_tool_rounds,
} => {
execute_lightweight(
&ctx,
&routine,
prompt,
context_paths,
*max_tokens,
*use_tools,
*max_tool_rounds,
)
.await
}
RoutineAction::FullJob { RoutineAction::FullJob {
title, title,
description, description,
@@ -683,8 +670,6 @@ async fn execute_lightweight(
prompt: &str, prompt: &str,
context_paths: &[String], context_paths: &[String],
max_tokens: u32, max_tokens: u32,
use_tools: bool,
max_tool_rounds: u32,
) -> Result<(RunStatus, Option<String>, Option<i32>), RoutineError> { ) -> Result<(RunStatus, Option<String>, Option<i32>), RoutineError> {
// Load context from workspace // Load context from workspace
let mut context_parts = Vec::new(); let mut context_parts = Vec::new();
@@ -747,15 +732,14 @@ async fn execute_lightweight(
Err(_) => max_tokens, Err(_) => max_tokens,
}; };
// If tools are enabled (both globally and per-routine), use the tool execution loop // If tools are enabled, use the tool execution loop; otherwise, single LLM call
if use_tools && ctx.config.lightweight_tools_enabled { if ctx.config.lightweight_tools_enabled {
execute_lightweight_with_tools( execute_lightweight_with_tools(
ctx, ctx,
routine, routine,
&system_prompt, &system_prompt,
&full_prompt, &full_prompt,
effective_max_tokens, effective_max_tokens,
max_tool_rounds,
) )
.await .await
} else { } else {
@@ -799,12 +783,24 @@ async fn execute_lightweight_no_tools(
reason: e.to_string(), reason: e.to_string(),
})?; })?;
handle_text_response( let content = response.content.trim();
&response.content, let tokens_used = Some((response.input_tokens + response.output_tokens) as i32);
response.finish_reason,
response.input_tokens, // Empty content guard
response.output_tokens, if content.is_empty() {
) return if response.finish_reason == FinishReason::Length {
Err(RoutineError::TruncatedResponse)
} else {
Err(RoutineError::EmptyResponse)
};
}
// Check for the "nothing to do" sentinel
if content == "ROUTINE_OK" || content.contains("ROUTINE_OK") {
return Ok((RunStatus::Ok, None, tokens_used));
}
Ok((RunStatus::Attention, Some(content.to_string()), tokens_used))
} }
/// Handle a text-only LLM response in lightweight routine execution. /// Handle a text-only LLM response in lightweight routine execution.
@@ -854,7 +850,6 @@ async fn execute_lightweight_with_tools(
system_prompt: &str, system_prompt: &str,
full_prompt: &str, full_prompt: &str,
effective_max_tokens: u32, effective_max_tokens: u32,
max_tool_rounds: u32,
) -> Result<(RunStatus, Option<String>, Option<i32>), RoutineError> { ) -> Result<(RunStatus, Option<String>, Option<i32>), RoutineError> {
let mut messages = if system_prompt.is_empty() { let mut messages = if system_prompt.is_empty() {
vec![ChatMessage::user(full_prompt)] vec![ChatMessage::user(full_prompt)]
@@ -865,9 +860,7 @@ async fn execute_lightweight_with_tools(
] ]
}; };
let max_iterations = max_tool_rounds let max_iterations = ctx.config.lightweight_max_iterations.min(5);
.min(ctx.config.lightweight_max_iterations)
.min(5);
let mut iteration = 0; let mut iteration = 0;
let mut total_input_tokens = 0; let mut total_input_tokens = 0;
let mut total_output_tokens = 0; let mut total_output_tokens = 0;
@@ -913,10 +906,7 @@ async fn execute_lightweight_with_tools(
); );
} else { } else {
// Tool-enabled iteration // Tool-enabled iteration
let tool_defs = ctx let tool_defs = ctx.tools.tool_definitions().await;
.tools
.tool_definitions_excluding(ROUTINE_TOOL_DENYLIST)
.await;
let request = ToolCompletionRequest::new(messages.clone(), tool_defs) let request = ToolCompletionRequest::new(messages.clone(), tool_defs)
.with_max_tokens(effective_max_tokens) .with_max_tokens(effective_max_tokens)
@@ -982,33 +972,12 @@ async fn execute_lightweight_with_tools(
} }
} }
/// Tools that must never be callable from lightweight routines.
///
/// These tools pose autonomy-escalation risks: a routine could self-replicate,
/// modify its own triggers/prompts, delete other routines, or restart the agent.
const ROUTINE_TOOL_DENYLIST: &[&str] = &[
"routine_create",
"routine_update",
"routine_delete",
"routine_fire",
"restart",
];
/// Execute a single tool for a lightweight routine. /// Execute a single tool for a lightweight routine.
async fn execute_routine_tool( async fn execute_routine_tool(
ctx: &EngineContext, ctx: &EngineContext,
job_ctx: &JobContext, job_ctx: &JobContext,
tc: &ToolCall, tc: &ToolCall,
) -> Result<String, Box<dyn std::error::Error + Send + Sync>> { ) -> Result<String, Box<dyn std::error::Error + Send + Sync>> {
// Block tools that pose autonomy-escalation risks
if ROUTINE_TOOL_DENYLIST.contains(&tc.name.as_str()) {
return Err(format!(
"Tool '{}' is not available in lightweight routines",
tc.name
)
.into());
}
// Check if tool exists // Check if tool exists
let tool = ctx let tool = ctx
.tools .tools
@@ -1314,36 +1283,6 @@ mod tests {
} }
} }
#[test]
fn test_routine_tool_denylist_blocks_self_management_tools() {
let denylisted = vec![
"routine_create",
"routine_update",
"routine_delete",
"routine_fire",
"restart",
];
for tool in &denylisted {
assert!(
super::ROUTINE_TOOL_DENYLIST.contains(tool),
"Tool '{}' should be in ROUTINE_TOOL_DENYLIST",
tool
);
}
}
#[test]
fn test_routine_tool_denylist_allows_safe_tools() {
let allowed = vec!["echo", "time", "json", "http", "memory_search", "shell"];
for tool in &allowed {
assert!(
!super::ROUTINE_TOOL_DENYLIST.contains(tool),
"Tool '{}' should NOT be in ROUTINE_TOOL_DENYLIST",
tool
);
}
}
#[test] #[test]
fn test_empty_response_handling() { fn test_empty_response_handling() {
// Simulate the empty content guard logic // Simulate the empty content guard logic
+3 -18
View File
@@ -486,12 +486,7 @@ impl Agent {
.channels .channels
.send_status( .send_status(
&message.channel, &message.channel,
StatusUpdate::ApprovalNeeded { StatusUpdate::Status("Awaiting approval".into()),
request_id: request_id.to_string(),
tool_name: tool_name.clone(),
description: description.clone(),
parameters: parameters.clone(),
},
&message.metadata, &message.metadata,
) )
.await; .await;
@@ -1302,12 +1297,7 @@ impl Agent {
.channels .channels
.send_status( .send_status(
&message.channel, &message.channel,
StatusUpdate::ApprovalNeeded { StatusUpdate::Status("Awaiting approval".into()),
request_id: request_id.to_string(),
tool_name: tool_name.clone(),
description: description.clone(),
parameters: parameters.clone(),
},
&message.metadata, &message.metadata,
) )
.await; .await;
@@ -1378,12 +1368,7 @@ impl Agent {
.channels .channels
.send_status( .send_status(
&message.channel, &message.channel,
StatusUpdate::ApprovalNeeded { StatusUpdate::Status("Awaiting approval".into()),
request_id: request_id.to_string(),
tool_name: tool_name.clone(),
description: description.clone(),
parameters: parameters.clone(),
},
&message.metadata, &message.metadata,
) )
.await; .await;
+1 -74
View File
@@ -9,7 +9,6 @@
use std::sync::Arc; use std::sync::Arc;
use crate::agent::SessionManager as AgentSessionManager;
use crate::channels::web::log_layer::LogBroadcaster; use crate::channels::web::log_layer::LogBroadcaster;
use crate::config::Config; use crate::config::Config;
use crate::context::ContextManager; use crate::context::ContextManager;
@@ -47,8 +46,6 @@ pub struct AppComponents {
pub log_broadcaster: Arc<LogBroadcaster>, pub log_broadcaster: Arc<LogBroadcaster>,
pub context_manager: Arc<ContextManager>, pub context_manager: Arc<ContextManager>,
pub hooks: Arc<HookRegistry>, pub hooks: Arc<HookRegistry>,
/// Shared thread/session manager used by the standard agent runtime.
pub agent_session_manager: Arc<AgentSessionManager>,
pub skill_registry: Option<Arc<std::sync::RwLock<SkillRegistry>>>, pub skill_registry: Option<Arc<std::sync::RwLock<SkillRegistry>>>,
pub skill_catalog: Option<Arc<SkillCatalog>>, pub skill_catalog: Option<Arc<SkillCatalog>>,
pub cost_guard: Arc<crate::agent::cost_guard::CostGuard>, pub cost_guard: Arc<crate::agent::cost_guard::CostGuard>,
@@ -303,8 +300,7 @@ impl AppBuilder {
// Register memory tools if database is available // Register memory tools if database is available
let workspace = if let Some(ref db) = self.db { let workspace = if let Some(ref db) = self.db {
let mut ws = Workspace::new_with_db("default", db.clone()) let mut ws = Workspace::new_with_db("default", db.clone());
.with_search_config(&self.config.search);
if let Some(ref emb) = embeddings { if let Some(ref emb) = embeddings {
ws = ws.with_embeddings(emb.clone()); ws = ws.with_embeddings(emb.clone());
} }
@@ -693,8 +689,6 @@ impl AppBuilder {
// Create hook registry early so runtime extension activation can register hooks. // Create hook registry early so runtime extension activation can register hooks.
let hooks = Arc::new(HookRegistry::new()); let hooks = Arc::new(HookRegistry::new());
let agent_session_manager =
Arc::new(AgentSessionManager::new().with_hooks(Arc::clone(&hooks)));
let ( let (
mcp_session_manager, mcp_session_manager,
@@ -801,7 +795,6 @@ impl AppBuilder {
log_broadcaster: self.log_broadcaster, log_broadcaster: self.log_broadcaster,
context_manager, context_manager,
hooks, hooks,
agent_session_manager,
skill_registry, skill_registry,
skill_catalog, skill_catalog,
cost_guard, cost_guard,
@@ -812,69 +805,3 @@ impl AppBuilder {
}) })
} }
} }
#[cfg(test)]
mod tests {
use std::sync::Arc;
use async_trait::async_trait;
use tokio::sync::mpsc;
use crate::agent::SessionManager as AgentSessionManager;
use crate::hooks::{
Hook, HookContext, HookError, HookEvent, HookOutcome, HookPoint, HookRegistry,
};
struct SessionStartHook {
tx: mpsc::UnboundedSender<(String, String)>,
}
#[async_trait]
impl Hook for SessionStartHook {
fn name(&self) -> &str {
"session-start-test"
}
fn hook_points(&self) -> &[HookPoint] {
&[HookPoint::OnSessionStart]
}
async fn execute(
&self,
event: &HookEvent,
_ctx: &HookContext,
) -> Result<HookOutcome, HookError> {
if let HookEvent::SessionStart {
user_id,
session_id,
} = event
{
self.tx
.send((user_id.clone(), session_id.clone()))
.expect("test channel receiver should be alive");
} else {
panic!("SessionStartHook received an unexpected event: {event:?}");
}
Ok(HookOutcome::ok())
}
}
#[tokio::test]
async fn agent_session_manager_runs_session_start_hooks() {
let hooks = Arc::new(HookRegistry::new());
let (tx, mut rx) = mpsc::unbounded_channel();
hooks.register(Arc::new(SessionStartHook { tx })).await;
let manager = AgentSessionManager::new().with_hooks(Arc::clone(&hooks));
manager.get_or_create_session("user-123").await;
let (user_id, session_id) =
tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv())
.await
.expect("session start hook should fire")
.expect("session start payload should be present");
assert_eq!(user_id, "user-123");
assert!(!session_id.is_empty());
}
}
-61
View File
@@ -807,67 +807,6 @@ mod tests {
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
} }
/// Regression test for issue #869: RwLock read guard was held across
/// tx.send(msg).await in `process_message()`, blocking shutdown() from
/// acquiring the write lock when the channel buffer was full.
///
/// This test exercises the actual production code path (`process_message`)
/// with a full channel buffer, then verifies shutdown() can still complete.
#[tokio::test]
async fn shutdown_completes_while_process_message_blocked() {
let channel = Arc::new(test_channel(Some("secret")));
let stream = channel.start().await.unwrap();
// Fill all 256 slots in the channel buffer
{
let tx = {
let guard = channel.state.tx.read().await;
guard.as_ref().unwrap().clone()
};
for i in 0..256 {
let msg = IncomingMessage::new("http", "user", format!("fill-{}", i));
tx.send(msg).await.unwrap();
}
}
// Signal so we know the spawned task has started and is about to
// call process_message (which will block on the full channel).
let started = Arc::new(tokio::sync::Notify::new());
let started_clone = started.clone();
// Spawn a task that calls the actual production code path.
// process_message() internally acquires the RwLock read guard and
// sends on the channel. With the fix, the guard is released before
// send().await; without the fix, shutdown() would deadlock.
let state = channel.state.clone();
let blocked_send = tokio::spawn(async move {
started_clone.notify_one();
let msg = IncomingMessage::new("http", "user", "blocked-257th");
let _ = process_message(state, msg, false).await;
});
// Wait for the spawned task to start, then give it time to reach
// the send().await and verify that it is still pending (i.e., blocked).
started.notified().await;
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
assert!(
!blocked_send.is_finished(),
"process_message task should still be pending before shutdown()"
);
// shutdown() must complete even though process_message is blocked on
// send(). Before the fix, the read guard held across send().await
// would prevent shutdown() from acquiring the write lock.
let result =
tokio::time::timeout(std::time::Duration::from_secs(2), channel.shutdown()).await;
assert!(result.is_ok(), "shutdown() must not deadlock");
assert!(result.unwrap().is_ok());
// Drop the stream (receiver) so the blocked send task can complete
drop(stream);
let _ = blocked_send.await;
}
#[tokio::test] #[tokio::test]
async fn webhook_missing_all_auth_returns_unauthorized() { async fn webhook_missing_all_auth_returns_unauthorized() {
let channel = test_channel(Some("correct-secret")); let channel = test_channel(Some("correct-secret"));
+3 -225
View File
@@ -408,120 +408,12 @@ impl Channel for RelayChannel {
Ok(()) Ok(())
} }
/// Status updates are not forwarded to messaging providers to avoid noise.
async fn send_status( async fn send_status(
&self, &self,
status: StatusUpdate, _status: StatusUpdate,
metadata: &serde_json::Value, _metadata: &serde_json::Value,
) -> Result<(), ChannelError> { ) -> Result<(), ChannelError> {
// Only handle ApprovalNeeded — all other variants are no-ops
let StatusUpdate::ApprovalNeeded {
request_id,
tool_name,
description,
parameters,
} = status
else {
return Ok(());
};
// Only send buttons in DMs (dispatcher gates upstream, but guard here too)
let event_type = metadata
.get("event_type")
.and_then(|v| v.as_str())
.unwrap_or("");
if event_type != "direct_message" {
tracing::warn!(
tool = %tool_name,
event_type,
"Approval requested in non-DM, skipping buttons"
);
return Ok(());
}
// Extract required metadata — error if missing
let channel_id = metadata
.get("channel_id")
.and_then(|v| v.as_str())
.ok_or_else(|| ChannelError::SendFailed {
name: self.name().to_string(),
reason: "Missing channel_id for approval buttons".into(),
})?;
let sender_id = metadata
.get("sender_id")
.and_then(|v| v.as_str())
.ok_or_else(|| ChannelError::SendFailed {
name: self.name().to_string(),
reason: "Missing sender_id for approval buttons".into(),
})?;
let thread_id = metadata.get("thread_id").and_then(|v| v.as_str());
let team_id = metadata
.get("team_id")
.and_then(|v| v.as_str())
.unwrap_or(&self.team_id);
// Button value payload (Slack limits button values to 2000 chars;
// safe with typical UUIDs but documented here as a constraint)
let value_payload = serde_json::json!({
"instance_id": self.instance_id,
"team_id": team_id,
"channel_id": channel_id,
"thread_ts": thread_id,
"request_id": request_id,
"sender_id": sender_id,
});
let value_str = value_payload.to_string();
// Parameters are already redacted via redact_params() in dispatcher.rs
let params_display =
serde_json::to_string_pretty(&parameters).unwrap_or_else(|_| parameters.to_string());
let blocks = serde_json::json!([
{
"type": "section",
"text": {
"type": "mrkdwn",
"text": format!(
"*Tool approval required*\n`{tool_name}`: {description}\n```{params_display}```"
)
}
},
{
"type": "actions",
"elements": [
{
"type": "button",
"text": { "type": "plain_text", "text": "Approve" },
"style": "primary",
"action_id": "approve_tool",
"value": value_str,
},
{
"type": "button",
"text": { "type": "plain_text", "text": "Deny" },
"style": "danger",
"action_id": "deny_tool",
"value": value_str,
}
]
}
]);
let mut body = serde_json::json!({
"channel": channel_id,
"text": format!("Tool approval required: {tool_name} - {description}"),
"blocks": blocks,
});
if let Some(tid) = thread_id {
body["thread_ts"] = serde_json::Value::String(tid.to_string());
}
self.proxy_send(team_id, "chat.postMessage", body)
.await
.map_err(|e| ChannelError::SendFailed {
name: self.name().to_string(),
reason: e.to_string(),
})?;
Ok(()) Ok(())
} }
@@ -747,118 +639,4 @@ mod tests {
// The reconnect loop now skips team validation when team_id is empty, // The reconnect loop now skips team validation when team_id is empty,
// so the channel remains alive. // so the channel remains alive.
} }
#[tokio::test]
async fn test_send_status_non_approval_is_noop() {
let channel = RelayChannel::new(
test_client(),
"token".into(),
"T123".into(),
"inst1".into(),
"user1".into(),
);
let metadata = serde_json::json!({});
let result = channel
.send_status(
StatusUpdate::ToolStarted {
name: "echo".into(),
},
&metadata,
)
.await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_send_status_approval_non_dm_skips() {
let channel = RelayChannel::new(
test_client(),
"token".into(),
"T123".into(),
"inst1".into(),
"user1".into(),
);
let metadata = serde_json::json!({
"event_type": "message",
"channel_id": "C456",
"sender_id": "U789",
});
let result = channel
.send_status(
StatusUpdate::ApprovalNeeded {
request_id: "req1".into(),
tool_name: "shell".into(),
description: "run command".into(),
parameters: serde_json::json!({}),
},
&metadata,
)
.await;
// Non-DM approval requests are silently skipped (no HTTP call)
assert!(result.is_ok());
}
#[tokio::test]
async fn test_send_status_approval_dm_missing_channel_id_errors() {
let channel = RelayChannel::new(
test_client(),
"token".into(),
"T123".into(),
"inst1".into(),
"user1".into(),
);
let metadata = serde_json::json!({
"event_type": "direct_message",
"sender_id": "U789",
});
let result = channel
.send_status(
StatusUpdate::ApprovalNeeded {
request_id: "req1".into(),
tool_name: "shell".into(),
description: "run command".into(),
parameters: serde_json::json!({}),
},
&metadata,
)
.await;
assert!(result.is_err());
let err = result.unwrap_err().to_string();
assert!(
err.contains("channel_id"),
"expected channel_id error, got: {err}"
);
}
#[tokio::test]
async fn test_send_status_approval_dm_missing_sender_id_errors() {
let channel = RelayChannel::new(
test_client(),
"token".into(),
"T123".into(),
"inst1".into(),
"user1".into(),
);
let metadata = serde_json::json!({
"event_type": "direct_message",
"channel_id": "C456",
});
let result = channel
.send_status(
StatusUpdate::ApprovalNeeded {
request_id: "req1".into(),
tool_name: "shell".into(),
description: "run command".into(),
parameters: serde_json::json!({}),
},
&metadata,
)
.await;
assert!(result.is_err());
let err = result.unwrap_err().to_string();
assert!(
err.contains("sender_id"),
"expected sender_id error, got: {err}"
);
}
} }
+2 -26
View File
@@ -63,11 +63,7 @@ const ALLOWED_MIME_PREFIXES: &[&str] = &[
"application/x-tar", "application/x-tar",
"application/octet-stream", "application/octet-stream",
]; ];
/// Truncate a string to at most `max_bytes` without splitting UTF-8 code points.
fn truncate_utf8(s: &str, max_bytes: usize) -> &str {
let end = crate::util::floor_char_boundary(s, max_bytes);
&s[..end]
}
/// A message emitted by a WASM channel to be sent to the agent. /// A message emitted by a WASM channel to be sent to the agent.
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub struct EmittedMessage { pub struct EmittedMessage {
@@ -268,7 +264,7 @@ impl ChannelHostState {
max = MAX_MESSAGE_CONTENT_SIZE, max = MAX_MESSAGE_CONTENT_SIZE,
"Message content too large, truncating" "Message content too large, truncating"
); );
let mut truncated = truncate_utf8(&msg.content, MAX_MESSAGE_CONTENT_SIZE).to_string(); let mut truncated = msg.content[..MAX_MESSAGE_CONTENT_SIZE].to_string();
truncated.push_str("... (truncated)"); truncated.push_str("... (truncated)");
let msg = EmittedMessage { let msg = EmittedMessage {
content: truncated, content: truncated,
@@ -635,7 +631,6 @@ mod tests {
use crate::channels::wasm::host::{ use crate::channels::wasm::host::{
Attachment, ChannelEmitRateLimiter, ChannelHostState, EmittedMessage, Attachment, ChannelEmitRateLimiter, ChannelHostState, EmittedMessage,
MAX_ATTACHMENT_TOTAL_SIZE, MAX_ATTACHMENTS_PER_MESSAGE, MAX_EMITS_PER_EXECUTION, MAX_ATTACHMENT_TOTAL_SIZE, MAX_ATTACHMENTS_PER_MESSAGE, MAX_EMITS_PER_EXECUTION,
MAX_MESSAGE_CONTENT_SIZE,
}; };
#[test] #[test]
@@ -694,25 +689,6 @@ mod tests {
assert_eq!(state.emits_dropped(), 1); assert_eq!(state.emits_dropped(), 1);
} }
#[test]
fn test_emit_message_truncates_utf8_safely() {
let caps = ChannelCapabilities::for_channel("test");
let mut state = ChannelHostState::new("test", caps);
let prefix = "a".repeat(MAX_MESSAGE_CONTENT_SIZE - 1);
let content = format!("{}🙂suffix", prefix);
let msg = EmittedMessage::new("user123", content);
state.emit_message(msg).unwrap();
let messages = state.take_emitted_messages();
assert_eq!(messages.len(), 1);
let emitted = &messages[0].content;
assert!(emitted.starts_with(&prefix));
assert!(emitted.ends_with("... (truncated)"));
assert!(!emitted.contains("🙂"));
}
#[test] #[test]
fn test_workspace_write_prefixing() { fn test_workspace_write_prefixing() {
let caps = ChannelCapabilities::for_channel("slack"); let caps = ChannelCapabilities::for_channel("slack");
+40 -50
View File
@@ -1994,33 +1994,28 @@ impl WasmChannel {
return Ok(()); return Ok(());
} }
// Clone sender to avoid holding RwLock read guard across send().await in the loop let tx_guard = self.message_tx.read().await;
let tx = { let Some(tx) = tx_guard.as_ref() else {
let tx_guard = self.message_tx.read().await; tracing::error!(
let Some(tx) = tx_guard.as_ref() else { channel = %self.name,
tracing::error!( count = messages.len(),
channel = %self.name, "Messages emitted but no sender available - channel may not be started!"
count = messages.len(), );
"Messages emitted but no sender available - channel may not be started!" return Ok(());
);
return Ok(());
};
tx.clone()
}; };
let mut rate_limiter = self.rate_limiter.write().await;
for emitted in messages { for emitted in messages {
// Check rate limit — acquire and release the write lock before send().await // Check rate limit
{ if !rate_limiter.check_and_record() {
let mut rate_limiter = self.rate_limiter.write().await; tracing::warn!(
if !rate_limiter.check_and_record() { channel = %self.name,
tracing::warn!( "Message emission rate limited"
channel = %self.name, );
"Message emission rate limited" return Err(WasmChannelError::EmitRateLimited {
); name: self.name.clone(),
return Err(WasmChannelError::EmitRateLimited { });
name: self.name.clone(),
});
}
} }
// Convert to IncomingMessage // Convert to IncomingMessage
@@ -2062,7 +2057,7 @@ impl WasmChannel {
self.update_broadcast_metadata(&emitted.metadata_json).await; self.update_broadcast_metadata(&emitted.metadata_json).await;
} }
// Send to stream — no locks held across this await // Send to stream
tracing::info!( tracing::info!(
channel = %self.name, channel = %self.name,
user_id = %emitted.user_id, user_id = %emitted.user_id,
@@ -2286,33 +2281,28 @@ impl WasmChannel {
"Processing emitted messages from polling callback" "Processing emitted messages from polling callback"
); );
// Clone sender to avoid holding RwLock read guard across send().await in the loop let tx_guard = message_tx.read().await;
let tx = { let Some(tx) = tx_guard.as_ref() else {
let tx_guard = message_tx.read().await; tracing::error!(
let Some(tx) = tx_guard.as_ref() else { channel = %channel_name,
tracing::error!( count = messages.len(),
channel = %channel_name, "Messages emitted but no sender available - channel may not be started!"
count = messages.len(), );
"Messages emitted but no sender available - channel may not be started!" return Ok(());
);
return Ok(());
};
tx.clone()
}; };
let mut limiter = rate_limiter.write().await;
for emitted in messages { for emitted in messages {
// Check rate limit — acquire and release the write lock before send().await // Check rate limit
{ if !limiter.check_and_record() {
let mut limiter = rate_limiter.write().await; tracing::warn!(
if !limiter.check_and_record() { channel = %channel_name,
tracing::warn!( "Message emission rate limited"
channel = %channel_name, );
"Message emission rate limited" return Err(WasmChannelError::EmitRateLimited {
); name: channel_name.to_string(),
return Err(WasmChannelError::EmitRateLimited { });
name: channel_name.to_string(),
});
}
} }
// Convert to IncomingMessage // Convert to IncomingMessage
@@ -2360,7 +2350,7 @@ impl WasmChannel {
.await; .await;
} }
// Send to stream — no locks held across this await // Send to stream
tracing::info!( tracing::info!(
channel = %channel_name, channel = %channel_name,
user_id = %emitted.user_id, user_id = %emitted.user_id,
+10 -22
View File
@@ -37,17 +37,11 @@ pub async fn chat_send_handler(
let msg_id = msg.id; let msg_id = msg.id;
let thread_id = msg.thread_id.clone(); let thread_id = msg.thread_id.clone();
// Clone sender to avoid holding RwLock read guard across send().await let tx_guard = state.msg_tx.read().await;
let tx = { let tx = tx_guard.as_ref().ok_or((
let tx_guard = state.msg_tx.read().await; StatusCode::SERVICE_UNAVAILABLE,
tx_guard "Channel not started".to_string(),
.as_ref() ))?;
.ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Channel not started".to_string(),
))?
.clone()
};
tx.send(msg).await.map_err(|_| { tx.send(msg).await.map_err(|_| {
( (
@@ -117,17 +111,11 @@ pub async fn chat_approval_handler(
let msg_id = msg.id; let msg_id = msg.id;
// Clone sender to avoid holding RwLock read guard across send().await let tx_guard = state.msg_tx.read().await;
let tx = { let tx = tx_guard.as_ref().ok_or((
let tx_guard = state.msg_tx.read().await; StatusCode::SERVICE_UNAVAILABLE,
tx_guard "Channel not started".to_string(),
.as_ref() ))?;
.ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Channel not started".to_string(),
))?
.clone()
};
tx.send(msg).await.map_err(|_| { tx.send(msg).await.map_err(|_| {
( (
+13 -74
View File
@@ -581,12 +581,7 @@ async fn oauth_callback_handler(
let exchange_proxy_url = std::env::var("IRONCLAW_OAUTH_EXCHANGE_URL").ok(); let exchange_proxy_url = std::env::var("IRONCLAW_OAUTH_EXCHANGE_URL").ok();
let result: Result<(), String> = async { let result: Result<(), String> = async {
let token_response = if let (Some(proxy_url), None) = (&exchange_proxy_url, &flow.resource) let token_response = if let Some(ref proxy_url) = exchange_proxy_url {
{
// Use the platform exchange proxy when configured and no resource
// parameter is needed. The proxy holds client_secret server-side so
// the container never sees it. MCP flows (resource.is_some()) bypass
// the proxy because it doesn't forward the RFC 8707 resource param.
let gateway_token = flow.gateway_token.as_deref().unwrap_or_default(); let gateway_token = flow.gateway_token.as_deref().unwrap_or_default();
oauth_defaults::exchange_via_proxy( oauth_defaults::exchange_via_proxy(
proxy_url, proxy_url,
@@ -599,10 +594,7 @@ async fn oauth_callback_handler(
.await .await
.map_err(|e| e.to_string())? .map_err(|e| e.to_string())?
} else { } else {
// Direct token exchange: uses exchange_oauth_code_with_resource so MCP oauth_defaults::exchange_oauth_code(
// flows can include the RFC 8707 `resource` parameter to scope the
// issued token to the specific MCP server.
oauth_defaults::exchange_oauth_code_with_resource(
&flow.token_url, &flow.token_url,
&flow.client_id, &flow.client_id,
flow.client_secret.as_deref(), flow.client_secret.as_deref(),
@@ -610,7 +602,6 @@ async fn oauth_callback_handler(
&flow.redirect_uri, &flow.redirect_uri,
flow.code_verifier.as_deref(), flow.code_verifier.as_deref(),
&flow.access_token_field, &flow.access_token_field,
flow.resource.as_deref(),
) )
.await .await
.map_err(|e| e.to_string())? .map_err(|e| e.to_string())?
@@ -637,19 +628,6 @@ async fn oauth_callback_handler(
.await .await
.map_err(|e| e.to_string())?; .map_err(|e| e.to_string())?;
// For MCP OAuth flows (identified by resource field), persist the
// client_id so token refresh works without re-authentication.
// The CLI flow stores this in authorize_mcp_server(); the gateway
// callback must do the same.
if let Some(ref client_id_secret) = flow.client_id_secret_name {
let params = crate::secrets::CreateSecretParams::new(client_id_secret, &flow.client_id)
.with_provider(flow.provider.as_ref().cloned().unwrap_or_default());
flow.secrets
.create(&flow.user_id, params)
.await
.map_err(|e| e.to_string())?;
}
Ok(()) Ok(())
} }
.await; .await;
@@ -681,35 +659,12 @@ async fn oauth_callback_handler(
} }
} }
// After successful OAuth, auto-activate the extension so it moves
// from "Installed (Authenticate)" → "Active" without a second click.
// OAuth success is independent of activation — tokens are already stored.
// Report auth as successful and attempt activation as a bonus step.
let final_message = if success {
match ext_mgr.activate(&flow.extension_name).await {
Ok(result) => result.message,
Err(e) => {
tracing::warn!(
extension = %flow.extension_name,
error = %e,
"Auto-activation after OAuth failed"
);
format!(
"{} authenticated successfully. Activation failed: {}. Try activating manually.",
flow.display_name, e
)
}
}
} else {
message
};
// Broadcast SSE event to notify the web UI // Broadcast SSE event to notify the web UI
if let Some(ref sender) = flow.sse_sender { if let Some(ref sender) = flow.sse_sender {
let _ = sender.send(SseEvent::AuthCompleted { let _ = sender.send(SseEvent::AuthCompleted {
extension_name: flow.extension_name, extension_name: flow.extension_name,
success, success,
message: final_message.clone(), message,
}); });
} }
@@ -1018,17 +973,11 @@ async fn chat_send_handler(
req.images.len() req.images.len()
); );
// Clone sender to avoid holding RwLock read guard across send().await let tx_guard = state.msg_tx.read().await;
let tx = { let tx = tx_guard.as_ref().ok_or((
let tx_guard = state.msg_tx.read().await; StatusCode::SERVICE_UNAVAILABLE,
tx_guard "Channel not started".to_string(),
.as_ref() ))?;
.ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Channel not started".to_string(),
))?
.clone()
};
tracing::debug!("[chat_send_handler] Sending message through channel"); tracing::debug!("[chat_send_handler] Sending message through channel");
tx.send(msg).await.map_err(|_| { tx.send(msg).await.map_err(|_| {
@@ -1094,17 +1043,11 @@ async fn chat_approval_handler(
let msg_id = msg.id; let msg_id = msg.id;
// Clone sender to avoid holding RwLock read guard across send().await let tx_guard = state.msg_tx.read().await;
let tx = { let tx = tx_guard.as_ref().ok_or((
let tx_guard = state.msg_tx.read().await; StatusCode::SERVICE_UNAVAILABLE,
tx_guard "Channel not started".to_string(),
.as_ref() ))?;
.ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Channel not started".to_string(),
))?
.clone()
};
tx.send(msg).await.map_err(|_| { tx.send(msg).await.map_err(|_| {
( (
@@ -3011,8 +2954,6 @@ mod tests {
secrets, secrets,
sse_sender: None, sse_sender: None,
gateway_token: None, gateway_token: None,
resource: None,
client_id_secret_name: None,
created_at: std::time::Instant::now() created_at: std::time::Instant::now()
.checked_sub(std::time::Duration::from_secs(600)) .checked_sub(std::time::Duration::from_secs(600))
.expect("System uptime is too low to run expired flow test"), .expect("System uptime is too low to run expired flow test"),
@@ -3122,8 +3063,6 @@ mod tests {
secrets, secrets,
sse_sender: None, sse_sender: None,
gateway_token: None, gateway_token: None,
resource: None,
client_id_secret_name: None,
// Expired — handler will reject after lookup (no network I/O) // Expired — handler will reject after lookup (no network I/O)
created_at: std::time::Instant::now() created_at: std::time::Instant::now()
.checked_sub(std::time::Duration::from_secs(600)) .checked_sub(std::time::Duration::from_secs(600))
+28 -149
View File
@@ -670,7 +670,7 @@ function renderMarkdown(text) {
// Sanitize HTML output to prevent XSS from tool output or LLM responses. // Sanitize HTML output to prevent XSS from tool output or LLM responses.
html = sanitizeRenderedHtml(html); html = sanitizeRenderedHtml(html);
// Inject copy buttons into <pre> blocks // Inject copy buttons into <pre> blocks
html = html.replace(/<pre>/g, '<pre class="code-block-wrapper"><button class="copy-btn" data-action="copy-code">Copy</button>'); html = html.replace(/<pre>/g, '<pre class="code-block-wrapper"><button class="copy-btn" onclick="copyCodeBlock(this)">Copy</button>');
return html; return html;
} }
return escapeHtml(text); return escapeHtml(text);
@@ -702,25 +702,16 @@ function copyCodeBlock(btn) {
}); });
} }
function copyMessage(btn) {
const message = btn.closest('.message');
if (!message) return;
const text = message.getAttribute('data-copy-text')
|| message.getAttribute('data-raw')
|| message.textContent
|| '';
navigator.clipboard.writeText(text).then(() => {
btn.textContent = 'Copied';
setTimeout(() => { btn.textContent = 'Copy'; }, 1200);
}).catch(() => {
btn.textContent = 'Failed';
setTimeout(() => { btn.textContent = 'Copy'; }, 1200);
});
}
function addMessage(role, content) { function addMessage(role, content) {
const container = document.getElementById('chat-messages'); const container = document.getElementById('chat-messages');
const div = createMessageElement(role, content); const div = document.createElement('div');
div.className = 'message ' + role;
if (role === 'user') {
div.textContent = content;
} else {
div.setAttribute('data-raw', content);
div.innerHTML = renderMarkdown(content);
}
container.appendChild(div); container.appendChild(div);
container.scrollTop = container.scrollHeight; container.scrollTop = container.scrollHeight;
} }
@@ -732,11 +723,7 @@ function appendToLastAssistant(chunk) {
const last = messages[messages.length - 1]; const last = messages[messages.length - 1];
const raw = (last.getAttribute('data-raw') || '') + chunk; const raw = (last.getAttribute('data-raw') || '') + chunk;
last.setAttribute('data-raw', raw); last.setAttribute('data-raw', raw);
last.setAttribute('data-copy-text', raw); last.innerHTML = renderMarkdown(raw);
const content = last.querySelector('.message-content');
if (content) {
content.innerHTML = renderMarkdown(raw);
}
container.scrollTop = container.scrollHeight; container.scrollTop = container.scrollHeight;
} else { } else {
addMessage('assistant', chunk); addMessage('assistant', chunk);
@@ -1323,31 +1310,12 @@ function loadHistory(before) {
function createMessageElement(role, content) { function createMessageElement(role, content) {
const div = document.createElement('div'); const div = document.createElement('div');
div.className = 'message ' + role; div.className = 'message ' + role;
if (role === 'user') {
if (role === 'assistant' || role === 'user') { div.textContent = content;
div.classList.add('has-copy');
div.setAttribute('data-copy-text', content);
const copyBtn = document.createElement('button');
copyBtn.className = 'message-copy-btn';
copyBtn.type = 'button';
copyBtn.setAttribute('aria-label', 'Copy message');
copyBtn.textContent = 'Copy';
copyBtn.addEventListener('click', (e) => {
e.stopPropagation();
copyMessage(copyBtn);
});
div.appendChild(copyBtn);
}
const body = document.createElement('div');
body.className = 'message-content';
if (role === 'user' || role === 'system') {
body.textContent = content;
} else { } else {
div.setAttribute('data-raw', content); div.setAttribute('data-raw', content);
body.innerHTML = renderMarkdown(content); div.innerHTML = renderMarkdown(content);
} }
div.appendChild(body);
return div; return div;
} }
@@ -1851,11 +1819,13 @@ function saveMemoryEdit() {
function buildBreadcrumb(path) { function buildBreadcrumb(path) {
const parts = path.split('/'); const parts = path.split('/');
let html = '<a data-action="breadcrumb-root" href="#">workspace</a>'; let html = '<a onclick="loadMemoryTree()">workspace</a>';
let current = ''; let current = '';
for (const part of parts) { for (const part of parts) {
current += (current ? '/' : '') + part; current += (current ? '/' : '') + part;
html += ' / <a data-action="breadcrumb-file" data-path="' + escapeHtml(current) + '" href="#">' + escapeHtml(part) + '</a>'; // Store the path in data-path (HTML-escaped) and read it back via this.dataset.path
// to avoid single-quote injection in inline JS string literals.
html += ' / <a onclick="readMemoryFile(this.dataset.path)" data-path="' + escapeHtml(current) + '">' + escapeHtml(part) + '</a>';
} }
return html; return html;
} }
@@ -2825,11 +2795,11 @@ function renderJobsList(jobs) {
let actionBtns = ''; let actionBtns = '';
if (job.state === 'pending' || job.state === 'in_progress') { if (job.state === 'pending' || job.state === 'in_progress') {
actionBtns = '<button class="btn-cancel" data-action="cancel-job" data-id="' + escapeHtml(job.id) + '">Cancel</button>'; actionBtns = '<button class="btn-cancel" onclick="event.stopPropagation(); cancelJob(\'' + job.id + '\')">Cancel</button>';
} }
// Retry is only shown in the detail view where can_restart is available. // Retry is only shown in the detail view where can_restart is available.
return '<tr class="job-row" data-action="open-job" data-id="' + escapeHtml(job.id) + '">' return '<tr class="job-row" onclick="openJobDetail(\'' + job.id + '\')">'
+ '<td title="' + escapeHtml(job.id) + '">' + shortId + '</td>' + '<td title="' + escapeHtml(job.id) + '">' + shortId + '</td>'
+ '<td>' + escapeHtml(job.title) + '</td>' + '<td>' + escapeHtml(job.title) + '</td>'
+ '<td><span class="badge ' + stateClass + '">' + escapeHtml(job.state) + '</span></td>' + '<td><span class="badge ' + stateClass + '">' + escapeHtml(job.state) + '</span></td>'
@@ -2892,12 +2862,12 @@ function renderJobDetail(job) {
const header = document.createElement('div'); const header = document.createElement('div');
header.className = 'job-detail-header'; header.className = 'job-detail-header';
let headerHtml = '<button class="btn-back" data-action="close-job-detail">&larr; Back</button>' let headerHtml = '<button class="btn-back" onclick="closeJobDetail()">&larr; Back</button>'
+ '<h2>' + escapeHtml(job.title) + '</h2>' + '<h2>' + escapeHtml(job.title) + '</h2>'
+ '<span class="badge ' + stateClass + '">' + escapeHtml(job.state) + '</span>'; + '<span class="badge ' + stateClass + '">' + escapeHtml(job.state) + '</span>';
if ((job.state === 'failed' || job.state === 'interrupted') && job.can_restart === true) { if ((job.state === 'failed' || job.state === 'interrupted') && job.can_restart === true) {
headerHtml += '<button class="btn-restart" data-action="restart-job" data-id="' + escapeHtml(job.id) + '">Retry</button>'; headerHtml += '<button class="btn-restart" onclick="restartJob(\'' + job.id + '\')">Retry</button>';
} }
if (job.browse_url) { if (job.browse_url) {
headerHtml += '<a class="btn-browse" href="' + escapeHtml(job.browse_url) + '" target="_blank">Browse Files</a>'; headerHtml += '<a class="btn-browse" href="' + escapeHtml(job.browse_url) + '" target="_blank">Browse Files</a>';
@@ -3354,7 +3324,7 @@ function renderRoutinesList(routines) {
const toggleLabel = r.enabled ? 'Disable' : 'Enable'; const toggleLabel = r.enabled ? 'Disable' : 'Enable';
const toggleClass = r.enabled ? 'btn-cancel' : 'btn-restart'; const toggleClass = r.enabled ? 'btn-cancel' : 'btn-restart';
return '<tr class="routine-row" data-action="open-routine" data-id="' + escapeHtml(r.id) + '">' return '<tr class="routine-row" onclick="openRoutineDetail(\'' + r.id + '\')">'
+ '<td>' + escapeHtml(r.name) + '</td>' + '<td>' + escapeHtml(r.name) + '</td>'
+ '<td>' + escapeHtml(r.trigger_summary) + '</td>' + '<td>' + escapeHtml(r.trigger_summary) + '</td>'
+ '<td>' + escapeHtml(r.action_type) + '</td>' + '<td>' + escapeHtml(r.action_type) + '</td>'
@@ -3363,9 +3333,9 @@ function renderRoutinesList(routines) {
+ '<td>' + r.run_count + '</td>' + '<td>' + r.run_count + '</td>'
+ '<td><span class="badge ' + statusClass + '">' + escapeHtml(r.status) + '</span></td>' + '<td><span class="badge ' + statusClass + '">' + escapeHtml(r.status) + '</span></td>'
+ '<td>' + '<td>'
+ '<button class="' + toggleClass + '" data-action="toggle-routine" data-id="' + escapeHtml(r.id) + '">' + toggleLabel + '</button> ' + '<button class="' + toggleClass + '" onclick="event.stopPropagation(); toggleRoutine(\'' + r.id + '\')">' + toggleLabel + '</button> '
+ '<button class="btn-restart" data-action="trigger-routine" data-id="' + escapeHtml(r.id) + '">Run</button> ' + '<button class="btn-restart" onclick="event.stopPropagation(); triggerRoutine(\'' + r.id + '\')">Run</button> '
+ '<button class="btn-cancel" data-action="delete-routine" data-id="' + escapeHtml(r.id) + '" data-name="' + escapeHtml(r.name) + '">Delete</button>' + '<button class="btn-cancel" onclick="event.stopPropagation(); deleteRoutine(\'' + r.id + '\', \'' + escapeHtml(r.name) + '\')">Delete</button>'
+ '</td>' + '</td>'
+ '</tr>'; + '</tr>';
}).join(''); }).join('');
@@ -3401,7 +3371,7 @@ function renderRoutineDetail(routine) {
: 'active'; : 'active';
let html = '<div class="job-detail-header">' let html = '<div class="job-detail-header">'
+ '<button class="btn-back" data-action="close-routine-detail">&larr; Back</button>' + '<button class="btn-back" onclick="closeRoutineDetail()">&larr; Back</button>'
+ '<h2>' + escapeHtml(routine.name) + '</h2>' + '<h2>' + escapeHtml(routine.name) + '</h2>'
+ '<span class="badge ' + statusClass + '">' + escapeHtml(statusLabel) + '</span>' + '<span class="badge ' + statusClass + '">' + escapeHtml(statusLabel) + '</span>'
+ '</div>'; + '</div>';
@@ -3448,7 +3418,7 @@ function renderRoutineDetail(routine) {
+ '<td>' + formatDate(run.completed_at) + '</td>' + '<td>' + formatDate(run.completed_at) + '</td>'
+ '<td><span class="badge ' + runStatusClass + '">' + escapeHtml(run.status) + '</span></td>' + '<td><span class="badge ' + runStatusClass + '">' + escapeHtml(run.status) + '</span></td>'
+ '<td>' + escapeHtml(run.result_summary || '-') + '<td>' + escapeHtml(run.result_summary || '-')
+ (run.job_id ? ' <a href="#" data-action="view-run-job" data-id="' + escapeHtml(run.job_id) + '">[view job]</a>' : '') + (run.job_id ? ' <a href="#" onclick="event.preventDefault(); switchTab(\'jobs\'); openJobDetail(\'' + run.job_id + '\')">[view job]</a>' : '')
+ '</td>' + '</td>'
+ '<td>' + (run.tokens_used != null ? run.tokens_used : '-') + '</td>' + '<td>' + (run.tokens_used != null ? run.tokens_used : '-') + '</td>'
+ '</tr>'; + '</tr>';
@@ -3691,7 +3661,7 @@ function renderTeePopover(report) {
+ '<div class="tee-field"><div class="tee-field-label">VM Config</div>' + '<div class="tee-field"><div class="tee-field-label">VM Config</div>'
+ '<div class="tee-field-value">' + escapeHtml(vmConfig) + '</div></div>' + '<div class="tee-field-value">' + escapeHtml(vmConfig) + '</div></div>'
+ '<div class="tee-popover-actions">' + '<div class="tee-popover-actions">'
+ '<button class="tee-btn-copy" data-action="copy-tee-report">Copy Full Report</button></div>'; + '<button class="tee-btn-copy" onclick="copyTeeReport()">Copy Full Report</button></div>';
} }
function copyTeeReport() { function copyTeeReport() {
@@ -4173,94 +4143,3 @@ function formatDate(isoString) {
const d = new Date(isoString); const d = new Date(isoString);
return d.toLocaleString(); return d.toLocaleString();
} }
// --- Event Listener Registration (CSP-safe, no inline handlers) ---
document.getElementById('auth-connect-btn').addEventListener('click', () => authenticate());
document.getElementById('restart-overlay').addEventListener('click', () => cancelRestart());
document.getElementById('restart-close-btn').addEventListener('click', () => cancelRestart());
document.getElementById('restart-cancel-btn').addEventListener('click', () => cancelRestart());
document.getElementById('restart-confirm-btn').addEventListener('click', () => confirmRestart());
document.getElementById('restart-btn').addEventListener('click', () => triggerRestart());
document.getElementById('thread-new-btn').addEventListener('click', () => createNewThread());
document.getElementById('thread-toggle-btn').addEventListener('click', () => toggleThreadSidebar());
document.getElementById('assistant-thread').addEventListener('click', () => switchToAssistant());
document.getElementById('send-btn').addEventListener('click', () => sendMessage());
document.getElementById('memory-edit-btn').addEventListener('click', () => startMemoryEdit());
document.getElementById('memory-save-btn').addEventListener('click', () => saveMemoryEdit());
document.getElementById('memory-cancel-btn').addEventListener('click', () => cancelMemoryEdit());
document.getElementById('logs-server-level').addEventListener('change', (e) => setServerLogLevel(e.target.value));
document.getElementById('logs-pause-btn').addEventListener('click', () => toggleLogsPause());
document.getElementById('logs-clear-btn').addEventListener('click', () => clearLogs());
document.getElementById('wasm-install-btn').addEventListener('click', () => installWasmExtension());
document.getElementById('mcp-add-btn').addEventListener('click', () => addMcpServer());
document.getElementById('skill-search-btn').addEventListener('click', () => searchClawHub());
document.getElementById('skill-install-btn').addEventListener('click', () => installSkillFromForm());
// --- Delegated Event Handlers (for dynamically generated HTML) ---
document.addEventListener('click', function(e) {
const el = e.target.closest('[data-action]');
if (!el) return;
const action = el.dataset.action;
switch (action) {
case 'copy-code':
copyCodeBlock(el);
break;
case 'breadcrumb-root':
e.preventDefault();
loadMemoryTree();
break;
case 'breadcrumb-file':
e.preventDefault();
readMemoryFile(el.dataset.path);
break;
case 'cancel-job':
e.stopPropagation();
cancelJob(el.dataset.id);
break;
case 'open-job':
openJobDetail(el.dataset.id);
break;
case 'close-job-detail':
closeJobDetail();
break;
case 'restart-job':
restartJob(el.dataset.id);
break;
case 'open-routine':
openRoutineDetail(el.dataset.id);
break;
case 'toggle-routine':
e.stopPropagation();
toggleRoutine(el.dataset.id);
break;
case 'trigger-routine':
e.stopPropagation();
triggerRoutine(el.dataset.id);
break;
case 'delete-routine':
e.stopPropagation();
deleteRoutine(el.dataset.id, el.dataset.name);
break;
case 'close-routine-detail':
closeRoutineDetail();
break;
case 'view-run-job':
e.preventDefault();
switchTab('jobs');
openJobDetail(el.dataset.id);
break;
case 'copy-tee-report':
copyTeeReport();
break;
case 'switch-language':
if (typeof switchLanguage === 'function') switchLanguage(el.dataset.lang);
break;
}
});
document.getElementById('language-btn').addEventListener('click', function() {
if (typeof toggleLanguageMenu === 'function') toggleLanguageMenu();
});
+23 -23
View File
@@ -37,7 +37,7 @@
<div class="auth-form"> <div class="auth-form">
<label for="token-input" data-i18n="auth.tokenLabel">Gateway Token</label> <label for="token-input" data-i18n="auth.tokenLabel">Gateway Token</label>
<input type="password" id="token-input" data-i18n="auth.tokenPlaceholder" data-i18n-attr="placeholder" placeholder="Paste your auth token" autofocus> <input type="password" id="token-input" data-i18n="auth.tokenPlaceholder" data-i18n-attr="placeholder" placeholder="Paste your auth token" autofocus>
<button id="auth-connect-btn" data-i18n="auth.connect">Connect</button> <button onclick="authenticate()" data-i18n="auth.connect">Connect</button>
</div> </div>
<div id="auth-error"></div> <div id="auth-error"></div>
<p class="auth-hint" data-i18n="auth.hint">Enter the GATEWAY_AUTH_TOKEN from your .env configuration.</p> <p class="auth-hint" data-i18n="auth.hint">Enter the GATEWAY_AUTH_TOKEN from your .env configuration.</p>
@@ -46,11 +46,11 @@
<!-- Restart Confirmation Modal --> <!-- Restart Confirmation Modal -->
<div id="restart-confirm-modal" class="restart-modal" style="display: none;"> <div id="restart-confirm-modal" class="restart-modal" style="display: none;">
<div class="restart-modal-overlay" id="restart-overlay"></div> <div class="restart-modal-overlay" onclick="cancelRestart()"></div>
<div class="restart-modal-content"> <div class="restart-modal-content">
<div class="restart-modal-header"> <div class="restart-modal-header">
<h2 data-i18n="restart.title">Restart IronClaw Instance</h2> <h2 data-i18n="restart.title">Restart IronClaw Instance</h2>
<button class="restart-modal-close" id="restart-close-btn" data-i18n="restart.closeTooltip" data-i18n-attr="title" <button class="restart-modal-close" onclick="cancelRestart()" data-i18n="restart.closeTooltip" data-i18n-attr="title"
title="Close">×</button> title="Close">×</button>
</div> </div>
<div class="restart-modal-body"> <div class="restart-modal-body">
@@ -63,8 +63,8 @@
</div> </div>
</div> </div>
<div class="restart-modal-footer"> <div class="restart-modal-footer">
<button class="restart-modal-btn cancel" id="restart-cancel-btn" data-i18n="restart.cancel">Cancel</button> <button class="restart-modal-btn cancel" onclick="cancelRestart()" data-i18n="restart.cancel">Cancel</button>
<button class="restart-modal-btn confirm" id="restart-confirm-btn" data-i18n="restart.confirm">Confirm Restart</button> <button class="restart-modal-btn confirm" onclick="confirmRestart()" data-i18n="restart.confirm">Confirm Restart</button>
</div> </div>
</div> </div>
</div> </div>
@@ -101,11 +101,11 @@
<!-- Language Switcher --> <!-- Language Switcher -->
<div class="language-switcher"> <div class="language-switcher">
<button class="language-btn" id="language-btn" type="button" title="Switch Language" <button class="language-btn" id="language-btn" type="button" onclick="toggleLanguageMenu()" title="Switch Language"
aria-label="Switch language" aria-haspopup="true" aria-expanded="false" aria-controls="language-menu">🌐</button> aria-label="Switch language" aria-haspopup="true" aria-expanded="false" aria-controls="language-menu">🌐</button>
<div class="language-menu" id="language-menu" style="display: none;"> <div class="language-menu" id="language-menu" style="display: none;">
<button type="button" class="language-option" data-action="switch-language" data-lang="en">English</button> <button type="button" class="language-option" onclick="switchLanguage('en')" data-lang="en">English</button>
<button type="button" class="language-option" data-action="switch-language" data-lang="zh-CN">简体中文</button> <button type="button" class="language-option" onclick="switchLanguage('zh-CN')" data-lang="zh-CN">简体中文</button>
</div> </div>
</div> </div>
@@ -122,7 +122,7 @@
<span id="sse-status" data-i18n="status.connected">Connected</span> <span id="sse-status" data-i18n="status.connected">Connected</span>
<div class="gateway-popover" id="gateway-popover"></div> <div class="gateway-popover" id="gateway-popover"></div>
</div> </div>
<button class="restart-btn" id="restart-btn" data-i18n="status.restartTooltip" <button class="restart-btn" id="restart-btn" onclick="triggerRestart()" data-i18n="status.restartTooltip"
data-i18n-attr="title" title="Gracefully restart the process" style="display: none;"> data-i18n-attr="title" title="Gracefully restart the process" style="display: none;">
<svg id="restart-icon" width="13" height="13" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round"> <svg id="restart-icon" width="13" height="13" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round">
<path d="M23 4v6h-6"></path> <path d="M23 4v6h-6"></path>
@@ -137,13 +137,13 @@
<div class="tab-panel active" id="tab-chat"> <div class="tab-panel active" id="tab-chat">
<div class="thread-sidebar" id="thread-sidebar"> <div class="thread-sidebar" id="thread-sidebar">
<div class="thread-sidebar-header"> <div class="thread-sidebar-header">
<button class="thread-new-btn" id="thread-new-btn" data-i18n="chat.newThread" data-i18n-attr="title" <button class="thread-new-btn" onclick="createNewThread()" data-i18n="chat.newThread" data-i18n-attr="title"
title="New thread (Ctrl/Cmd+N)">+</button> title="New thread (Ctrl/Cmd+N)">+</button>
<div class="spacer"></div> <div class="spacer"></div>
<button class="thread-toggle-btn" id="thread-toggle-btn" data-i18n="chat.toggleSidebar" <button class="thread-toggle-btn" id="thread-toggle-btn" onclick="toggleThreadSidebar()" data-i18n="chat.toggleSidebar"
data-i18n-attr="title" title="Toggle sidebar">&laquo;</button> data-i18n-attr="title" title="Toggle sidebar">&laquo;</button>
</div> </div>
<div class="assistant-item" id="assistant-thread"> <div class="assistant-item" id="assistant-thread" onclick="switchToAssistant()">
<span class="assistant-label" id="assistant-label" data-i18n="chat.assistant">Assistant</span> <span class="assistant-label" id="assistant-label" data-i18n="chat.assistant">Assistant</span>
<span class="assistant-meta" id="assistant-meta"></span> <span class="assistant-meta" id="assistant-meta"></span>
</div> </div>
@@ -161,7 +161,7 @@
<input type="file" id="image-file-input" accept="image/*" multiple style="display:none"> <input type="file" id="image-file-input" accept="image/*" multiple style="display:none">
<button id="attach-btn" class="attach-btn" data-i18n="chat.attachImages" data-i18n-attr="title" title="Attach images" <button id="attach-btn" class="attach-btn" data-i18n="chat.attachImages" data-i18n-attr="title" title="Attach images"
aria-label="Attach images">&#x1F4CE;</button> aria-label="Attach images">&#x1F4CE;</button>
<button id="send-btn" data-i18n="chat.send">Send</button> <button id="send-btn" onclick="sendMessage()" data-i18n="chat.send">Send</button>
</div> </div>
</div> </div>
</div> </div>
@@ -178,7 +178,7 @@
<div class="memory-content"> <div class="memory-content">
<div class="memory-breadcrumb" id="memory-breadcrumb"> <div class="memory-breadcrumb" id="memory-breadcrumb">
<span id="memory-breadcrumb-path">workspace /</span> <span id="memory-breadcrumb-path">workspace /</span>
<button class="memory-edit-btn" id="memory-edit-btn" style="display:none" data-i18n="memory.edit">Edit</button> <button class="memory-edit-btn" id="memory-edit-btn" style="display:none" onclick="startMemoryEdit()" data-i18n="memory.edit">Edit</button>
</div> </div>
<div class="memory-viewer" id="memory-viewer"> <div class="memory-viewer" id="memory-viewer">
<div class="empty" data-i18n="memory.selectFile">Select a file to view its contents</div> <div class="empty" data-i18n="memory.selectFile">Select a file to view its contents</div>
@@ -186,8 +186,8 @@
<div class="memory-editor" id="memory-editor" style="display:none"> <div class="memory-editor" id="memory-editor" style="display:none">
<textarea id="memory-edit-textarea"></textarea> <textarea id="memory-edit-textarea"></textarea>
<div class="memory-editor-actions"> <div class="memory-editor-actions">
<button class="btn-save" id="memory-save-btn" data-i18n="memory.save">Save</button> <button class="btn-save" onclick="saveMemoryEdit()" data-i18n="memory.save">Save</button>
<button class="btn-cancel-edit" id="memory-cancel-btn" data-i18n="memory.cancel">Cancel</button> <button class="btn-cancel-edit" onclick="cancelMemoryEdit()" data-i18n="memory.cancel">Cancel</button>
</div> </div>
</div> </div>
</div> </div>
@@ -219,7 +219,7 @@
<div class="tab-panel" id="tab-logs"> <div class="tab-panel" id="tab-logs">
<div class="logs-container"> <div class="logs-container">
<div class="logs-toolbar"> <div class="logs-toolbar">
<select id="logs-server-level" title="Server-side log level (changes what the server emits)"> <select id="logs-server-level" onchange="setServerLogLevel(this.value)" title="Server-side log level (changes what the server emits)">
<option value="error">Server: ERROR</option> <option value="error">Server: ERROR</option>
<option value="warn">Server: WARN</option> <option value="warn">Server: WARN</option>
<option value="info" selected>Server: INFO</option> <option value="info" selected>Server: INFO</option>
@@ -234,8 +234,8 @@
</select> </select>
<input type="text" id="logs-target-filter" placeholder="Filter by target..."> <input type="text" id="logs-target-filter" placeholder="Filter by target...">
<label class="logs-checkbox"><input type="checkbox" id="logs-autoscroll" checked> <span data-i18n="logs.autoScroll">Auto-scroll</span></label> <label class="logs-checkbox"><input type="checkbox" id="logs-autoscroll" checked> <span data-i18n="logs.autoScroll">Auto-scroll</span></label>
<button id="logs-pause-btn" data-i18n="logs.pause">Pause</button> <button id="logs-pause-btn" onclick="toggleLogsPause()" data-i18n="logs.pause">Pause</button>
<button id="logs-clear-btn" data-i18n="logs.clear">Clear</button> <button onclick="clearLogs()" data-i18n="logs.clear">Clear</button>
</div> </div>
<div class="logs-output" id="logs-output"></div> <div class="logs-output" id="logs-output"></div>
</div> </div>
@@ -287,7 +287,7 @@
<div class="ext-install-form"> <div class="ext-install-form">
<input type="text" id="wasm-install-name" data-i18n-placeholder="common.name" placeholder="Extension name"> <input type="text" id="wasm-install-name" data-i18n-placeholder="common.name" placeholder="Extension name">
<input type="text" id="wasm-install-url" placeholder="URL to .tar.gz bundle"> <input type="text" id="wasm-install-url" placeholder="URL to .tar.gz bundle">
<button id="wasm-install-btn" data-i18n="extensions.install">Install</button> <button onclick="installWasmExtension()" data-i18n="extensions.install">Install</button>
</div> </div>
</div> </div>
<div class="extensions-section"> <div class="extensions-section">
@@ -299,7 +299,7 @@
<div class="ext-install-form"> <div class="ext-install-form">
<input type="text" id="mcp-install-name" data-i18n-placeholder="common.name" placeholder="Server name"> <input type="text" id="mcp-install-name" data-i18n-placeholder="common.name" placeholder="Server name">
<input type="text" id="mcp-install-url" placeholder="MCP server URL (https://...)"> <input type="text" id="mcp-install-url" placeholder="MCP server URL (https://...)">
<button id="mcp-add-btn" data-i18n="mcp.add">Add</button> <button onclick="addMcpServer()" data-i18n="mcp.add">Add</button>
</div> </div>
</div> </div>
<div class="extensions-section"> <div class="extensions-section">
@@ -320,7 +320,7 @@
<h3 data-i18n="skills.searchClawHub">Search ClawHub</h3> <h3 data-i18n="skills.searchClawHub">Search ClawHub</h3>
<div class="skill-search-box"> <div class="skill-search-box">
<input type="text" id="skill-search-input" data-i18n-placeholder="skills.searchPlaceholder" placeholder="Search..."> <input type="text" id="skill-search-input" data-i18n-placeholder="skills.searchPlaceholder" placeholder="Search...">
<button id="skill-search-btn" data-i18n="skills.search">Search</button> <button onclick="searchClawHub()" data-i18n="skills.search">Search</button>
</div> </div>
<div class="extensions-list" id="skill-search-results"></div> <div class="extensions-list" id="skill-search-results"></div>
</div> </div>
@@ -335,7 +335,7 @@
<div class="ext-install-form"> <div class="ext-install-form">
<input type="text" id="skill-install-name" data-i18n-placeholder="skills.namePlaceholder" placeholder="Skill name or slug"> <input type="text" id="skill-install-name" data-i18n-placeholder="skills.namePlaceholder" placeholder="Skill name or slug">
<input type="text" id="skill-install-url" data-i18n-placeholder="skills.urlPlaceholder" placeholder="HTTPS URL to SKILL.md (optional)"> <input type="text" id="skill-install-url" data-i18n-placeholder="skills.urlPlaceholder" placeholder="HTTPS URL to SKILL.md (optional)">
<button id="skill-install-btn" data-i18n="extensions.install">Install</button> <button onclick="installSkillFromForm()" data-i18n="extensions.install">Install</button>
</div> </div>
</div> </div>
</div> </div>
-53
View File
@@ -666,7 +666,6 @@ body {
font-size: 14px; font-size: 14px;
line-height: 1.5; line-height: 1.5;
word-wrap: break-word; word-wrap: break-word;
position: relative;
} }
.message.user { .message.user {
@@ -687,58 +686,6 @@ body {
line-height: 1.6; line-height: 1.6;
} }
.message.has-copy {
padding-right: 52px;
}
.message-content {
min-width: 0;
}
.message-copy-btn {
position: absolute;
top: 8px;
right: 8px;
z-index: 2;
border: 1px solid var(--border);
background: var(--bg-primary);
color: var(--text-secondary);
border-radius: 8px;
font-size: 11px;
padding: 2px 8px;
opacity: 0;
pointer-events: none;
transition: opacity 0.15s ease;
}
.message.user:hover .message-copy-btn,
.message.assistant:hover .message-copy-btn,
.message.user:focus-within .message-copy-btn,
.message.assistant:focus-within .message-copy-btn {
opacity: 1;
pointer-events: auto;
}
.message-copy-btn:focus-visible {
opacity: 1;
pointer-events: auto;
outline: 2px solid var(--accent);
outline-offset: 1px;
}
.message-copy-btn:hover {
background: var(--bg-secondary);
color: var(--text-primary);
}
@media (hover: none) {
.message.user .message-copy-btn,
.message.assistant .message-copy-btn {
opacity: 1;
pointer-events: auto;
}
}
.message.system { .message.system {
align-self: center; align-self: center;
background: var(--bg-tertiary); background: var(--bg-tertiary);
+4 -12
View File
@@ -176,12 +176,8 @@ async fn handle_client_message(
incoming = incoming.with_attachments(attachments); incoming = incoming.with_attachments(attachments);
} }
// Clone sender to avoid holding RwLock read guard across send().await let tx_guard = state.msg_tx.read().await;
let tx = { if let Some(ref tx) = *tx_guard {
let tx_guard = state.msg_tx.read().await;
tx_guard.as_ref().cloned()
};
if let Some(tx) = tx {
if tx.send(incoming).await.is_err() { if tx.send(incoming).await.is_err() {
let _ = direct_tx let _ = direct_tx
.send(WsServerMessage::Error { .send(WsServerMessage::Error {
@@ -249,12 +245,8 @@ async fn handle_client_message(
if let Some(ref tid) = thread_id { if let Some(ref tid) = thread_id {
msg = msg.with_thread(tid); msg = msg.with_thread(tid);
} }
// Clone sender to avoid holding RwLock read guard across send().await let tx_guard = state.msg_tx.read().await;
let tx = { if let Some(ref tx) = *tx_guard {
let tx_guard = state.msg_tx.read().await;
tx_guard.as_ref().cloned()
};
if let Some(tx) = tx {
let _ = tx.send(msg).await; let _ = tx.send(msg).await;
} }
} }
-79
View File
@@ -172,35 +172,6 @@ pub async fn exchange_oauth_code(
redirect_uri: &str, redirect_uri: &str,
code_verifier: Option<&str>, code_verifier: Option<&str>,
access_token_field: &str, access_token_field: &str,
) -> Result<OAuthTokenResponse, OAuthCallbackError> {
// Delegates to exchange_oauth_code_with_resource with resource=None.
// Non-MCP OAuth flows don't need the RFC 8707 resource parameter.
exchange_oauth_code_with_resource(
token_url,
client_id,
client_secret,
code,
redirect_uri,
code_verifier,
access_token_field,
None,
)
.await
}
/// Exchange an OAuth authorization code for tokens, with optional RFC 8707 `resource` parameter.
///
/// The `resource` parameter scopes the issued token to a specific server (used by MCP OAuth).
#[allow(clippy::too_many_arguments)]
pub async fn exchange_oauth_code_with_resource(
token_url: &str,
client_id: &str,
client_secret: Option<&str>,
code: &str,
redirect_uri: &str,
code_verifier: Option<&str>,
access_token_field: &str,
resource: Option<&str>,
) -> Result<OAuthTokenResponse, OAuthCallbackError> { ) -> Result<OAuthTokenResponse, OAuthCallbackError> {
let client = reqwest::Client::new(); let client = reqwest::Client::new();
let mut token_params = vec![ let mut token_params = vec![
@@ -213,12 +184,6 @@ pub async fn exchange_oauth_code_with_resource(
token_params.push(("code_verifier", verifier.to_string())); token_params.push(("code_verifier", verifier.to_string()));
} }
// RFC 8707: include the `resource` parameter so the authorization server
// scopes the issued token to the specific MCP server (protected resource).
if let Some(resource) = resource {
token_params.push(("resource", resource.to_string()));
}
let mut request = client.post(token_url); let mut request = client.post(token_url);
if let Some(secret) = client_secret { if let Some(secret) = client_secret {
@@ -423,12 +388,6 @@ pub struct PendingOAuthFlow {
pub sse_sender: Option<tokio::sync::broadcast::Sender<crate::channels::web::types::SseEvent>>, pub sse_sender: Option<tokio::sync::broadcast::Sender<crate::channels::web::types::SseEvent>>,
/// Gateway auth token for authenticating with the platform token exchange proxy. /// Gateway auth token for authenticating with the platform token exchange proxy.
pub gateway_token: Option<String>, pub gateway_token: Option<String>,
/// RFC 8707 resource parameter (MCP OAuth only).
/// Sent during token exchange to scope the token to a specific MCP server.
pub resource: Option<String>,
/// Secret name for persisting the client ID (MCP OAuth only).
/// Needed so token refresh can find the client_id after the session ends.
pub client_id_secret_name: Option<String>,
/// When this flow was created (for expiry). /// When this flow was created (for expiry).
pub created_at: std::time::Instant, pub created_at: std::time::Instant,
} }
@@ -1016,42 +975,4 @@ mod tests {
assert_eq!(strip_instance_prefix("abc123"), "abc123"); assert_eq!(strip_instance_prefix("abc123"), "abc123");
assert_eq!(strip_instance_prefix(""), ""); assert_eq!(strip_instance_prefix(""), "");
} }
/// Verify that `build_oauth_url` includes the RFC 8707 `resource` parameter
/// when passed through `extra_params`, which is how MCP OAuth gateway mode
/// scopes tokens to a specific MCP server.
#[test]
fn test_build_oauth_url_includes_resource_via_extra_params() {
use std::collections::HashMap;
use crate::cli::oauth_defaults::build_oauth_url;
let mut extra = HashMap::new();
extra.insert(
"resource".to_string(),
"https://mcp.example.com".to_string(),
);
let result = build_oauth_url(
"https://auth.example.com/authorize",
"client-123",
"https://gateway.example.com/oauth/callback",
&["read".to_string()],
true,
&extra,
);
// The resource parameter should be URL-encoded in the auth URL
assert!(
result
.url
.contains("resource=https%3A%2F%2Fmcp.example.com"),
"Expected resource param in URL: {}",
result.url
);
// State and PKCE should be present
assert!(result.url.contains("state="));
assert!(result.url.contains("code_challenge="));
assert!(result.code_verifier.is_some());
}
} }
-2
View File
@@ -330,8 +330,6 @@ async fn create(
prompt: prompt.to_string(), prompt: prompt.to_string(),
context_paths: Vec::new(), context_paths: Vec::new(),
max_tokens: 4096, max_tokens: 4096,
use_tools: false,
max_tool_rounds: 0,
}, },
guardrails: RoutineGuardrails { guardrails: RoutineGuardrails {
cooldown: std::time::Duration::from_secs(cooldown_secs), cooldown: std::time::Duration::from_secs(cooldown_secs),
-5
View File
@@ -18,7 +18,6 @@ pub mod relay;
mod routines; mod routines;
mod safety; mod safety;
mod sandbox; mod sandbox;
mod search;
mod secrets; mod secrets;
mod skills; mod skills;
mod transcription; mod transcription;
@@ -45,7 +44,6 @@ pub use self::routines::RoutineConfig;
pub use self::safety::SafetyConfig; pub use self::safety::SafetyConfig;
use self::safety::resolve_safety_config; use self::safety::resolve_safety_config;
pub use self::sandbox::{ClaudeCodeConfig, SandboxModeConfig}; pub use self::sandbox::{ClaudeCodeConfig, SandboxModeConfig};
pub use self::search::WorkspaceSearchConfig;
pub use self::secrets::SecretsConfig; pub use self::secrets::SecretsConfig;
pub use self::skills::SkillsConfig; pub use self::skills::SkillsConfig;
pub use self::transcription::TranscriptionConfig; pub use self::transcription::TranscriptionConfig;
@@ -93,7 +91,6 @@ pub struct Config {
pub claude_code: ClaudeCodeConfig, pub claude_code: ClaudeCodeConfig,
pub skills: SkillsConfig, pub skills: SkillsConfig,
pub transcription: TranscriptionConfig, pub transcription: TranscriptionConfig,
pub search: WorkspaceSearchConfig,
pub observability: crate::observability::ObservabilityConfig, pub observability: crate::observability::ObservabilityConfig,
/// Channel-relay integration (Slack via external relay service). /// Channel-relay integration (Slack via external relay service).
/// Present only when both `CHANNEL_RELAY_URL` and `CHANNEL_RELAY_API_KEY` are set. /// Present only when both `CHANNEL_RELAY_URL` and `CHANNEL_RELAY_API_KEY` are set.
@@ -169,7 +166,6 @@ impl Config {
..SkillsConfig::default() ..SkillsConfig::default()
}, },
transcription: TranscriptionConfig::default(), transcription: TranscriptionConfig::default(),
search: WorkspaceSearchConfig::default(),
observability: crate::observability::ObservabilityConfig::default(), observability: crate::observability::ObservabilityConfig::default(),
relay: None, relay: None,
} }
@@ -322,7 +318,6 @@ impl Config {
claude_code: ClaudeCodeConfig::resolve()?, claude_code: ClaudeCodeConfig::resolve()?,
skills: SkillsConfig::resolve()?, skills: SkillsConfig::resolve()?,
transcription: TranscriptionConfig::resolve(settings)?, transcription: TranscriptionConfig::resolve(settings)?,
search: WorkspaceSearchConfig::resolve()?,
observability: crate::observability::ObservabilityConfig { observability: crate::observability::ObservabilityConfig {
backend: std::env::var("OBSERVABILITY_BACKEND").unwrap_or_else(|_| "none".into()), backend: std::env::var("OBSERVABILITY_BACKEND").unwrap_or_else(|_| "none".into()),
}, },
-211
View File
@@ -1,211 +0,0 @@
use crate::config::helpers::{optional_env, parse_optional_env};
use crate::error::ConfigError;
use crate::workspace::FusionStrategy;
/// Workspace search configuration resolved from environment variables.
#[derive(Debug, Clone)]
pub struct WorkspaceSearchConfig {
/// Fusion strategy: "rrf" or "weighted".
pub fusion_strategy: FusionStrategy,
/// RRF constant k (default 60).
pub rrf_k: u32,
/// FTS weight for fusion.
///
/// [`Default`] uses 0.5. When the configuration is resolved, per-strategy
/// defaults are applied: 0.5 (RRF) or 0.3 (weighted).
pub fts_weight: f32,
/// Vector weight for fusion.
///
/// [`Default`] uses 0.5. When the configuration is resolved, per-strategy
/// defaults are applied: 0.5 (RRF) or 0.7 (weighted).
pub vector_weight: f32,
}
impl Default for WorkspaceSearchConfig {
fn default() -> Self {
Self {
fusion_strategy: FusionStrategy::default(),
rrf_k: 60,
fts_weight: 0.5,
vector_weight: 0.5,
}
}
}
impl WorkspaceSearchConfig {
pub(crate) fn resolve() -> Result<Self, ConfigError> {
let fusion_strategy = match optional_env("SEARCH_FUSION_STRATEGY")? {
Some(s) => match s.to_lowercase().as_str() {
"rrf" => FusionStrategy::Rrf,
"weighted" => FusionStrategy::WeightedScore,
other => {
return Err(ConfigError::InvalidValue {
key: "SEARCH_FUSION_STRATEGY".to_string(),
message: format!("must be 'rrf' or 'weighted', got '{other}'"),
});
}
},
None => FusionStrategy::default(),
};
let rrf_k = parse_optional_env("SEARCH_RRF_K", 60u32)?;
// Per-strategy weight defaults: RRF uses 0.5/0.5, weighted uses 0.3/0.7 (vector-biased).
let (default_fts, default_vec) = match fusion_strategy {
FusionStrategy::Rrf => (0.5f32, 0.5f32),
FusionStrategy::WeightedScore => (0.3f32, 0.7f32),
};
let fts_weight = parse_optional_env("SEARCH_FTS_WEIGHT", default_fts)?;
let vector_weight = parse_optional_env("SEARCH_VECTOR_WEIGHT", default_vec)?;
if !fts_weight.is_finite() || fts_weight < 0.0 {
return Err(ConfigError::InvalidValue {
key: "SEARCH_FTS_WEIGHT".to_string(),
message: "must be a finite, non-negative float".to_string(),
});
}
if !vector_weight.is_finite() || vector_weight < 0.0 {
return Err(ConfigError::InvalidValue {
key: "SEARCH_VECTOR_WEIGHT".to_string(),
message: "must be a finite, non-negative float".to_string(),
});
}
if matches!(fusion_strategy, FusionStrategy::WeightedScore)
&& fts_weight == 0.0
&& vector_weight == 0.0
{
return Err(ConfigError::InvalidValue {
key: "SEARCH_FTS_WEIGHT/SEARCH_VECTOR_WEIGHT".to_string(),
message: "weighted fusion requires at least one non-zero weight".to_string(),
});
}
Ok(Self {
fusion_strategy,
rrf_k,
fts_weight,
vector_weight,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::helpers::ENV_MUTEX;
fn clear_search_env() {
// SAFETY: Only called under ENV_MUTEX in tests.
unsafe {
std::env::remove_var("SEARCH_FUSION_STRATEGY");
std::env::remove_var("SEARCH_RRF_K");
std::env::remove_var("SEARCH_FTS_WEIGHT");
std::env::remove_var("SEARCH_VECTOR_WEIGHT");
}
}
#[test]
fn defaults_when_no_env() {
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
clear_search_env();
let config = WorkspaceSearchConfig::resolve().expect("should resolve");
assert_eq!(config.fusion_strategy, FusionStrategy::Rrf);
assert_eq!(config.rrf_k, 60);
assert!((config.fts_weight - 0.5).abs() < 0.001);
assert!((config.vector_weight - 0.5).abs() < 0.001);
}
#[test]
fn env_overrides() {
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
clear_search_env();
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::set_var("SEARCH_FUSION_STRATEGY", "weighted");
std::env::set_var("SEARCH_RRF_K", "30");
std::env::set_var("SEARCH_FTS_WEIGHT", "0.9");
std::env::set_var("SEARCH_VECTOR_WEIGHT", "0.1");
}
let config = WorkspaceSearchConfig::resolve().expect("should resolve");
assert_eq!(config.fusion_strategy, FusionStrategy::WeightedScore);
assert_eq!(config.rrf_k, 30);
assert!((config.fts_weight - 0.9).abs() < 0.001);
assert!((config.vector_weight - 0.1).abs() < 0.001);
clear_search_env();
}
#[test]
fn invalid_strategy_rejected() {
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
clear_search_env();
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::set_var("SEARCH_FUSION_STRATEGY", "bm25");
}
let result = WorkspaceSearchConfig::resolve();
assert!(result.is_err());
clear_search_env();
}
#[test]
fn weighted_strategy_defaults() {
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
clear_search_env();
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::set_var("SEARCH_FUSION_STRATEGY", "weighted");
}
let config = WorkspaceSearchConfig::resolve().expect("should resolve");
assert_eq!(config.fusion_strategy, FusionStrategy::WeightedScore);
// Weighted mode should default to 0.3 FTS / 0.7 vector
assert!((config.fts_weight - 0.3).abs() < 0.001);
assert!((config.vector_weight - 0.7).abs() < 0.001);
clear_search_env();
}
#[test]
fn weighted_both_zero_rejected() {
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
clear_search_env();
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::set_var("SEARCH_FUSION_STRATEGY", "weighted");
std::env::set_var("SEARCH_FTS_WEIGHT", "0.0");
std::env::set_var("SEARCH_VECTOR_WEIGHT", "0.0");
}
let result = WorkspaceSearchConfig::resolve();
assert!(result.is_err());
clear_search_env();
}
#[test]
fn rrf_both_zero_allowed() {
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
clear_search_env();
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::set_var("SEARCH_FTS_WEIGHT", "0.0");
std::env::set_var("SEARCH_VECTOR_WEIGHT", "0.0");
}
// RRF ignores weights, so both=0 is fine
let config = WorkspaceSearchConfig::resolve().expect("should resolve");
assert_eq!(config.fusion_strategy, FusionStrategy::Rrf);
clear_search_env();
}
}
+2 -20
View File
@@ -169,7 +169,7 @@ pub(crate) fn parse_timestamp(s: &str) -> Result<DateTime<Utc>, String> {
} }
// Naive with fractional seconds (legacy or SQLite datetime() output) // Naive with fractional seconds (legacy or SQLite datetime() output)
if let Ok(ndt) = NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S%.f") { if let Ok(ndt) = NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S%.f") {
tracing::debug!( tracing::warn!(
timestamp = %s, timestamp = %s,
"parsed naive timestamp without timezone; assuming UTC for backward compatibility" "parsed naive timestamp without timezone; assuming UTC for backward compatibility"
); );
@@ -177,7 +177,7 @@ pub(crate) fn parse_timestamp(s: &str) -> Result<DateTime<Utc>, String> {
} }
// Naive without fractional seconds (legacy format) // Naive without fractional seconds (legacy format)
if let Ok(ndt) = NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S") { if let Ok(ndt) = NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S") {
tracing::debug!( tracing::warn!(
timestamp = %s, timestamp = %s,
"parsed naive timestamp without timezone; assuming UTC for backward compatibility" "parsed naive timestamp without timezone; assuming UTC for backward compatibility"
); );
@@ -326,24 +326,6 @@ impl Database for LibSqlBackend {
libsql_migrations::run_incremental(&conn).await?; libsql_migrations::run_incremental(&conn).await?;
Ok(()) Ok(())
} }
async fn shutdown(&self) -> Result<(), DatabaseError> {
match self.db.flush_replicator().await {
Ok(Some(frame_no)) => {
tracing::debug!("libSQL replicator flushed at frame {}", frame_no);
Ok(())
}
Ok(None) => {
tracing::debug!("No libSQL replicator to flush, skipping shutdown sync");
Ok(())
}
Err(libsql::Error::SyncNotSupported(_)) => {
tracing::debug!("libSQL sync not supported, skipping flush on shutdown");
Ok(())
}
Err(error) => Err(DatabaseError::from(error)),
}
}
} }
// ==================== Row conversion helpers ==================== // ==================== Row conversion helpers ====================
+2 -2
View File
@@ -14,7 +14,7 @@ use crate::db::WorkspaceStore;
use crate::error::WorkspaceError; use crate::error::WorkspaceError;
use crate::workspace::{ use crate::workspace::{
MemoryChunk, MemoryDocument, RankedResult, SearchConfig, SearchResult, WorkspaceEntry, MemoryChunk, MemoryDocument, RankedResult, SearchConfig, SearchResult, WorkspaceEntry,
fuse_results, reciprocal_rank_fusion,
}; };
use chrono::Utc; use chrono::Utc;
@@ -614,6 +614,6 @@ impl WorkspaceStore for LibSqlBackend {
); );
} }
Ok(fuse_results(fts_results, vector_results, config)) Ok(reciprocal_rank_fusion(fts_results, vector_results, config))
} }
} }
-7
View File
@@ -523,13 +523,6 @@ pub trait Database:
{ {
/// Run schema migrations for this backend. /// Run schema migrations for this backend.
async fn run_migrations(&self) -> Result<(), DatabaseError>; async fn run_migrations(&self) -> Result<(), DatabaseError>;
/// Shutdown hook for backend-specific drain/flush behavior.
///
/// Default implementation is a no-op so existing backends remain compatible.
async fn shutdown(&self) -> Result<(), DatabaseError> {
Ok(())
}
} }
#[cfg(test)] #[cfg(test)]
-5
View File
@@ -61,11 +61,6 @@ impl Database for PgBackend {
async fn run_migrations(&self) -> Result<(), DatabaseError> { async fn run_migrations(&self) -> Result<(), DatabaseError> {
self.store.run_migrations().await self.store.run_migrations().await
} }
async fn shutdown(&self) -> Result<(), DatabaseError> {
self.store.pool().close();
Ok(())
}
} }
// ==================== ConversationStore ==================== // ==================== ConversationStore ====================
+59 -447
View File
@@ -27,7 +27,7 @@ use crate::secrets::{CreateSecretParams, SecretsStore};
use crate::tools::ToolRegistry; use crate::tools::ToolRegistry;
use crate::tools::mcp::McpClient; use crate::tools::mcp::McpClient;
use crate::tools::mcp::auth::{ use crate::tools::mcp::auth::{
authorize_mcp_server, canonical_resource_uri, discover_full_oauth_metadata, PkceChallenge, authorize_mcp_server, build_authorization_url, discover_full_oauth_metadata,
find_available_port, is_authenticated, register_client, find_available_port, is_authenticated, register_client,
}; };
use crate::tools::mcp::config::McpServerConfig; use crate::tools::mcp::config::McpServerConfig;
@@ -108,13 +108,6 @@ pub struct ExtensionManager {
/// Relay config captured at startup. Used by `auth_channel_relay` and /// Relay config captured at startup. Used by `auth_channel_relay` and
/// `activate_channel_relay` instead of re-reading env vars. /// `activate_channel_relay` instead of re-reading env vars.
relay_config: Option<crate::config::RelayConfig>, relay_config: Option<crate::config::RelayConfig>,
/// When `true`, OAuth flows always return an auth URL to the caller
/// instead of opening a browser on the server via `open::that()`.
/// Set by the web gateway at startup via `enable_gateway_mode()`.
gateway_mode: std::sync::atomic::AtomicBool,
/// The gateway's own base URL for building OAuth redirect URIs.
/// Set by the web gateway at startup via `enable_gateway_mode()`.
gateway_base_url: RwLock<Option<String>>,
} }
/// Sanitize a URL for logging by removing query parameters and credentials. /// Sanitize a URL for logging by removing query parameters and credentials.
@@ -188,75 +181,9 @@ impl ExtensionManager {
pending_oauth_flows: crate::cli::oauth_defaults::new_pending_oauth_registry(), pending_oauth_flows: crate::cli::oauth_defaults::new_pending_oauth_registry(),
gateway_token: std::env::var("GATEWAY_AUTH_TOKEN").ok(), gateway_token: std::env::var("GATEWAY_AUTH_TOKEN").ok(),
relay_config: crate::config::RelayConfig::from_env(), relay_config: crate::config::RelayConfig::from_env(),
gateway_mode: std::sync::atomic::AtomicBool::new(false),
gateway_base_url: RwLock::new(None),
} }
} }
/// Enable gateway mode so OAuth flows return auth URLs to the frontend
/// instead of calling `open::that()` on the server.
///
/// `base_url` is the gateway's own public URL (e.g. `https://my-gateway.example.com`),
/// used to build OAuth redirect URIs when `IRONCLAW_OAUTH_CALLBACK_URL` is not set.
pub async fn enable_gateway_mode(&self, base_url: String) {
self.gateway_mode
.store(true, std::sync::atomic::Ordering::Release);
*self.gateway_base_url.write().await = Some(base_url);
}
/// Returns `true` if OAuth should use gateway mode (return auth URL to
/// frontend) rather than CLI mode (open browser on server via `open::that`).
///
/// Gateway mode is active when any of:
/// - `enable_gateway_mode()` was called (web gateway is running), OR
/// - `IRONCLAW_OAUTH_CALLBACK_URL` is set to a non-loopback URL, OR
/// - `self.tunnel_url` is set to a non-loopback URL
pub fn should_use_gateway_mode(&self) -> bool {
if self.gateway_mode.load(std::sync::atomic::Ordering::Acquire) {
return true;
}
if crate::cli::oauth_defaults::use_gateway_callback() {
return true;
}
self.tunnel_url
.as_ref()
.filter(|u| !u.is_empty())
.and_then(|raw| url::Url::parse(raw).ok())
.and_then(|u| u.host_str().map(String::from))
.map(|host| !crate::cli::oauth_defaults::is_loopback_host(&host))
.unwrap_or(false)
}
/// Returns the OAuth redirect URI for gateway mode, or `None` for local mode.
///
/// Priority:
/// 1. `IRONCLAW_OAUTH_CALLBACK_URL` env var (via `callback_url()`)
/// 2. `gateway_base_url` (set by `enable_gateway_mode()`)
/// 3. `tunnel_url` (from config)
/// 4. `None` (local/CLI mode)
async fn gateway_callback_redirect_uri(&self) -> Option<String> {
use crate::cli::oauth_defaults;
if oauth_defaults::use_gateway_callback() {
return Some(format!("{}/oauth/callback", oauth_defaults::callback_url()));
}
// Use gateway_base_url from enable_gateway_mode()
if let Some(ref base) = *self.gateway_base_url.read().await {
let base = base.trim_end_matches('/');
return Some(format!("{}/oauth/callback", base));
}
// Fall back to tunnel_url
self.tunnel_url
.as_ref()
.filter(|u| !u.is_empty())
.and_then(|raw| url::Url::parse(raw).ok())
.and_then(|u| u.host_str().map(String::from))
.filter(|host| !oauth_defaults::is_loopback_host(host))
.map(|_| {
let base = self.tunnel_url.as_ref().unwrap().trim_end_matches('/');
format!("{}/oauth/callback", base)
})
}
/// Get the relay config stored at startup. /// Get the relay config stored at startup.
fn relay_config(&self) -> Result<&crate::config::RelayConfig, ExtensionError> { fn relay_config(&self) -> Result<&crate::config::RelayConfig, ExtensionError> {
self.relay_config.as_ref().ok_or_else(|| { self.relay_config.as_ref().ok_or_else(|| {
@@ -266,12 +193,6 @@ impl ExtensionManager {
}) })
} }
/// Inject a registry entry for testing. The entry is added to the discovery
/// cache so it appears in search results alongside built-in entries.
pub async fn inject_registry_entry(&self, entry: crate::extensions::RegistryEntry) {
self.registry.cache_discovered(vec![entry]).await;
}
/// Configure the channel runtime infrastructure for hot-activating WASM channels. /// Configure the channel runtime infrastructure for hot-activating WASM channels.
/// ///
/// Call after construction (and after wrapping in `Arc`) once the channel /// Call after construction (and after wrapping in `Arc`) once the channel
@@ -1763,46 +1684,29 @@ impl ExtensionManager {
return Ok(AuthResult::authenticated(name, ExtensionKind::McpServer)); return Ok(AuthResult::authenticated(name, ExtensionKind::McpServer));
} }
// In gateway mode, build an auth URL and return it for the frontend to // Run the full OAuth flow (opens browser, waits for callback)
// open in the same browser. The gateway's /oauth/callback handler will
// complete the token exchange.
if self.should_use_gateway_mode() {
return match self.auth_mcp_build_url(name, &server).await {
Ok(result) => Ok(result),
Err(ExtensionError::AuthNotSupported(_)) => Ok(AuthResult::awaiting_token(
name,
ExtensionKind::McpServer,
format!(
"Server '{}' does not support OAuth. \
Please provide an API token/key for this server.",
name
),
None,
)),
Err(e) => Err(e),
};
}
// CLI/local mode: run the full blocking OAuth flow (opens browser, waits for callback)
match authorize_mcp_server(&server, &self.secrets, &self.user_id).await { match authorize_mcp_server(&server, &self.secrets, &self.user_id).await {
Ok(_token) => { Ok(_token) => {
tracing::info!("MCP server '{}' authenticated via OAuth", name); tracing::info!("MCP server '{}' authenticated via OAuth", name);
Ok(AuthResult::authenticated(name, ExtensionKind::McpServer)) Ok(AuthResult::authenticated(name, ExtensionKind::McpServer))
} }
Err(crate::tools::mcp::auth::AuthError::NotSupported) => { Err(crate::tools::mcp::auth::AuthError::NotSupported) => {
// Server doesn't support OAuth, try building a URL // Server doesn't support OAuth, try building a URL first
match self.auth_mcp_build_url(name, &server).await { match self.auth_mcp_build_url(name, &server).await {
Ok(result) => Ok(result), Ok(result) => Ok(result),
Err(_) => Ok(AuthResult::awaiting_token( Err(_) => {
name, // No OAuth, no DCR: fall back to manual token entry
ExtensionKind::McpServer, Ok(AuthResult::awaiting_token(
format!( name,
"Server '{}' does not support OAuth. \ ExtensionKind::McpServer,
Please provide an API token/key for this server.", format!(
name "Server '{}' does not support OAuth. \
), Please provide an API token/key for this server.",
None, name
)), ),
None,
))
}
} }
} }
Err(e) => { Err(e) => {
@@ -1821,12 +1725,8 @@ impl ExtensionManager {
} }
} }
/// Build an auth URL for MCP OAuth. /// Build an auth URL for cases where non-interactive auth is needed
/// /// (e.g., running via Telegram where we can't open a browser).
/// In gateway mode, stores a `PendingOAuthFlow` so the web gateway's
/// `/oauth/callback` handler can complete the token exchange — the auth
/// URL is sent to the frontend which opens it in the same browser.
/// In local/CLI mode, builds the URL for the user to open manually.
async fn auth_mcp_build_url( async fn auth_mcp_build_url(
&self, &self,
name: &str, name: &str,
@@ -1835,153 +1735,60 @@ impl ExtensionManager {
// Try to discover OAuth metadata and build a URL the user can open manually // Try to discover OAuth metadata and build a URL the user can open manually
let metadata = discover_full_oauth_metadata(&server.url) let metadata = discover_full_oauth_metadata(&server.url)
.await .await
.map_err(|e| match e { .map_err(|e| ExtensionError::AuthFailed(e.to_string()))?;
crate::tools::mcp::auth::AuthError::NotSupported => {
ExtensionError::AuthNotSupported(e.to_string())
}
_ => ExtensionError::AuthFailed(e.to_string()),
})?;
use crate::cli::oauth_defaults; // Try DCR if no client_id configured
let (client_id, redirect_uri) = if let Some(ref oauth) = server.oauth {
let is_gateway = self.should_use_gateway_mode();
// Build redirect URI: gateway uses the public callback URL,
// local mode binds a random port.
let redirect_uri = if let Some(uri) = self.gateway_callback_redirect_uri().await {
uri
} else {
let port = find_available_port() let port = find_available_port()
.await .await
.map_err(|e| ExtensionError::AuthFailed(e.to_string()))?; .map_err(|e| ExtensionError::AuthFailed(e.to_string()))?;
format!("http://localhost:{}/callback", port.1) let redirect = format!("http://localhost:{}/callback", port.1);
}; (oauth.client_id.clone(), redirect)
// Try DCR if no client_id configured
let (client_id, client_secret) = if let Some(ref oauth) = server.oauth {
(oauth.client_id.clone(), None)
} else if let Some(ref reg_endpoint) = metadata.registration_endpoint { } else if let Some(ref reg_endpoint) = metadata.registration_endpoint {
let registration = register_client(reg_endpoint, &redirect_uri) let port = find_available_port()
.await
.map_err(|e| ExtensionError::AuthFailed(e.to_string()))?;
let redirect = format!("http://localhost:{}/callback", port.1);
let registration = register_client(reg_endpoint, &redirect)
.await .await
.map_err(|e| ExtensionError::AuthFailed(e.to_string()))?; .map_err(|e| ExtensionError::AuthFailed(e.to_string()))?;
(registration.client_id, None) (registration.client_id, redirect)
} else { } else {
return Err(ExtensionError::AuthNotSupported( return Err(ExtensionError::AuthFailed(
"Server doesn't support OAuth or Dynamic Client Registration".to_string(), "Server doesn't support OAuth or Dynamic Client Registration".to_string(),
)); ));
}; };
// RFC 8707: resource parameter to scope the token to this MCP server let pkce = PkceChallenge::generate();
let resource = canonical_resource_uri(&server.url); let auth_url = build_authorization_url(
// Build authorization URL with CSRF state using the shared oauth_defaults
// builder, which generates PKCE + state for us.
let mut extra_params = server
.oauth
.as_ref()
.map(|o| o.extra_params.clone())
.unwrap_or_default();
extra_params.insert("resource".to_string(), resource.clone());
let scopes = server
.oauth
.as_ref()
.map(|o| o.scopes.clone())
.unwrap_or_else(|| metadata.scopes_supported.clone());
let oauth_result = oauth_defaults::build_oauth_url(
&metadata.authorization_endpoint, &metadata.authorization_endpoint,
&client_id, &client_id,
&redirect_uri, &redirect_uri,
&scopes, &metadata.scopes_supported,
true, // Always use PKCE for MCP Some(&pkce),
&extra_params, &std::collections::HashMap::new(),
None,
); );
let expected_state = oauth_result.state;
let code_verifier = oauth_result.code_verifier;
if is_gateway { // Store pending auth for later callback handling
// Gateway mode: store pending flow for the /oauth/callback handler. self.pending_auth.write().await.insert(
oauth_defaults::sweep_expired_flows(&self.pending_oauth_flows).await; name.to_string(),
PendingAuth {
// Platform routing: prepend instance name to state _name: name.to_string(),
let platform_state = oauth_defaults::build_platform_state(&expected_state); _kind: ExtensionKind::McpServer,
let auth_url = if platform_state != expected_state {
oauth_result.url.replace(
&format!("state={}", urlencoding::encode(&expected_state)),
&format!("state={}", urlencoding::encode(&platform_state)),
)
} else {
oauth_result.url
};
let flow = oauth_defaults::PendingOAuthFlow {
extension_name: name.to_string(),
display_name: server.name.clone(),
token_url: metadata.token_endpoint,
client_id,
client_secret,
redirect_uri,
code_verifier,
access_token_field: "access_token".to_string(),
secret_name: server.token_secret_name(),
provider: Some(format!("mcp:{}", name)),
validation_endpoint: None,
scopes,
user_id: self.user_id.clone(),
secrets: Arc::clone(&self.secrets),
sse_sender: self.sse_sender.read().await.clone(),
gateway_token: self.gateway_token.clone(),
resource: Some(resource),
client_id_secret_name: if server.oauth.is_none() {
Some(server.client_id_secret_name())
} else {
None
},
created_at: std::time::Instant::now(), created_at: std::time::Instant::now(),
}; task_handle: None,
},
);
self.pending_oauth_flows Ok(AuthResult::awaiting_authorization(
.write() name,
.await ExtensionKind::McpServer,
.insert(expected_state, flow); auth_url,
"local".to_string(),
self.pending_auth.write().await.insert( ))
name.to_string(),
PendingAuth {
_name: name.to_string(),
_kind: ExtensionKind::McpServer,
created_at: std::time::Instant::now(),
task_handle: None,
},
);
Ok(AuthResult::awaiting_authorization(
name,
ExtensionKind::McpServer,
auth_url,
"gateway".to_string(),
))
} else {
// Local mode: return URL for manual opening
self.pending_auth.write().await.insert(
name.to_string(),
PendingAuth {
_name: name.to_string(),
_kind: ExtensionKind::McpServer,
created_at: std::time::Instant::now(),
task_handle: None,
},
);
Ok(AuthResult::awaiting_authorization(
name,
ExtensionKind::McpServer,
oauth_result.url,
"local".to_string(),
))
}
} }
async fn auth_wasm_tool(&self, name: &str) -> Result<AuthResult, ExtensionError> { async fn auth_wasm_tool(&self, name: &str) -> Result<AuthResult, ExtensionError> {
@@ -2396,10 +2203,7 @@ impl ExtensionManager {
flows.retain(|_, flow| flow.extension_name != name); flows.retain(|_, flow| flow.extension_name != name);
} }
let redirect_uri = self let redirect_uri = format!("{}/callback", oauth_defaults::callback_url());
.gateway_callback_redirect_uri()
.await
.unwrap_or_else(|| format!("{}/callback", oauth_defaults::callback_url()));
// Merge scopes from all tools sharing this provider // Merge scopes from all tools sharing this provider
let merged_scopes = self let merged_scopes = self
@@ -2424,7 +2228,7 @@ impl ExtensionManager {
.clone() .clone()
.unwrap_or_else(|| name.to_string()); .unwrap_or_else(|| name.to_string());
if self.should_use_gateway_mode() { if oauth_defaults::use_gateway_callback() {
// Gateway mode: store pending flow state for the web gateway's // Gateway mode: store pending flow state for the web gateway's
// `/oauth/callback` handler to complete the exchange. No TCP listener // `/oauth/callback` handler to complete the exchange. No TCP listener
// needed — the OAuth provider redirects to the gateway URL. // needed — the OAuth provider redirects to the gateway URL.
@@ -2460,8 +2264,6 @@ impl ExtensionManager {
secrets: Arc::clone(&self.secrets), secrets: Arc::clone(&self.secrets),
sse_sender: self.sse_sender.read().await.clone(), sse_sender: self.sse_sender.read().await.clone(),
gateway_token: self.gateway_token.clone(), gateway_token: self.gateway_token.clone(),
resource: None,
client_id_secret_name: None,
created_at: std::time::Instant::now(), created_at: std::time::Instant::now(),
}; };
@@ -2803,17 +2605,11 @@ impl ExtensionManager {
.await .await
.map_err(|e| ExtensionError::ActivationFailed(e.to_string()))?; .map_err(|e| ExtensionError::ActivationFailed(e.to_string()))?;
// Try to list and create tools. // Try to list and create tools
// A 401/auth error means the server requires OAuth — surface as let mcp_tools = client
// AuthRequired so the activate handler triggers the OAuth flow. .list_tools()
let mcp_tools = client.list_tools().await.map_err(|e| { .await
let msg = e.to_string(); .map_err(|e| ExtensionError::ActivationFailed(e.to_string()))?;
if msg.contains("requires authentication") || msg.contains("401") {
ExtensionError::AuthRequired
} else {
ExtensionError::ActivationFailed(msg)
}
})?;
let tool_impls = client let tool_impls = client
.create_tools() .create_tools()
@@ -4970,190 +4766,6 @@ mod tests {
assert!(result.contains("/v1/users/123/profile")); assert!(result.contains("/v1/users/123/profile"));
} }
// ---- gateway mode detection tests ----
// Regression tests for a bug where MCP OAuth called `open::that()` on the
// server machine instead of returning an auth URL to the gateway frontend.
// The root cause was that `should_use_gateway_mode()` only checked the
// `IRONCLAW_OAUTH_CALLBACK_URL` env var, ignoring `self.tunnel_url`.
/// Serializes env-mutating tests to prevent parallel races.
static GATEWAY_ENV_MUTEX: std::sync::Mutex<()> = std::sync::Mutex::new(());
/// Build a minimal ExtensionManager with a custom tunnel_url.
fn make_manager_with_tunnel(tunnel_url: Option<String>) -> ExtensionManager {
use crate::secrets::{InMemorySecretsStore, SecretsCrypto};
use crate::tools::mcp::process::McpProcessManager;
use crate::tools::mcp::session::McpSessionManager;
let key = secrecy::SecretString::from(crate::secrets::keychain::generate_master_key_hex());
let crypto = Arc::new(SecretsCrypto::new(key).expect("crypto"));
let secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync> =
Arc::new(InMemorySecretsStore::new(crypto));
let tools = Arc::new(crate::tools::ToolRegistry::new());
let mcp = Arc::new(McpSessionManager::new());
let dir = std::env::temp_dir().join("ironclaw-test-gateway-mode");
ExtensionManager::new(
mcp,
Arc::new(McpProcessManager::new()),
secrets,
tools,
None,
None,
dir.clone(),
dir,
tunnel_url,
"test".to_string(),
None,
vec![],
)
}
#[test]
fn should_use_gateway_mode_true_for_tunnel_url() {
let _guard = GATEWAY_ENV_MUTEX.lock().expect("env mutex poisoned");
let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok();
// SAFETY: Under GATEWAY_ENV_MUTEX, no concurrent env access.
unsafe {
std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL");
}
let mgr = make_manager_with_tunnel(Some("https://my-gateway.example.com".into()));
assert!(
mgr.should_use_gateway_mode(),
"should detect gateway mode from tunnel_url"
);
unsafe {
if let Some(val) = original {
std::env::set_var("IRONCLAW_OAUTH_CALLBACK_URL", val);
}
}
}
#[test]
fn should_use_gateway_mode_false_without_tunnel() {
let _guard = GATEWAY_ENV_MUTEX.lock().expect("env mutex poisoned");
let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok();
unsafe {
std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL");
}
let mgr = make_manager_with_tunnel(None);
assert!(
!mgr.should_use_gateway_mode(),
"should not detect gateway mode without tunnel_url or env var"
);
unsafe {
if let Some(val) = original {
std::env::set_var("IRONCLAW_OAUTH_CALLBACK_URL", val);
}
}
}
#[test]
fn should_use_gateway_mode_false_for_loopback_tunnel() {
let _guard = GATEWAY_ENV_MUTEX.lock().expect("env mutex poisoned");
let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok();
unsafe {
std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL");
}
let mgr = make_manager_with_tunnel(Some("http://127.0.0.1:3001".into()));
assert!(
!mgr.should_use_gateway_mode(),
"should not detect gateway mode for loopback tunnel_url"
);
unsafe {
if let Some(val) = original {
std::env::set_var("IRONCLAW_OAUTH_CALLBACK_URL", val);
}
}
}
/// Helper to run an async test body while holding the env mutex.
/// Clears `IRONCLAW_OAUTH_CALLBACK_URL` for the duration, restoring on drop.
struct EnvGuard {
original: Option<String>,
_mutex: std::sync::MutexGuard<'static, ()>,
}
impl EnvGuard {
fn new() -> Self {
let guard = GATEWAY_ENV_MUTEX.lock().expect("env mutex poisoned");
let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok();
// SAFETY: Under GATEWAY_ENV_MUTEX, no concurrent env access.
unsafe {
std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL");
}
Self {
original,
_mutex: guard,
}
}
}
impl Drop for EnvGuard {
fn drop(&mut self) {
// SAFETY: Under GATEWAY_ENV_MUTEX (still held by _mutex), no concurrent env access.
unsafe {
if let Some(ref val) = self.original {
std::env::set_var("IRONCLAW_OAUTH_CALLBACK_URL", val);
} else {
std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL");
}
}
}
}
#[tokio::test]
async fn gateway_callback_redirect_uri_from_tunnel_url() {
let _env = EnvGuard::new();
let mgr = make_manager_with_tunnel(Some("https://my-gateway.example.com".into()));
assert_eq!(
mgr.gateway_callback_redirect_uri().await,
Some("https://my-gateway.example.com/oauth/callback".to_string()),
);
}
#[tokio::test]
async fn gateway_callback_redirect_uri_none_without_tunnel() {
let _env = EnvGuard::new();
let mgr = make_manager_with_tunnel(None);
assert_eq!(mgr.gateway_callback_redirect_uri().await, None);
}
#[tokio::test]
async fn gateway_callback_redirect_uri_trims_trailing_slash() {
let _env = EnvGuard::new();
let mgr = make_manager_with_tunnel(Some("https://my-gateway.example.com/".into()));
assert_eq!(
mgr.gateway_callback_redirect_uri().await,
Some("https://my-gateway.example.com/oauth/callback".to_string()),
);
}
#[tokio::test]
async fn gateway_mode_enabled_explicitly() {
let _env = EnvGuard::new();
let mgr = make_manager_with_tunnel(None);
assert!(!mgr.should_use_gateway_mode());
mgr.enable_gateway_mode("https://my-gateway.example.com".into())
.await;
assert!(mgr.should_use_gateway_mode());
assert_eq!(
mgr.gateway_callback_redirect_uri().await,
Some("https://my-gateway.example.com/oauth/callback".to_string()),
);
}
// ── Regression tests for PR #677 (unify-extension-lifecycle) ───────── // ── Regression tests for PR #677 (unify-extension-lifecycle) ─────────
#[tokio::test] #[tokio::test]
-3
View File
@@ -517,9 +517,6 @@ pub enum ExtensionError {
#[error("Authentication failed: {0}")] #[error("Authentication failed: {0}")]
AuthFailed(String), AuthFailed(String),
#[error("Server does not support OAuth: {0}")]
AuthNotSupported(String),
#[error("Activation failed: {0}")] #[error("Activation failed: {0}")]
ActivationFailed(String), ActivationFailed(String),
+5 -21
View File
@@ -9,8 +9,8 @@ use ironclaw::{
agent::{Agent, AgentDeps}, agent::{Agent, AgentDeps},
app::{AppBuilder, AppBuilderFlags}, app::{AppBuilder, AppBuilderFlags},
channels::{ channels::{
ChannelManager, GatewayChannel, HttpChannel, ReplChannel, SignalChannel, WebhookServer, ChannelManager, ChannelSecretUpdater, GatewayChannel, HttpChannel, ReplChannel,
WebhookServerConfig, SignalChannel, WebhookServer, WebhookServerConfig,
wasm::{WasmChannelRouter, WasmChannelRuntime}, wasm::{WasmChannelRouter, WasmChannelRuntime},
web::log_layer::LogBroadcaster, web::log_layer::LogBroadcaster,
}, },
@@ -433,8 +433,9 @@ async fn async_main() -> anyhow::Result<()> {
"Lifecycle hooks initialized" "Lifecycle hooks initialized"
); );
// Reuse the shared agent session manager prepared by AppBuilder. // Create session manager (shared between agent and web gateway)
let session_manager = Arc::clone(&components.agent_session_manager); let session_manager =
Arc::new(ironclaw::agent::SessionManager::new().with_hooks(components.hooks.clone()));
// Lazy scheduler slot — filled after Agent::new creates the Scheduler. // Lazy scheduler slot — filled after Agent::new creates the Scheduler.
// Allows CreateJobTool to dispatch local jobs via the Scheduler even though // Allows CreateJobTool to dispatch local jobs via the Scheduler even though
@@ -475,14 +476,6 @@ async fn async_main() -> anyhow::Result<()> {
gw = gw.with_log_level_handle(Arc::clone(&log_level_handle)); gw = gw.with_log_level_handle(Arc::clone(&log_level_handle));
gw = gw.with_tool_registry(Arc::clone(&components.tools)); gw = gw.with_tool_registry(Arc::clone(&components.tools));
if let Some(ref ext_mgr) = components.extension_manager { if let Some(ref ext_mgr) = components.extension_manager {
// Enable gateway mode so MCP OAuth returns auth URLs to the frontend
// instead of calling open::that() on the server.
let gw_base = config
.tunnel
.public_url
.clone()
.unwrap_or_else(|| format!("http://{}:{}", gw_config.host, gw_config.port));
ext_mgr.enable_gateway_mode(gw_base).await;
gw = gw.with_extension_manager(Arc::clone(ext_mgr)); gw = gw.with_extension_manager(Arc::clone(ext_mgr));
} }
if !components.catalog_entries.is_empty() { if !components.catalog_entries.is_empty() {
@@ -672,8 +665,6 @@ async fn async_main() -> anyhow::Result<()> {
.as_ref() .as_ref()
.map(|db| Arc::clone(db) as Arc<dyn ironclaw::db::SettingsStore>); .map(|db| Arc::clone(db) as Arc<dyn ironclaw::db::SettingsStore>);
let db_for_shutdown = components.db.clone();
let deps = AgentDeps { let deps = AgentDeps {
store: components.db, store: components.db,
llm: components.llm, llm: components.llm,
@@ -738,7 +729,6 @@ async fn async_main() -> anyhow::Result<()> {
#[cfg(unix)] #[cfg(unix)]
{ {
use ironclaw::channels::ChannelSecretUpdater;
// Collect all channels that support secret updates // Collect all channels that support secret updates
let mut secret_updaters: Vec<Arc<dyn ChannelSecretUpdater>> = Vec::new(); let mut secret_updaters: Vec<Arc<dyn ChannelSecretUpdater>> = Vec::new();
if let Some(ref state) = http_channel_state { if let Some(ref state) = http_channel_state {
@@ -932,12 +922,6 @@ async fn async_main() -> anyhow::Result<()> {
} }
} }
if let Some(db) = db_for_shutdown {
if let Err(e) = db.shutdown().await {
tracing::warn!("Failed to shutdown database cleanly: {}", e);
}
}
tracing::debug!("Agent shutdown complete"); tracing::debug!("Agent shutdown complete");
Ok(()) Ok(())
-4
View File
@@ -1067,8 +1067,6 @@ mod tests {
prompt: "Check status".to_string(), prompt: "Check status".to_string(),
context_paths: vec![], context_paths: vec![],
max_tokens: 500, max_tokens: 500,
use_tools: false,
max_tool_rounds: 3,
}, },
guardrails: RoutineGuardrails { guardrails: RoutineGuardrails {
cooldown: std::time::Duration::from_secs(60), cooldown: std::time::Duration::from_secs(60),
@@ -1200,8 +1198,6 @@ mod tests {
prompt: "test".to_string(), prompt: "test".to_string(),
context_paths: vec![], context_paths: vec![],
max_tokens: 100, max_tokens: 100,
use_tools: false,
max_tool_rounds: 3,
}, },
guardrails: RoutineGuardrails { guardrails: RoutineGuardrails {
cooldown: std::time::Duration::from_secs(0), cooldown: std::time::Duration::from_secs(0),
+1 -23
View File
@@ -256,13 +256,7 @@ impl Tool for ToolAuthTool {
} }
fn requires_approval(&self, _params: &serde_json::Value) -> ApprovalRequirement { fn requires_approval(&self, _params: &serde_json::Value) -> ApprovalRequirement {
// In gateway mode, tool_auth only returns an auth URL for the frontend ApprovalRequirement::UnlessAutoApproved
// to open — no browser is launched server-side, so no approval needed.
if self.manager.should_use_gateway_mode() {
ApprovalRequirement::Never
} else {
ApprovalRequirement::UnlessAutoApproved
}
} }
} }
@@ -739,22 +733,6 @@ mod tests {
} }
} }
#[tokio::test]
async fn tool_auth_no_approval_in_gateway_mode() {
let manager = test_manager_stub();
manager
.enable_gateway_mode("http://localhost:3000".to_string())
.await;
let tool = ToolAuthTool {
manager: manager.clone(),
};
assert_eq!(
tool.requires_approval(&serde_json::json!({})),
ApprovalRequirement::Never,
"tool_auth should not require approval in gateway mode"
);
}
#[test] #[test]
fn test_tool_upgrade_schema() { fn test_tool_upgrade_schema() {
use crate::tools::tool::ApprovalRequirement; use crate::tools::tool::ApprovalRequirement;
+37 -212
View File
@@ -31,12 +31,6 @@ const MAX_RESPONSE_SIZE: usize = 5 * 1024 * 1024;
/// in memory for LLM context. Matches the WASM attachment size cap. /// in memory for LLM context. Matches the WASM attachment size cap.
const MAX_SAVE_TO_SIZE: usize = 50 * 1024 * 1024; const MAX_SAVE_TO_SIZE: usize = 50 * 1024 * 1024;
/// Default request timeout when the caller does not provide one.
const DEFAULT_TIMEOUT_SECS: u64 = 30;
/// Maximum allowed request timeout to bound resource usage from LLM-controlled inputs.
const MAX_TIMEOUT_SECS: u64 = 300;
/// Maximum number of redirects to follow for simple GET requests. /// Maximum number of redirects to follow for simple GET requests.
const MAX_REDIRECTS: usize = 3; const MAX_REDIRECTS: usize = 3;
@@ -250,120 +244,43 @@ fn is_html_response(headers: &HashMap<String, String>) -> bool {
fn parse_headers_param( fn parse_headers_param(
headers: Option<&serde_json::Value>, headers: Option<&serde_json::Value>,
) -> Result<Vec<(String, String)>, ToolError> { ) -> Result<Vec<(String, String)>, ToolError> {
fn parse_header_object(
map: &serde_json::Map<String, serde_json::Value>,
) -> Result<Vec<(String, String)>, ToolError> {
let mut out = Vec::with_capacity(map.len());
for (k, v) in map {
let value = v.as_str().ok_or_else(|| {
ToolError::InvalidParameters(format!("header '{}' must have a string value", k))
})?;
out.push((k.clone(), value.to_string()));
}
Ok(out)
}
fn parse_header_array(items: &[serde_json::Value]) -> Result<Vec<(String, String)>, ToolError> {
let mut out = Vec::with_capacity(items.len());
for (idx, item) in items.iter().enumerate() {
let obj = item.as_object().ok_or_else(|| {
ToolError::InvalidParameters(format!(
"headers[{}] must be an object with 'name' and 'value'",
idx
))
})?;
let name = obj.get("name").and_then(|v| v.as_str()).ok_or_else(|| {
ToolError::InvalidParameters(format!("headers[{}].name must be a string", idx))
})?;
let value = obj.get("value").and_then(|v| v.as_str()).ok_or_else(|| {
ToolError::InvalidParameters(format!("headers[{}].value must be a string", idx))
})?;
out.push((name.to_string(), value.to_string()));
}
Ok(out)
}
match headers { match headers {
None => Ok(Vec::new()), None => Ok(Vec::new()),
Some(serde_json::Value::String(raw)) => { Some(serde_json::Value::Object(map)) => {
let trimmed = raw.trim(); let mut out = Vec::with_capacity(map.len());
if trimmed.is_empty() { for (k, v) in map {
return Ok(Vec::new()); let value = v.as_str().ok_or_else(|| {
} ToolError::InvalidParameters(format!("header '{}' must have a string value", k))
let parsed = serde_json::from_str::<serde_json::Value>(trimmed).map_err(|e| { })?;
ToolError::InvalidParameters(format!( out.push((k.clone(), value.to_string()));
"headers string must contain valid JSON object/array: {}",
e
))
})?;
match parsed {
serde_json::Value::Object(map) => parse_header_object(&map),
serde_json::Value::Array(items) => parse_header_array(&items),
_ => Err(ToolError::InvalidParameters(
"headers string must decode to a JSON object or array".to_string(),
)),
} }
Ok(out)
}
Some(serde_json::Value::Array(items)) => {
let mut out = Vec::with_capacity(items.len());
for (idx, item) in items.iter().enumerate() {
let obj = item.as_object().ok_or_else(|| {
ToolError::InvalidParameters(format!(
"headers[{}] must be an object with 'name' and 'value'",
idx
))
})?;
let name = obj.get("name").and_then(|v| v.as_str()).ok_or_else(|| {
ToolError::InvalidParameters(format!("headers[{}].name must be a string", idx))
})?;
let value = obj.get("value").and_then(|v| v.as_str()).ok_or_else(|| {
ToolError::InvalidParameters(format!("headers[{}].value must be a string", idx))
})?;
out.push((name.to_string(), value.to_string()));
}
Ok(out)
} }
Some(serde_json::Value::Object(map)) => parse_header_object(map),
Some(serde_json::Value::Array(items)) => parse_header_array(items),
Some(_) => Err(ToolError::InvalidParameters( Some(_) => Err(ToolError::InvalidParameters(
"'headers' must be an object or an array of {name, value}".to_string(), "'headers' must be an object or an array of {name, value}".to_string(),
)), )),
} }
} }
fn parse_timeout_secs_param(timeout: Option<&serde_json::Value>) -> Result<Option<u64>, ToolError> {
let parsed = match timeout {
None | Some(serde_json::Value::Null) => Ok(None),
Some(serde_json::Value::Number(n)) => n.as_u64().map(Some).ok_or_else(|| {
ToolError::InvalidParameters("timeout_secs must be a non-negative integer".to_string())
}),
Some(serde_json::Value::String(raw)) => {
let trimmed = raw.trim();
if trimmed.is_empty() {
return Ok(None);
}
let secs = trimmed.parse::<u64>().map_err(|_| {
ToolError::InvalidParameters(
"timeout_secs string must contain a non-negative integer".to_string(),
)
})?;
Ok(Some(secs))
}
Some(_) => Err(ToolError::InvalidParameters(
"timeout_secs must be an integer".to_string(),
)),
}?;
if let Some(secs) = parsed
&& secs > MAX_TIMEOUT_SECS
{
return Err(ToolError::InvalidParameters(format!(
"timeout_secs must be <= {}",
MAX_TIMEOUT_SECS
)));
}
Ok(parsed)
}
fn parse_save_to_param(save_to: Option<&serde_json::Value>) -> Result<Option<String>, ToolError> {
match save_to {
None | Some(serde_json::Value::Null) => Ok(None),
Some(serde_json::Value::String(path)) => {
let trimmed = path.trim();
if trimmed.is_empty() {
Ok(None)
} else {
Ok(Some(trimmed.to_string()))
}
}
Some(_) => Err(ToolError::InvalidParameters(
"save_to must be a string".to_string(),
)),
}
}
/// Extract host from URL in params (for approval checks). /// Extract host from URL in params (for approval checks).
fn extract_host_from_params(params: &serde_json::Value) -> Option<String> { fn extract_host_from_params(params: &serde_json::Value) -> Option<String> {
params params
@@ -441,7 +358,6 @@ impl Tool for HttpTool {
let start = std::time::Instant::now(); let start = std::time::Instant::now();
let method = require_str(&params, "method")?; let method = require_str(&params, "method")?;
let method_upper = method.to_uppercase();
let url = require_str(&params, "url")?; let url = require_str(&params, "url")?;
let mut parsed_url = validate_url(url)?; let mut parsed_url = validate_url(url)?;
@@ -463,9 +379,6 @@ impl Tool for HttpTool {
// Parse headers // Parse headers
let mut headers_vec = parse_headers_param(params.get("headers"))?; let mut headers_vec = parse_headers_param(params.get("headers"))?;
let timeout_secs = parse_timeout_secs_param(params.get("timeout_secs"))?;
let save_to = parse_save_to_param(params.get("save_to"))?;
let effective_timeout = Duration::from_secs(timeout_secs.unwrap_or(DEFAULT_TIMEOUT_SECS));
// Build request // Build request
let mut request = match method.to_uppercase().as_str() { let mut request = match method.to_uppercase().as_str() {
@@ -482,8 +395,6 @@ impl Tool for HttpTool {
} }
}; };
request = request.timeout(effective_timeout);
// Add headers // Add headers
for (key, value) in &headers_vec { for (key, value) in &headers_vec {
request = request.header(key.as_str(), value.as_str()); request = request.header(key.as_str(), value.as_str());
@@ -492,9 +403,7 @@ impl Tool for HttpTool {
// Add body if present // Add body if present
let body_bytes = if let Some(body) = params.get("body") { let body_bytes = if let Some(body) = params.get("body") {
if let Some(body_str) = body.as_str() { if let Some(body_str) = body.as_str() {
if body_str.is_empty() { if let Ok(json_body) = serde_json::from_str::<serde_json::Value>(body_str) {
None
} else if let Ok(json_body) = serde_json::from_str::<serde_json::Value>(body_str) {
let bytes = serde_json::to_vec(&json_body).map_err(|e| { let bytes = serde_json::to_vec(&json_body).map_err(|e| {
ToolError::InvalidParameters(format!("invalid body JSON: {}", e)) ToolError::InvalidParameters(format!("invalid body JSON: {}", e))
})?; })?;
@@ -559,7 +468,7 @@ impl Tool for HttpTool {
// Build the interceptor request descriptor for recording/replay // Build the interceptor request descriptor for recording/replay
let intercept_req = crate::llm::recording::HttpExchangeRequest { let intercept_req = crate::llm::recording::HttpExchangeRequest {
method: method_upper, method: method.to_uppercase(),
url: parsed_url.to_string(), url: parsed_url.to_string(),
headers: headers_vec.clone(), headers: headers_vec.clone(),
body: body_bytes body: body_bytes
@@ -601,7 +510,7 @@ impl Tool for HttpTool {
let hop_client = build_pinned_client( let hop_client = build_pinned_client(
&hop_host, &hop_host,
&hop_addrs, &hop_addrs,
effective_timeout, Duration::from_secs(30),
reqwest::redirect::Policy::none(), reqwest::redirect::Policy::none(),
)?; )?;
@@ -615,7 +524,7 @@ impl Tool for HttpTool {
.await .await
.map_err(|e| { .map_err(|e| {
if e.is_timeout() { if e.is_timeout() {
ToolError::Timeout(effective_timeout) ToolError::Timeout(Duration::from_secs(30))
} else { } else {
ToolError::ExternalService(e.to_string()) ToolError::ExternalService(e.to_string())
} }
@@ -679,7 +588,7 @@ impl Tool for HttpTool {
} else { } else {
let resp = request.send().await.map_err(|e| { let resp = request.send().await.map_err(|e| {
if e.is_timeout() { if e.is_timeout() {
ToolError::Timeout(effective_timeout) ToolError::Timeout(Duration::from_secs(30))
} else { } else {
ToolError::ExternalService(e.to_string()) ToolError::ExternalService(e.to_string())
} }
@@ -707,7 +616,7 @@ impl Tool for HttpTool {
.collect(); .collect();
// Use a larger size limit when saving to disk (file downloads) // Use a larger size limit when saving to disk (file downloads)
let saving_to_disk = save_to.is_some(); let saving_to_disk = params.get("save_to").is_some();
let max_size = if saving_to_disk { let max_size = if saving_to_disk {
MAX_SAVE_TO_SIZE MAX_SAVE_TO_SIZE
} else { } else {
@@ -752,11 +661,11 @@ impl Tool for HttpTool {
let body_bytes = bytes::Bytes::from(body); let body_bytes = bytes::Bytes::from(body);
// If save_to is specified, write raw bytes to file and return metadata. // If save_to is specified, write raw bytes to file and return metadata.
if let Some(save_to) = save_to { if let Some(save_to) = params.get("save_to").and_then(|v| v.as_str()) {
let saved_to = save_to.clone(); let save_to_owned = save_to.to_string();
let bytes_clone = body_bytes.clone(); let bytes_clone = body_bytes.clone();
tokio::task::spawn_blocking(move || { tokio::task::spawn_blocking(move || {
let canonical = validate_save_to_path(&save_to)?; let canonical = validate_save_to_path(&save_to_owned)?;
std::fs::write(&canonical, &bytes_clone).map_err(|e| { std::fs::write(&canonical, &bytes_clone).map_err(|e| {
ToolError::ExecutionFailed(format!("failed to write file: {}", e)) ToolError::ExecutionFailed(format!("failed to write file: {}", e))
})?; })?;
@@ -767,7 +676,7 @@ impl Tool for HttpTool {
.map_err(|e: ToolError| e)?; .map_err(|e: ToolError| e)?;
let result = serde_json::json!({ let result = serde_json::json!({
"status": status, "status": status,
"saved_to": saved_to, "saved_to": save_to,
"size_bytes": body_bytes.len(), "size_bytes": body_bytes.len(),
"headers": headers, "headers": headers,
}); });
@@ -978,71 +887,6 @@ mod tests {
); );
} }
#[test]
fn test_parse_headers_param_accepts_stringified_array() {
let headers =
serde_json::json!("[{\"name\":\"Authorization\",\"value\":\"Bearer token\"}]");
let parsed = parse_headers_param(Some(&headers)).unwrap();
assert_eq!(
parsed,
vec![("Authorization".to_string(), "Bearer token".to_string())]
);
}
#[test]
fn test_parse_headers_param_rejects_double_string_encoding() {
let headers = serde_json::json!("\"hello\"");
let err = parse_headers_param(Some(&headers)).unwrap_err();
assert!(
err.to_string()
.contains("headers string must decode to a JSON object or array"),
"unexpected error: {}",
err
);
}
#[test]
fn test_parse_timeout_secs_param_accepts_string_integer() {
let timeout = serde_json::json!("30");
assert_eq!(parse_timeout_secs_param(Some(&timeout)).unwrap(), Some(30));
}
#[test]
fn test_parse_timeout_secs_param_treats_empty_string_as_none() {
let timeout = serde_json::json!("");
assert_eq!(parse_timeout_secs_param(Some(&timeout)).unwrap(), None);
}
#[test]
fn test_parse_timeout_secs_param_rejects_value_above_cap() {
let timeout = serde_json::json!(MAX_TIMEOUT_SECS + 1);
let err = parse_timeout_secs_param(Some(&timeout)).unwrap_err();
assert!(
err.to_string()
.contains(&format!("timeout_secs must be <= {}", MAX_TIMEOUT_SECS)),
"unexpected error: {}",
err
);
}
#[test]
fn test_parse_timeout_secs_param_rejects_string_value_above_cap() {
let timeout = serde_json::json!((MAX_TIMEOUT_SECS + 1).to_string());
let err = parse_timeout_secs_param(Some(&timeout)).unwrap_err();
assert!(
err.to_string()
.contains(&format!("timeout_secs must be <= {}", MAX_TIMEOUT_SECS)),
"unexpected error: {}",
err
);
}
#[test]
fn test_parse_save_to_param_treats_empty_string_as_none() {
let save_to = serde_json::json!("");
assert_eq!(parse_save_to_param(Some(&save_to)).unwrap(), None);
}
#[test] #[test]
fn test_http_tool_schema_body_is_freeform() { fn test_http_tool_schema_body_is_freeform() {
let schema = HttpTool::new().parameters_schema(); let schema = HttpTool::new().parameters_schema();
@@ -1275,25 +1119,6 @@ mod tests {
assert_eq!(extract_host_from_params(&params), None); assert_eq!(extract_host_from_params(&params), None);
} }
#[test]
fn test_requires_approval_with_stringified_http_params() {
use crate::tools::wasm::SharedCredentialRegistry;
let tool = HttpTool::new().with_credentials(
Arc::new(SharedCredentialRegistry::new()),
Arc::new(test_secrets_store()),
);
let req = serde_json::json!({
"body": "",
"headers": "[]",
"method": "GET",
"save_to": "",
"timeout_secs": "30",
"url": "https://r.jina.ai/http://news.baidu.com/"
});
let _ = tool.requires_approval(&req);
}
// ── DNS pinning tests ───────────────────────────────────────────── // ── DNS pinning tests ─────────────────────────────────────────────
#[tokio::test] #[tokio::test]
+2 -63
View File
@@ -12,7 +12,6 @@
//! Use `memory_write` to persist important facts that should be remembered //! Use `memory_write` to persist important facts that should be remembered
//! across sessions. //! across sessions.
use std::path::Path;
use std::sync::Arc; use std::sync::Arc;
use async_trait::async_trait; use async_trait::async_trait;
@@ -27,28 +26,6 @@ use crate::workspace::{Workspace, paths};
const PROTECTED_IDENTITY_FILES: &[&str] = const PROTECTED_IDENTITY_FILES: &[&str] =
&[paths::IDENTITY, paths::SOUL, paths::AGENTS, paths::USER]; &[paths::IDENTITY, paths::SOUL, paths::AGENTS, paths::USER];
/// Detect paths that are clearly local filesystem references, not workspace-memory docs.
///
/// Examples:
/// - `/Users/.../file.md` (Unix absolute)
/// - `C:\Users\...` or `D:/work/...` (Windows absolute)
/// - `~/notes.md` (home expansion shorthand)
fn looks_like_filesystem_path(path: &str) -> bool {
if path.is_empty() {
return false;
}
if Path::new(path).is_absolute() || path.starts_with("~/") {
return true;
}
let bytes = path.as_bytes();
bytes.len() >= 3
&& bytes[0].is_ascii_alphabetic()
&& bytes[1] == b':'
&& (bytes[2] == b'\\' || bytes[2] == b'/')
}
/// Tool for searching workspace memory. /// Tool for searching workspace memory.
/// ///
/// Performs hybrid search (FTS + semantic) across all memory documents. /// Performs hybrid search (FTS + semantic) across all memory documents.
@@ -166,8 +143,7 @@ impl Tool for MemoryWriteTool {
be remembered across sessions. Targets: 'memory' for curated long-term facts, \ be remembered across sessions. Targets: 'memory' for curated long-term facts, \
'daily_log' for timestamped session notes, 'heartbeat' for the periodic \ 'daily_log' for timestamped session notes, 'heartbeat' for the periodic \
checklist (HEARTBEAT.md), 'bootstrap' to clear the first-run ritual file, \ checklist (HEARTBEAT.md), 'bootstrap' to clear the first-run ritual file, \
or provide a custom workspace path for arbitrary file creation. \ or provide a custom path for arbitrary file creation."
Never pass absolute filesystem paths like '/Users/...' or 'C:\\...'."
} }
fn parameters_schema(&self) -> serde_json::Value { fn parameters_schema(&self) -> serde_json::Value {
@@ -207,14 +183,6 @@ impl Tool for MemoryWriteTool {
.and_then(|v| v.as_str()) .and_then(|v| v.as_str())
.unwrap_or("daily_log"); .unwrap_or("daily_log");
if looks_like_filesystem_path(target) {
return Err(ToolError::InvalidParameters(format!(
"'{}' looks like a local filesystem path. memory_write only works with workspace-memory paths. \
Use write_file for filesystem writes. For opening files in an editor, use shell with: open \"<absolute_path>\".",
target
)));
}
// Bootstrap target: clear BOOTSTRAP.md to mark first-run ritual complete. // Bootstrap target: clear BOOTSTRAP.md to mark first-run ritual complete.
// Handled early because it accepts empty content (unlike other targets). // Handled early because it accepts empty content (unlike other targets).
if target == "bootstrap" { if target == "bootstrap" {
@@ -364,8 +332,7 @@ impl Tool for MemoryReadTool {
fn description(&self) -> &str { fn description(&self) -> &str {
"Read a file from the workspace memory (database-backed storage). \ "Read a file from the workspace memory (database-backed storage). \
Use this to read files shown by memory_tree. NOT for local filesystem files \ Use this to read files shown by memory_tree. NOT for local filesystem files \
(use read_file for those). Do not pass absolute paths like '/Users/...' or 'C:\\...'. \ (use read_file for those). Works with identity files, heartbeat checklist, \
Works with identity files, heartbeat checklist, \
memory, daily logs, or any custom workspace path." memory, daily logs, or any custom workspace path."
} }
@@ -391,14 +358,6 @@ impl Tool for MemoryReadTool {
let path = require_str(&params, "path")?; let path = require_str(&params, "path")?;
if looks_like_filesystem_path(path) {
return Err(ToolError::InvalidParameters(format!(
"'{}' looks like a local filesystem path. memory_read only works with workspace-memory paths. \
Use read_file for filesystem reads. For opening files in an editor, use shell with: open \"<absolute_path>\".",
path
)));
}
let doc = self let doc = self
.workspace .workspace
.read(path) .read(path)
@@ -539,26 +498,6 @@ impl Tool for MemoryTreeTool {
} }
} }
#[cfg(test)]
mod path_routing_tests {
use super::looks_like_filesystem_path;
#[test]
fn detects_filesystem_paths() {
assert!(looks_like_filesystem_path("/Users/nige/file.md"));
assert!(looks_like_filesystem_path("C:\\Users\\nige\\file.md"));
assert!(looks_like_filesystem_path("D:/work/file.md"));
assert!(looks_like_filesystem_path("~/notes.md"));
}
#[test]
fn allows_workspace_memory_paths() {
assert!(!looks_like_filesystem_path("MEMORY.md"));
assert!(!looks_like_filesystem_path("daily/2026-03-11.md"));
assert!(!looks_like_filesystem_path("projects/alpha/notes.md"));
}
}
#[cfg(all(test, feature = "postgres"))] #[cfg(all(test, feature = "postgres"))]
mod tests { mod tests {
use super::*; use super::*;
-21
View File
@@ -104,14 +104,6 @@ impl Tool for RoutineCreateTool {
"enum": ["lightweight", "full_job"], "enum": ["lightweight", "full_job"],
"description": "Execution mode: 'lightweight' (single LLM call, default) or 'full_job' (multi-turn with tools)" "description": "Execution mode: 'lightweight' (single LLM call, default) or 'full_job' (multi-turn with tools)"
}, },
"use_tools": {
"type": "boolean",
"description": "Enable tool access in lightweight mode (default: false). Only safe tools (no approval required) are available. Ignored for full_job mode."
},
"max_tool_rounds": {
"type": "integer",
"description": "Max tool call rounds in lightweight mode (default: 3). Only used when use_tools is true."
},
"cooldown_secs": { "cooldown_secs": {
"type": "integer", "type": "integer",
"description": "Minimum seconds between fires (default: 300)" "description": "Minimum seconds between fires (default: 300)"
@@ -270,24 +262,11 @@ impl Tool for RoutineCreateTool {
}) })
.unwrap_or_default(); .unwrap_or_default();
let use_tools = params
.get("use_tools")
.and_then(|v| v.as_bool())
.unwrap_or(false);
let max_tool_rounds = params
.get("max_tool_rounds")
.and_then(|v| v.as_u64())
.map(|v| v.clamp(1, crate::agent::routine::MAX_TOOL_ROUNDS_LIMIT as u64) as u32)
.unwrap_or(3);
let action = match action_type { let action = match action_type {
"lightweight" => RoutineAction::Lightweight { "lightweight" => RoutineAction::Lightweight {
prompt: prompt.to_string(), prompt: prompt.to_string(),
context_paths, context_paths,
max_tokens: 4096, max_tokens: 4096,
use_tools,
max_tool_rounds,
}, },
"full_job" => { "full_job" => {
let tool_permissions = crate::agent::routine::parse_tool_permissions(&params); let tool_permissions = crate::agent::routine::parse_tool_permissions(&params);
+2 -77
View File
@@ -669,7 +669,7 @@ pub async fn authorize_mcp_server(
} }
// Determine client_id and endpoints // Determine client_id and endpoints
let (client_id, authorization_url, token_url, use_pkce, scopes, mut extra_params) = let (client_id, authorization_url, token_url, use_pkce, scopes, extra_params) =
if let Some(oauth) = &server_config.oauth { if let Some(oauth) = &server_config.oauth {
// Pre-configured OAuth // Pre-configured OAuth
let (auth_url, tok_url) = discover_oauth_endpoints(server_config).await?; let (auth_url, tok_url) = discover_oauth_endpoints(server_config).await?;
@@ -711,13 +711,6 @@ pub async fn authorize_mcp_server(
None None
}; };
// Generate OAuth state parameter. While optional in OAuth 2.1 with PKCE,
// some MCP servers (e.g. Attio) require it.
let mut state_bytes = [0u8; 16];
rand::rngs::OsRng.fill_bytes(&mut state_bytes);
let state = URL_SAFE_NO_PAD.encode(state_bytes);
extra_params.insert("state".to_string(), state);
// Compute canonical resource URI for RFC 8707 // Compute canonical resource URI for RFC 8707
let resource = canonical_resource_uri(&server_config.url); let resource = canonical_resource_uri(&server_config.url);
@@ -748,10 +741,7 @@ pub async fn authorize_mcp_server(
println!(" Waiting for authorization..."); println!(" Waiting for authorization...");
// Wait for callback. State is sent in the URL for servers that require it // Wait for callback
// (e.g. Attio), but we don't enforce validation on the callback because MCP
// servers use PKCE which already binds the request to the token exchange,
// and some servers may not echo state back.
let code = wait_for_authorization_callback(listener, &server_config.name).await?; let code = wait_for_authorization_callback(listener, &server_config.name).await?;
println!(" Exchanging code for token..."); println!(" Exchanging code for token...");
@@ -1721,69 +1711,4 @@ mod tests {
assert!(!url.contains("resource=")); assert!(!url.contains("resource="));
} }
/// Regression test: MCP OAuth authorization URLs must include a `state`
/// parameter. While OAuth 2.1 makes `state` optional when PKCE is used,
/// some MCP servers (e.g. Attio) require it and reject requests without it:
/// {"error":"invalid_request","error_description":"Invalid value provided
/// for: state"}
///
/// Including `state` is harmless for servers that don't require it, since
/// it is a standard OAuth parameter that compliant servers will echo back
/// or ignore.
///
/// The state is generated in `authorize_mcp_server` and injected into
/// `extra_params` before `build_authorization_url` is called. This test
/// verifies that `build_authorization_url` correctly propagates state from
/// extra_params into the URL, and that each generated state is unique.
#[test]
fn test_authorization_url_includes_state_parameter() {
// Simulate what authorize_mcp_server does: generate state and
// insert it into extra_params.
let mut extra_params = HashMap::new();
let mut state_bytes = [0u8; 16];
rand::rngs::OsRng.fill_bytes(&mut state_bytes);
let state = URL_SAFE_NO_PAD.encode(state_bytes);
extra_params.insert("state".to_string(), state.clone());
let pkce = PkceChallenge::generate();
let url = build_authorization_url(
"https://app.attio.com/oidc/authorize",
"test-client",
"http://127.0.0.1:9876/callback",
&[
"mcp".to_string(),
"offline_access".to_string(),
"openid".to_string(),
],
Some(&pkce),
&extra_params,
Some("https://mcp.attio.com/mcp"),
);
// State must be present in the URL
assert!(
url.contains(&format!("state={}", state)),
"Authorization URL must include the state parameter, got: {}",
url,
);
// State must be base64url-encoded (no padding, no +/)
assert!(!state.contains('+'), "State must be base64url-safe");
assert!(!state.contains('/'), "State must be base64url-safe");
assert!(!state.contains('='), "State must not have padding");
// State must have sufficient entropy (16 bytes -> 22 base64url chars)
assert!(
state.len() >= 22,
"State must have at least 128 bits of entropy, got {} chars",
state.len(),
);
// Two generated states must differ
let mut state_bytes_2 = [0u8; 16];
rand::rngs::OsRng.fill_bytes(&mut state_bytes_2);
let state_2 = URL_SAFE_NO_PAD.encode(state_bytes_2);
assert_ne!(state, state_2, "State must be unique per request");
}
} }
-1
View File
@@ -12,7 +12,6 @@ pub mod builtin;
pub mod execute; pub mod execute;
pub mod mcp; pub mod mcp;
pub mod rate_limiter; pub mod rate_limiter;
pub mod redaction;
pub mod schema_validator; pub mod schema_validator;
pub mod wasm; pub mod wasm;
-251
View File
@@ -1,251 +0,0 @@
use serde_json::{Map, Value};
const REDACTED: &str = "[REDACTED]";
const SENSITIVE_EXACT: &[&str] = &[
"authorization",
"proxy-authorization",
"cookie",
"set-cookie",
"x-api-key",
"api-key",
"api_key",
"access_token",
"refresh_token",
"session_token",
"id_token",
"token",
"password",
"passwd",
"secret",
"client_secret",
"private_key",
"apikey",
"apisecret",
];
const SENSITIVE_PARTS: &[&str] = &[
"password",
"passwd",
"secret",
"credential",
"authorization",
"cookie",
"apikey",
"apisecret",
];
const TOKEN_PARTS: &[&str] = &["token", "jwt"];
const KEY_PARTS: &[&str] = &["key"];
const CONTEXT_PARTS: &[&str] = &[
"auth",
"oauth",
"authorization",
"api",
"access",
"refresh",
"session",
"bearer",
"private",
"client",
"id",
"app",
"user",
"application",
"account",
];
fn split_camel_case_key_parts(key: &str) -> Vec<String> {
if key.is_empty() {
return Vec::new();
}
let chars: Vec<char> = key.chars().collect();
let mut parts = Vec::new();
let mut start = 0;
for i in 1..chars.len() {
let prev = chars[i - 1];
let cur = chars[i];
let next = chars.get(i + 1).copied();
let boundary = (prev.is_ascii_lowercase() && cur.is_ascii_uppercase())
|| (prev.is_ascii_alphabetic() && cur.is_ascii_digit())
|| (prev.is_ascii_digit() && cur.is_ascii_alphabetic())
|| (prev.is_ascii_uppercase()
&& cur.is_ascii_uppercase()
&& next.map(|n| n.is_ascii_lowercase()).unwrap_or(false));
if boundary {
parts.push(chars[start..i].iter().collect::<String>());
start = i;
}
}
parts.push(chars[start..].iter().collect::<String>());
parts
}
fn tokenize_key_parts(key: &str) -> Vec<String> {
let mut parts = Vec::new();
for segment in key.split(|c: char| !c.is_ascii_alphanumeric()) {
if segment.is_empty() {
continue;
}
parts.extend(split_camel_case_key_parts(segment));
}
parts.into_iter().map(|p| p.to_ascii_lowercase()).collect()
}
fn has_exact(parts: &[String], candidates: &[&str]) -> bool {
parts
.iter()
.any(|part| candidates.iter().any(|candidate| part == candidate))
}
fn has_candidate_or_numbered_variant(parts: &[String], candidates: &[&str]) -> bool {
parts.iter().any(|part| {
candidates.iter().any(|candidate| {
if part == candidate {
return true;
}
let Some(suffix) = part.strip_prefix(candidate) else {
return false;
};
!suffix.is_empty() && suffix.chars().all(|c| c.is_ascii_digit())
})
})
}
fn has_contextual_suffix(parts: &[String], candidates: &[&str]) -> bool {
parts.iter().any(|part| {
candidates.iter().any(|candidate| {
let Some(prefix) = part.strip_suffix(candidate) else {
return false;
};
!prefix.is_empty() && CONTEXT_PARTS.contains(&prefix)
})
})
}
fn is_sensitive_key(key: &str) -> bool {
let lower = key.to_ascii_lowercase();
if SENSITIVE_EXACT.contains(&lower.as_str()) {
return true;
}
let parts = tokenize_key_parts(key);
if parts.is_empty() {
return false;
}
if has_candidate_or_numbered_variant(&parts, SENSITIVE_PARTS) {
return true;
}
let has_token = has_candidate_or_numbered_variant(&parts, TOKEN_PARTS);
let has_key = has_candidate_or_numbered_variant(&parts, KEY_PARTS);
if has_token && has_key {
return true;
}
if has_contextual_suffix(&parts, TOKEN_PARTS) || has_contextual_suffix(&parts, KEY_PARTS) {
return true;
}
let has_context = has_exact(&parts, CONTEXT_PARTS);
has_context && (has_token || has_key)
}
fn redact_in_place(value: &mut Value) {
match value {
Value::Object(map) => redact_object(map),
Value::Array(items) => {
for item in items {
redact_in_place(item);
}
}
_ => {}
}
}
fn redact_object(map: &mut Map<String, Value>) {
for (key, val) in map {
if is_sensitive_key(key) {
*val = Value::String(REDACTED.to_string());
} else {
redact_in_place(val);
}
}
}
pub fn redact_sensitive_json(value: &Value) -> Value {
let mut cloned = value.clone();
redact_in_place(&mut cloned);
cloned
}
#[cfg(test)]
mod tests {
use super::{is_sensitive_key, redact_sensitive_json};
#[test]
fn redacts_exact_sensitive_keys() {
let input = serde_json::json!({
"headers": {
"Authorization": "Bearer abc",
"x-api-key": "k-123",
"content-type": "application/json"
},
"password": "p@ss"
});
let out = redact_sensitive_json(&input);
assert_eq!(out["headers"]["Authorization"], "[REDACTED]");
assert_eq!(out["headers"]["x-api-key"], "[REDACTED]");
assert_eq!(out["headers"]["content-type"], "application/json");
assert_eq!(out["password"], "[REDACTED]");
}
#[test]
fn redacts_nested_sensitive_keys() {
let input = serde_json::json!({
"body": {
"clientSecret": "xyz",
"nested": [{"authToken": "123"}, {"query": "ok"}]
}
});
let out = redact_sensitive_json(&input);
assert_eq!(out["body"]["clientSecret"], "[REDACTED]");
assert_eq!(out["body"]["nested"][0]["authToken"], "[REDACTED]");
assert_eq!(out["body"]["nested"][1]["query"], "ok");
}
#[test]
fn does_not_over_redact_common_non_sensitive_keys() {
assert!(!is_sensitive_key("author"));
assert!(!is_sensitive_key("authorize_user"));
assert!(!is_sensitive_key("token_count"));
assert!(!is_sensitive_key("tokenize"));
assert!(!is_sensitive_key("oauth_redirect_uri"));
}
#[test]
fn still_redacts_expected_token_keys() {
assert!(is_sensitive_key("auth_token"));
assert!(is_sensitive_key("oauth_token"));
assert!(is_sensitive_key("accessToken"));
assert!(is_sensitive_key("apiKey"));
assert!(is_sensitive_key("token_key"));
assert!(is_sensitive_key("appTokenKey"));
assert!(is_sensitive_key("userJwt"));
}
#[test]
fn redacts_lowercase_digit_suffix_segments() {
assert!(is_sensitive_key("password123"));
assert!(is_sensitive_key("secret99"));
assert!(is_sensitive_key("accounttoken2"));
}
}
+1 -33
View File
@@ -23,7 +23,7 @@ use crate::tools::builtin::{
ToolUpgradeTool, WriteFileTool, ToolUpgradeTool, WriteFileTool,
}; };
use crate::tools::rate_limiter::RateLimiter; use crate::tools::rate_limiter::RateLimiter;
use crate::tools::tool::{ApprovalRequirement, Tool, ToolDomain}; use crate::tools::tool::{Tool, ToolDomain};
use crate::tools::wasm::{ use crate::tools::wasm::{
Capabilities, OAuthRefreshConfig, ResourceLimits, SharedCredentialRegistry, WasmError, Capabilities, OAuthRefreshConfig, ResourceLimits, SharedCredentialRegistry, WasmError,
WasmStorageError, WasmToolRuntime, WasmToolStore, WasmToolWrapper, WasmStorageError, WasmToolRuntime, WasmToolStore, WasmToolWrapper,
@@ -278,38 +278,6 @@ impl ToolRegistry {
.collect() .collect()
} }
/// Get tool definitions excluding specific tools by name.
///
/// Used by lightweight routines to filter out denylisted and approval-gated tools
/// so the LLM only sees tools it is actually allowed to call.
pub async fn tool_definitions_excluding(&self, deny: &[&str]) -> Vec<ToolDefinition> {
let empty_params = serde_json::Value::Object(serde_json::Map::new());
let mut defs: Vec<ToolDefinition> = self
.tools
.read()
.await
.values()
.filter(|tool| {
// Exclude denylisted tools
if deny.contains(&tool.name()) {
return false;
}
// Exclude tools that require approval
matches!(
tool.requires_approval(&empty_params),
ApprovalRequirement::Never
)
})
.map(|tool| ToolDefinition {
name: tool.name().to_string(),
description: tool.description().to_string(),
parameters: tool.parameters_schema(),
})
.collect();
defs.sort_unstable_by(|a, b| a.name.cmp(&b.name));
defs
}
/// Register development tools for building software. /// Register development tools for building software.
/// ///
/// These tools provide shell access, file operations, and code editing /// These tools provide shell access, file operations, and code editing
+3 -19
View File
@@ -55,9 +55,7 @@ pub use embeddings::{
}; };
#[cfg(feature = "postgres")] #[cfg(feature = "postgres")]
pub use repository::Repository; pub use repository::Repository;
pub use search::{ pub use search::{RankedResult, SearchConfig, SearchResult, reciprocal_rank_fusion};
FusionStrategy, RankedResult, SearchConfig, SearchResult, fuse_results, reciprocal_rank_fusion,
};
use std::sync::Arc; use std::sync::Arc;
@@ -334,8 +332,6 @@ pub struct Workspace {
storage: WorkspaceStorage, storage: WorkspaceStorage,
/// Embedding provider for semantic search. /// Embedding provider for semantic search.
embeddings: Option<Arc<dyn EmbeddingProvider>>, embeddings: Option<Arc<dyn EmbeddingProvider>>,
/// Default search configuration applied to all queries.
search_defaults: SearchConfig,
} }
impl Workspace { impl Workspace {
@@ -347,7 +343,6 @@ impl Workspace {
agent_id: None, agent_id: None,
storage: WorkspaceStorage::Repo(Repository::new(pool)), storage: WorkspaceStorage::Repo(Repository::new(pool)),
embeddings: None, embeddings: None,
search_defaults: SearchConfig::default(),
} }
} }
@@ -360,7 +355,6 @@ impl Workspace {
agent_id: None, agent_id: None,
storage: WorkspaceStorage::Db(db), storage: WorkspaceStorage::Db(db),
embeddings: None, embeddings: None,
search_defaults: SearchConfig::default(),
} }
} }
@@ -376,16 +370,6 @@ impl Workspace {
self self
} }
/// Set the default search configuration from workspace search config.
pub fn with_search_config(mut self, config: &crate::config::WorkspaceSearchConfig) -> Self {
self.search_defaults = SearchConfig::default()
.with_fusion_strategy(config.fusion_strategy)
.with_rrf_k(config.rrf_k)
.with_fts_weight(config.fts_weight)
.with_vector_weight(config.vector_weight);
self
}
/// Get the user ID. /// Get the user ID.
pub fn user_id(&self) -> &str { pub fn user_id(&self) -> &str {
&self.user_id &self.user_id
@@ -725,13 +709,13 @@ impl Workspace {
/// Hybrid search across all memory documents. /// Hybrid search across all memory documents.
/// ///
/// Combines full-text search (BM25) with semantic search (vector similarity) /// Combines full-text search (BM25) with semantic search (vector similarity)
/// using the configured fusion strategy. /// using Reciprocal Rank Fusion (RRF).
pub async fn search( pub async fn search(
&self, &self,
query: &str, query: &str,
limit: usize, limit: usize,
) -> Result<Vec<SearchResult>, WorkspaceError> { ) -> Result<Vec<SearchResult>, WorkspaceError> {
self.search_with_config(query, self.search_defaults.clone().with_limit(limit)) self.search_with_config(query, SearchConfig::default().with_limit(limit))
.await .await
} }
+2 -2
View File
@@ -12,7 +12,7 @@ use uuid::Uuid;
use crate::error::WorkspaceError; use crate::error::WorkspaceError;
use crate::workspace::document::{MemoryChunk, MemoryDocument, WorkspaceEntry}; use crate::workspace::document::{MemoryChunk, MemoryDocument, WorkspaceEntry};
use crate::workspace::search::{RankedResult, SearchConfig, SearchResult, fuse_results}; use crate::workspace::search::{RankedResult, SearchConfig, SearchResult, reciprocal_rank_fusion};
/// Database repository for workspace operations. /// Database repository for workspace operations.
pub struct Repository { pub struct Repository {
@@ -415,7 +415,7 @@ impl Repository {
Vec::new() Vec::new()
}; };
Ok(fuse_results(fts_results, vector_results, config)) Ok(reciprocal_rank_fusion(fts_results, vector_results, config))
} }
/// Full-text search using PostgreSQL ts_rank_cd. /// Full-text search using PostgreSQL ts_rank_cd.
+7 -314
View File
@@ -1,30 +1,17 @@
//! Hybrid search combining full-text and semantic search. //! Hybrid search combining full-text and semantic search.
//! //!
//! Supports two fusion strategies: //! Uses Reciprocal Rank Fusion (RRF) to combine results from:
//! 1. **RRF** (Reciprocal Rank Fusion) — the default, rank-based method. //! 1. PostgreSQL full-text search (ts_rank_cd)
//! `score = sum(1 / (k + rank))` for each retrieval method. //! 2. pgvector cosine similarity search
//! 2. **WeightedScore** — converts ranks to scores via `1/rank`, combines with
//! configurable weights (`fts_weight * fts_score + vector_weight * vector_score`),
//! then normalizes to \[0,1\] by dividing by the maximum combined score.
//! //!
//! Both strategies combine results from: //! RRF formula: score = sum(1 / (k + rank)) for each retrieval method
//! - PostgreSQL / libSQL full-text search //! This is robust to different score scales and produces better results
//! - pgvector / libsql_vector cosine similarity search //! than simple score averaging.
use std::collections::HashMap; use std::collections::HashMap;
use uuid::Uuid; use uuid::Uuid;
/// Strategy used to fuse FTS and vector search results.
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub enum FusionStrategy {
/// Reciprocal Rank Fusion (default). Ignores `fts_weight`/`vector_weight`.
#[default]
Rrf,
/// Weighted score fusion using normalized rank-derived scores.
WeightedScore,
}
/// Configuration for hybrid search. /// Configuration for hybrid search.
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub struct SearchConfig { pub struct SearchConfig {
@@ -40,16 +27,6 @@ pub struct SearchConfig {
pub min_score: f32, pub min_score: f32,
/// Maximum results to fetch from each method before fusion. /// Maximum results to fetch from each method before fusion.
pub pre_fusion_limit: usize, pub pre_fusion_limit: usize,
/// Fusion strategy to use when combining results.
pub fusion_strategy: FusionStrategy,
/// Weight for FTS results in `WeightedScore` fusion (default 0.5).
/// Ignored by `Rrf` fusion. For env-based config via
/// `WorkspaceSearchConfig::resolve`, defaults are per-strategy.
pub fts_weight: f32,
/// Weight for vector results in `WeightedScore` fusion (default 0.5).
/// Ignored by `Rrf` fusion. For env-based config via
/// `WorkspaceSearchConfig::resolve`, defaults are per-strategy.
pub vector_weight: f32,
} }
impl Default for SearchConfig { impl Default for SearchConfig {
@@ -61,9 +38,6 @@ impl Default for SearchConfig {
use_vector: true, use_vector: true,
min_score: 0.0, min_score: 0.0,
pre_fusion_limit: 50, pre_fusion_limit: 50,
fusion_strategy: FusionStrategy::default(),
fts_weight: 0.5,
vector_weight: 0.5,
} }
} }
} }
@@ -100,32 +74,6 @@ impl SearchConfig {
self.min_score = score.clamp(0.0, 1.0); self.min_score = score.clamp(0.0, 1.0);
self self
} }
/// Set the fusion strategy.
pub fn with_fusion_strategy(mut self, strategy: FusionStrategy) -> Self {
self.fusion_strategy = strategy;
self
}
/// Set the FTS weight for `WeightedScore` fusion.
///
/// Non-finite (NaN, ±inf) or negative values are ignored.
pub fn with_fts_weight(mut self, weight: f32) -> Self {
if weight.is_finite() && weight >= 0.0 {
self.fts_weight = weight;
}
self
}
/// Set the vector weight for `WeightedScore` fusion.
///
/// Non-finite (NaN, ±inf) or negative values are ignored.
pub fn with_vector_weight(mut self, weight: f32) -> Self {
if weight.is_finite() && weight >= 0.0 {
self.vector_weight = weight;
}
self
}
} }
/// A search result with hybrid scoring. /// A search result with hybrid scoring.
@@ -139,7 +87,7 @@ pub struct SearchResult {
pub chunk_id: Uuid, pub chunk_id: Uuid,
/// Chunk content. /// Chunk content.
pub content: String, pub content: String,
/// Combined fusion score (0.0-1.0 normalized). Strategy-dependent (RRF or WeightedScore). /// Combined RRF score (0.0-1.0 normalized).
pub score: f32, pub score: f32,
/// Rank in FTS results (1-based, None if not in FTS results). /// Rank in FTS results (1-based, None if not in FTS results).
pub fts_rank: Option<u32>, pub fts_rank: Option<u32>,
@@ -175,22 +123,6 @@ pub struct RankedResult {
pub rank: u32, // 1-based rank pub rank: u32, // 1-based rank
} }
/// Fuse FTS and vector search results using the strategy specified in `config`.
///
/// This is the primary entry point for result fusion. Delegates to
/// [`reciprocal_rank_fusion`] or [`weighted_score_fusion`] based on
/// `config.fusion_strategy`.
pub fn fuse_results(
fts_results: Vec<RankedResult>,
vector_results: Vec<RankedResult>,
config: &SearchConfig,
) -> Vec<SearchResult> {
match config.fusion_strategy {
FusionStrategy::Rrf => reciprocal_rank_fusion(fts_results, vector_results, config),
FusionStrategy::WeightedScore => weighted_score_fusion(fts_results, vector_results, config),
}
}
/// Reciprocal Rank Fusion algorithm. /// Reciprocal Rank Fusion algorithm.
/// ///
/// Combines ranked results from multiple retrieval methods using the formula: /// Combines ranked results from multiple retrieval methods using the formula:
@@ -303,109 +235,6 @@ pub fn reciprocal_rank_fusion(
results results
} }
/// Weighted score fusion.
///
/// Converts ranks from each method into scores using `1/rank`
/// (so rank 1 → 1.0, rank N → 1/N), then combines them with
/// configurable weights: `fts_weight * fts_score + vector_weight * vector_score`.
///
/// The combined scores are then normalized to [0,1] by dividing by the
/// maximum score; post-processing (normalization, min_score filter, sort,
/// truncate) matches RRF.
pub fn weighted_score_fusion(
fts_results: Vec<RankedResult>,
vector_results: Vec<RankedResult>,
config: &SearchConfig,
) -> Vec<SearchResult> {
struct ChunkInfo {
document_id: Uuid,
document_path: String,
content: String,
score: f32,
fts_rank: Option<u32>,
vector_rank: Option<u32>,
}
let mut chunk_scores: HashMap<Uuid, ChunkInfo> = HashMap::new();
// Process FTS results: score = fts_weight * (1 / rank)
for result in fts_results {
let score = config.fts_weight * (1.0 / result.rank as f32);
chunk_scores
.entry(result.chunk_id)
.and_modify(|info| {
info.score += score;
info.fts_rank = Some(result.rank);
})
.or_insert(ChunkInfo {
document_id: result.document_id,
document_path: result.document_path,
content: result.content,
score,
fts_rank: Some(result.rank),
vector_rank: None,
});
}
// Process vector results: score = vector_weight * (1 / rank)
for result in vector_results {
let score = config.vector_weight * (1.0 / result.rank as f32);
chunk_scores
.entry(result.chunk_id)
.and_modify(|info| {
info.score += score;
info.vector_rank = Some(result.rank);
})
.or_insert(ChunkInfo {
document_id: result.document_id,
document_path: result.document_path,
content: result.content,
score,
fts_rank: None,
vector_rank: Some(result.rank),
});
}
let mut results: Vec<SearchResult> = chunk_scores
.into_iter()
.map(|(chunk_id, info)| SearchResult {
document_id: info.document_id,
document_path: info.document_path,
chunk_id,
content: info.content,
score: info.score,
fts_rank: info.fts_rank,
vector_rank: info.vector_rank,
})
.collect();
// Normalize scores to 0-1 range
if let Some(max_score) = results.iter().map(|r| r.score).reduce(f32::max)
&& max_score > 0.0
{
for result in &mut results {
result.score /= max_score;
}
}
// Filter by minimum score
if config.min_score > 0.0 {
results.retain(|r| r.score >= config.min_score);
}
// Sort by score descending
results.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
// Limit results
results.truncate(config.limit);
results
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
@@ -628,142 +457,6 @@ mod tests {
let vector_only = SearchConfig::default().vector_only(); let vector_only = SearchConfig::default().vector_only();
assert!(!vector_only.use_fts); assert!(!vector_only.use_fts);
assert!(vector_only.use_vector); assert!(vector_only.use_vector);
let weighted = SearchConfig::default()
.with_fusion_strategy(FusionStrategy::WeightedScore)
.with_fts_weight(0.8)
.with_vector_weight(0.2);
assert_eq!(weighted.fusion_strategy, FusionStrategy::WeightedScore);
assert!((weighted.fts_weight - 0.8).abs() < 0.001);
assert!((weighted.vector_weight - 0.2).abs() < 0.001);
}
#[test]
fn test_weighted_fusion_basic() {
// With equal weights, a hybrid match should still rank highest.
let config = SearchConfig::default()
.with_fusion_strategy(FusionStrategy::WeightedScore)
.with_fts_weight(1.0)
.with_vector_weight(1.0)
.with_limit(10);
let chunk1 = Uuid::new_v4(); // In both
let chunk2 = Uuid::new_v4(); // FTS only
let chunk3 = Uuid::new_v4(); // Vector only
let doc = Uuid::new_v4();
let fts = vec![make_result(chunk1, doc, 1), make_result(chunk2, doc, 2)];
let vec_results = vec![make_result(chunk1, doc, 1), make_result(chunk3, doc, 2)];
let results = weighted_score_fusion(fts, vec_results, &config);
assert_eq!(results.len(), 3);
// Hybrid match (chunk1) should be first — it gets score from both
assert_eq!(results[0].chunk_id, chunk1);
assert!(results[0].is_hybrid());
assert!(results[0].score > results[1].score);
}
#[test]
fn test_weighted_fusion_fts_boost() {
// High FTS weight should elevate FTS-only results above vector-only.
let config = SearchConfig::default()
.with_fusion_strategy(FusionStrategy::WeightedScore)
.with_fts_weight(2.0)
.with_vector_weight(0.5)
.with_limit(10);
let chunk_fts = Uuid::new_v4(); // FTS only, rank 2
let chunk_vec = Uuid::new_v4(); // Vector only, rank 2
let doc = Uuid::new_v4();
let fts = vec![make_result(chunk_fts, doc, 2)];
let vec_results = vec![make_result(chunk_vec, doc, 2)];
let results = weighted_score_fusion(fts, vec_results, &config);
assert_eq!(results.len(), 2);
// FTS result should rank higher because of the 2.0 weight vs 0.5
assert_eq!(results[0].chunk_id, chunk_fts);
assert!(results[0].from_fts());
assert!(!results[0].from_vector());
}
#[test]
fn test_weighted_fusion_single_source() {
// Only FTS results — should still work correctly.
let config = SearchConfig::default()
.with_fusion_strategy(FusionStrategy::WeightedScore)
.with_limit(10);
let chunk1 = Uuid::new_v4();
let chunk2 = Uuid::new_v4();
let doc = Uuid::new_v4();
let fts = vec![make_result(chunk1, doc, 1), make_result(chunk2, doc, 3)];
let results = weighted_score_fusion(fts, Vec::new(), &config);
assert_eq!(results.len(), 2);
assert_eq!(results[0].chunk_id, chunk1);
assert!(results[0].score > results[1].score);
// Top result should be normalized to 1.0
assert!((results[0].score - 1.0).abs() < 0.001);
}
#[test]
fn test_weight_setters_reject_invalid() {
let config = SearchConfig::default();
let original_fts = config.fts_weight;
let original_vec = config.vector_weight;
// NaN is ignored
let c = config.clone().with_fts_weight(f32::NAN);
assert!((c.fts_weight - original_fts).abs() < 0.001);
// Infinity is ignored
let c = config.clone().with_vector_weight(f32::INFINITY);
assert!((c.vector_weight - original_vec).abs() < 0.001);
// Negative is ignored
let c = config.clone().with_fts_weight(-1.0);
assert!((c.fts_weight - original_fts).abs() < 0.001);
// Negative infinity is ignored
let c = config.clone().with_vector_weight(f32::NEG_INFINITY);
assert!((c.vector_weight - original_vec).abs() < 0.001);
// Valid values > 1.0 are accepted (weights don't need to sum to 1.0)
let c = config.clone().with_fts_weight(2.0);
assert!((c.fts_weight - 2.0).abs() < 0.001);
// Zero is valid
let c = config.clone().with_vector_weight(0.0);
assert!(c.vector_weight.abs() < 0.001);
}
#[test]
fn test_fuse_results_dispatches_correctly() {
let chunk1 = Uuid::new_v4();
let doc = Uuid::new_v4();
let fts = vec![make_result(chunk1, doc, 1)];
// RRF strategy
let rrf_config = SearchConfig::default().with_limit(10);
let rrf_results = fuse_results(fts.clone(), Vec::new(), &rrf_config);
assert_eq!(rrf_results.len(), 1);
// Weighted strategy
let weighted_config = SearchConfig::default()
.with_fusion_strategy(FusionStrategy::WeightedScore)
.with_limit(10);
let weighted_results = fuse_results(fts, Vec::new(), &weighted_config);
assert_eq!(weighted_results.len(), 1);
// Both should normalize single result to 1.0
assert!((rrf_results[0].score - 1.0).abs() < 0.001);
assert!((weighted_results[0].score - 1.0).abs() < 0.001);
} }
// --- Edge case tests --- // --- Edge case tests ---
-99
View File
@@ -1,99 +0,0 @@
"""Scenario: Content Security Policy compliance.
Detects CSP violations (inline scripts, blocked resources) that would
break the gateway JS. This test catches regressions like adding
inline onclick handlers while a script-src CSP is active.
"""
from helpers import SEL
async def test_no_csp_violations_on_load(page):
"""Page load must produce zero CSP violation reports."""
violations = []
page.on("console", lambda msg: (
violations.append(msg.text)
if "content security policy" in msg.text.lower()
or msg.type == "error" and "refused" in msg.text.lower()
else None
))
# Reload the page to catch violations from initial load.
# Use "load" (not "networkidle") because the SSE stream keeps the
# connection open indefinitely, preventing networkidle from firing.
await page.reload(wait_until="load")
# Wait a moment for any deferred script execution
await page.wait_for_timeout(2000)
assert violations == [], (
f"CSP violations detected on page load:\n" + "\n".join(violations)
)
async def test_no_inline_event_handlers_in_html(page):
"""Static HTML must not contain any inline event handler attributes."""
inline_handlers = await page.evaluate("""() => {
const allElements = document.querySelectorAll('*');
const found = [];
const handlerAttrs = [
'onclick', 'onchange', 'onsubmit', 'onload', 'onerror',
'onmouseover', 'onfocus', 'onblur', 'onkeydown', 'onkeyup',
'oninput', 'onmousedown', 'onmouseup'
];
for (const el of allElements) {
for (const attr of handlerAttrs) {
if (el.hasAttribute(attr)) {
const tag = el.tagName.toLowerCase();
const id = el.id ? '#' + el.id : '';
const cls = el.className ? '.' + el.className.split(' ')[0] : '';
found.push(tag + id + cls + '[' + attr + ']');
}
}
}
return found;
}""")
assert inline_handlers == [], (
f"Found inline event handlers (CSP-incompatible):\n"
+ "\n".join(f" - {h}" for h in inline_handlers)
)
async def test_no_js_errors_on_page_load(page):
"""No JavaScript errors should occur on page load."""
errors = []
page.on("pageerror", lambda err: errors.append(str(err)))
await page.reload(wait_until="load")
await page.wait_for_timeout(2000)
assert errors == [], (
f"JavaScript errors on page load:\n" + "\n".join(errors)
)
async def test_buttons_still_functional_after_csp_migration(page):
"""Core buttons must still be wired up via addEventListener."""
# Verify that key buttons have click handlers attached (not inline)
# by checking that clicking them doesn't throw and they exist in the DOM
button_ids = [
'send-btn',
'thread-new-btn',
'thread-toggle-btn',
'restart-btn',
'memory-edit-btn',
'logs-pause-btn',
'logs-clear-btn',
]
for btn_id in button_ids:
exists = await page.evaluate(
"id => document.getElementById(id) !== null", btn_id
)
assert exists, f"Button #{btn_id} not found in DOM"
# Verify the assistant thread div is clickable (has no onclick but
# should be handled by delegation or direct addEventListener)
assistant_el = page.locator(SEL["chat_input"])
await assistant_el.wait_for(state="visible", timeout=5000)
-132
View File
@@ -403,136 +403,4 @@ mod advanced {
rig.verify_trace_expects(&trace, &responses); rig.verify_trace_expects(&trace, &responses);
rig.shutdown(); rig.shutdown();
} }
// -----------------------------------------------------------------------
// 8. MCP extension lifecycle (search → install → activate → use)
//
// Exercises the MCP extension flow with a mock MCP server:
// Turn 1: tool_search → tool_install → text
// (inject token + activate between turns)
// Turn 2: mock-notion_notion-search → mock-notion_notion-fetch → text
// -----------------------------------------------------------------------
#[tokio::test]
async fn mcp_extension_lifecycle() {
use crate::support::mock_mcp_server::{MockToolResponse, start_mock_mcp_server};
use ironclaw::extensions::{AuthHint, ExtensionKind, ExtensionSource, RegistryEntry};
// 1. Start mock MCP server with pre-configured tool responses.
let mock_server = start_mock_mcp_server(vec![
MockToolResponse {
name: "notion-search".into(),
content: serde_json::json!({
"results": [
{"id": "page-001", "title": "Project Alpha", "type": "page"},
{"id": "page-002", "title": "Sprint Planning", "type": "page"}
]
}),
},
MockToolResponse {
name: "notion-fetch".into(),
content: serde_json::json!({
"id": "page-001",
"title": "Project Alpha",
"content": "Status: In Progress\n- Sprint planning on March 15\n- API redesign review pending"
}),
},
])
.await;
// 2. Load trace fixture.
let trace =
LlmTrace::from_file(format!("{FIXTURES}/mcp_extension_lifecycle.json")).unwrap();
// 3. Build rig with auto-approve (so tool_install doesn't block).
let rig = TestRigBuilder::new()
.with_trace(trace.clone())
.with_auto_approve_tools(true)
.with_max_tool_iterations(15)
.build()
.await;
// 4. Inject mock-notion registry entry pointing to the mock server.
let ext_mgr = rig
.extension_manager()
.expect("test rig must expose extension manager");
ext_mgr
.inject_registry_entry(RegistryEntry {
name: "mock-notion".to_string(),
display_name: "Mock Notion".to_string(),
kind: ExtensionKind::McpServer,
description: "Test MCP server for E2E lifecycle test".to_string(),
keywords: vec!["mock-notion".into(), "notion".into()],
source: ExtensionSource::McpUrl {
url: mock_server.mcp_url(),
},
fallback_source: None,
auth_hint: AuthHint::Dcr,
version: None,
})
.await;
// 5. Turn 1: "setup mock-notion" → search → install → text.
rig.send_message("setup mock-notion").await;
let r1 = rig.wait_for_responses(1, TIMEOUT).await;
assert!(!r1.is_empty(), "Turn 1: no response");
// 6. Simulate OAuth completion: inject token + activate.
// This mirrors what the gateway's oauth_callback_handler does after
// the user completes the OAuth flow in their browser.
let secret_name = "mcp_mock-notion_access_token";
ext_mgr
.secrets()
.create(
"default",
ironclaw::secrets::CreateSecretParams::new(secret_name, "mock-access-token")
.with_provider("mcp:mock-notion".to_string()),
)
.await
.expect("failed to inject test token");
let activate_result = ext_mgr.activate("mock-notion").await;
assert!(
activate_result.is_ok(),
"activation failed: {:?}",
activate_result.err()
);
// 7. Turn 2: "check what's in my notion" → notion-search → notion-fetch → text.
// Wait for r1.len() + 1 to ensure we observe at least one new turn-2 response.
let turn1_count = r1.len();
rig.send_message("it's done, check what's in my notion")
.await;
let r2 = rig.wait_for_responses(turn1_count + 1, TIMEOUT).await;
assert!(
r2.len() > turn1_count,
"Turn 2: expected new responses beyond turn 1's {turn1_count}, got {}",
r2.len()
);
// 8. Verify tool calls across both turns.
let started = rig.tool_calls_started();
assert!(
started.iter().any(|s| s == "tool_search"),
"tool_search not called: {started:?}"
);
assert!(
started.iter().any(|s| s == "tool_install"),
"tool_install not called: {started:?}"
);
// Verify MCP tools were called in turn 2.
assert!(
started.iter().any(|s| s.starts_with("mock-notion_")),
"No mock-notion MCP tools called: {started:?}"
);
// Verify all tools that completed did so successfully.
let completed = rig.tool_calls_completed();
let failed: Vec<_> = completed.iter().filter(|(_, success)| !success).collect();
assert!(failed.is_empty(), "Tools failed: {failed:?}");
mock_server.shutdown().await;
rig.shutdown();
}
} }
-2
View File
@@ -61,8 +61,6 @@ mod tests {
prompt: prompt.to_string(), prompt: prompt.to_string(),
context_paths: vec![], context_paths: vec![],
max_tokens: 1000, max_tokens: 1000,
use_tools: false,
max_tool_rounds: 3,
}, },
guardrails: RoutineGuardrails { guardrails: RoutineGuardrails {
cooldown: Duration::from_secs(0), cooldown: Duration::from_secs(0),
@@ -1,98 +0,0 @@
{
"model_name": "advanced-mcp-extension-lifecycle",
"expects": {
"tools_used": ["tool_search", "tool_install"],
"tools_order": ["tool_search", "tool_install"],
"all_tools_succeeded": true,
"min_responses": 2
},
"turns": [
{
"user_input": "setup mock-notion",
"steps": [
{
"request_hint": { "last_user_message_contains": "setup mock-notion" },
"response": {
"type": "tool_calls",
"tool_calls": [
{
"id": "call_search_1",
"name": "tool_search",
"arguments": { "query": "mock-notion" }
}
],
"input_tokens": 500,
"output_tokens": 30
}
},
{
"request_hint": { "last_user_message_contains": "setup mock-notion", "min_message_count": 4 },
"response": {
"type": "tool_calls",
"tool_calls": [
{
"id": "call_install_1",
"name": "tool_install",
"arguments": { "name": "mock-notion" }
}
],
"input_tokens": 600,
"output_tokens": 30
}
},
{
"request_hint": { "last_user_message_contains": "setup mock-notion", "min_message_count": 6 },
"response": {
"type": "text",
"content": "I've installed Mock Notion. Please authenticate to connect your account — once done, tell me and I'll load the MCP tools.",
"input_tokens": 700,
"output_tokens": 35
}
}
]
},
{
"user_input": "it's done, check what's in my notion",
"steps": [
{
"request_hint": { "last_user_message_contains": "notion" },
"response": {
"type": "tool_calls",
"tool_calls": [
{
"id": "call_ns_1",
"name": "mock-notion_notion-search",
"arguments": { "query": "recent notes" }
}
],
"input_tokens": 900,
"output_tokens": 30
}
},
{
"request_hint": { "min_message_count": 4 },
"response": {
"type": "tool_calls",
"tool_calls": [
{
"id": "call_nf_1",
"name": "mock-notion_notion-fetch",
"arguments": { "query": "page-001" }
}
],
"input_tokens": 1000,
"output_tokens": 30
}
},
{
"response": {
"type": "text",
"content": "Here's what I found in your Notion:\n\n**Project Alpha** — Status: In Progress\n- Sprint planning on March 15\n- API redesign review pending\n\nLet me know if you want more details on any item.",
"input_tokens": 1100,
"output_tokens": 50
}
}
]
}
]
}
-340
View File
@@ -1,340 +0,0 @@
//! Mock MCP server for E2E testing of the extension lifecycle.
//!
//! Provides a minimal HTTP server with:
//! - OAuth 2.1 discovery (`.well-known/oauth-protected-resource`, `.well-known/oauth-authorization-server`)
//! - Dynamic Client Registration (`/register`)
//! - Token exchange (`/token`)
//! - MCP JSON-RPC endpoint (`/mcp`) with `initialize`, `tools/list`, `tools/call`
//!
//! Tool call responses are pre-configured via `MockToolResponse`.
#![allow(dead_code)]
use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::Arc;
use axum::extract::State;
use axum::http::{HeaderMap, StatusCode};
use axum::response::IntoResponse;
use axum::routing::{get, post};
use axum::{Json, Router};
use serde::{Deserialize, Serialize};
use tokio::sync::oneshot;
/// A pre-configured response for a specific MCP tool call.
#[derive(Clone, Debug)]
pub struct MockToolResponse {
/// Tool name (e.g., "notion-search").
pub name: String,
/// JSON response content for `tools/call`.
pub content: serde_json::Value,
}
/// A running mock MCP server.
pub struct MockMcpServer {
/// Base URL including port (e.g., "http://127.0.0.1:12345").
pub base_url: String,
/// Shutdown signal sender.
shutdown_tx: Option<oneshot::Sender<()>>,
/// Server task handle.
handle: Option<tokio::task::JoinHandle<()>>,
}
impl MockMcpServer {
/// The MCP endpoint URL for use in registry entries.
pub fn mcp_url(&self) -> String {
format!("{}/mcp", self.base_url)
}
/// Shut down the server.
pub async fn shutdown(mut self) {
if let Some(tx) = self.shutdown_tx.take() {
let _ = tx.send(());
}
if let Some(h) = self.handle.take() {
let _ = h.await;
}
}
}
impl Drop for MockMcpServer {
fn drop(&mut self) {
if let Some(tx) = self.shutdown_tx.take() {
let _ = tx.send(());
}
if let Some(h) = self.handle.take() {
h.abort();
}
}
}
/// Shared state for the mock server handlers.
struct MockState {
/// Base URL (filled after bind).
base_url: String,
/// Tool definitions served by tools/list.
tools: Vec<McpToolDef>,
/// Pre-configured tool call responses keyed by tool name.
/// Multiple calls to the same tool return responses in order.
tool_responses: HashMap<String, Vec<serde_json::Value>>,
/// Counter for tool_responses consumption (per tool name).
tool_response_idx: std::sync::Mutex<HashMap<String, usize>>,
}
#[derive(Clone, Serialize)]
struct McpToolDef {
name: String,
description: String,
#[serde(rename = "inputSchema")]
input_schema: serde_json::Value,
}
/// Start a mock MCP server on a random port.
///
/// `tool_responses` configures what `tools/call` returns for each tool name.
/// Multiple responses for the same tool are returned in order.
pub async fn start_mock_mcp_server(tool_responses: Vec<MockToolResponse>) -> MockMcpServer {
// Build tool definitions and response map.
let mut tools = Vec::new();
let mut response_map: HashMap<String, Vec<serde_json::Value>> = HashMap::new();
let mut seen_tools = std::collections::HashSet::new();
for tr in &tool_responses {
if seen_tools.insert(tr.name.clone()) {
tools.push(McpToolDef {
name: tr.name.clone(),
description: format!("Mock tool: {}", tr.name),
input_schema: serde_json::json!({"type": "object", "properties": {}}),
});
}
response_map
.entry(tr.name.clone())
.or_default()
.push(tr.content.clone());
}
// Bind to a random port.
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("failed to bind mock MCP server");
let addr: SocketAddr = listener.local_addr().expect("no local addr");
let base_url = format!("http://127.0.0.1:{}", addr.port());
let state = Arc::new(MockState {
base_url: base_url.clone(),
tools,
tool_responses: response_map,
tool_response_idx: std::sync::Mutex::new(HashMap::new()),
});
let app = Router::new()
.route(
"/.well-known/oauth-protected-resource/mcp",
get(handle_protected_resource),
)
.route(
"/.well-known/oauth-authorization-server",
get(handle_auth_server_metadata),
)
.route("/register", post(handle_register))
.route("/authorize", get(handle_authorize))
.route("/token", post(handle_token))
.route("/mcp", post(handle_mcp))
.with_state(state);
let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>();
let handle = tokio::spawn(async move {
axum::serve(listener, app)
.with_graceful_shutdown(async {
let _ = shutdown_rx.await;
})
.await
.expect("mock MCP server failed");
});
// Wait briefly for the server to start accepting.
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
MockMcpServer {
base_url,
shutdown_tx: Some(shutdown_tx),
handle: Some(handle),
}
}
// ── OAuth discovery endpoints ───────────────────────────────────────────
async fn handle_protected_resource(State(state): State<Arc<MockState>>) -> impl IntoResponse {
Json(serde_json::json!({
"resource": format!("{}/mcp", state.base_url),
"authorization_servers": [state.base_url],
"scopes_supported": ["read", "write"]
}))
}
async fn handle_auth_server_metadata(State(state): State<Arc<MockState>>) -> impl IntoResponse {
Json(serde_json::json!({
"issuer": state.base_url,
"authorization_endpoint": format!("{}/authorize", state.base_url),
"token_endpoint": format!("{}/token", state.base_url),
"registration_endpoint": format!("{}/register", state.base_url),
"response_types_supported": ["code"],
"grant_types_supported": ["authorization_code"],
"code_challenge_methods_supported": ["S256"],
"scopes_supported": ["read", "write"]
}))
}
// ── OAuth DCR ───────────────────────────────────────────────────────────
async fn handle_register() -> impl IntoResponse {
Json(serde_json::json!({
"client_id": "mock-client-id",
"client_name": "ironclaw-test",
"redirect_uris": [],
"grant_types": ["authorization_code"],
"response_types": ["code"],
"token_endpoint_auth_method": "none"
}))
}
// ── OAuth authorize (auto-approve) ──────────────────────────────────────
/// In a real flow, this would show a consent screen. For testing, we just
/// need the endpoint to exist. The test will bypass OAuth by injecting
/// tokens directly.
async fn handle_authorize() -> impl IntoResponse {
// Return a simple HTML page; in practice the test injects tokens directly.
axum::response::Html(
"<html><body>Mock OAuth: authorize endpoint. Tests bypass this.</body></html>",
)
}
// ── OAuth token exchange ────────────────────────────────────────────────
async fn handle_token() -> impl IntoResponse {
Json(serde_json::json!({
"access_token": "mock-access-token",
"token_type": "Bearer",
"expires_in": 3600,
"refresh_token": "mock-refresh-token"
}))
}
// ── MCP JSON-RPC endpoint ───────────────────────────────────────────────
#[derive(Deserialize)]
struct JsonRpcRequest {
jsonrpc: String,
id: Option<serde_json::Value>,
method: String,
#[serde(default)]
params: Option<serde_json::Value>,
}
async fn handle_mcp(
State(state): State<Arc<MockState>>,
headers: HeaderMap,
Json(req): Json<JsonRpcRequest>,
) -> impl IntoResponse {
// Check for auth header.
let auth = headers
.get("authorization")
.and_then(|v| v.to_str().ok())
.unwrap_or("");
if !auth.starts_with("Bearer ") || &auth[7..] != "mock-access-token" {
// Return 401 with WWW-Authenticate header per MCP OAuth spec.
let www_auth = format!(
"Bearer resource_metadata=\"{}/.well-known/oauth-protected-resource/mcp\"",
state.base_url
);
return (
StatusCode::UNAUTHORIZED,
[("www-authenticate", www_auth.as_str())],
Json(serde_json::json!({
"jsonrpc": "2.0",
"id": req.id,
"error": {"code": -32000, "message": "Unauthorized"}
})),
)
.into_response();
}
// Handle notifications (no id) silently.
if req.id.is_none() {
return StatusCode::OK.into_response();
}
let response = match req.method.as_str() {
"initialize" => serde_json::json!({
"jsonrpc": "2.0",
"id": req.id,
"result": {
"protocolVersion": "2024-11-05",
"serverInfo": {
"name": "mock-mcp-server",
"version": "1.0.0"
},
"capabilities": {
"tools": {}
}
}
}),
"tools/list" => {
let tools: Vec<serde_json::Value> = state
.tools
.iter()
.map(|t| serde_json::to_value(t).unwrap())
.collect();
serde_json::json!({
"jsonrpc": "2.0",
"id": req.id,
"result": {
"tools": tools
}
})
}
"tools/call" => {
let tool_name = req
.params
.as_ref()
.and_then(|p| p.get("name"))
.and_then(|n| n.as_str())
.unwrap_or("unknown");
let content = {
let mut idx_map = state.tool_response_idx.lock().unwrap();
let idx = idx_map.entry(tool_name.to_string()).or_insert(0);
let responses = state.tool_responses.get(tool_name);
let result = responses
.and_then(|r| r.get(*idx))
.cloned()
.unwrap_or_else(|| serde_json::json!({"error": "no mock response configured"}));
*idx += 1;
result
};
serde_json::json!({
"jsonrpc": "2.0",
"id": req.id,
"result": {
"content": [
{
"type": "text",
"text": serde_json::to_string(&content).unwrap_or_default()
}
]
}
})
}
_ => serde_json::json!({
"jsonrpc": "2.0",
"id": req.id,
"error": {"code": -32601, "message": format!("Method not found: {}", req.method)}
}),
};
Json(response).into_response()
}
-1
View File
@@ -4,7 +4,6 @@ pub mod cleanup;
pub mod gateway_workflow_harness; pub mod gateway_workflow_harness;
pub mod instrumented_llm; pub mod instrumented_llm;
pub mod metrics; pub mod metrics;
pub mod mock_mcp_server;
pub mod mock_openai_server; pub mod mock_openai_server;
pub mod test_channel; pub mod test_channel;
pub mod test_rig; pub mod test_rig;
-10
View File
@@ -50,9 +50,6 @@ pub struct TestRig {
/// The underlying TraceLlm for inspecting captured requests. /// The underlying TraceLlm for inspecting captured requests.
#[cfg(feature = "libsql")] #[cfg(feature = "libsql")]
trace_llm: Option<Arc<TraceLlm>>, trace_llm: Option<Arc<TraceLlm>>,
/// Extension manager for direct extension operations in tests.
#[cfg(feature = "libsql")]
extension_manager: Option<Arc<ironclaw::extensions::ExtensionManager>>,
/// Temp directory guard -- keeps the libSQL database file alive. /// Temp directory guard -- keeps the libSQL database file alive.
#[cfg(feature = "libsql")] #[cfg(feature = "libsql")]
_temp_dir: tempfile::TempDir, _temp_dir: tempfile::TempDir,
@@ -79,11 +76,6 @@ impl TestRig {
.unwrap_or_default() .unwrap_or_default()
} }
/// Return the extension manager for direct extension operations in tests.
pub fn extension_manager(&self) -> Option<&Arc<ironclaw::extensions::ExtensionManager>> {
self.extension_manager.as_ref()
}
/// Wait until at least `n` responses have been captured, or `timeout` elapses. /// Wait until at least `n` responses have been captured, or `timeout` elapses.
pub async fn wait_for_responses(&self, n: usize, timeout: Duration) -> Vec<OutgoingResponse> { pub async fn wait_for_responses(&self, n: usize, timeout: Duration) -> Vec<OutgoingResponse> {
self.channel.wait_for_responses(n, timeout).await self.channel.wait_for_responses(n, timeout).await
@@ -608,7 +600,6 @@ impl TestRigBuilder {
// Save references for test accessors. // Save references for test accessors.
let db_ref = components.db.clone().expect("test rig requires a database"); let db_ref = components.db.clone().expect("test rig requires a database");
let workspace_ref = components.workspace.clone(); let workspace_ref = components.workspace.clone();
let ext_mgr_ref = components.extension_manager.clone();
// 7. Construct AgentDeps from AppComponents (mirrors main.rs). // 7. Construct AgentDeps from AppComponents (mirrors main.rs).
let deps = AgentDeps { let deps = AgentDeps {
@@ -704,7 +695,6 @@ impl TestRigBuilder {
db: db_ref, db: db_ref,
workspace: workspace_ref, workspace: workspace_ref,
trace_llm: trace_llm_ref, trace_llm: trace_llm_ref,
extension_manager: ext_mgr_ref,
_temp_dir: temp_dir, _temp_dir: temp_dir,
} }
} }
-23
View File
@@ -1,23 +0,0 @@
[package]
name = "llm-context-tool"
version = "0.1.0"
edition = "2021"
description = "Brave Search LLM Context tool for IronClaw (WASM component)"
license = "MIT OR Apache-2.0"
publish = false
[dependencies]
serde = { version = "1.0", features = ["derive"] }
serde_json = "1.0"
wit-bindgen = "0.41.0"
[lib]
crate-type = ["cdylib"]
[profile.release]
opt-level = "s"
lto = true
strip = true
codegen-units = 1
[workspace]
@@ -1,53 +0,0 @@
{
"version": "0.1.0",
"wit_version": "0.3.0",
"capabilities": {
"http": {
"allowlist": [
{
"host": "api.search.brave.com",
"path_prefix": "/res/v1/llm/context",
"methods": [
"POST"
]
}
],
"credentials": {
"brave_api_key": {
"secret_name": "brave_api_key",
"location": {
"type": "header",
"name": "X-Subscription-Token"
},
"host_patterns": [
"api.search.brave.com"
]
}
},
"rate_limit": {
"requests_per_minute": 30,
"requests_per_hour": 500
}
},
"secrets": {
"allowed_names": [
"brave_api_key"
]
}
},
"auth": {
"secret_name": "brave_api_key",
"display_name": "Brave Search",
"instructions": "Get a free API key at brave.com/search/api/ (Free tier: 2,000 queries/month). Same key as Web Search.",
"setup_url": "https://brave.com/search/api/",
"env_var": "BRAVE_API_KEY"
},
"setup": {
"required_secrets": [
{
"name": "brave_api_key",
"prompt": "Brave Search API key (from brave.com/search/api)"
}
]
}
}
File diff suppressed because it is too large Load Diff