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
117 changed files with 1121 additions and 9915 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_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 ===
# LLM_BACKEND=anthropic
# 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."
-59
View File
@@ -1,59 +0,0 @@
#!/usr/bin/env bash
load_commit_summary() {
local range="$1"
local max_commits="${2:-50}"
local commit_list overflow
commit_list="$(git log --oneline --no-merges --reverse "${range}" 2>/dev/null || echo "")"
if [ -n "${commit_list}" ]; then
COMMIT_COUNT="$(printf '%s\n' "${commit_list}" | wc -l | tr -d ' ')"
if [ "${COMMIT_COUNT}" -gt "${max_commits}" ]; then
COMMIT_MD="$(printf '%s\n' "${commit_list}" | head -n "${max_commits}" | sed 's/^/- /')"
overflow=$((COMMIT_COUNT - max_commits))
COMMIT_MD+=$'\n'"- ... and ${overflow} more (see compare view)"
else
COMMIT_MD="$(printf '%s\n' "${commit_list}" | sed 's/^/- /')"
fi
else
COMMIT_COUNT=0
COMMIT_MD="- (no non-merge commits in range)"
fi
}
replace_marked_section() {
local body_file="$1"
local section_file="$2"
local section_start="$3"
local section_end="$4"
local output_file="$5"
if grep -qF "${section_start}" "${body_file}" && grep -qF "${section_end}" "${body_file}"; then
awk -v start="${section_start}" -v end="${section_end}" -v replacement_file="${section_file}" '
BEGIN {
while ((getline line < replacement_file) > 0) {
replacement = replacement line ORS
}
in_block = 0
}
$0 == start {
printf "%s", replacement
in_block = 1
next
}
$0 == end {
in_block = 0
next
}
!in_block {
print
}
' "${body_file}" > "${output_file}"
else
cp "${body_file}" "${output_file}"
if [ -s "${output_file}" ]; then
printf '\n\n' >> "${output_file}"
fi
cat "${section_file}" >> "${output_file}"
fi
}
-101
View File
@@ -1,101 +0,0 @@
#!/usr/bin/env bash
set -euo pipefail
: "${PR_NUMBER:?PR_NUMBER is required}"
: "${REPO:?REPO is required}"
MAIN_BRANCH="${MAIN_BRANCH:-main}"
DRY_RUN="${DRY_RUN:-false}"
SECTION_START="<!-- staging-promotion-release-summary:start -->"
SECTION_END="<!-- staging-promotion-release-summary:end -->"
TMP_DIR="$(mktemp -d)"
trap 'rm -rf "${TMP_DIR}"' EXIT
# shellcheck source=.github/scripts/pr-body-utils.sh
source "$(dirname "$0")/pr-body-utils.sh"
gh pr view "${PR_NUMBER}" --repo "${REPO}" --json body > "${TMP_DIR}/pr.json"
jq -r '.body // ""' < "${TMP_DIR}/pr.json" > "${TMP_DIR}/body.md"
git fetch origin "${MAIN_BRANCH}"
git fetch origin "+refs/tags/v*:refs/tags/v*"
LAST_TAG="$(git describe --tags --match 'v*' --abbrev=0 "origin/${MAIN_BRANCH}" 2>/dev/null || true)"
if [ -n "${LAST_TAG}" ]; then
RANGE="${LAST_TAG}..origin/${MAIN_BRANCH}"
HEADER="## Staging promotion batches since ${LAST_TAG}"
EMPTY_MESSAGE="_No structured staging promotion merges found since ${LAST_TAG}._"
else
RANGE="origin/${MAIN_BRANCH}"
HEADER="## Staging promotion batches on ${MAIN_BRANCH}"
EMPTY_MESSAGE="_No structured staging promotion merges found on ${MAIN_BRANCH}._"
fi
{
echo "${SECTION_START}"
echo "${HEADER}"
echo
} > "${TMP_DIR}/section.md"
FOUND_SUMMARY=false
while IFS= read -r sha; do
[ -n "${sha}" ] || continue
BODY="$(git show -s --format=%b "${sha}")"
if ! printf '%s\n' "${BODY}" | grep -q '^staging-promotion-summary-v1$'; then
continue
fi
FOUND_SUMMARY=true
SUBJECT="$(git show -s --format=%s "${sha}")"
PR_REF="$(printf '%s\n' "${BODY}" | sed -n 's/^promotion-pr: //p' | head -n 1)"
COMMIT_COUNT="$(printf '%s\n' "${BODY}" | sed -n 's/^current-commit-count: //p' | head -n 1)"
CURRENT_RANGE="$(printf '%s\n' "${BODY}" | sed -n 's/^current-range: //p' | head -n 1)"
COMMIT_BLOCK="$(printf '%s\n' "${BODY}" | awk 'capture { print } /^Current commits in this promotion \([0-9]+\):$/ { capture = 1 }')"
{
echo "### ${SUBJECT}"
echo
if [ -n "${PR_REF}" ]; then
echo "**Promotion PR:** ${PR_REF}"
fi
if [ -n "${COMMIT_COUNT}" ]; then
echo "**Commit count:** ${COMMIT_COUNT}"
fi
if [ -n "${CURRENT_RANGE}" ]; then
echo "**Range:** \`${CURRENT_RANGE}\`"
fi
echo
if [ -n "${COMMIT_BLOCK}" ]; then
echo "${COMMIT_BLOCK}"
else
echo "- (no commit summary found)"
fi
echo
} >> "${TMP_DIR}/section.md"
done < <(git log --merges --reverse --format='%H' "${RANGE}")
if [ "${FOUND_SUMMARY}" = false ]; then
{
echo "${EMPTY_MESSAGE}"
echo
} >> "${TMP_DIR}/section.md"
fi
{
echo "*Auto-updated from structured staging promotion merge bodies on ${MAIN_BRANCH}.*"
echo "${SECTION_END}"
} >> "${TMP_DIR}/section.md"
replace_marked_section \
"${TMP_DIR}/body.md" \
"${TMP_DIR}/section.md" \
"${SECTION_START}" \
"${SECTION_END}" \
"${TMP_DIR}/new-body.md"
if [ "${DRY_RUN}" = "true" ]; then
echo "Dry run enabled. Computed PR body for #${PR_NUMBER}:"
cat "${TMP_DIR}/new-body.md"
else
gh pr edit "${PR_NUMBER}" --repo "${REPO}" --body-file "${TMP_DIR}/new-body.md"
fi
@@ -1,53 +0,0 @@
#!/usr/bin/env bash
set -euo pipefail
: "${PR_NUMBER:?PR_NUMBER is required}"
: "${REPO:?REPO is required}"
MAX_COMMITS="${MAX_COMMITS:-50}"
DRY_RUN="${DRY_RUN:-false}"
SECTION_START="<!-- staging-ci-current:start -->"
SECTION_END="<!-- staging-ci-current:end -->"
TMP_DIR="$(mktemp -d)"
trap 'rm -rf "${TMP_DIR}"' EXIT
# shellcheck source=.github/scripts/pr-body-utils.sh
source "$(dirname "$0")/pr-body-utils.sh"
gh pr view "${PR_NUMBER}" --repo "${REPO}" --json body,baseRefName,headRefName > "${TMP_DIR}/pr.json"
jq -r '.body // ""' < "${TMP_DIR}/pr.json" > "${TMP_DIR}/body.md"
BASE="$(jq -r '.baseRefName' < "${TMP_DIR}/pr.json")"
HEAD="$(jq -r '.headRefName' < "${TMP_DIR}/pr.json")"
RANGE="origin/${BASE}..origin/${HEAD}"
git fetch origin "${BASE}" "${HEAD}"
load_commit_summary "${RANGE}" "${MAX_COMMITS}"
{
echo "${SECTION_START}"
echo "### Current commits in this promotion (${COMMIT_COUNT})"
echo
echo "**Current base:** \`${BASE}\`"
echo "**Current head:** \`${HEAD}\`"
echo "**Current range:** \`${RANGE}\`"
echo
echo "${COMMIT_MD}"
echo
echo "*Auto-updated by staging promotion metadata workflow*"
echo "${SECTION_END}"
} > "${TMP_DIR}/section.md"
replace_marked_section \
"${TMP_DIR}/body.md" \
"${TMP_DIR}/section.md" \
"${SECTION_START}" \
"${SECTION_END}" \
"${TMP_DIR}/new-body.md"
if [ "${DRY_RUN}" = "true" ]; then
echo "Dry run enabled. Computed PR body for #${PR_NUMBER}:"
cat "${TMP_DIR}/new-body.md"
else
gh pr edit "${PR_NUMBER}" --repo "${REPO}" --body-file "${TMP_DIR}/new-body.md"
fi
+2 -2
View File
@@ -48,11 +48,11 @@ jobs:
matrix:
include:
- 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
files: "tests/e2e/scenarios/test_skills.py tests/e2e/scenarios/test_tool_approval.py"
- group: extensions
files: "tests/e2e/scenarios/test_extensions.py tests/e2e/scenarios/test_extension_oauth.py tests/e2e/scenarios/test_wasm_lifecycle.py tests/e2e/scenarios/test_tool_execution.py tests/e2e/scenarios/test_pairing.py"
files: "tests/e2e/scenarios/test_extensions.py"
steps:
- uses: actions/checkout@v6
+5 -12
View File
@@ -13,11 +13,6 @@ jobs:
with:
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
env:
PR_TITLE: ${{ github.event.pull_request.title }}
@@ -26,8 +21,6 @@ jobs:
set -euo pipefail
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 ---
IS_FIX=false
@@ -37,7 +30,7 @@ jobs:
fi
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
IS_FIX=true
fi
@@ -56,14 +49,14 @@ jobs:
exit 0
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
echo "[skip-regression-check] found in commit message — skipping."
exit 0
fi
# --- 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
echo "No changed files — skipping."
@@ -87,13 +80,13 @@ jobs:
# --- 4. Look for test changes ---
# 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."
exit 0
fi
# 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 }
/^ .*#\[test\]/ || /^ .*#\[tokio::test\]/ || /^ .*#\[cfg\(test\)\]/ || /^ .*mod tests/ { has_test=1 }
/^\+.*#\[test\]/ || /^\+.*#\[tokio::test\]/ || /^\+.*#\[cfg\(test\)\]/ || /^\+.*mod tests/ { has_test=1 }
@@ -1,44 +0,0 @@
name: Release-plz Batch Summary
on:
workflow_dispatch:
inputs:
pr_number:
description: "release-plz PR number to refresh"
required: true
type: string
dry_run:
description: "Compute the body update without editing the PR"
required: false
type: boolean
default: true
pull_request_target:
types: [opened, synchronize, reopened]
permissions:
contents: read
pull-requests: write
jobs:
update-release-pr:
if: >
(github.event_name == 'pull_request_target' &&
github.event.pull_request.head.repo.full_name == github.repository &&
startsWith(github.event.pull_request.head.ref, 'release-plz-')) ||
github.event_name == 'workflow_dispatch'
runs-on: ubuntu-latest
steps:
- name: Checkout base branch
uses: actions/checkout@v6
with:
ref: ${{ github.event_name == 'workflow_dispatch' && 'main' || github.event.pull_request.base.ref }}
fetch-depth: 0
fetch-tags: true
- name: Update release-plz PR body with staging batch summary
env:
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
PR_NUMBER: ${{ github.event_name == 'workflow_dispatch' && inputs.pr_number || github.event.pull_request.number }}
REPO: ${{ github.repository }}
DRY_RUN: ${{ github.event_name == 'workflow_dispatch' && inputs.dry_run || 'false' }}
run: bash .github/scripts/update-release-plz-body.sh
+1 -7
View File
@@ -58,16 +58,10 @@ jobs:
- *checkout
- *install-rust
- uses: Swatinem/rust-cache@v2
- name: Generate GitHub token
uses: actions/create-github-app-token@v2
id: generate-token
with:
app-id: ${{ secrets.GH_RELEASES_MANAGER_APP_ID }}
private-key: ${{ secrets.GH_RELEASES_MANAGER_APP_PRIVATE_KEY }}
- name: Run release-plz
uses: release-plz/[email protected]
with:
command: release-pr
env:
GITHUB_TOKEN: ${{ steps.generate-token.outputs.token }}
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
CARGO_REGISTRY_TOKEN: ${{ secrets.CARGO_REGISTRY_TOKEN }}
+62 -113
View File
@@ -25,35 +25,9 @@ concurrency:
cancel-in-progress: false # Let running suites finish
jobs:
# ── Resolve promotion base branch ───────────────────────────────
resolve-promotion-base:
name: Resolve promotion base
runs-on: ubuntu-latest
outputs:
promotion_base: ${{ steps.resolve.outputs.promotion_base }}
steps:
- name: Resolve promotion base
id: resolve
env:
GH_TOKEN: ${{ github.token }}
FALLBACK_BRANCH: main
REPO: ${{ github.repository }}
run: |
LATEST=$(gh pr list --repo "${REPO}" --label staging-promotion --state open \
--json headRefName,createdAt \
--jq '[.[] | select(.headRefName | startswith("staging-promote/"))] | sort_by(.createdAt) | last | .headRefName // empty')
if [ -n "$LATEST" ]; then
echo "promotion_base=${LATEST}" >> "$GITHUB_OUTPUT"
echo "Using open promotion branch as base: ${LATEST}"
else
echo "promotion_base=${FALLBACK_BRANCH}" >> "$GITHUB_OUTPUT"
echo "No open promotion branch found. Using ${FALLBACK_BRANCH}."
fi
# ── Check for new commits ──────────────────────────────────────
check-changes:
name: Check for new commits
needs: resolve-promotion-base
runs-on: ubuntu-latest
outputs:
has_changes: ${{ steps.check.outputs.has_changes }}
@@ -70,7 +44,6 @@ jobs:
id: check
env:
FORCE_RUN: ${{ inputs.force }}
PROMOTION_BASE: ${{ needs.resolve-promotion-base.outputs.promotion_base }}
run: |
CURRENT_HEAD=$(git rev-parse HEAD)
echo "current_head=${CURRENT_HEAD}" >> "$GITHUB_OUTPUT"
@@ -92,9 +65,9 @@ jobs:
echo "Found ${COMMIT_COUNT} new commit(s) since last tested"
DIFF_RANGE="${LAST_TESTED}..${CURRENT_HEAD}"
else
git fetch origin "${PROMOTION_BASE}"
MERGE_BASE=$(git merge-base "origin/${PROMOTION_BASE}" HEAD)
echo "First run -- reviewing from merge-base ${MERGE_BASE} against ${PROMOTION_BASE}"
git fetch origin main
MERGE_BASE=$(git merge-base origin/main HEAD)
echo "First run -- reviewing from merge-base ${MERGE_BASE}"
DIFF_RANGE="${MERGE_BASE}..${CURRENT_HEAD}"
fi
fi
@@ -128,7 +101,7 @@ jobs:
# ── Create promotion PR (triggers claude-review.yml on the PR) ──
create-promotion-pr:
name: Create Promotion PR
needs: [resolve-promotion-base, check-changes]
needs: check-changes
if: needs.check-changes.outputs.has_changes == 'true'
runs-on: ubuntu-latest
outputs:
@@ -156,19 +129,18 @@ jobs:
echo "token=${{ github.token }}" >> "$GITHUB_OUTPUT"
fi
- name: Check if staging is ahead of target branch
- name: Check if staging is ahead of main
id: ahead-check
env:
GH_TOKEN: ${{ steps.token.outputs.token }}
PROMOTION_BASE: ${{ needs.resolve-promotion-base.outputs.promotion_base }}
run: |
git fetch origin "${PROMOTION_BASE}"
AHEAD=$(git rev-list --count "origin/${PROMOTION_BASE}..origin/staging")
git fetch origin main
AHEAD=$(git rev-list --count origin/main..origin/staging)
echo "commits_ahead=${AHEAD}" >> "$GITHUB_OUTPUT"
if [ "$AHEAD" -eq 0 ]; then
echo "Staging is not ahead of ${PROMOTION_BASE}. Nothing to promote."
echo "Staging is not ahead of main. Nothing to promote."
else
echo "Staging is ${AHEAD} commits ahead of ${PROMOTION_BASE}."
echo "Staging is ${AHEAD} commits ahead of main."
fi
- name: Create promotion branch
@@ -182,53 +154,53 @@ jobs:
echo "branch=${BRANCH}" >> "$GITHUB_OUTPUT"
echo "Created promotion branch: ${BRANCH}"
- name: Find base branch
id: find-base
if: steps.ahead-check.outputs.commits_ahead != '0'
env:
GH_TOKEN: ${{ steps.token.outputs.token }}
run: |
# Find the newest open promotion PR with a staging-promote/* head branch
LATEST=$(gh pr list --label staging-promotion --state open \
--json headRefName,createdAt \
--jq '[.[] | select(.headRefName | startswith("staging-promote/"))] | sort_by(.createdAt) | last | .headRefName // empty')
if [ -n "$LATEST" ]; then
echo "base=${LATEST}" >> "$GITHUB_OUTPUT"
echo "Chaining onto existing promotion branch: ${LATEST}"
else
echo "base=main" >> "$GITHUB_OUTPUT"
echo "No existing promotion PR — targeting main"
fi
- name: Create promotion PR
id: create-pr
if: steps.ahead-check.outputs.commits_ahead != '0'
env:
GH_TOKEN: ${{ steps.token.outputs.token }}
run: |
source .github/scripts/pr-body-utils.sh
RANGE="${{ needs.check-changes.outputs.diff_range }}"
TIMESTAMP=$(date -u +"%Y-%m-%d %H:%M UTC")
BRANCH="${{ steps.branch.outputs.branch }}"
BASE="${{ needs.resolve-promotion-base.outputs.promotion_base }}"
MAX_COMMITS=50
load_commit_summary "${RANGE}" "${MAX_COMMITS}"
# Build PR body via concatenation to avoid heredoc shell expansion
# (commit messages in COMMIT_MD may contain $, backticks, or backslashes)
PR_BODY="## Auto-promotion from staging CI"
PR_BODY+=$'\n\n'"**Batch range:** \`${RANGE}\`"
PR_BODY+=$'\n'"**Promotion branch:** \`${BRANCH}\`"
PR_BODY+=$'\n'"**Base:** \`${BASE}\`"
PR_BODY+=$'\n'"**Triggered by:** Staging CI batch at ${TIMESTAMP}"
PR_BODY+=$'\n\n'"### Commits in this batch (${COMMIT_COUNT}):"
PR_BODY+=$'\n'"${COMMIT_MD}"
PR_BODY+=$'\n\n'"<!-- staging-ci-current:start -->"
PR_BODY+=$'\n'"### Current commits in this promotion (${COMMIT_COUNT})"
PR_BODY+=$'\n'
PR_BODY+=$'\n'"**Current base:** \`${BASE}\`"
PR_BODY+=$'\n'"**Current head:** \`${BRANCH}\`"
PR_BODY+=$'\n'"**Current range:** \`origin/${BASE}..origin/${BRANCH}\`"
PR_BODY+=$'\n'
PR_BODY+=$'\n'"${COMMIT_MD}"
PR_BODY+=$'\n'
PR_BODY+=$'\n'"*Auto-updated by staging promotion metadata workflow*"
PR_BODY+=$'\n'"<!-- staging-ci-current:end -->"
PR_BODY+=$'\n\n'"Waiting for gates:"
PR_BODY+=$'\n'"- Tests: pending"
PR_BODY+=$'\n'"- E2E: pending"
PR_BODY+=$'\n'"- Claude Code review: pending (will post comments on this PR)"
PR_BODY+=$'\n\n'"---"
PR_BODY+=$'\n'"*Auto-created by staging-ci workflow*"
BASE="${{ steps.find-base.outputs.base }}"
PR_URL=$(gh pr create \
--base "$BASE" \
--head "$BRANCH" \
--title "chore: promote staging to ${BASE} (${TIMESTAMP})" \
--body "$PR_BODY" \
--title "chore: promote staging to main (${TIMESTAMP})" \
--body "## Auto-promotion from staging CI
**Batch range:** \`${RANGE}\`
**Promotion branch:** \`${BRANCH}\`
**Base:** \`${BASE}\`
**Triggered by:** Staging CI batch at ${TIMESTAMP}
Waiting for gates:
- Tests: pending
- E2E: pending
- Claude Code review: pending (will post comments on this PR)
---
*Auto-created by staging-ci workflow*" \
--label "staging-promotion")
PR_NUM=$(echo "$PR_URL" | grep -oE '[0-9]+$')
@@ -253,8 +225,7 @@ jobs:
- uses: actions/checkout@v6
with:
ref: staging
# Need full history to recompute the final promoted range before merge.
fetch-depth: 0
fetch-depth: 1
- name: Generate GitHub App token
id: app-token
@@ -353,10 +324,8 @@ jobs:
# Use process substitution so variables propagate to parent shell
while read -r line; do
TAG=$(echo "$line" | grep -oE '^\[(CRITICAL|HIGH|MEDIUM|LOW):[0-9]+\]')
SEVERITY="${TAG#\[}"
SEVERITY="${SEVERITY%%:*}"
CONFIDENCE="${TAG##*:}"
CONFIDENCE="${CONFIDENCE%\]}"
SEVERITY=$(echo "$TAG" | sed 's/\[\(.*\):\(.*\)\]/\1/')
CONFIDENCE=$(echo "$TAG" | sed 's/\[\(.*\):\(.*\)\]/\2/')
DESC=$(echo "$line" | sed "s/\[${SEVERITY}:${CONFIDENCE}\] *//" | head -1)
echo "Found: [${SEVERITY}:${CONFIDENCE}] ${DESC}"
@@ -448,29 +417,11 @@ jobs:
GH_TOKEN: ${{ steps.token.outputs.token }}
PR_NUMBER: ${{ needs.create-promotion-pr.outputs.pr_number }}
run: |
source .github/scripts/pr-body-utils.sh
if [ -n "$PR_NUMBER" ]; then
BASE=$(gh pr view "$PR_NUMBER" --json baseRefName --jq '.baseRefName')
if [ "$BASE" = "main" ]; then
echo "Merging promotion PR #${PR_NUMBER} (targets main)"
TITLE=$(gh pr view "$PR_NUMBER" --json title --jq '.title')
HEAD_BRANCH=$(gh pr view "$PR_NUMBER" --json headRefName --jq '.headRefName')
git fetch origin "${BASE}" "${HEAD_BRANCH}"
CURRENT_RANGE="origin/${BASE}..origin/${HEAD_BRANCH}"
MAX_COMMITS=50
load_commit_summary "${CURRENT_RANGE}" "${MAX_COMMITS}"
{
echo "staging-promotion-summary-v1"
echo "promotion-pr: #${PR_NUMBER}"
echo "base: ${BASE}"
echo "head: ${HEAD_BRANCH}"
echo "current-range: ${CURRENT_RANGE}"
echo "current-commit-count: ${COMMIT_COUNT}"
echo ""
echo "Current commits in this promotion (${COMMIT_COUNT}):"
echo "${COMMIT_MD}"
} > /tmp/staging-promotion-merge-body.md
gh pr merge "$PR_NUMBER" --merge --subject "#${PR_NUMBER} $TITLE" --body-file /tmp/staging-promotion-merge-body.md
gh pr merge "$PR_NUMBER" --merge
echo "merged=true" >> "$GITHUB_OUTPUT"
else
echo "PR #${PR_NUMBER} targets '${BASE}' (not main) — leaving open for chain resolution"
@@ -510,20 +461,18 @@ jobs:
steps:
- name: Summary
run: |
{
echo "## Staging CI Batch Results"
echo ""
echo "| Check | Result |"
echo "|-------|--------|"
echo "| Tests | ${{ needs.tests.result }} |"
echo "| E2E | ${{ needs.e2e.result }} |"
echo "| Promotion PR | ${{ needs.create-promotion-pr.result }} |"
echo "| Gate | ${{ needs.gate.result }} |"
echo "| Tag Updated | ${{ needs.update-tag.result }} |"
echo ""
echo "Range: ${{ needs.check-changes.outputs.diff_range }}"
PR_NUM="${{ needs.create-promotion-pr.outputs.pr_number }}"
if [ -n "$PR_NUM" ]; then
echo "Promotion PR: #${PR_NUM}"
fi
} >> "$GITHUB_STEP_SUMMARY"
echo "## Staging CI Batch Results" >> "$GITHUB_STEP_SUMMARY"
echo "" >> "$GITHUB_STEP_SUMMARY"
echo "| Check | Result |" >> "$GITHUB_STEP_SUMMARY"
echo "|-------|--------|" >> "$GITHUB_STEP_SUMMARY"
echo "| Tests | ${{ needs.tests.result }} |" >> "$GITHUB_STEP_SUMMARY"
echo "| E2E | ${{ needs.e2e.result }} |" >> "$GITHUB_STEP_SUMMARY"
echo "| Promotion PR | ${{ needs.create-promotion-pr.result }} |" >> "$GITHUB_STEP_SUMMARY"
echo "| Gate | ${{ needs.gate.result }} |" >> "$GITHUB_STEP_SUMMARY"
echo "| Tag Updated | ${{ needs.update-tag.result }} |" >> "$GITHUB_STEP_SUMMARY"
echo "" >> "$GITHUB_STEP_SUMMARY"
echo "Range: ${{ needs.check-changes.outputs.diff_range }}" >> "$GITHUB_STEP_SUMMARY"
PR_NUM="${{ needs.create-promotion-pr.outputs.pr_number }}"
if [ -n "$PR_NUM" ]; then
echo "Promotion PR: #${PR_NUM}" >> "$GITHUB_STEP_SUMMARY"
fi
@@ -1,78 +0,0 @@
name: Staging Promotion Metadata
on:
workflow_dispatch:
inputs:
pr_number:
description: "Staging promotion PR number to refresh"
required: true
type: string
dry_run:
description: "Compute the body update without editing the PR"
required: false
type: boolean
default: true
pull_request_target:
types: [opened, synchronize, reopened]
push:
branches:
- main
permissions:
contents: read
pull-requests: write
jobs:
refresh-single-pr:
if: >
(github.event_name == 'pull_request_target' &&
github.event.pull_request.head.repo.full_name == github.repository &&
startsWith(github.event.pull_request.head.ref, 'staging-promote/')) ||
github.event_name == 'workflow_dispatch'
runs-on: ubuntu-latest
steps:
- name: Checkout workflow source
uses: actions/checkout@v6
with:
# For chained promotion PRs, the script lives on the trusted PR head,
# not necessarily on the older promotion branch used as the PR base.
ref: ${{ github.event_name == 'workflow_dispatch' && 'main' || github.event.pull_request.head.sha }}
fetch-depth: 0
fetch-tags: true
- name: Refresh staging promotion PR body
env:
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
PR_NUMBER: ${{ github.event_name == 'workflow_dispatch' && inputs.pr_number || github.event.pull_request.number }}
REPO: ${{ github.repository }}
DRY_RUN: ${{ github.event_name == 'workflow_dispatch' && inputs.dry_run || 'false' }}
run: bash .github/scripts/update-staging-promotion-body.sh
refresh-open-prs-after-main-push:
if: github.event_name == 'push'
runs-on: ubuntu-latest
steps:
- name: Checkout main
uses: actions/checkout@v6
with:
ref: main
fetch-depth: 0
fetch-tags: true
- name: Refresh all open staging promotion PR bodies
env:
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
REPO: ${{ github.repository }}
run: |
# ubuntu-latest uses bash 5.x, so mapfile is available here.
mapfile -t prs < <(gh pr list --repo "${REPO}" --label staging-promotion --state open \
--json number,headRefName \
--jq '.[] | select(.headRefName | startswith("staging-promote/")) | .number')
if [ "${#prs[@]}" -eq 0 ]; then
echo "No open staging promotion PRs to refresh."
exit 0
fi
for pr in "${prs[@]}"; do
echo "Refreshing staging promotion PR #${pr}"
PR_NUMBER="${pr}" bash .github/scripts/update-staging-promotion-body.sh
done
-9
View File
@@ -7,15 +7,6 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
## [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
### Added
Generated
+1 -1
View File
@@ -3350,7 +3350,7 @@ dependencies = [
[[package]]
name = "ironclaw"
version = "0.18.0"
version = "0.17.0"
dependencies = [
"aes-gcm",
"aho-corasick",
+1 -1
View File
@@ -20,7 +20,7 @@ exclude = [
[package]
name = "ironclaw"
version = "0.18.0"
version = "0.17.0"
edition = "2024"
rust-version = "1.92"
description = "Secure personal AI assistant that protects your data and expands its capabilities on the fly"
-1
View File
@@ -19,7 +19,6 @@ WORKDIR /app
# Copy manifests first for layer caching
COPY Cargo.toml Cargo.lock ./
COPY crates/ crates/
# Copy source, build script, tests, and supporting directories
COPY build.rs build.rs
-1
View File
@@ -20,7 +20,6 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
WORKDIR /app
COPY Cargo.toml Cargo.lock ./
COPY crates/ crates/
COPY build.rs build.rs
COPY src/ src/
COPY tests/ tests/
+1 -206
View File
@@ -20,162 +20,33 @@ version = "1.0.102"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7f202df86484c868dbad7eaa557ef785d5c66295e41b460ef922eca0723b842c"
[[package]]
name = "base64ct"
version = "1.8.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2af50177e190e07a26ab74f8b1efbfe2ef87da2116221318cb1c2e82baf7de06"
[[package]]
name = "bitflags"
version = "2.11.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
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]]
name = "cfg-if"
version = "1.0.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
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]]
name = "discord-channel"
version = "0.2.0"
version = "0.1.0"
dependencies = [
"ed25519-dalek",
"hex",
"serde",
"serde_json",
"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]]
name = "equivalent"
version = "1.0.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
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]]
name = "hashbrown"
version = "0.14.5"
@@ -197,12 +68,6 @@ version = "0.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea"
[[package]]
name = "hex"
version = "0.4.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70"
[[package]]
name = "id-arena"
version = "2.3.0"
@@ -233,12 +98,6 @@ version = "0.2.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "884e2677b40cc8c339eaefcb701c32ef1fd2493d71118dc0ca4b6a736c93bd67"
[[package]]
name = "libc"
version = "0.2.182"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6800badb6cb2082ffd7b6a67e6125bb39f18782f793520caee8cb8846be06112"
[[package]]
name = "log"
version = "0.4.29"
@@ -257,16 +116,6 @@ version = "1.21.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
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]]
name = "prettyplease"
version = "0.2.37"
@@ -295,15 +144,6 @@ dependencies = [
"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]]
name = "semver"
version = "1.0.27"
@@ -353,23 +193,6 @@ dependencies = [
"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]]
name = "smallvec"
version = "1.15.1"
@@ -385,22 +208,6 @@ dependencies = [
"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]]
name = "syn"
version = "2.0.117"
@@ -412,12 +219,6 @@ dependencies = [
"unicode-ident",
]
[[package]]
name = "typenum"
version = "1.19.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "562d481066bde0658276a35467c4af00bdc6ee726305698a55b86e61d7ad82bb"
[[package]]
name = "unicode-ident"
version = "1.0.24"
@@ -593,12 +394,6 @@ dependencies = [
"syn",
]
[[package]]
name = "zeroize"
version = "1.8.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b97154e67e32c85465826e8bcc1c59429aaaf107c1e4a9e53c8d8ccd5eff88d0"
[[package]]
name = "zmij"
version = "1.0.21"
-2
View File
@@ -10,8 +10,6 @@ publish = false
serde = { version = "1.0", features = ["derive"] }
serde_json = "1.0"
wit-bindgen = "0.36"
ed25519-dalek = { version = "2", default-features = false, features = ["alloc", "fast", "zeroize"] }
hex = "0.4"
[lib]
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
```
**Note:** The `discord_bot_token` secret is used for Discord REST API calls.
Interaction signature verification is performed inside the Discord channel
module and uses the channel config field `webhook_secret` (set this to your
Discord app public key hex).
**Note:** The `discord_bot_token` secret is the only value read directly by this
Discord channel WASM component. The `discord_app_id` and `discord_public_key`
secrets are used by the IronClaw host (for example, to verify Discord
interaction signatures and manage slash command registration) and are not
accessed from the WASM module itself.
## 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.
## 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
@@ -119,11 +96,8 @@ To send embeds, include an `embeds` array in the `metadata_json` field of the ag
### "Invalid Signature"
- Check that `webhook_secret` is set to your Discord app public key hex in the
Discord channel config.
- 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.
- Check that `discord_public_key` is set correctly in IronClaw secrets.
- This validation happens on the host before reaching the WASM.
### "401 Unauthorized"
@@ -3,7 +3,7 @@
"wit_version": "0.3.0",
"type": "channel",
"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": {
"required_secrets": [
{
@@ -41,7 +41,7 @@
},
"channel": {
"allowed_paths": ["/webhook/discord"],
"allow_polling": true,
"allow_polling": false,
"callback_timeout_secs": 45,
"workspace_prefix": "channels/discord/",
"emit_rate_limit": {
@@ -55,12 +55,8 @@
},
"config": {
"require_signature_verification": true,
"webhook_secret": null,
"polling_enabled": false,
"poll_interval_ms": 30000,
"mention_channel_ids": [],
"owner_id": null,
"dm_policy": "pairing",
"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.
# 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
# NEAR AI Cloud (API key auth, Chat Completions API)
+5 -9
View File
@@ -5,17 +5,13 @@ Requires=cloud-sql-proxy.service
[Service]
Type=simple
EnvironmentFile=/opt/ironclaw/.env
# Pin to a specific version tag or digest instead of :latest to prevent
# 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 \
ExecStartPre=/usr/bin/docker pull us-central1-docker.pkg.dev/ironclaw-prod/ironclaw/agent:latest
ExecStart=/usr/bin/docker run --rm \
--name ironclaw \
--env-file /opt/ironclaw/.env \
-p 3000:3000 \
us-central1-docker.pkg.dev/ironclaw-prod/ironclaw/agent:${IRONCLAW_VERSION:-latest} \
--no-onboard'
--network=host \
us-central1-docker.pkg.dev/ironclaw-prod/ironclaw/agent:latest \
--no-onboard
ExecStop=/usr/bin/docker stop ironclaw
Restart=always
RestartSec=10
+1 -8
View File
@@ -24,15 +24,8 @@ systemctl enable docker
systemctl start docker
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 \
"https://storage.googleapis.com/cloud-sql-connectors/cloud-sql-proxy/${CLOUD_SQL_PROXY_VERSION}/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
}
https://storage.googleapis.com/cloud-sql-connectors/cloud-sql-proxy/v2.14.3/cloud-sql-proxy.linux.amd64
chmod +x /usr/local/bin/cloud-sql-proxy
echo "==> Installing systemd services"
-20
View File
@@ -15,7 +15,6 @@ configurations.
| io.net | `ionet` | `IONET_API_KEY` | Intelligence API |
| Mistral | `mistral` | `MISTRAL_API_KEY` | Mistral 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 |
| Ollama | `ollama` | No | Local inference |
| 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`)
Uses the native AWS Converse API via `aws-sdk-bedrockruntime`. Supports standard AWS
-21
View File
@@ -382,27 +382,6 @@
"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",
"aliases": [
+3 -3
View File
@@ -2,7 +2,7 @@
"name": "discord",
"display_name": "Discord Channel",
"kind": "channel",
"version": "0.2.1",
"version": "0.2.0",
"wit_version": "0.3.0",
"description": "Talk to your agent in Discord",
"keywords": [
@@ -18,8 +18,8 @@
},
"artifacts": {
"wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/discord-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "efa1b9019fa33e243f8db1e1fcc732731d45836336bdd26ca19b6fe227ca8b69"
"url": "https://github.com/nearai/ironclaw/releases/latest/download/discord-wasm32-wasip2.tar.gz",
"sha256": null
}
},
"auth_summary": {
+2 -2
View File
@@ -18,8 +18,8 @@
},
"artifacts": {
"wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/slack-0.2.1-wasm32-wasip2.tar.gz",
"sha256": "d4667e35126986509d862bc3a0088777305d8f41c75de83c1e223b42312ede48"
"url": "https://github.com/nearai/ironclaw/releases/latest/download/slack-wasm32-wasip2.tar.gz",
"sha256": null
}
},
"auth_summary": {
+3 -3
View File
@@ -2,7 +2,7 @@
"name": "telegram",
"display_name": "Telegram Channel",
"kind": "channel",
"version": "0.2.3",
"version": "0.2.2",
"wit_version": "0.3.0",
"description": "Talk to your agent through a Telegram bot",
"keywords": [
@@ -18,8 +18,8 @@
},
"artifacts": {
"wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/telegram-0.2.3-wasm32-wasip2.tar.gz",
"sha256": "b9a83d5a2d1285ce0ec116b354336a1f245f893291ccb01dffbcaccf89d72aed"
"url": "https://github.com/nearai/ironclaw/releases/latest/download/telegram-wasm32-wasip2.tar.gz",
"sha256": null
}
},
"auth_summary": {
+2 -2
View File
@@ -18,8 +18,8 @@
},
"artifacts": {
"wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/whatsapp-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "feb9194719d9bed796b070ab4dc30348dbfb5d3dec56f9f21e02d14137abab01"
"url": "https://github.com/nearai/ironclaw/releases/latest/download/whatsapp-wasm32-wasip2.tar.gz",
"sha256": null
}
},
"auth_summary": {
+2 -2
View File
@@ -19,8 +19,8 @@
},
"artifacts": {
"wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/github-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "da9fac56b6f20197a415489bbaec9fefb085a5cf6324cab79ea48a47eb19c13b"
"url": "https://github.com/nearai/ironclaw/releases/latest/download/github-wasm32-wasip2.tar.gz",
"sha256": null
}
},
"auth_summary": {
+2 -2
View File
@@ -18,8 +18,8 @@
},
"artifacts": {
"wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/gmail-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "ee9574e02e92bc1d481f1310eb88afd99ee52bf6971074ab33bd76bf99b34b1d"
"url": "https://github.com/nearai/ironclaw/releases/latest/download/gmail-wasm32-wasip2.tar.gz",
"sha256": null
}
},
"auth_summary": {
+2 -2
View File
@@ -18,8 +18,8 @@
},
"artifacts": {
"wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-calendar-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "2fa47150ea222e787c122182ad6f4dfa2ffaf5fe490d05e8de887a76445f8d2d"
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-calendar-wasm32-wasip2.tar.gz",
"sha256": null
}
},
"auth_summary": {
+2 -2
View File
@@ -18,8 +18,8 @@
},
"artifacts": {
"wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-docs-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "40e134a1c1564f832ca861c3396895d4e33ec67b99313fc1f97baf8d971423a9"
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-docs-wasm32-wasip2.tar.gz",
"sha256": null
}
},
"auth_summary": {
+2 -2
View File
@@ -18,8 +18,8 @@
},
"artifacts": {
"wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-drive-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "002a341a1d58125563a7c69561b26fbc2629b04ea723cade744102bdc0fbb71f"
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-drive-wasm32-wasip2.tar.gz",
"sha256": null
}
},
"auth_summary": {
+2 -2
View File
@@ -18,8 +18,8 @@
},
"artifacts": {
"wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-sheets-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "8aa2c9d52f033edea3a6c2311b0ec694ccb6d0a54ef07e94d72bf8be1ce8009a"
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-sheets-wasm32-wasip2.tar.gz",
"sha256": null
}
},
"auth_summary": {
+2 -2
View File
@@ -17,8 +17,8 @@
},
"artifacts": {
"wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-slides-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "e931a97d4fd0b0b938e464dc7c7f2be6ea6b4d1508f5ea3cd931d44db23f05f5"
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-slides-wasm32-wasip2.tar.gz",
"sha256": null
}
},
"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": {
"wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/slack-0.2.1-wasm32-wasip2.tar.gz",
"sha256": "d4667e35126986509d862bc3a0088777305d8f41c75de83c1e223b42312ede48"
"url": "https://github.com/nearai/ironclaw/releases/latest/download/slack-tool-wasm32-wasip2.tar.gz",
"sha256": null
}
},
"auth_summary": {
+2 -2
View File
@@ -18,8 +18,8 @@
},
"artifacts": {
"wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/telegram-0.2.2-wasm32-wasip2.tar.gz",
"sha256": "b9a83d5a2d1285ce0ec116b354336a1f245f893291ccb01dffbcaccf89d72aed"
"url": "https://github.com/nearai/ironclaw/releases/latest/download/telegram-mtproto-wasm32-wasip2.tar.gz",
"sha256": null
}
},
"auth_summary": {
+3 -3
View File
@@ -2,7 +2,7 @@
"name": "web-search",
"display_name": "Web Search",
"kind": "tool",
"version": "0.2.1",
"version": "0.2.0",
"wit_version": "0.3.0",
"description": "Search the web using Brave Search API",
"keywords": [
@@ -18,8 +18,8 @@
},
"artifacts": {
"wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/web-search-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "56834573c54ea2a33cea1eb0f04bbdf59f1ef8d8702995cf431b0921302eeccc"
"url": "https://github.com/nearai/ironclaw/releases/latest/download/web-search-wasm32-wasip2.tar.gz",
"sha256": null
}
},
"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::submission::{Submission, SubmissionParser, SubmissionResult};
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::context::ContextManager;
use crate::db::Database;
@@ -936,10 +936,29 @@ impl Agent {
SubmissionResult::Ok { message } => Ok(message),
SubmissionResult::Error { message } => Ok(Some(format!("Error: {}", message))),
SubmissionResult::Interrupted => Ok(Some("Interrupted.".into())),
SubmissionResult::NeedApproval { .. } => {
// ApprovalNeeded status was already sent by thread_ops.rs before
// returning this result. Empty string signals the caller to skip
// respond() (no duplicate text).
SubmissionResult::NeedApproval {
request_id,
tool_name,
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()))
}
}
-72
View File
@@ -554,31 +554,6 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
};
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));
break;
}
@@ -2260,51 +2235,4 @@ mod tests {
"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;
pub mod undo;
pub use crate::worker::{Worker, WorkerDeps};
pub(crate) use agent_loop::truncate_for_preview;
pub use agent_loop::{Agent, AgentDeps};
pub use compaction::{CompactionResult, ContextCompactor};
+3 -116
View File
@@ -207,7 +207,7 @@ impl Trigger {
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum RoutineAction {
/// Single LLM call (optionally with tools). Cheap and fast.
/// Single LLM call, no tools. Cheap and fast.
Lightweight {
/// The prompt sent to the LLM.
prompt: String,
@@ -217,14 +217,6 @@ pub enum RoutineAction {
/// Max output tokens (default: 4096).
#[serde(default = "default_max_tokens")]
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.
FullJob {
@@ -251,19 +243,6 @@ fn default_max_iterations() -> u32 {
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>`.
pub fn parse_tool_permissions(value: &serde_json::Value) -> Vec<String> {
value
@@ -311,22 +290,10 @@ impl RoutineAction {
.get("max_tokens")
.and_then(|v| v.as_u64())
.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 {
prompt,
context_paths,
max_tokens,
use_tools,
max_tool_rounds,
})
}
"full_job" => {
@@ -372,14 +339,10 @@ impl RoutineAction {
prompt,
context_paths,
max_tokens,
use_tools,
max_tool_rounds,
} => serde_json::json!({
"prompt": prompt,
"context_paths": context_paths,
"max_tokens": max_tokens,
"use_tools": use_tools,
"max_tool_rounds": max_tool_rounds,
}),
RoutineAction::FullJob {
title,
@@ -541,8 +504,7 @@ pub fn next_cron_fire(
#[cfg(test)]
mod tests {
use crate::agent::routine::{
MAX_TOOL_ROUNDS_LIMIT, RoutineAction, RoutineGuardrails, RunStatus, Trigger, content_hash,
next_cron_fire,
RoutineAction, RoutineGuardrails, RunStatus, Trigger, content_hash, next_cron_fire,
};
#[test]
@@ -592,13 +554,11 @@ mod tests {
prompt: "Check PRs".to_string(),
context_paths: vec!["context/priorities.md".to_string()],
max_tokens: 2048,
use_tools: false,
max_tool_rounds: 3,
};
let json = action.to_config_json();
let parsed = RoutineAction::from_db("lightweight", json).expect("parse lightweight");
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)
);
}
@@ -735,77 +695,4 @@ mod tests {
);
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"),
}
}
}
+25 -95
View File
@@ -459,20 +459,7 @@ async fn execute_routine(ctx: EngineContext, routine: Routine, run: RoutineRun)
prompt,
context_paths,
max_tokens,
use_tools,
max_tool_rounds,
} => {
execute_lightweight(
&ctx,
&routine,
prompt,
context_paths,
*max_tokens,
*use_tools,
*max_tool_rounds,
)
.await
}
} => execute_lightweight(&ctx, &routine, prompt, context_paths, *max_tokens).await,
RoutineAction::FullJob {
title,
description,
@@ -683,8 +670,6 @@ async fn execute_lightweight(
prompt: &str,
context_paths: &[String],
max_tokens: u32,
use_tools: bool,
max_tool_rounds: u32,
) -> Result<(RunStatus, Option<String>, Option<i32>), RoutineError> {
// Load context from workspace
let mut context_parts = Vec::new();
@@ -747,15 +732,14 @@ async fn execute_lightweight(
Err(_) => max_tokens,
};
// If tools are enabled (both globally and per-routine), use the tool execution loop
if use_tools && ctx.config.lightweight_tools_enabled {
// If tools are enabled, use the tool execution loop; otherwise, single LLM call
if ctx.config.lightweight_tools_enabled {
execute_lightweight_with_tools(
ctx,
routine,
&system_prompt,
&full_prompt,
effective_max_tokens,
max_tool_rounds,
)
.await
} else {
@@ -799,12 +783,24 @@ async fn execute_lightweight_no_tools(
reason: e.to_string(),
})?;
handle_text_response(
&response.content,
response.finish_reason,
response.input_tokens,
response.output_tokens,
)
let content = response.content.trim();
let tokens_used = Some((response.input_tokens + response.output_tokens) as i32);
// Empty content guard
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.
@@ -854,7 +850,6 @@ async fn execute_lightweight_with_tools(
system_prompt: &str,
full_prompt: &str,
effective_max_tokens: u32,
max_tool_rounds: u32,
) -> Result<(RunStatus, Option<String>, Option<i32>), RoutineError> {
let mut messages = if system_prompt.is_empty() {
vec![ChatMessage::user(full_prompt)]
@@ -865,9 +860,7 @@ async fn execute_lightweight_with_tools(
]
};
let max_iterations = max_tool_rounds
.min(ctx.config.lightweight_max_iterations)
.min(5);
let max_iterations = ctx.config.lightweight_max_iterations.min(5);
let mut iteration = 0;
let mut total_input_tokens = 0;
let mut total_output_tokens = 0;
@@ -913,10 +906,7 @@ async fn execute_lightweight_with_tools(
);
} else {
// Tool-enabled iteration
let tool_defs = ctx
.tools
.tool_definitions_excluding(ROUTINE_TOOL_DENYLIST)
.await;
let tool_defs = ctx.tools.tool_definitions().await;
let request = ToolCompletionRequest::new(messages.clone(), tool_defs)
.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.
async fn execute_routine_tool(
ctx: &EngineContext,
job_ctx: &JobContext,
tc: &ToolCall,
) -> 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
let tool = ctx
.tools
@@ -1150,11 +1119,9 @@ pub fn spawn_cron_ticker(
interval: Duration,
) -> tokio::task::JoinHandle<()> {
tokio::spawn(async move {
// Run one check immediately so routines due at startup don't wait
// an extra full polling interval.
engine.check_cron_triggers().await;
let mut ticker = tokio::time::interval(interval);
// Skip immediate first tick
ticker.tick().await;
loop {
ticker.tick().await;
@@ -1316,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]
fn test_empty_response_handling() {
// Simulate the empty content guard logic
@@ -1360,11 +1297,4 @@ mod tests {
assert_eq!(finish_reason_length, crate::llm::FinishReason::Length);
assert_eq!(finish_reason_stop, crate::llm::FinishReason::Stop);
}
#[test]
fn test_truncate_adds_ellipsis_when_over_limit() {
let input = "abcdefghijk";
let out = super::truncate(input, 5);
assert_eq!(out, "abcde...");
}
}
+3 -18
View File
@@ -486,12 +486,7 @@ impl Agent {
.channels
.send_status(
&message.channel,
StatusUpdate::ApprovalNeeded {
request_id: request_id.to_string(),
tool_name: tool_name.clone(),
description: description.clone(),
parameters: parameters.clone(),
},
StatusUpdate::Status("Awaiting approval".into()),
&message.metadata,
)
.await;
@@ -1302,12 +1297,7 @@ impl Agent {
.channels
.send_status(
&message.channel,
StatusUpdate::ApprovalNeeded {
request_id: request_id.to_string(),
tool_name: tool_name.clone(),
description: description.clone(),
parameters: parameters.clone(),
},
StatusUpdate::Status("Awaiting approval".into()),
&message.metadata,
)
.await;
@@ -1378,12 +1368,7 @@ impl Agent {
.channels
.send_status(
&message.channel,
StatusUpdate::ApprovalNeeded {
request_id: request_id.to_string(),
tool_name: tool_name.clone(),
description: description.clone(),
parameters: parameters.clone(),
},
StatusUpdate::Status("Awaiting approval".into()),
&message.metadata,
)
.await;
+1 -75
View File
@@ -9,7 +9,6 @@
use std::sync::Arc;
use crate::agent::SessionManager as AgentSessionManager;
use crate::channels::web::log_layer::LogBroadcaster;
use crate::config::Config;
use crate::context::ContextManager;
@@ -47,8 +46,6 @@ pub struct AppComponents {
pub log_broadcaster: Arc<LogBroadcaster>,
pub context_manager: Arc<ContextManager>,
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_catalog: Option<Arc<SkillCatalog>>,
pub cost_guard: Arc<crate::agent::cost_guard::CostGuard>,
@@ -290,7 +287,6 @@ impl AppBuilder {
Arc::new(ToolRegistry::new())
};
tools.register_builtin_tools();
tools.register_tool_info();
if let Some(ref ss) = self.secrets_store {
tools.register_secrets_tools(Arc::clone(ss));
@@ -304,8 +300,7 @@ impl AppBuilder {
// Register memory tools if database is available
let workspace = if let Some(ref db) = self.db {
let mut ws = Workspace::new_with_db("default", db.clone())
.with_search_config(&self.config.search);
let mut ws = Workspace::new_with_db("default", db.clone());
if let Some(ref emb) = embeddings {
ws = ws.with_embeddings(emb.clone());
}
@@ -694,8 +689,6 @@ impl AppBuilder {
// Create hook registry early so runtime extension activation can register hooks.
let hooks = Arc::new(HookRegistry::new());
let agent_session_manager =
Arc::new(AgentSessionManager::new().with_hooks(Arc::clone(&hooks)));
let (
mcp_session_manager,
@@ -802,7 +795,6 @@ impl AppBuilder {
log_broadcaster: self.log_broadcaster,
context_manager,
hooks,
agent_session_manager,
skill_registry,
skill_catalog,
cost_guard,
@@ -813,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());
}
}
+76 -173
View File
@@ -269,105 +269,95 @@ async fn webhook_handler(
let mut fallback_req = None;
{
let webhook_secret = state.webhook_secret.read().await;
let Some(expected_secret) = webhook_secret.as_ref() else {
return (
StatusCode::UNAUTHORIZED,
Json(WebhookResponse {
message_id: Uuid::nil(),
status: "error".to_string(),
response: Some(
"Webhook authentication required: HTTP webhook secret is not configured."
.to_string(),
),
}),
)
.into_response();
};
let expected_secret = expected_secret.expose_secret();
if let Some(expected_secret) = webhook_secret.as_ref() {
let expected_secret = expected_secret.expose_secret();
match headers.get("x-ironclaw-signature") {
Some(raw_signature) => match raw_signature.to_str() {
Ok(signature) => {
if !verify_hmac_signature(expected_secret, &body, signature) {
return (
StatusCode::UNAUTHORIZED,
Json(WebhookResponse {
message_id: Uuid::nil(),
status: "error".to_string(),
response: Some("Invalid webhook signature".to_string()),
}),
)
.into_response();
match headers.get("x-ironclaw-signature") {
Some(raw_signature) => match raw_signature.to_str() {
Ok(signature) => {
if !verify_hmac_signature(expected_secret, &body, signature) {
return (
StatusCode::UNAUTHORIZED,
Json(WebhookResponse {
message_id: Uuid::nil(),
status: "error".to_string(),
response: Some("Invalid webhook signature".to_string()),
}),
)
.into_response();
}
}
}
Err(_) => {
return (
StatusCode::UNAUTHORIZED,
Json(WebhookResponse {
message_id: Uuid::nil(),
status: "error".to_string(),
response: Some("Invalid signature header encoding".to_string()),
}),
)
.into_response();
}
},
None => {
let req: WebhookRequest = match serde_json::from_slice(&body) {
Ok(req) => req,
Err(_) => {
return (
StatusCode::UNAUTHORIZED,
Json(WebhookResponse {
message_id: Uuid::nil(),
status: "error".to_string(),
response: Some(
"Webhook authentication required. Provide X-IronClaw-Signature header \
(preferred) or 'secret' field in body (deprecated)."
.to_string(),
),
response: Some("Invalid signature header encoding".to_string()),
}),
)
.into_response();
}
};
},
None => {
let req: WebhookRequest = match serde_json::from_slice(&body) {
Ok(req) => req,
Err(_) => {
return (
StatusCode::UNAUTHORIZED,
Json(WebhookResponse {
message_id: Uuid::nil(),
status: "error".to_string(),
response: Some(
"Webhook authentication required. Provide X-IronClaw-Signature header \
(preferred) or 'secret' field in body (deprecated)."
.to_string(),
),
}),
)
.into_response();
}
};
match &req.secret {
Some(provided)
if bool::from(provided.as_bytes().ct_eq(expected_secret.as_bytes())) =>
{
tracing::warn!(
"Webhook authenticated via deprecated 'secret' field in request body. \
Migrate to X-IronClaw-Signature header (HMAC-SHA256). \
Body secret support will be removed in a future release."
);
fallback_req = Some(req);
}
Some(_) => {
return (
StatusCode::UNAUTHORIZED,
Json(WebhookResponse {
message_id: Uuid::nil(),
status: "error".to_string(),
response: Some("Invalid webhook secret".to_string()),
}),
)
.into_response();
}
None => {
return (
StatusCode::UNAUTHORIZED,
Json(WebhookResponse {
message_id: Uuid::nil(),
status: "error".to_string(),
response: Some(
"Webhook authentication required. Provide X-IronClaw-Signature header \
(preferred) or 'secret' field in body (deprecated)."
.to_string(),
),
}),
)
.into_response();
match &req.secret {
Some(provided)
if bool::from(
provided.as_bytes().ct_eq(expected_secret.as_bytes()),
) =>
{
tracing::warn!(
"Webhook authenticated via deprecated 'secret' field in request body. \
Migrate to X-IronClaw-Signature header (HMAC-SHA256). \
Body secret support will be removed in a future release."
);
fallback_req = Some(req);
}
Some(_) => {
return (
StatusCode::UNAUTHORIZED,
Json(WebhookResponse {
message_id: Uuid::nil(),
status: "error".to_string(),
response: Some("Invalid webhook secret".to_string()),
}),
)
.into_response();
}
None => {
return (
StatusCode::UNAUTHORIZED,
Json(WebhookResponse {
message_id: Uuid::nil(),
status: "error".to_string(),
response: Some(
"Webhook authentication required. Provide X-IronClaw-Signature header \
(preferred) or 'secret' field in body (deprecated)."
.to_string(),
),
}),
)
.into_response();
}
}
}
}
@@ -817,67 +807,6 @@ mod tests {
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]
async fn webhook_missing_all_auth_returns_unauthorized() {
let channel = test_channel(Some("correct-secret"));
@@ -1062,32 +991,6 @@ mod tests {
);
}
#[tokio::test]
async fn webhook_rejects_requests_after_secret_is_cleared() {
let secret = "test-secret-123";
let channel = test_channel(Some(secret));
let _stream = channel.start().await.unwrap();
let app = channel.routes();
channel.update_secret(None).await;
let body = serde_json::json!({
"content": "hello"
});
let body_bytes = serde_json::to_vec(&body).unwrap();
let signature = compute_signature(secret, &body_bytes);
let req = Request::builder()
.method("POST")
.uri("/webhook")
.header("content-type", "application/json")
.header("x-ironclaw-signature", signature)
.body(Body::from(body_bytes))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test]
async fn test_concurrent_requests_during_secret_update() {
use std::sync::Arc as StdArc;
+3 -225
View File
@@ -408,120 +408,12 @@ impl Channel for RelayChannel {
Ok(())
}
/// Status updates are not forwarded to messaging providers to avoid noise.
async fn send_status(
&self,
status: StatusUpdate,
metadata: &serde_json::Value,
_status: StatusUpdate,
_metadata: &serde_json::Value,
) -> 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(())
}
@@ -747,118 +639,4 @@ mod tests {
// The reconnect loop now skips team validation when team_id is empty,
// 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/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.
#[derive(Debug, Clone)]
pub struct EmittedMessage {
@@ -268,7 +264,7 @@ impl ChannelHostState {
max = MAX_MESSAGE_CONTENT_SIZE,
"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)");
let msg = EmittedMessage {
content: truncated,
@@ -635,7 +631,6 @@ mod tests {
use crate::channels::wasm::host::{
Attachment, ChannelEmitRateLimiter, ChannelHostState, EmittedMessage,
MAX_ATTACHMENT_TOTAL_SIZE, MAX_ATTACHMENTS_PER_MESSAGE, MAX_EMITS_PER_EXECUTION,
MAX_MESSAGE_CONTENT_SIZE,
};
#[test]
@@ -694,25 +689,6 @@ mod tests {
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]
fn test_workspace_write_prefixing() {
let caps = ChannelCapabilities::for_channel("slack");
+40 -50
View File
@@ -1994,33 +1994,28 @@ impl WasmChannel {
return Ok(());
}
// Clone sender to avoid holding RwLock read guard across send().await in the loop
let tx = {
let tx_guard = self.message_tx.read().await;
let Some(tx) = tx_guard.as_ref() else {
tracing::error!(
channel = %self.name,
count = messages.len(),
"Messages emitted but no sender available - channel may not be started!"
);
return Ok(());
};
tx.clone()
let tx_guard = self.message_tx.read().await;
let Some(tx) = tx_guard.as_ref() else {
tracing::error!(
channel = %self.name,
count = messages.len(),
"Messages emitted but no sender available - channel may not be started!"
);
return Ok(());
};
let mut rate_limiter = self.rate_limiter.write().await;
for emitted in messages {
// Check rate limit — acquire and release the write lock before send().await
{
let mut rate_limiter = self.rate_limiter.write().await;
if !rate_limiter.check_and_record() {
tracing::warn!(
channel = %self.name,
"Message emission rate limited"
);
return Err(WasmChannelError::EmitRateLimited {
name: self.name.clone(),
});
}
// Check rate limit
if !rate_limiter.check_and_record() {
tracing::warn!(
channel = %self.name,
"Message emission rate limited"
);
return Err(WasmChannelError::EmitRateLimited {
name: self.name.clone(),
});
}
// Convert to IncomingMessage
@@ -2062,7 +2057,7 @@ impl WasmChannel {
self.update_broadcast_metadata(&emitted.metadata_json).await;
}
// Send to stream — no locks held across this await
// Send to stream
tracing::info!(
channel = %self.name,
user_id = %emitted.user_id,
@@ -2286,33 +2281,28 @@ impl WasmChannel {
"Processing emitted messages from polling callback"
);
// Clone sender to avoid holding RwLock read guard across send().await in the loop
let tx = {
let tx_guard = message_tx.read().await;
let Some(tx) = tx_guard.as_ref() else {
tracing::error!(
channel = %channel_name,
count = messages.len(),
"Messages emitted but no sender available - channel may not be started!"
);
return Ok(());
};
tx.clone()
let tx_guard = message_tx.read().await;
let Some(tx) = tx_guard.as_ref() else {
tracing::error!(
channel = %channel_name,
count = messages.len(),
"Messages emitted but no sender available - channel may not be started!"
);
return Ok(());
};
let mut limiter = rate_limiter.write().await;
for emitted in messages {
// Check rate limit — acquire and release the write lock before send().await
{
let mut limiter = rate_limiter.write().await;
if !limiter.check_and_record() {
tracing::warn!(
channel = %channel_name,
"Message emission rate limited"
);
return Err(WasmChannelError::EmitRateLimited {
name: channel_name.to_string(),
});
}
// Check rate limit
if !limiter.check_and_record() {
tracing::warn!(
channel = %channel_name,
"Message emission rate limited"
);
return Err(WasmChannelError::EmitRateLimited {
name: channel_name.to_string(),
});
}
// Convert to IncomingMessage
@@ -2360,7 +2350,7 @@ impl WasmChannel {
.await;
}
// Send to stream — no locks held across this await
// Send to stream
tracing::info!(
channel = %channel_name,
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 thread_id = msg.thread_id.clone();
// Clone sender to avoid holding RwLock read guard across send().await
let tx = {
let tx_guard = state.msg_tx.read().await;
tx_guard
.as_ref()
.ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Channel not started".to_string(),
))?
.clone()
};
let tx_guard = state.msg_tx.read().await;
let tx = tx_guard.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Channel not started".to_string(),
))?;
tx.send(msg).await.map_err(|_| {
(
@@ -117,17 +111,11 @@ pub async fn chat_approval_handler(
let msg_id = msg.id;
// Clone sender to avoid holding RwLock read guard across send().await
let tx = {
let tx_guard = state.msg_tx.read().await;
tx_guard
.as_ref()
.ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Channel not started".to_string(),
))?
.clone()
};
let tx_guard = state.msg_tx.read().await;
let tx = tx_guard.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Channel not started".to_string(),
))?;
tx.send(msg).await.map_err(|_| {
(
-10
View File
@@ -10,7 +10,6 @@ use axum::{
use serde::Deserialize;
use uuid::Uuid;
use crate::agent::routine::{Trigger, next_cron_fire};
use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*;
use crate::error::RoutineError;
@@ -183,21 +182,12 @@ pub async fn routines_toggle_handler(
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
let was_enabled = routine.enabled;
// If a specific value was provided, use it; otherwise toggle.
routine.enabled = match body {
Some(Json(req)) => req.enabled.unwrap_or(!routine.enabled),
None => !routine.enabled,
};
if routine.enabled
&& !was_enabled
&& let Trigger::Cron { schedule, timezone } = &routine.trigger
{
routine.next_fire_at = next_cron_fire(schedule, timezone.as_deref())
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
}
store
.update_routine(&routine)
.await
+85 -204
View File
@@ -26,7 +26,6 @@ use tower_http::set_header::SetResponseHeaderLayer;
use uuid::Uuid;
use crate::agent::SessionManager;
use crate::agent::routine::{Trigger, next_cron_fire};
use crate::bootstrap::ironclaw_base_dir;
use crate::channels::IncomingMessage;
use crate::channels::relay::DEFAULT_RELAY_NAME;
@@ -573,14 +572,6 @@ async fn oauth_callback_handler(
extension = %flow.extension_name,
"OAuth flow expired"
);
// Notify UI so auth card can show error instead of staying stuck
if let Some(ref sender) = flow.sse_sender {
let _ = sender.send(SseEvent::AuthCompleted {
extension_name: flow.extension_name.clone(),
success: false,
message: "OAuth flow expired. Please try again.".to_string(),
});
}
return oauth_error_page(&flow.display_name);
}
@@ -590,12 +581,7 @@ async fn oauth_callback_handler(
let exchange_proxy_url = std::env::var("IRONCLAW_OAUTH_EXCHANGE_URL").ok();
let result: Result<(), String> = async {
let token_response = if let (Some(proxy_url), None) = (&exchange_proxy_url, &flow.resource)
{
// 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 token_response = if let Some(ref proxy_url) = exchange_proxy_url {
let gateway_token = flow.gateway_token.as_deref().unwrap_or_default();
oauth_defaults::exchange_via_proxy(
proxy_url,
@@ -608,10 +594,7 @@ async fn oauth_callback_handler(
.await
.map_err(|e| e.to_string())?
} else {
// Direct token exchange: uses exchange_oauth_code_with_resource so MCP
// 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(
oauth_defaults::exchange_oauth_code(
&flow.token_url,
&flow.client_id,
flow.client_secret.as_deref(),
@@ -619,7 +602,6 @@ async fn oauth_callback_handler(
&flow.redirect_uri,
flow.code_verifier.as_deref(),
&flow.access_token_field,
flow.resource.as_deref(),
)
.await
.map_err(|e| e.to_string())?
@@ -646,19 +628,6 @@ async fn oauth_callback_handler(
.await
.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(())
}
.await;
@@ -690,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
if let Some(ref sender) = flow.sse_sender {
let _ = sender.send(SseEvent::AuthCompleted {
extension_name: flow.extension_name,
success,
message: final_message.clone(),
message,
});
}
@@ -1027,17 +973,11 @@ async fn chat_send_handler(
req.images.len()
);
// Clone sender to avoid holding RwLock read guard across send().await
let tx = {
let tx_guard = state.msg_tx.read().await;
tx_guard
.as_ref()
.ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Channel not started".to_string(),
))?
.clone()
};
let tx_guard = state.msg_tx.read().await;
let tx = tx_guard.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Channel not started".to_string(),
))?;
tracing::debug!("[chat_send_handler] Sending message through channel");
tx.send(msg).await.map_err(|_| {
@@ -1103,17 +1043,11 @@ async fn chat_approval_handler(
let msg_id = msg.id;
// Clone sender to avoid holding RwLock read guard across send().await
let tx = {
let tx_guard = state.msg_tx.read().await;
tx_guard
.as_ref()
.ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Channel not started".to_string(),
))?
.clone()
};
let tx_guard = state.msg_tx.read().await;
let tx = tx_guard.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Channel not started".to_string(),
))?;
tx.send(msg).await.map_err(|_| {
(
@@ -2425,21 +2359,12 @@ async fn routines_toggle_handler(
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
let was_enabled = routine.enabled;
// If a specific value was provided, use it; otherwise toggle.
routine.enabled = match body {
Some(Json(req)) => req.enabled.unwrap_or(!routine.enabled),
None => !routine.enabled,
};
if routine.enabled
&& !was_enabled
&& let Trigger::Cron { schedule, timezone } = &routine.trigger
{
routine.next_fire_at = next_cron_fire(schedule, timezone.as_deref())
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
}
store
.update_routine(&routine)
.await
@@ -2714,7 +2639,6 @@ struct GatewayStatusResponse {
#[cfg(test)]
mod tests {
use super::*;
use crate::cli::oauth_defaults;
use crate::testing::credentials::TEST_GATEWAY_CRYPTO_KEY;
#[test]
@@ -2832,11 +2756,6 @@ mod tests {
.with_state(state)
}
fn expired_flow_created_at() -> Option<std::time::Instant> {
std::time::Instant::now()
.checked_sub(oauth_defaults::OAUTH_FLOW_EXPIRY + std::time::Duration::from_secs(1))
}
#[tokio::test]
async fn test_csp_header_present_on_responses() {
use std::net::SocketAddr;
@@ -2943,14 +2862,29 @@ mod tests {
use tower::ServiceExt;
// Build an ExtensionManager so the handler can look up flows
let secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync> =
Arc::new(crate::secrets::InMemorySecretsStore::new(Arc::new(
crate::secrets::SecretsCrypto::new(secrecy::SecretString::from(
TEST_GATEWAY_CRYPTO_KEY.to_string(),
))
.expect("crypto"),
)));
let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(secrets);
let secrets = Arc::new(crate::secrets::InMemorySecretsStore::new(Arc::new(
crate::secrets::SecretsCrypto::new(secrecy::SecretString::from(
TEST_GATEWAY_CRYPTO_KEY.to_string(),
))
.expect("crypto"),
)));
let tool_registry = Arc::new(ToolRegistry::new());
let mcp_sm = Arc::new(crate::tools::mcp::session::McpSessionManager::new());
let ext_mgr = Arc::new(ExtensionManager::new(
mcp_sm,
Arc::new(crate::tools::mcp::process::McpProcessManager::new()),
secrets,
tool_registry,
None,
None,
std::path::PathBuf::from("/tmp/wasm_tools"),
std::path::PathBuf::from("/tmp/wasm_channels"),
None,
"test".to_string(),
None,
vec![],
));
let state = test_gateway_state(Some(ext_mgr));
let app = test_oauth_router(state);
@@ -2984,13 +2918,25 @@ mod tests {
))
.expect("crypto"),
)));
let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(secrets.clone());
let Some(created_at) = expired_flow_created_at() else {
eprintln!("Skipping expired OAuth flow test: monotonic uptime below expiry window");
return;
};
let tool_registry = Arc::new(ToolRegistry::new());
let mcp_sm = Arc::new(crate::tools::mcp::session::McpSessionManager::new());
// Insert an expired flow.
let ext_mgr = Arc::new(ExtensionManager::new(
mcp_sm,
Arc::new(crate::tools::mcp::process::McpProcessManager::new()),
secrets.clone(),
tool_registry,
None,
None,
std::path::PathBuf::from("/tmp/wasm_tools"),
std::path::PathBuf::from("/tmp/wasm_channels"),
None,
"test".to_string(),
None,
vec![],
));
// Insert an expired flow (created 10 minutes ago)
let flow = crate::cli::oauth_defaults::PendingOAuthFlow {
extension_name: "test_tool".to_string(),
display_name: "Test Tool".to_string(),
@@ -3008,9 +2954,9 @@ mod tests {
secrets,
sse_sender: None,
gateway_token: None,
resource: None,
client_id_secret_name: None,
created_at,
created_at: std::time::Instant::now()
.checked_sub(std::time::Duration::from_secs(600))
.expect("System uptime is too low to run expired flow test"),
};
ext_mgr
@@ -3040,80 +2986,6 @@ mod tests {
assert!(html.contains("Authorization Failed"));
}
#[tokio::test]
async fn test_oauth_callback_expired_flow_broadcasts_auth_completed_failure() {
use axum::body::Body;
use tower::ServiceExt;
let secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync> =
Arc::new(crate::secrets::InMemorySecretsStore::new(Arc::new(
crate::secrets::SecretsCrypto::new(secrecy::SecretString::from(
TEST_GATEWAY_CRYPTO_KEY.to_string(),
))
.expect("crypto"),
)));
let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(secrets.clone());
let (sender, mut receiver) = tokio::sync::broadcast::channel(4);
let Some(created_at) = expired_flow_created_at() else {
eprintln!("Skipping expired OAuth flow SSE test: monotonic uptime below expiry window");
return;
};
let flow = crate::cli::oauth_defaults::PendingOAuthFlow {
extension_name: "test_tool".to_string(),
display_name: "Test Tool".to_string(),
token_url: "https://example.com/token".to_string(),
client_id: "client123".to_string(),
client_secret: None,
redirect_uri: "https://example.com/oauth/callback".to_string(),
code_verifier: None,
access_token_field: "access_token".to_string(),
secret_name: "test_token".to_string(),
provider: None,
validation_endpoint: None,
scopes: vec![],
user_id: "test".to_string(),
secrets,
sse_sender: Some(sender),
gateway_token: None,
resource: None,
client_id_secret_name: None,
created_at,
};
ext_mgr
.pending_oauth_flows()
.write()
.await
.insert("expired_state".to_string(), flow);
let state = test_gateway_state(Some(ext_mgr));
let app = test_oauth_router(state);
let req = axum::http::Request::builder()
.uri("/oauth/callback?code=test_code&state=expired_state")
.body(Body::empty())
.expect("request");
let resp = ServiceExt::<axum::http::Request<Body>>::oneshot(app, req)
.await
.expect("response");
assert_eq!(resp.status(), StatusCode::OK);
match receiver.recv().await.expect("auth_completed event") {
crate::channels::web::types::SseEvent::AuthCompleted {
extension_name,
success,
message,
} => {
assert_eq!(extension_name, "test_tool");
assert!(!success, "expired OAuth flow should broadcast failure");
assert_eq!(message, "OAuth flow expired. Please try again.");
}
event => panic!("expected AuthCompleted event, got {event:?}"),
}
}
#[tokio::test]
async fn test_oauth_callback_no_extension_manager() {
use axum::body::Body;
@@ -3152,16 +3024,28 @@ mod tests {
))
.expect("crypto"),
)));
let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(secrets.clone());
let tool_registry = Arc::new(ToolRegistry::new());
let mcp_sm = Arc::new(crate::tools::mcp::session::McpSessionManager::new());
let ext_mgr = Arc::new(ExtensionManager::new(
mcp_sm,
Arc::new(crate::tools::mcp::process::McpProcessManager::new()),
secrets.clone(),
tool_registry,
None,
None,
std::path::PathBuf::from("/tmp/wasm_tools"),
std::path::PathBuf::from("/tmp/wasm_channels"),
None,
"test".to_string(),
None,
vec![],
));
// Insert a flow keyed by raw nonce "test_nonce" (without instance prefix).
// Use an expired flow so the handler exits before attempting a real HTTP
// token exchange — we only need to verify that the instance prefix was
// stripped and the flow was found by the raw nonce.
let Some(created_at) = expired_flow_created_at() else {
eprintln!("Skipping OAuth state-prefix test: monotonic uptime below expiry window");
return;
};
let flow = crate::cli::oauth_defaults::PendingOAuthFlow {
extension_name: "test_tool".to_string(),
display_name: "Test Tool".to_string(),
@@ -3179,10 +3063,10 @@ mod tests {
secrets,
sse_sender: None,
gateway_token: None,
resource: None,
client_id_secret_name: None,
// Expired — handler will reject after lookup (no network I/O)
created_at,
created_at: std::time::Instant::now()
.checked_sub(std::time::Duration::from_secs(600))
.expect("System uptime is too low to run expired flow test"),
};
ext_mgr
@@ -3253,27 +3137,24 @@ mod tests {
fn test_ext_mgr(
secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync>,
) -> (Arc<ExtensionManager>, tempfile::TempDir, tempfile::TempDir) {
) -> Arc<ExtensionManager> {
let tool_registry = Arc::new(ToolRegistry::new());
let mcp_sm = Arc::new(crate::tools::mcp::session::McpSessionManager::new());
let mcp_pm = Arc::new(crate::tools::mcp::process::McpProcessManager::new());
let wasm_tools_dir = tempfile::tempdir().expect("temp wasm tools dir");
let wasm_channels_dir = tempfile::tempdir().expect("temp wasm channels dir");
let ext_mgr = Arc::new(ExtensionManager::new(
Arc::new(ExtensionManager::new(
mcp_sm,
mcp_pm,
secrets,
tool_registry,
None,
None,
wasm_tools_dir.path().to_path_buf(),
wasm_channels_dir.path().to_path_buf(),
std::path::PathBuf::from("/tmp/wasm_tools"),
std::path::PathBuf::from("/tmp/wasm_channels"),
None,
"test".to_string(),
None,
vec![],
));
(ext_mgr, wasm_tools_dir, wasm_channels_dir)
))
}
#[tokio::test]
@@ -3282,7 +3163,7 @@ mod tests {
use tower::ServiceExt;
let secrets = test_secrets_store();
let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(secrets);
let ext_mgr = test_ext_mgr(secrets);
let state = test_gateway_state(Some(ext_mgr));
let app = test_relay_oauth_router(state);
@@ -3326,7 +3207,7 @@ mod tests {
.await
.expect("store nonce");
let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(secrets);
let ext_mgr = test_ext_mgr(secrets);
let state = test_gateway_state(Some(ext_mgr));
let app = test_relay_oauth_router(state);
@@ -3371,7 +3252,7 @@ mod tests {
.await
.expect("store nonce");
let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(secrets.clone());
let ext_mgr = test_ext_mgr(secrets.clone());
let state = test_gateway_state(Some(ext_mgr));
let app = test_relay_oauth_router(state);
+64 -286
View File
@@ -342,27 +342,31 @@ function connectSSE() {
eventSource.addEventListener('approval_needed', (e) => {
const data = JSON.parse(e.data);
const hasThread = !!data.thread_id;
const forCurrentThread = !hasThread || isCurrentThread(data.thread_id);
if (forCurrentThread) {
showApproval(data);
} else {
// Keep thread list fresh when approval is requested in a background thread.
unreadThreads.set(data.thread_id, (unreadThreads.get(data.thread_id) || 0) + 1);
debouncedLoadThreads();
}
// Extension setup flows can surface approvals while user is on Extensions tab.
if (currentTab === 'extensions') loadExtensions();
if (!isCurrentThread(data.thread_id)) return;
showApproval(data);
});
eventSource.addEventListener('auth_required', (e) => {
handleAuthRequired(JSON.parse(e.data));
const data = JSON.parse(e.data);
if (data.auth_url) {
// OAuth flow: show the auth card with an OAuth button + optional token paste field.
showAuthCard(data);
} else {
// Setup flow: fetch the extension's credential schema and show the multi-field
// configure modal (the same UI used by the Extensions tab "Setup" button).
showConfigureModal(data.extension_name);
}
});
eventSource.addEventListener('auth_completed', (e) => {
handleAuthCompleted(JSON.parse(e.data));
const data = JSON.parse(e.data);
// Dismiss whichever UI path was active: auth card (OAuth) or configure modal (setup).
removeAuthCard(data.extension_name);
closeConfigureModal();
showToast(data.message, data.success ? 'success' : 'error');
// Refresh extensions list so status indicators update
if (currentTab === 'extensions') loadExtensions();
enableChatInput();
});
eventSource.addEventListener('extension_status', (e) => {
@@ -666,7 +670,7 @@ function renderMarkdown(text) {
// Sanitize HTML output to prevent XSS from tool output or LLM responses.
html = sanitizeRenderedHtml(html);
// 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 escapeHtml(text);
@@ -698,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) {
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.scrollTop = container.scrollHeight;
}
@@ -728,11 +723,7 @@ function appendToLastAssistant(chunk) {
const last = messages[messages.length - 1];
const raw = (last.getAttribute('data-raw') || '') + chunk;
last.setAttribute('data-raw', raw);
last.setAttribute('data-copy-text', raw);
const content = last.querySelector('.message-content');
if (content) {
content.innerHTML = renderMarkdown(raw);
}
last.innerHTML = renderMarkdown(raw);
container.scrollTop = container.scrollHeight;
} else {
addMessage('assistant', chunk);
@@ -986,26 +977,7 @@ function finalizeActivityGroup() {
_activeToolCards = {};
}
function humanizeToolName(rawName) {
if (!rawName) return '';
return String(rawName)
.replace(/[_-]+/g, ' ')
.replace(/([a-z0-9])([A-Z])/g, '$1 $2')
.replace(/^tool([a-zA-Z])/, 'tool $1')
.replace(/\s+/g, ' ')
.trim();
}
function shouldShowChannelConnectedMessage(extensionName, success) {
if (!success || !extensionName) return false;
return String(extensionName).toLowerCase().includes('telegram');
}
function showApproval(data) {
// Avoid duplicate cards on reconnect/history refresh.
const existing = document.querySelector('.approval-card[data-request-id="' + CSS.escape(data.request_id) + '"]');
if (existing) return;
const container = document.getElementById('chat-messages');
const card = document.createElement('div');
card.className = 'approval-card';
@@ -1018,7 +990,7 @@ function showApproval(data) {
const toolName = document.createElement('div');
toolName.className = 'approval-tool-name';
toolName.textContent = humanizeToolName(data.tool_name);
toolName.textContent = data.tool_name;
card.appendChild(toolName);
if (data.description) {
@@ -1121,71 +1093,13 @@ function showJobCard(data) {
// --- Auth card ---
function handleAuthRequired(data) {
if (data.auth_url) {
// OAuth flow: show the global auth prompt with an OAuth button + optional token paste field.
showAuthCard(data);
} else {
// Setup flow: fetch the extension's credential schema and show the multi-field
// configure modal (the same UI used by the Extensions tab "Setup" button).
showConfigureModal(data.extension_name);
}
}
function handleAuthCompleted(data) {
// Dismiss only the matching extension's UI so unrelated setup work is not interrupted.
removeAuthCard(data.extension_name);
closeConfigureModal(data.extension_name);
showToast(data.message, data.success ? 'success' : 'error');
if (shouldShowChannelConnectedMessage(data.extension_name, data.success)) {
addMessage('system', 'Telegram is now connected. You can message me there and I can send you notifications.');
}
if (currentTab === 'extensions') loadExtensions();
enableChatInput();
}
function queryByDataAttribute(selector, attributeName, attributeValue) {
if (typeof attributeValue !== 'string') return document.querySelector(selector);
if (window.CSS && typeof window.CSS.escape === 'function') {
return document.querySelector(
selector + '[' + attributeName + '="' + window.CSS.escape(attributeValue) + '"]'
);
}
const candidates = document.querySelectorAll(selector);
for (const candidate of candidates) {
if (candidate.getAttribute(attributeName) === attributeValue) return candidate;
}
return null;
}
function getAuthOverlay(extensionName) {
return queryByDataAttribute('.auth-overlay', 'data-extension-name', extensionName);
}
function getAuthCard(extensionName) {
return queryByDataAttribute('.auth-card', 'data-extension-name', extensionName);
}
function getConfigureOverlay(extensionName) {
return queryByDataAttribute('.configure-overlay', 'data-extension-name', extensionName);
}
function showAuthCard(data) {
// Keep a single global auth prompt so the experience is consistent across tabs.
const existing = getAuthOverlay();
if (existing) existing.remove();
const overlay = document.createElement('div');
overlay.className = 'auth-overlay';
overlay.setAttribute('data-extension-name', data.extension_name);
overlay.addEventListener('click', (e) => {
if (e.target === overlay) cancelAuth(data.extension_name);
});
// Remove any existing card for this extension first
removeAuthCard(data.extension_name);
const container = document.getElementById('chat-messages');
const card = document.createElement('div');
card.className = 'auth-card auth-modal';
card.className = 'auth-card';
card.setAttribute('data-extension-name', data.extension_name);
const header = document.createElement('div');
@@ -1264,30 +1178,21 @@ function showAuthCard(data) {
actions.appendChild(cancelBtn);
card.appendChild(actions);
overlay.appendChild(card);
document.body.appendChild(overlay);
container.appendChild(card);
container.scrollTop = container.scrollHeight;
tokenInput.focus();
}
function removeAuthCard(extensionName) {
const overlay = getAuthOverlay(extensionName);
if (overlay) {
overlay.remove();
return;
}
const card = getAuthCard(extensionName);
if (card) {
const parentOverlay = card.closest('.auth-overlay');
if (parentOverlay) parentOverlay.remove();
else card.remove();
}
const card = document.querySelector('.auth-card[data-extension-name="' + extensionName + '"]');
if (card) card.remove();
}
function submitAuthToken(extensionName, tokenValue) {
if (!tokenValue || !tokenValue.trim()) return;
// Disable submit button while in flight
const card = getAuthCard(extensionName);
const card = document.querySelector('.auth-card[data-extension-name="' + extensionName + '"]');
if (card) {
const btns = card.querySelectorAll('button');
btns.forEach((b) => { b.disabled = true; });
@@ -1298,10 +1203,8 @@ function submitAuthToken(extensionName, tokenValue) {
body: { extension_name: extensionName, token: tokenValue.trim() },
}).then((result) => {
if (result.success) {
// Close immediately for responsiveness; the authoritative success UX
// (toast + extensions refresh) still comes from auth_completed SSE.
removeAuthCard(extensionName);
enableChatInput();
addMessage('system', result.message);
} else {
showAuthCardError(extensionName, result.message);
}
@@ -1320,7 +1223,7 @@ function cancelAuth(extensionName) {
}
function showAuthCardError(extensionName, message) {
const card = getAuthCard(extensionName);
const card = document.querySelector('.auth-card[data-extension-name="' + extensionName + '"]');
if (!card) return;
// Re-enable buttons
const btns = card.querySelectorAll('button');
@@ -1407,31 +1310,12 @@ function loadHistory(before) {
function createMessageElement(role, content) {
const div = document.createElement('div');
div.className = 'message ' + role;
if (role === 'assistant' || role === 'user') {
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;
if (role === 'user') {
div.textContent = content;
} else {
div.setAttribute('data-raw', content);
body.innerHTML = renderMarkdown(content);
div.innerHTML = renderMarkdown(content);
}
div.appendChild(body);
return div;
}
@@ -1935,11 +1819,13 @@ function saveMemoryEdit() {
function buildBreadcrumb(path) {
const parts = path.split('/');
let html = '<a data-action="breadcrumb-root" href="#">workspace</a>';
let html = '<a onclick="loadMemoryTree()">workspace</a>';
let current = '';
for (const part of parts) {
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;
}
@@ -2250,10 +2136,6 @@ function renderAvailableExtensionCard(entry) {
showToast(I18n.t('extensions.installedSuccess', {name: entry.display_name}), 'success');
// OAuth popup if auth started during install (builtin creds)
if (res.auth_url) {
showAuthCard({
extension_name: entry.name,
auth_url: res.auth_url,
});
showToast('Opening authentication for ' + entry.display_name, 'info');
openOAuthUrl(res.auth_url);
}
@@ -2519,10 +2401,6 @@ function activateExtension(name) {
if (res.success) {
// Even on success, the tool may need OAuth (e.g., WASM loaded but no token yet)
if (res.auth_url) {
showAuthCard({
extension_name: name,
auth_url: res.auth_url,
});
showToast('Opening authentication for ' + name, 'info');
openOAuthUrl(res.auth_url);
}
@@ -2531,10 +2409,6 @@ function activateExtension(name) {
}
if (res.auth_url) {
showAuthCard({
extension_name: name,
auth_url: res.auth_url,
});
showToast('Opening authentication for ' + name, 'info');
openOAuthUrl(res.auth_url);
} else if (res.awaiting_token) {
@@ -2577,7 +2451,6 @@ function renderConfigureModal(name, secrets) {
closeConfigureModal();
const overlay = document.createElement('div');
overlay.className = 'configure-overlay';
overlay.setAttribute('data-extension-name', name);
overlay.addEventListener('click', (e) => {
if (e.target === overlay) closeConfigureModal();
});
@@ -2671,8 +2544,7 @@ function submitConfigureModal(name, fields) {
}
// Disable buttons to prevent double-submit
const overlay = getConfigureOverlay(name) || document.querySelector('.configure-overlay');
var btns = overlay ? overlay.querySelectorAll('.configure-actions button') : [];
var btns = document.querySelectorAll('.configure-actions button');
btns.forEach(function(b) { b.disabled = true; });
apiFetch('/api/extensions/' + encodeURIComponent(name) + '/setup', {
@@ -2683,10 +2555,8 @@ function submitConfigureModal(name, fields) {
if (res.success) {
closeConfigureModal();
if (res.auth_url) {
showAuthCard({
extension_name: name,
auth_url: res.auth_url,
});
// OAuth flow started — open consent popup. The auth_completed SSE will
// not arrive immediately (it fires after OAuth callback), so show a toast now.
showToast('Opening OAuth authorization for ' + name, 'info');
openOAuthUrl(res.auth_url);
loadExtensions();
@@ -2705,9 +2575,8 @@ function submitConfigureModal(name, fields) {
});
}
function closeConfigureModal(extensionName) {
if (typeof extensionName !== 'string') extensionName = null;
const existing = getConfigureOverlay(extensionName);
function closeConfigureModal() {
const existing = document.querySelector('.configure-overlay');
if (existing) existing.remove();
}
@@ -2926,11 +2795,11 @@ function renderJobsList(jobs) {
let actionBtns = '';
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.
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>' + escapeHtml(job.title) + '</td>'
+ '<td><span class="badge ' + stateClass + '">' + escapeHtml(job.state) + '</span></td>'
@@ -2993,12 +2862,12 @@ function renderJobDetail(job) {
const header = document.createElement('div');
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>'
+ '<span class="badge ' + stateClass + '">' + escapeHtml(job.state) + '</span>';
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) {
headerHtml += '<a class="btn-browse" href="' + escapeHtml(job.browse_url) + '" target="_blank">Browse Files</a>';
@@ -3455,7 +3324,7 @@ function renderRoutinesList(routines) {
const toggleLabel = r.enabled ? 'Disable' : 'Enable';
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.trigger_summary) + '</td>'
+ '<td>' + escapeHtml(r.action_type) + '</td>'
@@ -3464,9 +3333,9 @@ function renderRoutinesList(routines) {
+ '<td>' + r.run_count + '</td>'
+ '<td><span class="badge ' + statusClass + '">' + escapeHtml(r.status) + '</span></td>'
+ '<td>'
+ '<button class="' + toggleClass + '" data-action="toggle-routine" data-id="' + escapeHtml(r.id) + '">' + toggleLabel + '</button> '
+ '<button class="btn-restart" data-action="trigger-routine" data-id="' + escapeHtml(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="' + toggleClass + '" onclick="event.stopPropagation(); toggleRoutine(\'' + r.id + '\')">' + toggleLabel + '</button> '
+ '<button class="btn-restart" onclick="event.stopPropagation(); triggerRoutine(\'' + r.id + '\')">Run</button> '
+ '<button class="btn-cancel" onclick="event.stopPropagation(); deleteRoutine(\'' + r.id + '\', \'' + escapeHtml(r.name) + '\')">Delete</button>'
+ '</td>'
+ '</tr>';
}).join('');
@@ -3502,7 +3371,7 @@ function renderRoutineDetail(routine) {
: 'active';
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>'
+ '<span class="badge ' + statusClass + '">' + escapeHtml(statusLabel) + '</span>'
+ '</div>';
@@ -3549,7 +3418,7 @@ function renderRoutineDetail(routine) {
+ '<td>' + formatDate(run.completed_at) + '</td>'
+ '<td><span class="badge ' + runStatusClass + '">' + escapeHtml(run.status) + '</span></td>'
+ '<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>' + (run.tokens_used != null ? run.tokens_used : '-') + '</td>'
+ '</tr>';
@@ -3792,7 +3661,7 @@ function renderTeePopover(report) {
+ '<div class="tee-field"><div class="tee-field-label">VM Config</div>'
+ '<div class="tee-field-value">' + escapeHtml(vmConfig) + '</div></div>'
+ '<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() {
@@ -4274,94 +4143,3 @@ function formatDate(isoString) {
const d = new Date(isoString);
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();
});
+27 -27
View File
@@ -9,12 +9,12 @@
<link rel="preconnect" href="https://fonts.gstatic.com" crossorigin>
<link href="https://fonts.googleapis.com/css2?family=DM+Sans:wght@400;500;600;700&family=IBM+Plex+Mono:wght@400;500;600&display=swap" rel="stylesheet">
<link rel="stylesheet" href="/style.css">
<!-- i18n Modules -->
<script src="/i18n/index.js"></script>
<script src="/i18n/en.js"></script>
<script src="/i18n/zh-CN.js"></script>
<script
src="https://cdnjs.cloudflare.com/ajax/libs/dompurify/3.2.3/purify.min.js"
integrity="sha384-osZDKVu4ipZP703HmPOhWdyBajcFyjX2Psjk//TG1Rc0AdwEtuToaylrmcK3LdAl"
@@ -37,7 +37,7 @@
<div class="auth-form">
<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>
<button id="auth-connect-btn" data-i18n="auth.connect">Connect</button>
<button onclick="authenticate()" data-i18n="auth.connect">Connect</button>
</div>
<div id="auth-error"></div>
<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 -->
<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-header">
<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>
</div>
<div class="restart-modal-body">
@@ -63,8 +63,8 @@
</div>
</div>
<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 confirm" id="restart-confirm-btn" data-i18n="restart.confirm">Confirm Restart</button>
<button class="restart-modal-btn cancel" onclick="cancelRestart()" data-i18n="restart.cancel">Cancel</button>
<button class="restart-modal-btn confirm" onclick="confirmRestart()" data-i18n="restart.confirm">Confirm Restart</button>
</div>
</div>
</div>
@@ -98,17 +98,17 @@
<button data-tab="extensions" data-i18n="tab.extensions">Extensions</button>
<button data-tab="skills" data-i18n="tab.skills">Skills</button>
<div class="spacer"></div>
<!-- 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>
<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" data-action="switch-language" data-lang="zh-CN">简体中文</button>
<button type="button" class="language-option" onclick="switchLanguage('en')" data-lang="en">English</button>
<button type="button" class="language-option" onclick="switchLanguage('zh-CN')" data-lang="zh-CN">简体中文</button>
</div>
</div>
<button class="status-logs-btn" data-tab="logs" data-i18n="tab.logs" title="Logs">Logs</button>
<div class="tee-shield" id="tee-shield" style="display:none" title="Running in a Trusted Execution Environment">
<svg width="16" height="16" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round">
@@ -122,7 +122,7 @@
<span id="sse-status" data-i18n="status.connected">Connected</span>
<div class="gateway-popover" id="gateway-popover"></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;">
<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>
@@ -137,13 +137,13 @@
<div class="tab-panel active" id="tab-chat">
<div class="thread-sidebar" id="thread-sidebar">
<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>
<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>
</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-meta" id="assistant-meta"></span>
</div>
@@ -161,7 +161,7 @@
<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"
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>
@@ -178,7 +178,7 @@
<div class="memory-content">
<div class="memory-breadcrumb" id="memory-breadcrumb">
<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 class="memory-viewer" id="memory-viewer">
<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">
<textarea id="memory-edit-textarea"></textarea>
<div class="memory-editor-actions">
<button class="btn-save" id="memory-save-btn" data-i18n="memory.save">Save</button>
<button class="btn-cancel-edit" id="memory-cancel-btn" data-i18n="memory.cancel">Cancel</button>
<button class="btn-save" onclick="saveMemoryEdit()" data-i18n="memory.save">Save</button>
<button class="btn-cancel-edit" onclick="cancelMemoryEdit()" data-i18n="memory.cancel">Cancel</button>
</div>
</div>
</div>
@@ -219,7 +219,7 @@
<div class="tab-panel" id="tab-logs">
<div class="logs-container">
<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="warn">Server: WARN</option>
<option value="info" selected>Server: INFO</option>
@@ -234,8 +234,8 @@
</select>
<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>
<button id="logs-pause-btn" data-i18n="logs.pause">Pause</button>
<button id="logs-clear-btn" data-i18n="logs.clear">Clear</button>
<button id="logs-pause-btn" onclick="toggleLogsPause()" data-i18n="logs.pause">Pause</button>
<button onclick="clearLogs()" data-i18n="logs.clear">Clear</button>
</div>
<div class="logs-output" id="logs-output"></div>
</div>
@@ -287,7 +287,7 @@
<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-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 class="extensions-section">
@@ -299,7 +299,7 @@
<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-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 class="extensions-section">
@@ -320,7 +320,7 @@
<h3 data-i18n="skills.searchClawHub">Search ClawHub</h3>
<div class="skill-search-box">
<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 class="extensions-list" id="skill-search-results"></div>
</div>
@@ -335,7 +335,7 @@
<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-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>
+1 -78
View File
@@ -666,7 +666,6 @@ body {
font-size: 14px;
line-height: 1.5;
word-wrap: break-word;
position: relative;
}
.message.user {
@@ -687,58 +686,6 @@ body {
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 {
align-self: center;
background: var(--bg-tertiary);
@@ -1219,21 +1166,7 @@ body {
color: var(--danger);
}
/* Auth prompt */
.auth-overlay {
position: fixed;
top: 0;
left: 0;
width: 100%;
height: 100%;
background: rgba(0, 0, 0, 0.6);
z-index: 1001;
display: flex;
align-items: center;
justify-content: center;
padding: 16px;
}
/* Auth card (inline in chat) */
.auth-card {
align-self: flex-start;
max-width: 80%;
@@ -1248,16 +1181,6 @@ body {
transition: border-color 0.2s;
}
.auth-overlay .auth-card {
width: 460px;
max-width: min(460px, 90vw);
margin: 0;
align-self: auto;
background: var(--bg);
border-color: rgba(52, 211, 153, 0.35);
box-shadow: 0 24px 48px rgba(0, 0, 0, 0.35);
}
.auth-card .auth-header {
font-weight: 600;
color: var(--accent);
+4 -12
View File
@@ -176,12 +176,8 @@ async fn handle_client_message(
incoming = incoming.with_attachments(attachments);
}
// Clone sender to avoid holding RwLock read guard across send().await
let tx = {
let tx_guard = state.msg_tx.read().await;
tx_guard.as_ref().cloned()
};
if let Some(tx) = tx {
let tx_guard = state.msg_tx.read().await;
if let Some(ref tx) = *tx_guard {
if tx.send(incoming).await.is_err() {
let _ = direct_tx
.send(WsServerMessage::Error {
@@ -249,12 +245,8 @@ async fn handle_client_message(
if let Some(ref tid) = thread_id {
msg = msg.with_thread(tid);
}
// Clone sender to avoid holding RwLock read guard across send().await
let tx = {
let tx_guard = state.msg_tx.read().await;
tx_guard.as_ref().cloned()
};
if let Some(tx) = tx {
let tx_guard = state.msg_tx.read().await;
if let Some(ref tx) = *tx_guard {
let _ = tx.send(msg).await;
}
}
-79
View File
@@ -172,35 +172,6 @@ pub async fn exchange_oauth_code(
redirect_uri: &str,
code_verifier: Option<&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> {
let client = reqwest::Client::new();
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()));
}
// 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);
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>>,
/// Gateway auth token for authenticating with the platform token exchange proxy.
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).
pub created_at: std::time::Instant,
}
@@ -1016,42 +975,4 @@ mod tests {
assert_eq!(strip_instance_prefix("abc123"), "abc123");
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(),
context_paths: Vec::new(),
max_tokens: 4096,
use_tools: false,
max_tool_rounds: 0,
},
guardrails: RoutineGuardrails {
cooldown: std::time::Duration::from_secs(cooldown_secs),
+7 -63
View File
@@ -23,9 +23,6 @@ pub struct EmbeddingsConfig {
pub ollama_base_url: String,
/// Embedding vector dimension. Inferred from the model name when not set explicitly.
pub dimension: usize,
/// Custom base URL for OpenAI-compatible embedding providers.
/// When set, overrides the default `https://api.openai.com`.
pub openai_base_url: Option<String>,
}
impl Default for EmbeddingsConfig {
@@ -39,7 +36,6 @@ impl Default for EmbeddingsConfig {
model,
ollama_base_url: "http://localhost:11434".to_string(),
dimension,
openai_base_url: None,
}
}
}
@@ -78,8 +74,6 @@ impl EmbeddingsConfig {
let enabled = parse_bool_env("EMBEDDING_ENABLED", settings.embeddings.enabled)?;
let openai_base_url = optional_env("EMBEDDING_BASE_URL")?;
Ok(Self {
enabled,
provider,
@@ -87,7 +81,6 @@ impl EmbeddingsConfig {
model,
ollama_base_url,
dimension,
openai_base_url,
})
}
@@ -137,27 +130,16 @@ impl EmbeddingsConfig {
}
_ => {
if let Some(api_key) = self.openai_api_key() {
let mut provider = crate::workspace::OpenAiEmbeddings::with_model(
tracing::debug!(
"Embeddings enabled via OpenAI (model: {}, dim: {})",
self.model,
self.dimension,
);
Some(Arc::new(crate::workspace::OpenAiEmbeddings::with_model(
api_key,
&self.model,
self.dimension,
);
if let Some(ref base_url) = self.openai_base_url {
tracing::debug!(
"Embeddings enabled via OpenAI (model: {}, base_url: {}, dim: {})",
self.model,
base_url,
self.dimension,
);
provider = provider.with_base_url(base_url);
} else {
tracing::debug!(
"Embeddings enabled via OpenAI (model: {}, dim: {})",
self.model,
self.dimension,
);
}
Some(Arc::new(provider))
)))
} else {
tracing::warn!("Embeddings configured but OPENAI_API_KEY not set");
None
@@ -182,7 +164,6 @@ mod tests {
std::env::remove_var("EMBEDDING_PROVIDER");
std::env::remove_var("EMBEDDING_MODEL");
std::env::remove_var("OPENAI_API_KEY");
std::env::remove_var("EMBEDDING_BASE_URL");
}
}
@@ -266,41 +247,4 @@ mod tests {
std::env::remove_var("EMBEDDING_ENABLED");
}
}
#[test]
fn embedding_base_url_parsed_from_env() {
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
clear_embedding_env();
// SAFETY: Under ENV_MUTEX, no concurrent env access.
unsafe {
std::env::set_var("EMBEDDING_BASE_URL", "https://custom.example.com");
}
let settings = Settings::default();
let config = EmbeddingsConfig::resolve(&settings).expect("resolve should succeed");
assert_eq!(
config.openai_base_url.as_deref(),
Some("https://custom.example.com"),
"EMBEDDING_BASE_URL env var should be parsed into openai_base_url"
);
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::remove_var("EMBEDDING_BASE_URL");
}
}
#[test]
fn embedding_base_url_defaults_to_none() {
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
clear_embedding_env();
let settings = Settings::default();
let config = EmbeddingsConfig::resolve(&settings).expect("resolve should succeed");
assert!(
config.openai_base_url.is_none(),
"openai_base_url should be None when EMBEDDING_BASE_URL is not set"
);
}
}
-5
View File
@@ -18,7 +18,6 @@ pub mod relay;
mod routines;
mod safety;
mod sandbox;
mod search;
mod secrets;
mod skills;
mod transcription;
@@ -45,7 +44,6 @@ pub use self::routines::RoutineConfig;
pub use self::safety::SafetyConfig;
use self::safety::resolve_safety_config;
pub use self::sandbox::{ClaudeCodeConfig, SandboxModeConfig};
pub use self::search::WorkspaceSearchConfig;
pub use self::secrets::SecretsConfig;
pub use self::skills::SkillsConfig;
pub use self::transcription::TranscriptionConfig;
@@ -93,7 +91,6 @@ pub struct Config {
pub claude_code: ClaudeCodeConfig,
pub skills: SkillsConfig,
pub transcription: TranscriptionConfig,
pub search: WorkspaceSearchConfig,
pub observability: crate::observability::ObservabilityConfig,
/// Channel-relay integration (Slack via external relay service).
/// Present only when both `CHANNEL_RELAY_URL` and `CHANNEL_RELAY_API_KEY` are set.
@@ -169,7 +166,6 @@ impl Config {
..SkillsConfig::default()
},
transcription: TranscriptionConfig::default(),
search: WorkspaceSearchConfig::default(),
observability: crate::observability::ObservabilityConfig::default(),
relay: None,
}
@@ -322,7 +318,6 @@ impl Config {
claude_code: ClaudeCodeConfig::resolve()?,
skills: SkillsConfig::resolve()?,
transcription: TranscriptionConfig::resolve(settings)?,
search: WorkspaceSearchConfig::resolve()?,
observability: crate::observability::ObservabilityConfig {
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();
}
}
+8 -14
View File
@@ -16,7 +16,6 @@ mod workspace;
use std::path::Path;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use async_trait::async_trait;
use chrono::{DateTime, NaiveDateTime, Utc};
@@ -33,8 +32,6 @@ use crate::workspace::MemoryDocument;
use crate::db::libsql_migrations;
static NAIVE_TIMESTAMP_LOGGED: AtomicBool = AtomicBool::new(false);
/// Explicit column list for routines table (matches positional access in `row_to_routine_libsql`).
pub(crate) const ROUTINE_COLUMNS: &str = "\
id, name, description, user_id, enabled, \
@@ -166,27 +163,24 @@ impl LibSqlBackend {
///
/// Returns an error if none of the formats match.
pub(crate) fn parse_timestamp(s: &str) -> Result<DateTime<Utc>, String> {
let log_naive_timestamp_once = || {
if !NAIVE_TIMESTAMP_LOGGED.swap(true, Ordering::Relaxed) {
tracing::debug!(
timestamp = %s,
"parsed naive timestamp without timezone; assuming UTC for backward compatibility"
);
}
};
// RFC 3339 (our canonical write format)
if let Ok(dt) = DateTime::parse_from_rfc3339(s) {
return Ok(dt.with_timezone(&Utc));
}
// 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") {
log_naive_timestamp_once();
tracing::warn!(
timestamp = %s,
"parsed naive timestamp without timezone; assuming UTC for backward compatibility"
);
return Ok(ndt.and_utc());
}
// Naive without fractional seconds (legacy format)
if let Ok(ndt) = NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S") {
log_naive_timestamp_once();
tracing::warn!(
timestamp = %s,
"parsed naive timestamp without timezone; assuming UTC for backward compatibility"
);
return Ok(ndt.and_utc());
}
Err(format!("unparseable timestamp: {:?}", s))
+2 -2
View File
@@ -14,7 +14,7 @@ use crate::db::WorkspaceStore;
use crate::error::WorkspaceError;
use crate::workspace::{
MemoryChunk, MemoryDocument, RankedResult, SearchConfig, SearchResult, WorkspaceEntry,
fuse_results,
reciprocal_rank_fusion,
};
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))
}
}
+64 -635
View File
@@ -27,7 +27,7 @@ use crate::secrets::{CreateSecretParams, SecretsStore};
use crate::tools::ToolRegistry;
use crate::tools::mcp::McpClient;
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,
};
use crate::tools::mcp::config::McpServerConfig;
@@ -108,13 +108,6 @@ pub struct ExtensionManager {
/// Relay config captured at startup. Used by `auth_channel_relay` and
/// `activate_channel_relay` instead of re-reading env vars.
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.
@@ -188,75 +181,9 @@ impl ExtensionManager {
pending_oauth_flows: crate::cli::oauth_defaults::new_pending_oauth_registry(),
gateway_token: std::env::var("GATEWAY_AUTH_TOKEN").ok(),
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.
fn relay_config(&self) -> Result<&crate::config::RelayConfig, ExtensionError> {
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.
///
/// Call after construction (and after wrapping in `Arc`) once the channel
@@ -786,19 +707,6 @@ impl ExtensionManager {
Self::validate_extension_name(name)?;
let kind = self.determine_installed_kind(name).await?;
// Clean up any in-progress OAuth flows for this extension.
// TCP mode: abort the listener task so port 9876 is freed immediately.
// Gateway mode: remove stale pending flow entries.
if let Some(pending) = self.pending_auth.write().await.remove(name)
&& let Some(handle) = pending.task_handle
{
handle.abort();
}
self.pending_oauth_flows
.write()
.await
.retain(|_, flow| flow.extension_name != name);
match kind {
ExtensionKind::McpServer => {
// Unregister tools with this server's prefix
@@ -832,14 +740,6 @@ impl ExtensionManager {
// Unregister from tool registry
self.tool_registry.unregister(name).await;
// Evict compiled module from runtime cache so reinstall uses fresh binary
if let Some(ref rt) = self.wasm_tool_runtime {
rt.remove(name).await;
}
// Clear stale activation errors so reinstall starts clean
self.activation_errors.write().await.remove(name);
// Revoke credential mappings from the shared registry
let cap_path = self
.wasm_tools_dir
@@ -880,9 +780,6 @@ impl ExtensionManager {
self.active_channel_names.write().await.remove(name);
self.persist_active_channels().await;
// Clear stale activation errors so reinstall starts clean
self.activation_errors.write().await.remove(name);
// Delete channel files
let wasm_path = self.wasm_channels_dir.join(format!("{}.wasm", name));
let cap_path = self
@@ -1787,46 +1684,29 @@ impl ExtensionManager {
return Ok(AuthResult::authenticated(name, ExtensionKind::McpServer));
}
// In gateway mode, build an auth URL and return it for the frontend to
// 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)
// Run the full OAuth flow (opens browser, waits for callback)
match authorize_mcp_server(&server, &self.secrets, &self.user_id).await {
Ok(_token) => {
tracing::info!("MCP server '{}' authenticated via OAuth", name);
Ok(AuthResult::authenticated(name, ExtensionKind::McpServer))
}
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 {
Ok(result) => Ok(result),
Err(_) => 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(_) => {
// No OAuth, no DCR: fall back to manual token entry
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) => {
@@ -1845,12 +1725,8 @@ impl ExtensionManager {
}
}
/// Build an auth URL for MCP OAuth.
///
/// 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.
/// Build an auth URL for cases where non-interactive auth is needed
/// (e.g., running via Telegram where we can't open a browser).
async fn auth_mcp_build_url(
&self,
name: &str,
@@ -1859,153 +1735,60 @@ impl ExtensionManager {
// Try to discover OAuth metadata and build a URL the user can open manually
let metadata = discover_full_oauth_metadata(&server.url)
.await
.map_err(|e| match e {
crate::tools::mcp::auth::AuthError::NotSupported => {
ExtensionError::AuthNotSupported(e.to_string())
}
_ => ExtensionError::AuthFailed(e.to_string()),
})?;
.map_err(|e| ExtensionError::AuthFailed(e.to_string()))?;
use crate::cli::oauth_defaults;
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 {
// Try DCR if no client_id configured
let (client_id, redirect_uri) = if let Some(ref oauth) = server.oauth {
let port = find_available_port()
.await
.map_err(|e| ExtensionError::AuthFailed(e.to_string()))?;
format!("http://localhost:{}/callback", port.1)
};
// Try DCR if no client_id configured
let (client_id, client_secret) = if let Some(ref oauth) = server.oauth {
(oauth.client_id.clone(), None)
let redirect = format!("http://localhost:{}/callback", port.1);
(oauth.client_id.clone(), redirect)
} 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
.map_err(|e| ExtensionError::AuthFailed(e.to_string()))?;
(registration.client_id, None)
(registration.client_id, redirect)
} else {
return Err(ExtensionError::AuthNotSupported(
return Err(ExtensionError::AuthFailed(
"Server doesn't support OAuth or Dynamic Client Registration".to_string(),
));
};
// RFC 8707: resource parameter to scope the token to this MCP server
let resource = canonical_resource_uri(&server.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(
let pkce = PkceChallenge::generate();
let auth_url = build_authorization_url(
&metadata.authorization_endpoint,
&client_id,
&redirect_uri,
&scopes,
true, // Always use PKCE for MCP
&extra_params,
&metadata.scopes_supported,
Some(&pkce),
&std::collections::HashMap::new(),
None,
);
let expected_state = oauth_result.state;
let code_verifier = oauth_result.code_verifier;
if is_gateway {
// Gateway mode: store pending flow for the /oauth/callback handler.
oauth_defaults::sweep_expired_flows(&self.pending_oauth_flows).await;
// Platform routing: prepend instance name to state
let platform_state = oauth_defaults::build_platform_state(&expected_state);
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
},
// Store pending auth for later callback handling
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,
},
);
self.pending_oauth_flows
.write()
.await
.insert(expected_state, flow);
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(),
))
}
Ok(AuthResult::awaiting_authorization(
name,
ExtensionKind::McpServer,
auth_url,
"local".to_string(),
))
}
async fn auth_wasm_tool(&self, name: &str) -> Result<AuthResult, ExtensionError> {
@@ -2420,10 +2203,7 @@ impl ExtensionManager {
flows.retain(|_, flow| flow.extension_name != name);
}
let redirect_uri = self
.gateway_callback_redirect_uri()
.await
.unwrap_or_else(|| format!("{}/callback", oauth_defaults::callback_url()));
let redirect_uri = format!("{}/callback", oauth_defaults::callback_url());
// Merge scopes from all tools sharing this provider
let merged_scopes = self
@@ -2448,7 +2228,7 @@ impl ExtensionManager {
.clone()
.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
// `/oauth/callback` handler to complete the exchange. No TCP listener
// needed — the OAuth provider redirects to the gateway URL.
@@ -2484,8 +2264,6 @@ impl ExtensionManager {
secrets: Arc::clone(&self.secrets),
sse_sender: self.sse_sender.read().await.clone(),
gateway_token: self.gateway_token.clone(),
resource: None,
client_id_secret_name: None,
created_at: std::time::Instant::now(),
};
@@ -2827,17 +2605,11 @@ impl ExtensionManager {
.await
.map_err(|e| ExtensionError::ActivationFailed(e.to_string()))?;
// Try to list and create tools.
// A 401/auth error means the server requires OAuth — surface as
// AuthRequired so the activate handler triggers the OAuth flow.
let mcp_tools = client.list_tools().await.map_err(|e| {
let msg = e.to_string();
if msg.contains("requires authentication") || msg.contains("401") {
ExtensionError::AuthRequired
} else {
ExtensionError::ActivationFailed(msg)
}
})?;
// Try to list and create tools
let mcp_tools = client
.list_tools()
.await
.map_err(|e| ExtensionError::ActivationFailed(e.to_string()))?;
let tool_impls = client
.create_tools()
@@ -2884,17 +2656,6 @@ impl ExtensionManager {
});
}
// Check auth status — block activation if required secrets are missing.
// NeedsAuth (OAuth not yet completed) is allowed because configure() loads
// the tool first, then starts the OAuth flow to obtain the token.
let auth_state = self.check_tool_auth_status(name).await;
if auth_state == ToolAuthState::NeedsSetup {
return Err(ExtensionError::ActivationFailed(format!(
"Tool '{}' requires configuration. Use the setup form to provide credentials.",
name
)));
}
let runtime = self.wasm_tool_runtime.as_ref().ok_or_else(|| {
ExtensionError::ActivationFailed("WASM runtime not available".to_string())
})?;
@@ -4530,18 +4291,14 @@ mod tests {
// available" because the ExtensionManager had `wasm_tool_runtime: None`.
/// Build a minimal ExtensionManager suitable for unit tests.
fn make_test_manager_with_dirs(
fn make_test_manager(
wasm_runtime: Option<Arc<crate::tools::wasm::WasmToolRuntime>>,
tools_dir: std::path::PathBuf,
channels_dir: std::path::PathBuf,
) -> crate::extensions::manager::ExtensionManager {
use crate::secrets::{InMemorySecretsStore, SecretsCrypto};
use crate::tools::mcp::process::McpProcessManager;
use crate::tools::mcp::session::McpSessionManager;
std::fs::create_dir_all(&tools_dir).ok();
std::fs::create_dir_all(&channels_dir).ok();
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> =
@@ -4556,22 +4313,15 @@ mod tests {
tools,
None, // hooks
wasm_runtime,
tools_dir,
channels_dir,
None, // tunnel_url
tools_dir.clone(),
tools_dir, // channels dir (unused here)
None, // tunnel_url
"test".to_string(),
None, // db
vec![],
)
}
fn make_test_manager(
wasm_runtime: Option<Arc<crate::tools::wasm::WasmToolRuntime>>,
tools_dir: std::path::PathBuf,
) -> crate::extensions::manager::ExtensionManager {
make_test_manager_with_dirs(wasm_runtime, tools_dir.clone(), tools_dir)
}
#[tokio::test]
async fn test_activate_wasm_tool_with_runtime_passes_runtime_check() {
// When the ExtensionManager has a WASM runtime, activation should get
@@ -4924,145 +4674,6 @@ mod tests {
);
}
#[tokio::test]
async fn test_remove_wasm_tool_clears_pending_oauth_state_and_activation_error() {
let dir = tempfile::tempdir().expect("temp dir");
let mgr = make_test_manager(None, dir.path().to_path_buf());
std::fs::write(dir.path().join("gmail.wasm"), b"fake-tool").expect("write tool");
let listener = tokio::spawn(async {
std::future::pending::<()>().await;
});
let abort_handle = listener.abort_handle();
mgr.pending_auth.write().await.insert(
"gmail".to_string(),
super::PendingAuth {
_name: "gmail".to_string(),
_kind: ExtensionKind::WasmTool,
created_at: std::time::Instant::now(),
task_handle: Some(listener),
},
);
mgr.activation_errors
.write()
.await
.insert("gmail".to_string(), "cached failure".to_string());
let secrets = Arc::clone(&mgr.secrets);
mgr.pending_oauth_flows().write().await.insert(
"gmail-state".to_string(),
crate::cli::oauth_defaults::PendingOAuthFlow {
extension_name: "gmail".to_string(),
display_name: "Gmail".to_string(),
token_url: "https://example.com/token".to_string(),
client_id: "client123".to_string(),
client_secret: None,
redirect_uri: "https://example.com/oauth/callback".to_string(),
code_verifier: None,
access_token_field: "access_token".to_string(),
secret_name: "google_oauth_token".to_string(),
provider: None,
validation_endpoint: None,
scopes: vec![],
user_id: "test".to_string(),
secrets: Arc::clone(&secrets),
sse_sender: None,
gateway_token: None,
resource: None,
client_id_secret_name: None,
created_at: std::time::Instant::now(),
},
);
mgr.pending_oauth_flows().write().await.insert(
"other-state".to_string(),
crate::cli::oauth_defaults::PendingOAuthFlow {
extension_name: "web-search".to_string(),
display_name: "Web Search".to_string(),
token_url: "https://example.com/token".to_string(),
client_id: "client456".to_string(),
client_secret: None,
redirect_uri: "https://example.com/oauth/callback".to_string(),
code_verifier: None,
access_token_field: "access_token".to_string(),
secret_name: "other_token".to_string(),
provider: None,
validation_endpoint: None,
scopes: vec![],
user_id: "test".to_string(),
secrets,
sse_sender: None,
gateway_token: None,
resource: None,
client_id_secret_name: None,
created_at: std::time::Instant::now(),
},
);
let result = mgr.remove("gmail").await;
assert!(result.is_ok(), "remove should succeed: {:?}", result.err());
tokio::task::yield_now().await;
assert!(
mgr.pending_auth.read().await.get("gmail").is_none(),
"pending auth entry should be removed"
);
assert!(
abort_handle.is_finished(),
"pending auth listener should be aborted"
);
assert!(
!mgr.activation_errors.read().await.contains_key("gmail"),
"stale activation error should be cleared"
);
let flows = mgr.pending_oauth_flows().read().await;
assert!(
!flows.contains_key("gmail-state"),
"gateway OAuth flow for removed extension should be cleared"
);
assert!(
flows.contains_key("other-state"),
"unrelated pending OAuth flows should be retained"
);
}
#[tokio::test]
async fn test_remove_wasm_channel_clears_activation_error_and_deletes_files() {
let dir = tempfile::tempdir().expect("temp dir");
let tools_dir = dir.path().join("tools");
let channels_dir = dir.path().join("channels");
let mgr = make_test_manager_with_dirs(None, tools_dir, channels_dir.clone());
let wasm_path = channels_dir.join("telegram.wasm");
let cap_path = channels_dir.join("telegram.capabilities.json");
std::fs::write(&wasm_path, b"fake-channel").expect("write channel");
std::fs::write(&cap_path, b"{}").expect("write capabilities");
mgr.activation_errors
.write()
.await
.insert("telegram".to_string(), "channel failed".to_string());
let result = mgr.remove("telegram").await;
assert!(result.is_ok(), "remove should succeed: {:?}", result.err());
assert!(
!mgr.activation_errors.read().await.contains_key("telegram"),
"channel activation error should be cleared on remove"
);
assert!(
!wasm_path.exists(),
"channel wasm file should be deleted on remove"
);
assert!(
!cap_path.exists(),
"channel capabilities file should be deleted on remove"
);
}
#[test]
fn test_sanitize_url_with_query_params() {
let url = "https://api.example.com/path?api_key=secret123&token=abc";
@@ -5155,189 +4766,6 @@ mod tests {
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) ─────────
#[tokio::test]
@@ -5487,6 +4915,7 @@ mod tests {
"configure should have stored the relay stream token"
);
}
#[test]
fn test_validation_failed_is_distinct_error_variant() {
// Regression: ValidationFailed must be a distinct error variant so
-3
View File
@@ -517,9 +517,6 @@ pub enum ExtensionError {
#[error("Authentication failed: {0}")]
AuthFailed(String),
#[error("Server does not support OAuth: {0}")]
AuthNotSupported(String),
#[error("Activation failed: {0}")]
ActivationFailed(String),
+5 -13
View File
@@ -9,8 +9,8 @@ use ironclaw::{
agent::{Agent, AgentDeps},
app::{AppBuilder, AppBuilderFlags},
channels::{
ChannelManager, GatewayChannel, HttpChannel, ReplChannel, SignalChannel, WebhookServer,
WebhookServerConfig,
ChannelManager, ChannelSecretUpdater, GatewayChannel, HttpChannel, ReplChannel,
SignalChannel, WebhookServer, WebhookServerConfig,
wasm::{WasmChannelRouter, WasmChannelRuntime},
web::log_layer::LogBroadcaster,
},
@@ -433,8 +433,9 @@ async fn async_main() -> anyhow::Result<()> {
"Lifecycle hooks initialized"
);
// Reuse the shared agent session manager prepared by AppBuilder.
let session_manager = Arc::clone(&components.agent_session_manager);
// Create session manager (shared between agent and web gateway)
let session_manager =
Arc::new(ironclaw::agent::SessionManager::new().with_hooks(components.hooks.clone()));
// Lazy scheduler slot — filled after Agent::new creates the Scheduler.
// 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_tool_registry(Arc::clone(&components.tools));
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));
}
if !components.catalog_entries.is_empty() {
@@ -736,7 +729,6 @@ async fn async_main() -> anyhow::Result<()> {
#[cfg(unix)]
{
use ironclaw::channels::ChannelSecretUpdater;
// Collect all channels that support secret updates
let mut secret_updaters: Vec<Arc<dyn ChannelSecretUpdater>> = Vec::new();
if let Some(ref state) = http_channel_state {
+10 -30
View File
@@ -65,20 +65,7 @@ fn install_macos() -> Result<()> {
let stdout = logs_dir.join("daemon.stdout.log");
let stderr = logs_dir.join("daemon.stderr.log");
let plist = macos_plist_content(
&exe.display().to_string(),
&stdout.display().to_string(),
&stderr.display().to_string(),
);
std::fs::write(&file, plist)?;
println!("Installed launchd service: {}", file.display());
println!(" Start with: ironclaw service start");
Ok(())
}
fn macos_plist_content(exe: &str, stdout: &str, stderr: &str) -> String {
format!(
let plist = format!(
r#"<?xml version="1.0" encoding="UTF-8"?>
<!DOCTYPE plist PUBLIC "-//Apple//DTD PLIST 1.0//EN" "http://www.apple.com/DTDs/PropertyList-1.0.dtd">
<plist version="1.0">
@@ -94,11 +81,6 @@ fn macos_plist_content(exe: &str, stdout: &str, stderr: &str) -> String {
<true/>
<key>KeepAlive</key>
<true/>
<key>EnvironmentVariables</key>
<dict>
<key>CLI_ENABLED</key>
<string>false</string>
</dict>
<key>StandardOutPath</key>
<string>{stdout}</string>
<key>StandardErrorPath</key>
@@ -107,10 +89,15 @@ fn macos_plist_content(exe: &str, stdout: &str, stderr: &str) -> String {
</plist>
"#,
label = SERVICE_LABEL,
exe = xml_escape(exe),
stdout = xml_escape(stdout),
stderr = xml_escape(stderr),
)
exe = xml_escape(&exe.display().to_string()),
stdout = xml_escape(&stdout.display().to_string()),
stderr = xml_escape(&stderr.display().to_string()),
);
std::fs::write(&file, plist)?;
println!("Installed launchd service: {}", file.display());
println!(" Start with: ironclaw service start");
Ok(())
}
fn install_linux() -> Result<()> {
@@ -369,11 +356,4 @@ mod tests {
let s = path.to_string_lossy();
assert!(s.ends_with(".ironclaw/logs"), "unexpected path: {s}");
}
#[test]
fn macos_plist_sets_cli_enabled_false() {
let plist = macos_plist_content("/tmp/ironclaw", "/tmp/stdout.log", "/tmp/stderr.log");
assert!(plist.contains("<key>EnvironmentVariables</key>"));
assert!(plist.contains(" <key>CLI_ENABLED</key>\n <string>false</string>"));
}
}
-4
View File
@@ -1067,8 +1067,6 @@ mod tests {
prompt: "Check status".to_string(),
context_paths: vec![],
max_tokens: 500,
use_tools: false,
max_tool_rounds: 3,
},
guardrails: RoutineGuardrails {
cooldown: std::time::Duration::from_secs(60),
@@ -1200,8 +1198,6 @@ mod tests {
prompt: "test".to_string(),
context_paths: vec![],
max_tokens: 100,
use_tools: false,
max_tool_rounds: 3,
},
guardrails: RoutineGuardrails {
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 {
// In gateway mode, tool_auth only returns an auth URL for the frontend
// to open — no browser is launched server-side, so no approval needed.
if self.manager.should_use_gateway_mode() {
ApprovalRequirement::Never
} else {
ApprovalRequirement::UnlessAutoApproved
}
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]
fn test_tool_upgrade_schema() {
use crate::tools::tool::ApprovalRequirement;
+4
View File
@@ -397,6 +397,10 @@ impl Tool for ListDirTool {
false // Directory listings are safe
}
fn requires_approval(&self, _params: &serde_json::Value) -> ApprovalRequirement {
ApprovalRequirement::UnlessAutoApproved
}
fn domain(&self) -> ToolDomain {
ToolDomain::Container
}
+69 -246
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.
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.
const MAX_REDIRECTS: usize = 3;
@@ -250,120 +244,43 @@ fn is_html_response(headers: &HashMap<String, String>) -> bool {
fn parse_headers_param(
headers: Option<&serde_json::Value>,
) -> 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 {
None => Ok(Vec::new()),
Some(serde_json::Value::String(raw)) => {
let trimmed = raw.trim();
if trimmed.is_empty() {
return Ok(Vec::new());
}
let parsed = serde_json::from_str::<serde_json::Value>(trimmed).map_err(|e| {
ToolError::InvalidParameters(format!(
"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(),
)),
Some(serde_json::Value::Object(map)) => {
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)
}
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(
"'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).
fn extract_host_from_params(params: &serde_json::Value) -> Option<String> {
params
@@ -398,7 +315,7 @@ impl Tool for HttpTool {
"method": {
"type": "string",
"enum": ["GET", "POST", "PUT", "DELETE", "PATCH"],
"description": "HTTP method (default: GET)"
"description": "HTTP method"
},
"url": {
"type": "string",
@@ -429,7 +346,7 @@ impl Tool for HttpTool {
"description": "Save response body as raw bytes to this file path instead of returning it. Use for binary downloads (images, PDFs, etc.). The path must be under /tmp/."
}
},
"required": ["url"]
"required": ["method", "url"]
})
}
@@ -440,8 +357,7 @@ impl Tool for HttpTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let method = params["method"].as_str().unwrap_or("GET");
let method_upper = method.to_uppercase();
let method = require_str(&params, "method")?;
let url = require_str(&params, "url")?;
let mut parsed_url = validate_url(url)?;
@@ -463,9 +379,6 @@ impl Tool for HttpTool {
// Parse 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
let mut request = match method.to_uppercase().as_str() {
@@ -482,8 +395,6 @@ impl Tool for HttpTool {
}
};
request = request.timeout(effective_timeout);
// Add headers
for (key, value) in &headers_vec {
request = request.header(key.as_str(), value.as_str());
@@ -492,9 +403,7 @@ impl Tool for HttpTool {
// Add body if present
let body_bytes = if let Some(body) = params.get("body") {
if let Some(body_str) = body.as_str() {
if body_str.is_empty() {
None
} else if let Ok(json_body) = serde_json::from_str::<serde_json::Value>(body_str) {
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| {
ToolError::InvalidParameters(format!("invalid body JSON: {}", e))
})?;
@@ -559,7 +468,7 @@ impl Tool for HttpTool {
// Build the interceptor request descriptor for recording/replay
let intercept_req = crate::llm::recording::HttpExchangeRequest {
method: method_upper,
method: method.to_uppercase(),
url: parsed_url.to_string(),
headers: headers_vec.clone(),
body: body_bytes
@@ -601,7 +510,7 @@ impl Tool for HttpTool {
let hop_client = build_pinned_client(
&hop_host,
&hop_addrs,
effective_timeout,
Duration::from_secs(30),
reqwest::redirect::Policy::none(),
)?;
@@ -615,7 +524,7 @@ impl Tool for HttpTool {
.await
.map_err(|e| {
if e.is_timeout() {
ToolError::Timeout(effective_timeout)
ToolError::Timeout(Duration::from_secs(30))
} else {
ToolError::ExternalService(e.to_string())
}
@@ -679,7 +588,7 @@ impl Tool for HttpTool {
} else {
let resp = request.send().await.map_err(|e| {
if e.is_timeout() {
ToolError::Timeout(effective_timeout)
ToolError::Timeout(Duration::from_secs(30))
} else {
ToolError::ExternalService(e.to_string())
}
@@ -707,7 +616,7 @@ impl Tool for HttpTool {
.collect();
// 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 {
MAX_SAVE_TO_SIZE
} else {
@@ -752,11 +661,11 @@ impl Tool for HttpTool {
let body_bytes = bytes::Bytes::from(body);
// If save_to is specified, write raw bytes to file and return metadata.
if let Some(save_to) = save_to {
let saved_to = save_to.clone();
if let Some(save_to) = params.get("save_to").and_then(|v| v.as_str()) {
let save_to_owned = save_to.to_string();
let bytes_clone = body_bytes.clone();
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| {
ToolError::ExecutionFailed(format!("failed to write file: {}", e))
})?;
@@ -767,7 +676,7 @@ impl Tool for HttpTool {
.map_err(|e: ToolError| e)?;
let result = serde_json::json!({
"status": status,
"saved_to": saved_to,
"saved_to": save_to,
"size_bytes": body_bytes.len(),
"headers": headers,
});
@@ -829,22 +738,18 @@ impl Tool for HttpTool {
}
fn requires_approval(&self, params: &serde_json::Value) -> ApprovalRequirement {
let has_credentials = crate::safety::params_contain_manual_credentials(params)
|| (self.credential_registry.as_ref().is_some_and(|registry| {
extract_host_from_params(params)
.is_some_and(|host| registry.has_credentials_for_host(&host))
}));
if has_credentials {
// 1. Manual auth headers/query params in LLM params
if crate::safety::params_contain_manual_credentials(params) {
return ApprovalRequirement::Always;
}
// GET requests (or missing method, since GET is the default) are low-risk
let method = params["method"].as_str().unwrap_or("GET");
if method.eq_ignore_ascii_case("GET") {
return ApprovalRequirement::Never;
// 2. Target host has credential mappings (will be auto-injected)
if let Some(ref registry) = self.credential_registry
&& let Some(host) = extract_host_from_params(params)
&& registry.has_credentials_for_host(&host)
{
return ApprovalRequirement::Always;
}
// Default: outbound HTTP still needs approval unless auto-approved
ApprovalRequirement::UnlessAutoApproved
}
@@ -982,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]
fn test_http_tool_schema_body_is_freeform() {
let schema = HttpTool::new().parameters_schema();
@@ -1067,22 +907,12 @@ mod tests {
// ── Approval requirement tests ──────────────────────────────────────
#[test]
fn test_get_no_auth_headers_returns_never() {
fn test_no_auth_headers_returns_unless_auto_approved() {
let tool = HttpTool::new();
let params = serde_json::json!({
"method": "GET",
"url": "https://api.example.com/data"
});
assert_eq!(tool.requires_approval(&params), ApprovalRequirement::Never);
}
#[test]
fn test_post_no_auth_headers_returns_unless_auto_approved() {
let tool = HttpTool::new();
let params = serde_json::json!({
"method": "POST",
"url": "https://api.example.com/data"
});
assert_eq!(
tool.requires_approval(&params),
ApprovalRequirement::UnlessAutoApproved
@@ -1166,18 +996,21 @@ mod tests {
}
#[test]
fn test_get_non_auth_headers_return_never() {
fn test_non_auth_headers_return_unless_auto_approved() {
let tool = HttpTool::new();
let params = serde_json::json!({
"method": "GET",
"url": "https://example.com",
"headers": {"Content-Type": "application/json", "Accept": "text/html"}
});
assert_eq!(tool.requires_approval(&params), ApprovalRequirement::Never);
assert_eq!(
tool.requires_approval(&params),
ApprovalRequirement::UnlessAutoApproved
);
}
#[test]
fn test_get_empty_headers_return_never() {
fn test_empty_headers_return_unless_auto_approved() {
let tool = HttpTool::new();
// Empty object
@@ -1186,7 +1019,10 @@ mod tests {
"url": "https://example.com",
"headers": {}
});
assert_eq!(tool.requires_approval(&params), ApprovalRequirement::Never);
assert_eq!(
tool.requires_approval(&params),
ApprovalRequirement::UnlessAutoApproved
);
// Empty array
let params = serde_json::json!({
@@ -1194,7 +1030,10 @@ mod tests {
"url": "https://example.com",
"headers": []
});
assert_eq!(tool.requires_approval(&params), ApprovalRequirement::Never);
assert_eq!(
tool.requires_approval(&params),
ApprovalRequirement::UnlessAutoApproved
);
}
// ── Credential registry approval tests ─────────────────────────────
@@ -1224,7 +1063,7 @@ mod tests {
}
#[test]
fn test_get_host_without_credential_mapping_returns_never() {
fn test_host_without_credential_mapping_returns_unless_auto_approved() {
use crate::tools::wasm::SharedCredentialRegistry;
let registry = Arc::new(SharedCredentialRegistry::new());
@@ -1236,7 +1075,10 @@ mod tests {
"method": "GET",
"url": "https://api.example.com/data"
});
assert_eq!(tool.requires_approval(&params), ApprovalRequirement::Never);
assert_eq!(
tool.requires_approval(&params),
ApprovalRequirement::UnlessAutoApproved
);
}
#[test]
@@ -1277,25 +1119,6 @@ mod tests {
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 ─────────────────────────────────────────────
#[tokio::test]
+7 -4
View File
@@ -8,7 +8,7 @@ use secrecy::{ExposeSecret, SecretString};
use crate::context::JobContext;
use crate::tools::builtin::path_utils::validate_path;
use crate::tools::tool::{Tool, ToolError, ToolOutput};
use crate::tools::tool::{ApprovalRequirement, Tool, ToolError, ToolOutput};
/// Tool for analyzing images using a vision-capable model.
pub struct ImageAnalyzeTool {
@@ -86,6 +86,10 @@ impl Tool for ImageAnalyzeTool {
})
}
fn requires_approval(&self, _params: &serde_json::Value) -> ApprovalRequirement {
ApprovalRequirement::UnlessAutoApproved
}
fn requires_sanitization(&self) -> bool {
true
}
@@ -181,7 +185,6 @@ impl Tool for ImageAnalyzeTool {
mod tests {
use super::super::media_type_from_path;
use super::*;
use crate::tools::tool::ApprovalRequirement;
use tempfile::TempDir;
#[test]
@@ -196,7 +199,7 @@ mod tests {
}
#[test]
fn test_requires_approval_returns_never() {
fn test_requires_approval_returns_unless_auto_approved() {
let tool = ImageAnalyzeTool::new(
"https://api.example.com".to_string(),
"test-key".to_string(),
@@ -205,7 +208,7 @@ mod tests {
);
assert_eq!(
tool.requires_approval(&serde_json::json!({})),
ApprovalRequirement::Never
ApprovalRequirement::UnlessAutoApproved
);
}
+6 -3
View File
@@ -7,7 +7,7 @@ use secrecy::{ExposeSecret, SecretString};
use crate::context::JobContext;
use crate::tools::builtin::path_utils::validate_path;
use crate::tools::tool::{Tool, ToolError, ToolOutput};
use crate::tools::tool::{ApprovalRequirement, Tool, ToolError, ToolOutput};
/// Tool for editing images using an AI image editing API.
pub struct ImageEditTool {
@@ -85,6 +85,10 @@ impl Tool for ImageEditTool {
})
}
fn requires_approval(&self, _params: &serde_json::Value) -> ApprovalRequirement {
ApprovalRequirement::UnlessAutoApproved
}
fn requires_sanitization(&self) -> bool {
false
}
@@ -262,7 +266,6 @@ impl ImageEditTool {
#[cfg(test)]
mod tests {
use super::*;
use crate::tools::tool::ApprovalRequirement;
use tempfile::TempDir;
#[test]
@@ -277,7 +280,7 @@ mod tests {
assert!(!tool.requires_sanitization());
assert_eq!(
tool.requires_approval(&serde_json::json!({})),
ApprovalRequirement::Never
ApprovalRequirement::UnlessAutoApproved
);
}
+6 -2
View File
@@ -5,6 +5,7 @@ use secrecy::{ExposeSecret, SecretString};
use serde::{Deserialize, Serialize};
use crate::context::JobContext;
use crate::tools::tool::ApprovalRequirement;
use crate::tools::{Tool, ToolError, ToolOutput};
/// Tool for generating images using FLUX or compatible image generation APIs.
@@ -86,6 +87,10 @@ impl Tool for ImageGenerateTool {
})
}
fn requires_approval(&self, _params: &serde_json::Value) -> ApprovalRequirement {
ApprovalRequirement::UnlessAutoApproved
}
fn requires_sanitization(&self) -> bool {
false
}
@@ -181,7 +186,6 @@ impl Tool for ImageGenerateTool {
#[cfg(test)]
mod tests {
use super::*;
use crate::tools::tool::ApprovalRequirement;
#[test]
fn test_tool_metadata() {
@@ -193,7 +197,7 @@ mod tests {
assert_eq!(tool.name(), "image_generate");
assert_eq!(
tool.requires_approval(&serde_json::json!({})),
ApprovalRequirement::Never
ApprovalRequirement::UnlessAutoApproved
);
let schema = tool.parameters_schema();
+61 -122
View File
@@ -12,7 +12,6 @@
//! Use `memory_write` to persist important facts that should be remembered
//! across sessions.
use std::path::Path;
use std::sync::Arc;
use async_trait::async_trait;
@@ -27,28 +26,6 @@ use crate::workspace::{Workspace, paths};
const PROTECTED_IDENTITY_FILES: &[&str] =
&[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.
///
/// 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, \
'daily_log' for timestamped session notes, 'heartbeat' for the periodic \
checklist (HEARTBEAT.md), 'bootstrap' to clear the first-run ritual file, \
or provide a custom workspace path for arbitrary file creation. \
Never pass absolute filesystem paths like '/Users/...' or 'C:\\...'."
or provide a custom path for arbitrary file creation."
}
fn parameters_schema(&self) -> serde_json::Value {
@@ -207,14 +183,6 @@ impl Tool for MemoryWriteTool {
.and_then(|v| v.as_str())
.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.
// Handled early because it accepts empty content (unlike other targets).
if target == "bootstrap" {
@@ -364,8 +332,7 @@ impl Tool for MemoryReadTool {
fn description(&self) -> &str {
"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 read_file for those). Do not pass absolute paths like '/Users/...' or 'C:\\...'. \
Works with identity files, heartbeat checklist, \
(use read_file for those). Works with identity files, heartbeat checklist, \
memory, daily logs, or any custom workspace path."
}
@@ -391,14 +358,6 @@ impl Tool for MemoryReadTool {
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
.workspace
.read(path)
@@ -539,100 +498,80 @@ impl Tool for MemoryTreeTool {
}
}
#[cfg(test)]
#[cfg(all(test, feature = "postgres"))]
mod tests {
use super::*;
#[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(feature = "postgres")]
mod postgres_schema_tests {
use super::*;
fn make_test_workspace() -> Arc<Workspace> {
Arc::new(Workspace::new(
"test_user",
deadpool_postgres::Pool::builder(deadpool_postgres::Manager::new(
tokio_postgres::Config::new(),
tokio_postgres::NoTls,
))
.build()
.unwrap(),
fn make_test_workspace() -> Arc<Workspace> {
Arc::new(Workspace::new(
"test_user",
deadpool_postgres::Pool::builder(deadpool_postgres::Manager::new(
tokio_postgres::Config::new(),
tokio_postgres::NoTls,
))
}
.build()
.unwrap(),
))
}
#[test]
fn test_memory_search_schema() {
let workspace = make_test_workspace();
let tool = MemorySearchTool::new(workspace);
#[test]
fn test_memory_search_schema() {
let workspace = make_test_workspace();
let tool = MemorySearchTool::new(workspace);
assert_eq!(tool.name(), "memory_search");
assert!(!tool.requires_sanitization());
assert_eq!(tool.name(), "memory_search");
assert!(!tool.requires_sanitization());
let schema = tool.parameters_schema();
assert!(schema["properties"]["query"].is_object());
assert!(
schema["required"]
.as_array()
.unwrap()
.contains(&"query".into())
);
}
let schema = tool.parameters_schema();
assert!(schema["properties"]["query"].is_object());
assert!(
schema["required"]
.as_array()
.unwrap()
.contains(&"query".into())
);
}
#[test]
fn test_memory_write_schema() {
let workspace = make_test_workspace();
let tool = MemoryWriteTool::new(workspace);
#[test]
fn test_memory_write_schema() {
let workspace = make_test_workspace();
let tool = MemoryWriteTool::new(workspace);
assert_eq!(tool.name(), "memory_write");
assert_eq!(tool.name(), "memory_write");
let schema = tool.parameters_schema();
assert!(schema["properties"]["content"].is_object());
assert!(schema["properties"]["target"].is_object());
assert!(schema["properties"]["append"].is_object());
}
let schema = tool.parameters_schema();
assert!(schema["properties"]["content"].is_object());
assert!(schema["properties"]["target"].is_object());
assert!(schema["properties"]["append"].is_object());
}
#[test]
fn test_memory_read_schema() {
let workspace = make_test_workspace();
let tool = MemoryReadTool::new(workspace);
#[test]
fn test_memory_read_schema() {
let workspace = make_test_workspace();
let tool = MemoryReadTool::new(workspace);
assert_eq!(tool.name(), "memory_read");
assert_eq!(tool.name(), "memory_read");
let schema = tool.parameters_schema();
assert!(schema["properties"]["path"].is_object());
assert!(
schema["required"]
.as_array()
.unwrap()
.contains(&"path".into())
);
}
let schema = tool.parameters_schema();
assert!(schema["properties"]["path"].is_object());
assert!(
schema["required"]
.as_array()
.unwrap()
.contains(&"path".into())
);
}
#[test]
fn test_memory_tree_schema() {
let workspace = make_test_workspace();
let tool = MemoryTreeTool::new(workspace);
#[test]
fn test_memory_tree_schema() {
let workspace = make_test_workspace();
let tool = MemoryTreeTool::new(workspace);
assert_eq!(tool.name(), "memory_tree");
assert_eq!(tool.name(), "memory_tree");
let schema = tool.parameters_schema();
assert!(schema["properties"]["path"].is_object());
assert!(schema["properties"]["depth"].is_object());
assert_eq!(schema["properties"]["depth"]["default"], 1);
}
let schema = tool.parameters_schema();
assert!(schema["properties"]["path"].is_object());
assert!(schema["properties"]["depth"].is_object());
assert_eq!(schema["properties"]["depth"]["default"], 1);
}
}
-2
View File
@@ -15,7 +15,6 @@ pub mod secrets_tools;
pub(crate) mod shell;
pub mod skill_tools;
mod time;
mod tool_info;
pub use echo::EchoTool;
pub use extension_tools::{
@@ -40,7 +39,6 @@ pub use secrets_tools::{SecretDeleteTool, SecretListTool};
pub use shell::ShellTool;
pub use skill_tools::{SkillInstallTool, SkillListTool, SkillRemoveTool, SkillSearchTool};
pub use time::TimeTool;
pub use tool_info::ToolInfoTool;
mod html_converter;
pub mod image_analyze;
pub mod image_edit;
-21
View File
@@ -104,14 +104,6 @@ impl Tool for RoutineCreateTool {
"enum": ["lightweight", "full_job"],
"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": {
"type": "integer",
"description": "Minimum seconds between fires (default: 300)"
@@ -270,24 +262,11 @@ impl Tool for RoutineCreateTool {
})
.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 {
"lightweight" => RoutineAction::Lightweight {
prompt: prompt.to_string(),
context_paths,
max_tokens: 4096,
use_tools,
max_tool_rounds,
},
"full_job" => {
let tool_permissions = crate::agent::routine::parse_tool_permissions(&params);
-183
View File
@@ -1,183 +0,0 @@
//! On-demand tool discovery (like CLI `--help`).
//!
//! Two levels of detail:
//! - Default: name, description, parameter names (compact ~150 bytes)
//! - `include_schema: true`: adds the full typed JSON Schema
//!
//! Keeps the tools array compact (WASM tools use permissive schemas)
//! while allowing precise discovery when needed.
use std::sync::Weak;
use async_trait::async_trait;
use crate::context::JobContext;
use crate::tools::registry::ToolRegistry;
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
pub struct ToolInfoTool {
registry: Weak<ToolRegistry>,
}
impl ToolInfoTool {
pub fn new(registry: Weak<ToolRegistry>) -> Self {
Self { registry }
}
}
#[async_trait]
impl Tool for ToolInfoTool {
fn name(&self) -> &str {
"tool_info"
}
fn description(&self) -> &str {
"Get info about any tool: description and parameter names. \
Set include_schema to true for the full typed parameter schema."
}
fn parameters_schema(&self) -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": {
"name": {
"type": "string",
"description": "Name of the tool to get info about"
},
"include_schema": {
"type": "boolean",
"description": "If true, include the full typed JSON Schema for parameters (larger response). Default: false.",
"default": false
}
},
"required": ["name"]
})
}
async fn execute(
&self,
params: serde_json::Value,
_ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let name = require_str(&params, "name")?;
let include_schema = params
.get("include_schema")
.and_then(|v| v.as_bool())
.unwrap_or(false);
let registry = self.registry.upgrade().ok_or_else(|| {
ToolError::ExecutionFailed(
"tool registry is no longer available for tool_info".to_string(),
)
})?;
let tool = registry.get(name).await.ok_or_else(|| {
ToolError::InvalidParameters(format!("No tool named '{name}' is registered"))
})?;
let schema = tool.discovery_schema();
// Extract just param names from the schema's "properties" keys
let param_names: Vec<&str> = schema
.get("properties")
.and_then(|p| p.as_object())
.map(|props| props.keys().map(|k| k.as_str()).collect())
.unwrap_or_default();
let mut info = serde_json::json!({
"name": tool.name(),
"description": tool.description(),
"parameters": param_names,
});
if include_schema {
info["schema"] = schema;
}
Ok(ToolOutput::success(info, start.elapsed()))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::tools::builtin::EchoTool;
use std::sync::Arc;
#[tokio::test]
async fn test_tool_info_default_returns_param_names() {
let registry = Arc::new(ToolRegistry::new());
registry.register(Arc::new(EchoTool)).await;
let tool = ToolInfoTool::new(Arc::downgrade(&registry));
let ctx = JobContext::default();
let result = tool
.execute(serde_json::json!({"name": "echo"}), &ctx)
.await
.unwrap();
let info = &result.result;
assert_eq!(info["name"], "echo");
assert!(!info["description"].as_str().unwrap().is_empty());
// Default: parameters is an array of names, not the full schema
assert!(info["parameters"].is_array());
assert!(
info["parameters"]
.as_array()
.unwrap()
.iter()
.any(|v| v.as_str() == Some("message")),
"echo tool should have 'message' parameter: {:?}",
info["parameters"]
);
// No schema field by default
assert!(info.get("schema").is_none());
}
#[tokio::test]
async fn test_tool_info_with_schema() {
let registry = Arc::new(ToolRegistry::new());
registry.register(Arc::new(EchoTool)).await;
let tool = ToolInfoTool::new(Arc::downgrade(&registry));
let ctx = JobContext::default();
let result = tool
.execute(
serde_json::json!({"name": "echo", "include_schema": true}),
&ctx,
)
.await
.unwrap();
let info = &result.result;
assert_eq!(info["name"], "echo");
// With include_schema: true, schema field should be present
assert!(info["schema"].is_object());
assert!(info["schema"]["properties"].is_object());
}
#[tokio::test]
async fn test_tool_info_unknown_tool() {
let registry = Arc::new(ToolRegistry::new());
let tool = ToolInfoTool::new(Arc::downgrade(&registry));
let ctx = JobContext::default();
let result = tool
.execute(serde_json::json!({"name": "nonexistent"}), &ctx)
.await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_tool_info_registry_dropped() {
let registry = Arc::new(ToolRegistry::new());
let tool = ToolInfoTool::new(Arc::downgrade(&registry));
drop(registry);
let ctx = JobContext::default();
let result = tool
.execute(serde_json::json!({"name": "echo"}), &ctx)
.await;
assert!(matches!(result, Err(ToolError::ExecutionFailed(_))));
}
}
+2 -77
View File
@@ -669,7 +669,7 @@ pub async fn authorize_mcp_server(
}
// 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 {
// Pre-configured OAuth
let (auth_url, tok_url) = discover_oauth_endpoints(server_config).await?;
@@ -711,13 +711,6 @@ pub async fn authorize_mcp_server(
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
let resource = canonical_resource_uri(&server_config.url);
@@ -748,10 +741,7 @@ pub async fn authorize_mcp_server(
println!(" Waiting for authorization...");
// Wait for callback. State is sent in the URL for servers that require it
// (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.
// Wait for callback
let code = wait_for_authorization_callback(listener, &server_config.name).await?;
println!(" Exchanging code for token...");
@@ -1721,69 +1711,4 @@ mod tests {
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 mcp;
pub mod rate_limiter;
pub mod redaction;
pub mod schema_validator;
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 -45
View File
@@ -23,7 +23,7 @@ use crate::tools::builtin::{
ToolUpgradeTool, WriteFileTool,
};
use crate::tools::rate_limiter::RateLimiter;
use crate::tools::tool::{ApprovalRequirement, Tool, ToolDomain};
use crate::tools::tool::{Tool, ToolDomain};
use crate::tools::wasm::{
Capabilities, OAuthRefreshConfig, ResourceLimits, SharedCredentialRegistry, WasmError,
WasmStorageError, WasmToolRuntime, WasmToolStore, WasmToolWrapper,
@@ -75,7 +75,6 @@ const PROTECTED_TOOL_NAMES: &[&str] = &[
"image_generate",
"image_edit",
"image_analyze",
"tool_info",
];
/// Registry of available tools.
@@ -246,17 +245,6 @@ impl ToolRegistry {
tracing::debug!("Registered {} built-in tools", self.count());
}
/// Register the `tool_info` discovery tool.
///
/// Requires `Arc<Self>` so the tool can query the registry for other tools'
/// schemas at runtime. Call after `register_builtin_tools()`.
pub fn register_tool_info(self: &Arc<Self>) {
use crate::tools::builtin::ToolInfoTool;
let tool = ToolInfoTool::new(Arc::downgrade(self));
self.register_sync(Arc::new(tool));
tracing::debug!("Registered tool_info discovery tool");
}
/// Register only orchestrator-domain tools (safe for the main process).
///
/// This registers tools that don't touch the filesystem or run shell commands:
@@ -290,38 +278,6 @@ impl ToolRegistry {
.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.
///
/// These tools provide shell access, file operations, and code editing
-11
View File
@@ -336,17 +336,6 @@ pub trait Tool: Send + Sync {
None
}
/// Full parameter schema for discovery and coercion purposes.
///
/// Unlike `parameters_schema()` (which may be permissive to keep the tools
/// array compact), this returns the complete typed schema. Used by the
/// `tool_info` built-in and by WASM parameter coercion.
///
/// Default: delegates to `parameters_schema()`.
fn discovery_schema(&self) -> serde_json::Value {
self.parameters_schema()
}
/// Get the tool schema for LLM function calling.
fn schema(&self) -> ToolSchema {
ToolSchema {
+84 -6
View File
@@ -1,5 +1,7 @@
//! WASM sandbox error types.
use std::fmt;
use thiserror::Error;
/// Errors that can occur during WASM tool execution.
@@ -66,13 +68,13 @@ pub enum WasmError {
Timeout(std::time::Duration),
/// Component returned an error response.
/// When `hint` is non-empty it points the LLM to `tool_info` so it can
/// fetch the tool's full parameter schema on demand.
/// When `hint` is non-empty it carries the tool's description and parameter
/// schema so the LLM can retry with correct arguments.
#[error("Tool error: {message}{}", if hint.is_empty() { String::new() } else { format!("\n\nTool usage hint:\n{hint}") })]
ToolReturnedError {
/// The error message from the WASM tool.
message: String,
/// Optional retry hint (empty when unavailable).
/// Optional description + schema hint (empty when unavailable).
hint: String,
},
@@ -97,9 +99,73 @@ impl From<WasmError> for crate::tools::ToolError {
}
}
/// Details about a trap that occurred during execution.
#[derive(Debug, Clone)]
pub struct TrapInfo {
/// Human-readable trap message.
pub message: String,
/// Trap code if available.
pub code: Option<TrapCode>,
}
impl fmt::Display for TrapInfo {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match &self.code {
Some(code) => write!(f, "{}: {}", code, self.message),
None => write!(f, "{}", self.message),
}
}
}
/// Known trap codes from Wasmtime.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TrapCode {
/// Out of bounds memory access.
MemoryOutOfBounds,
/// Out of bounds table access.
TableOutOfBounds,
/// Indirect call type mismatch.
IndirectCallToNull,
/// Signature mismatch on indirect call.
BadSignature,
/// Integer overflow.
IntegerOverflow,
/// Integer division by zero.
IntegerDivisionByZero,
/// Invalid conversion to integer.
BadConversionToInteger,
/// Unreachable instruction executed.
UnreachableCodeReached,
/// Call stack exhausted.
StackOverflow,
/// Out of fuel.
OutOfFuel,
/// Unknown trap code.
Unknown,
}
impl fmt::Display for TrapCode {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let s = match self {
TrapCode::MemoryOutOfBounds => "memory out of bounds",
TrapCode::TableOutOfBounds => "table out of bounds",
TrapCode::IndirectCallToNull => "indirect call to null",
TrapCode::BadSignature => "bad signature",
TrapCode::IntegerOverflow => "integer overflow",
TrapCode::IntegerDivisionByZero => "integer division by zero",
TrapCode::BadConversionToInteger => "bad conversion to integer",
TrapCode::UnreachableCodeReached => "unreachable code reached",
TrapCode::StackOverflow => "stack overflow",
TrapCode::OutOfFuel => "out of fuel",
TrapCode::Unknown => "unknown trap",
};
write!(f, "{}", s)
}
}
#[cfg(test)]
mod tests {
use crate::tools::wasm::error::WasmError;
use crate::tools::wasm::error::{TrapCode, TrapInfo, WasmError};
#[test]
fn test_error_display() {
@@ -114,6 +180,17 @@ mod tests {
assert!(err.to_string().contains("10000000"));
}
#[test]
fn test_trap_info_display() {
let info = TrapInfo {
message: "access at offset 0x1000".to_string(),
code: Some(TrapCode::MemoryOutOfBounds),
};
let s = info.to_string();
assert!(s.contains("memory out of bounds"));
assert!(s.contains("access at offset"));
}
#[test]
fn test_conversion_to_tool_error() {
let wasm_err = WasmError::Trapped("test trap".to_string());
@@ -141,11 +218,12 @@ mod tests {
fn test_tool_returned_error_with_hint() {
let err = WasmError::ToolReturnedError {
message: "unknown action: foobar".to_string(),
hint: "Tip: call tool_info(name: \"gmail\", include_schema: true) for the full parameter schema.".to_string(),
hint: "Description: Gmail tool\nParameters schema: {\"type\":\"object\"}".to_string(),
};
let display = err.to_string();
assert!(display.contains("unknown action: foobar"));
assert!(display.contains("Tool usage hint"));
assert!(display.contains("tool_info"));
assert!(display.contains("Gmail tool"));
assert!(display.contains("Parameters schema"));
}
}
+8
View File
@@ -67,8 +67,14 @@ pub struct WasmResourceLimiter {
memory_used: u64,
/// Maximum tables allowed.
max_tables: u32,
/// Current table count.
#[allow(dead_code)] // Reserved for table limit enforcement
tables_created: u32,
/// Maximum instances allowed.
max_instances: u32,
/// Current instance count.
#[allow(dead_code)] // Reserved for instance limit enforcement
instances_created: u32,
}
impl WasmResourceLimiter {
@@ -81,7 +87,9 @@ impl WasmResourceLimiter {
memory_limit,
memory_used: 0,
max_tables: 10,
tables_created: 0,
max_instances: 10, // Component model needs multiple instances for WASI
instances_created: 0,
}
}
+1 -1
View File
@@ -96,7 +96,7 @@ pub(crate) mod storage;
mod wrapper;
// Core types
pub use error::WasmError;
pub use error::{TrapCode, TrapInfo, WasmError};
pub use host::{HostState, LogEntry, LogLevel};
pub use limits::{
DEFAULT_FUEL_LIMIT, DEFAULT_MEMORY_LIMIT, DEFAULT_TIMEOUT, FuelConfig, ResourceLimits,
+41 -26
View File
@@ -123,9 +123,7 @@ pub struct PreparedModule {
pub name: String,
/// Tool description (cached from component).
pub description: String,
/// Full parameter schema JSON extracted from the component.
/// Used for discovery and coercion, not necessarily for the compact
/// schema advertised in the main tools array.
/// Parameter schema JSON (cached from component).
pub schema: serde_json::Value,
/// Pre-compiled component (cheaply cloneable via internal Arc).
component: wasmtime::component::Component,
@@ -267,29 +265,11 @@ impl WasmToolRuntime {
let component = wasmtime::component::Component::new(&engine, &wasm_bytes)
.map_err(|e| WasmError::CompilationFailed(e.to_string()))?;
// Briefly instantiate to extract metadata (description + schema)
// from the tool's exports, analogous to MCP's list_tools().
let effective_limits = limits.clone().unwrap_or(default_limits.clone());
let (description, schema) = crate::tools::wasm::wrapper::extract_wasm_metadata(
&engine,
&component,
&effective_limits,
)
.unwrap_or_else(|e| {
tracing::warn!(
name = %name,
error = %e,
"WASM metadata extraction failed, using fallbacks"
);
(
"WASM sandboxed tool".to_string(),
serde_json::json!({
"type": "object",
"properties": {},
"additionalProperties": true
}),
)
});
// We need to instantiate briefly to extract metadata.
// In a full implementation, we'd use WIT bindgen to get typed access.
// For now, we extract what we can from the component.
let description = extract_tool_description(&engine, &component)?;
let schema = extract_tool_schema(&engine, &component)?;
Ok::<_, WasmError>(PreparedModule {
name: name.clone(),
@@ -341,6 +321,41 @@ impl WasmToolRuntime {
}
}
/// Extract tool description from a compiled component.
///
/// Returns a generic fallback. Callers should prefer loading the description
/// from the sidecar `*.capabilities.json` file and overriding via
/// `WasmToolWrapper::with_description()` or the `WasmToolRegistration::description` field.
fn extract_tool_description(
_engine: &Engine,
_component: &wasmtime::component::Component,
) -> Result<String, WasmError> {
// WIT bindgen extraction is not yet implemented (see TODO #4 in CLAUDE.md).
// Real descriptions come from the capabilities.json sidecar file, which is
// loaded by the WasmToolLoader and passed as an override at registration time.
Ok("WASM sandboxed tool".to_string())
}
/// Extract tool parameter schema from a compiled component.
///
/// Returns a permissive fallback that accepts any JSON object. Callers should
/// prefer loading the schema from the sidecar `*.capabilities.json` file and
/// overriding via `WasmToolWrapper::with_schema()` or the
/// `WasmToolRegistration::schema` field.
fn extract_tool_schema(
_engine: &Engine,
_component: &wasmtime::component::Component,
) -> Result<serde_json::Value, WasmError> {
// WIT bindgen extraction is not yet implemented (see TODO #4 in CLAUDE.md).
// Real schemas come from the capabilities.json sidecar file, which is
// loaded by the WasmToolLoader and passed as an override at registration time.
Ok(serde_json::json!({
"type": "object",
"properties": {},
"additionalProperties": true
}))
}
impl std::fmt::Debug for WasmToolRuntime {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("WasmToolRuntime")
+53 -271
View File
@@ -464,10 +464,9 @@ pub struct WasmToolWrapper {
/// Capabilities to grant to this tool.
capabilities: Capabilities,
/// Cached description (from PreparedModule or override).
/// Stored without any tool_info hints — hints are composed at display time.
description: String,
/// Compact and discovery schemas for this tool.
schemas: WasmToolSchemas,
/// Cached schema (from PreparedModule or override).
schema: serde_json::Value,
/// Injected credentials for HTTP requests (e.g., OAuth tokens).
/// Keys are placeholder names like "GOOGLE_ACCESS_TOKEN".
credentials: HashMap<String, String>,
@@ -478,84 +477,6 @@ pub struct WasmToolWrapper {
oauth_refresh: Option<OAuthRefreshConfig>,
}
#[derive(Debug, Clone)]
struct WasmToolSchemas {
/// Compact schema advertised in the main tools array.
///
/// This stays permissive by default to avoid serializing full exported
/// WASM schemas on every LLM call. Sidecars can override it explicitly.
advertised: serde_json::Value,
/// Full schema available for discovery and coercion.
///
/// Seeded from the WASM `schema()` export at registration time, unless a
/// sidecar explicitly overrides it.
discovery: serde_json::Value,
}
impl WasmToolSchemas {
fn permissive_schema() -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": {},
"additionalProperties": true
})
}
fn is_permissive_schema(schema: &serde_json::Value) -> bool {
schema
.get("properties")
.and_then(|p| p.as_object())
.is_none_or(|p| p.is_empty())
}
fn new(discovery: serde_json::Value) -> Self {
Self {
advertised: Self::permissive_schema(),
discovery,
}
}
fn with_override(&self, schema: serde_json::Value) -> Self {
Self {
advertised: schema.clone(),
discovery: schema,
}
}
fn is_advertised_permissive(&self) -> bool {
Self::is_permissive_schema(&self.advertised)
}
fn advertised(&self) -> serde_json::Value {
self.advertised.clone()
}
fn discovery(&self) -> serde_json::Value {
self.discovery.clone()
}
/// Return the best schema available for type coercion.
///
/// Prefers the discovery schema when it has typed properties. Falls back
/// to the `PreparedModule` schema extracted at load time rather than
/// re-calling the WASM `schema()` export mid-execution, which could
/// interact with mutable linear memory state.
fn effective_for_coercion(&self, prepared_schema: &serde_json::Value) -> serde_json::Value {
if !Self::is_permissive_schema(&self.discovery) {
return self.discovery.clone();
}
// Fall back to the load-time extracted schema from PreparedModule.
// This avoids calling schema() on the already-running WASM instance
// where mutable state could produce inconsistent results.
if !Self::is_permissive_schema(prepared_schema) {
return prepared_schema.clone();
}
self.discovery.clone()
}
}
impl WasmToolWrapper {
/// Create a new WASM tool wrapper.
pub fn new(
@@ -565,7 +486,7 @@ impl WasmToolWrapper {
) -> Self {
Self {
description: prepared.description.clone(),
schemas: WasmToolSchemas::new(prepared.schema.clone()),
schema: prepared.schema.clone(),
runtime,
prepared,
capabilities,
@@ -583,7 +504,7 @@ impl WasmToolWrapper {
/// Override the parameter schema.
pub fn with_schema(mut self, schema: serde_json::Value) -> Self {
self.schemas = self.schemas.with_override(schema);
self.schema = schema;
self
}
@@ -694,18 +615,9 @@ impl WasmToolWrapper {
}
})?;
// Get typed interface — used for execute.
let tool_iface = instance.near_agent_tool();
// Determine effective schema for type coercion.
// Prefer the discovery schema when typed; fall back to the load-time
// extracted schema from PreparedModule rather than re-calling the WASM
// export on the already-running instance.
let effective_schema = self.schemas.effective_for_coercion(&self.prepared.schema);
// Coerce string-encoded values to their schema-declared types.
// LLMs frequently pass numeric values as strings (e.g. "5" instead of 5).
let params = coerce_params_to_schema(params, &effective_schema);
let params = coerce_params_to_schema(params, &self.schema);
// Prepare the request
let params_json = serde_json::to_string(&params)
@@ -717,6 +629,7 @@ impl WasmToolWrapper {
};
// Call execute using the generated typed interface
let tool_iface = instance.near_agent_tool();
let response = tool_iface.call_execute(&mut store, &request).map_err(|e| {
let error_str = e.to_string();
if error_str.contains("out of fuel") {
@@ -731,13 +644,12 @@ impl WasmToolWrapper {
// Get logs from host state
let logs = store.data_mut().host_state.take_logs();
// Check for tool-level error — point the LLM to tool_info for the
// full schema instead of dumping ~3.5KB inline.
// Check for tool-level error — on failure, call the WASM module's
// description() and schema() exports so the LLM can retry with the
// correct parameters without us having to include the (large) schema
// in every request's tools array.
if let Some(err) = response.error {
let hint = format!(
"Tip: call tool_info(name: \"{}\", include_schema: true) for the full parameter schema.",
self.prepared.name
);
let hint = build_tool_hint(tool_iface, &mut store);
return Err(WasmError::ToolReturnedError { message: err, hint });
}
@@ -746,55 +658,47 @@ impl WasmToolWrapper {
}
}
/// Extract metadata (description + schema) from a WASM tool by briefly
/// instantiating it and calling its `description()` and `schema()` exports.
/// Analogous to MCP's `list_tools()` — discovers tool capabilities at load time.
///
/// Falls back to generic description and permissive schema on failure.
pub(super) fn extract_wasm_metadata(
engine: &wasmtime::Engine,
component: &wasmtime::component::Component,
limits: &ResourceLimits,
) -> Result<(String, serde_json::Value), WasmError> {
let store_data = StoreData::new(
limits.memory_bytes,
Capabilities::default(),
HashMap::new(),
vec![],
);
let mut store = Store::new(engine, store_data);
/// Maximum characters for the description portion of a tool hint.
const HINT_DESC_MAX: usize = 500;
/// Maximum characters for the schema portion of a tool hint.
const HINT_SCHEMA_MAX: usize = 3000;
// Configure fuel + epoch deadline so extraction can't hang
if let Err(e) = store.set_fuel(limits.fuel) {
tracing::debug!("Fuel not enabled for metadata extraction: {e}");
}
store.epoch_deadline_trap();
let ticks = (limits.timeout.as_millis() / EPOCH_TICK_INTERVAL.as_millis()).max(1) as u64;
store.set_epoch_deadline(ticks);
store.limiter(|data| &mut data.limiter);
// Instantiate with minimal linker
let mut linker = Linker::new(engine);
WasmToolWrapper::add_host_functions(&mut linker)?;
let instance = SandboxedTool::instantiate(&mut store, component, &linker)
.map_err(|e| WasmError::InstantiationFailed(e.to_string()))?;
let tool_iface = instance.near_agent_tool();
// Extract description (fall back to generic)
let description = tool_iface
.call_description(&mut store)
.unwrap_or_else(|_| "WASM sandboxed tool".to_string());
// Extract and parse schema (fall back to permissive)
let schema = tool_iface
.call_schema(&mut store)
/// Call the WASM module's `description()` and `schema()` exports to build a
/// hint string. Returns an empty string if both calls fail or return empty.
/// Description is capped at [`HINT_DESC_MAX`] chars, schema at
/// [`HINT_SCHEMA_MAX`] chars.
fn build_tool_hint(tool_iface: &wit_tool::Guest, store: &mut Store<StoreData>) -> String {
let desc = tool_iface
.call_description(&mut *store)
.ok()
.and_then(|s| serde_json::from_str::<serde_json::Value>(&s).ok())
.unwrap_or_else(|| {
serde_json::json!({"type": "object", "properties": {}, "additionalProperties": true})
});
Ok((description, schema))
.unwrap_or_default();
let schema = tool_iface.call_schema(&mut *store).ok().unwrap_or_default();
if desc.is_empty() && schema.is_empty() {
return String::new();
}
let mut hint = String::new();
if !desc.is_empty() {
hint.push_str("Description: ");
if desc.len() > HINT_DESC_MAX {
let end = crate::util::floor_char_boundary(&desc, HINT_DESC_MAX);
hint.push_str(&desc[..end]);
hint.push('…');
} else {
hint.push_str(&desc);
}
hint.push('\n');
}
if !schema.is_empty() {
hint.push_str("Parameters schema: ");
if schema.len() > HINT_SCHEMA_MAX {
let end = crate::util::floor_char_boundary(&schema, HINT_SCHEMA_MAX);
hint.push_str(&schema[..end]);
hint.push('…');
} else {
hint.push_str(&schema);
}
}
hint
}
#[async_trait]
@@ -808,33 +712,7 @@ impl Tool for WasmToolWrapper {
}
fn parameters_schema(&self) -> serde_json::Value {
self.schemas.advertised()
}
fn discovery_schema(&self) -> serde_json::Value {
self.schemas.discovery()
}
/// Compose the tool schema for LLM function calling.
///
/// When the advertised schema is permissive (no typed properties), appends
/// a hint to the description directing the LLM to call `tool_info` for the
/// full parameter schema. This keeps the raw description clean while still
/// guiding the LLM.
fn schema(&self) -> crate::tools::tool::ToolSchema {
let description = if self.schemas.is_advertised_permissive() {
format!(
"{} (call tool_info(name: \"{}\", include_schema: true) for parameter schema)",
self.description, self.prepared.name
)
} else {
self.description.clone()
};
crate::tools::tool::ToolSchema {
name: self.prepared.name.clone(),
description,
parameters: self.schemas.advertised(),
}
self.schema.clone()
}
async fn execute(
@@ -871,7 +749,7 @@ impl Tool for WasmToolWrapper {
let prepared = Arc::clone(&self.prepared);
let capabilities = self.capabilities.clone();
let description = self.description.clone();
let schemas = self.schemas.clone();
let schema = self.schema.clone();
let credentials = self.credentials.clone();
// Execute in blocking task with timeout
@@ -881,7 +759,7 @@ impl Tool for WasmToolWrapper {
prepared,
capabilities,
description,
schemas,
schema,
credentials,
secrets_store: None, // Not needed in blocking task
oauth_refresh: None, // Already used above for pre-refresh
@@ -1354,7 +1232,6 @@ mod tests {
TEST_GOOGLE_OAUTH_TOKEN, TEST_OAUTH_CLIENT_ID, TEST_OAUTH_CLIENT_SECRET,
test_secrets_store,
};
use crate::tools::tool::Tool;
use crate::tools::wasm::capabilities::Capabilities;
use crate::tools::wasm::runtime::{WasmRuntimeConfig, WasmToolRuntime};
@@ -1369,84 +1246,6 @@ mod tests {
assert!(runtime.config().fuel_config.enabled);
}
#[tokio::test]
async fn test_advertised_schema_stays_permissive_until_sidecar_override() {
let discovery_schema = serde_json::json!({
"type": "object",
"properties": {
"query": { "type": "string" },
"limit": { "type": "integer" }
},
"required": ["query"]
});
let runtime = Arc::new(WasmToolRuntime::new(WasmRuntimeConfig::for_testing()).unwrap());
let prepared = runtime
.prepare("search", b"\0asm\x0d\0\x01\0", None)
.await
.unwrap();
let mut wrapper =
super::WasmToolWrapper::new(Arc::clone(&runtime), prepared, Capabilities::default());
wrapper.schemas = super::WasmToolSchemas::new(discovery_schema.clone());
wrapper.description = "Search documents".to_string();
// Advertised schema stays permissive; discovery holds the typed schema
assert_eq!(
wrapper.parameters_schema(),
serde_json::json!({
"type": "object",
"properties": {},
"additionalProperties": true
})
);
assert_eq!(wrapper.discovery_schema(), discovery_schema);
// Raw description is clean — no tool_info hint baked in
assert!(!wrapper.description().contains("tool_info"));
// But schema() composes the hint at display time when advertised is permissive
let schema = wrapper.schema();
assert!(
schema.description.contains("tool_info"),
"schema().description should contain tool_info hint: {}",
schema.description
);
assert!(
schema.description.contains("include_schema: true"),
"hint should mention include_schema: true: {}",
schema.description
);
// After sidecar override, both schemas match and hint disappears
let wrapper = wrapper.with_schema(serde_json::json!({
"type": "object",
"properties": {
"query": { "type": "string" }
},
"required": ["query"]
}));
assert_eq!(
wrapper.parameters_schema(),
serde_json::json!({
"type": "object",
"properties": {
"query": { "type": "string" }
},
"required": ["query"]
})
);
assert_eq!(wrapper.discovery_schema(), wrapper.parameters_schema());
// With typed schema, schema() should NOT include tool_info hint
let schema = wrapper.schema();
assert!(
!schema.description.contains("tool_info"),
"schema().description should not contain tool_info hint when typed: {}",
schema.description
);
}
#[test]
fn test_capabilities_default() {
let caps = Capabilities::default();
@@ -1989,23 +1788,6 @@ mod tests {
assert_eq!(result["count"], serde_json::json!("not-a-number"));
}
/// Regression: permissive fallback schema (empty properties) must NOT coerce.
/// This documents the bug where WASM tools with no sidecar `parameters` field
/// got the permissive fallback, causing coercion to be a no-op and LLM-provided
/// string integers to reach the WASM tool un-coerced.
#[test]
fn test_coerce_noop_with_permissive_schema() {
let permissive = serde_json::json!({
"type": "object",
"properties": {},
"additionalProperties": true
});
let params = serde_json::json!({"query": "test", "count": "10"});
let result = super::coerce_params_to_schema(params, &permissive);
// With empty properties, no coercion happens — string stays string
assert_eq!(result["count"], serde_json::json!("10"));
}
/// Regression test: leak scan must run on raw headers (before credential
/// injection), not after. If it ran post-injection, the host-injected
/// Slack bot token (`xoxb-...`) would trigger a Block and reject the
+1 -67
View File
@@ -60,18 +60,12 @@ pub trait EmbeddingProvider: Send + Sync {
}
}
/// Default base URL for the OpenAI API.
const OPENAI_API_BASE_URL: &str = "https://api.openai.com";
/// OpenAI embedding provider using text-embedding-ada-002 or text-embedding-3-small.
///
/// Supports any OpenAI-compatible embedding endpoint via [`with_base_url`](Self::with_base_url).
pub struct OpenAiEmbeddings {
client: reqwest::Client,
api_key: String,
model: String,
dimension: usize,
base_url: String,
}
impl OpenAiEmbeddings {
@@ -84,7 +78,6 @@ impl OpenAiEmbeddings {
api_key: api_key.into(),
model: "text-embedding-3-small".to_string(),
dimension: 1536,
base_url: OPENAI_API_BASE_URL.to_string(),
}
}
@@ -95,7 +88,6 @@ impl OpenAiEmbeddings {
api_key: api_key.into(),
model: "text-embedding-ada-002".to_string(),
dimension: 1536,
base_url: OPENAI_API_BASE_URL.to_string(),
}
}
@@ -106,7 +98,6 @@ impl OpenAiEmbeddings {
api_key: api_key.into(),
model: "text-embedding-3-large".to_string(),
dimension: 3072,
base_url: OPENAI_API_BASE_URL.to_string(),
}
}
@@ -121,35 +112,8 @@ impl OpenAiEmbeddings {
api_key: api_key.into(),
model: model.into(),
dimension,
base_url: OPENAI_API_BASE_URL.to_string(),
}
}
/// Set a custom base URL for OpenAI-compatible embedding providers.
///
/// The URL must use `http://` or `https://` scheme. If no scheme is present,
/// `https://` is prepended automatically. Trailing slashes are stripped.
pub fn with_base_url(mut self, base_url: &str) -> Self {
let url = base_url.trim();
// Auto-prepend https:// if no scheme is present.
let mut url = if !url.starts_with("http://") && !url.starts_with("https://") {
tracing::debug!(
"No scheme in embedding base URL '{}', prepending https://",
url
);
format!("https://{url}")
} else {
url.to_string()
};
while url.ends_with('/') {
url.pop();
}
self.base_url = url;
self
}
}
#[derive(Debug, Serialize)]
@@ -209,11 +173,9 @@ impl EmbeddingProvider for OpenAiEmbeddings {
input: texts,
};
let url = format!("{}/v1/embeddings", self.base_url);
let response = self
.client
.post(&url)
.post("https://api.openai.com/v1/embeddings")
.header("Authorization", format!("Bearer {}", self.api_key))
.json(&request)
.send()
@@ -613,37 +575,9 @@ mod tests {
let provider = OpenAiEmbeddings::new("test-key");
assert_eq!(provider.dimension(), 1536);
assert_eq!(provider.model_name(), "text-embedding-3-small");
assert_eq!(provider.base_url, OPENAI_API_BASE_URL);
let provider = OpenAiEmbeddings::large("test-key");
assert_eq!(provider.dimension(), 3072);
assert_eq!(provider.model_name(), "text-embedding-3-large");
assert_eq!(provider.base_url, OPENAI_API_BASE_URL);
}
#[test]
fn test_openai_with_base_url_valid() {
let provider =
OpenAiEmbeddings::new("test-key").with_base_url("https://custom.example.com");
assert_eq!(provider.base_url, "https://custom.example.com");
}
#[test]
fn test_openai_with_base_url_strips_trailing_slashes() {
let provider =
OpenAiEmbeddings::new("test-key").with_base_url("https://custom.example.com///");
assert_eq!(provider.base_url, "https://custom.example.com");
}
#[test]
fn test_openai_with_base_url_http_scheme() {
let provider = OpenAiEmbeddings::new("test-key").with_base_url("http://localhost:8080");
assert_eq!(provider.base_url, "http://localhost:8080");
}
#[test]
fn test_openai_with_base_url_schemeless_prepends_https() {
let provider = OpenAiEmbeddings::new("test-key").with_base_url("custom.example.com/v1");
assert_eq!(provider.base_url, "https://custom.example.com/v1");
}
}
+3 -19
View File
@@ -55,9 +55,7 @@ pub use embeddings::{
};
#[cfg(feature = "postgres")]
pub use repository::Repository;
pub use search::{
FusionStrategy, RankedResult, SearchConfig, SearchResult, fuse_results, reciprocal_rank_fusion,
};
pub use search::{RankedResult, SearchConfig, SearchResult, reciprocal_rank_fusion};
use std::sync::Arc;
@@ -334,8 +332,6 @@ pub struct Workspace {
storage: WorkspaceStorage,
/// Embedding provider for semantic search.
embeddings: Option<Arc<dyn EmbeddingProvider>>,
/// Default search configuration applied to all queries.
search_defaults: SearchConfig,
}
impl Workspace {
@@ -347,7 +343,6 @@ impl Workspace {
agent_id: None,
storage: WorkspaceStorage::Repo(Repository::new(pool)),
embeddings: None,
search_defaults: SearchConfig::default(),
}
}
@@ -360,7 +355,6 @@ impl Workspace {
agent_id: None,
storage: WorkspaceStorage::Db(db),
embeddings: None,
search_defaults: SearchConfig::default(),
}
}
@@ -376,16 +370,6 @@ impl Workspace {
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.
pub fn user_id(&self) -> &str {
&self.user_id
@@ -725,13 +709,13 @@ impl Workspace {
/// Hybrid search across all memory documents.
///
/// Combines full-text search (BM25) with semantic search (vector similarity)
/// using the configured fusion strategy.
/// using Reciprocal Rank Fusion (RRF).
pub async fn search(
&self,
query: &str,
limit: usize,
) -> 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
}
+2 -2
View File
@@ -12,7 +12,7 @@ use uuid::Uuid;
use crate::error::WorkspaceError;
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.
pub struct Repository {
@@ -415,7 +415,7 @@ impl Repository {
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.
+7 -314
View File
@@ -1,30 +1,17 @@
//! Hybrid search combining full-text and semantic search.
//!
//! Supports two fusion strategies:
//! 1. **RRF** (Reciprocal Rank Fusion) — the default, rank-based method.
//! `score = sum(1 / (k + rank))` for each retrieval method.
//! 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.
//! Uses Reciprocal Rank Fusion (RRF) to combine results from:
//! 1. PostgreSQL full-text search (ts_rank_cd)
//! 2. pgvector cosine similarity search
//!
//! Both strategies combine results from:
//! - PostgreSQL / libSQL full-text search
//! - pgvector / libsql_vector cosine similarity search
//! RRF formula: score = sum(1 / (k + rank)) for each retrieval method
//! This is robust to different score scales and produces better results
//! than simple score averaging.
use std::collections::HashMap;
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.
#[derive(Debug, Clone)]
pub struct SearchConfig {
@@ -40,16 +27,6 @@ pub struct SearchConfig {
pub min_score: f32,
/// Maximum results to fetch from each method before fusion.
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 {
@@ -61,9 +38,6 @@ impl Default for SearchConfig {
use_vector: true,
min_score: 0.0,
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
}
/// 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.
@@ -139,7 +87,7 @@ pub struct SearchResult {
pub chunk_id: Uuid,
/// Chunk content.
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,
/// Rank in FTS results (1-based, None if not in FTS results).
pub fts_rank: Option<u32>,
@@ -175,22 +123,6 @@ pub struct RankedResult {
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.
///
/// Combines ranked results from multiple retrieval methods using the formula:
@@ -303,109 +235,6 @@ pub fn reciprocal_rank_fusion(
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)]
mod tests {
use super::*;
@@ -628,142 +457,6 @@ mod tests {
let vector_only = SearchConfig::default().vector_only();
assert!(!vector_only.use_fts);
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 ---
+1 -77
View File
@@ -20,31 +20,9 @@ from helpers import AUTH_TOKEN, wait_for_port_line, wait_for_ready
# Project root (two levels up from tests/e2e/)
ROOT = Path(__file__).resolve().parent.parent.parent
# Git main repo root (for worktree support — WASM build artifacts live
# in the main repo's tools-src/*/target/ and aren't shared across worktrees)
_MAIN_ROOT = None
try:
import subprocess as _sp
_common = _sp.check_output(
["git", "worktree", "list", "--porcelain"],
cwd=ROOT, text=True, stderr=_sp.DEVNULL,
)
for line in _common.splitlines():
if line.startswith("worktree "):
_MAIN_ROOT = Path(line.split(" ", 1)[1])
break # first entry is always the main worktree
except Exception:
pass
# Temp directory for the libSQL database file (cleaned up automatically)
_DB_TMPDIR = tempfile.TemporaryDirectory(prefix="ironclaw-e2e-")
# Temp directories for WASM extensions. These start empty and are populated by
# the install pipeline during tests; fixtures do not pre-populate dev build
# artifacts into them.
_WASM_TOOLS_TMPDIR = tempfile.TemporaryDirectory(prefix="ironclaw-e2e-wasm-tools-")
_WASM_CHANNELS_TMPDIR = tempfile.TemporaryDirectory(prefix="ironclaw-e2e-wasm-channels-")
def _find_free_port() -> int:
"""Bind to port 0 and return the OS-assigned port."""
@@ -92,53 +70,7 @@ async def mock_llm_server():
@pytest.fixture(scope="session")
def wasm_tools_dir(_wasm_build_symlinks):
"""Empty temp dir for WASM tools.
Starts empty so the server has no pre-loaded extensions at boot.
The install API (POST /api/extensions/install) downloads and writes
WASM files here; tests exercise the full install pipeline.
NOTE on capabilities file naming: Cargo builds with underscored stems
(web_search_tool.wasm) but capabilities use hyphens (web-search-tool.
capabilities.json). The loader expects matching stems. If you pre-load
files, rename caps: web-search-tool web_search_tool.
"""
return str(Path(_WASM_TOOLS_TMPDIR.name))
@pytest.fixture(scope="session", autouse=True)
def _wasm_build_symlinks():
"""Symlink WASM build artifacts from the main repo into the worktree.
In a git worktree, tools-src/*/target/ directories don't exist because
Cargo build artifacts aren't shared. The install API's source fallback
checks these paths. Symlinking makes the fallback work without rebuilding.
"""
if _MAIN_ROOT is None or _MAIN_ROOT == ROOT:
yield
return
created = []
tools_src = ROOT / "tools-src"
main_tools_src = _MAIN_ROOT / "tools-src"
if tools_src.is_dir() and main_tools_src.is_dir():
for tool_dir in tools_src.iterdir():
if not tool_dir.is_dir():
continue
target = tool_dir / "target"
main_target = main_tools_src / tool_dir.name / "target"
if not target.exists() and main_target.is_dir():
target.symlink_to(main_target)
created.append(target)
yield
for link in created:
if link.is_symlink():
link.unlink()
@pytest.fixture(scope="session")
async def ironclaw_server(ironclaw_binary, mock_llm_server, wasm_tools_dir):
async def ironclaw_server(ironclaw_binary, mock_llm_server):
"""Start the ironclaw gateway. Yields the base URL."""
gateway_port = _find_free_port()
env = {
@@ -163,16 +95,8 @@ async def ironclaw_server(ironclaw_binary, mock_llm_server, wasm_tools_dir):
"ROUTINES_ENABLED": "false",
"HEARTBEAT_ENABLED": "false",
"EMBEDDING_ENABLED": "false",
# WASM tool/channel support
"WASM_ENABLED": "true",
"WASM_TOOLS_DIR": wasm_tools_dir,
"WASM_CHANNELS_DIR": _WASM_CHANNELS_TMPDIR.name,
# Prevent onboarding wizard from triggering
"ONBOARD_COMPLETED": "true",
# Force gateway OAuth callback mode (non-loopback URL) and point
# token exchange at mock_llm.py so OAuth tests work without Google.
"IRONCLAW_OAUTH_CALLBACK_URL": "https://oauth.test.example/oauth/callback",
"IRONCLAW_OAUTH_EXCHANGE_URL": mock_llm_server,
}
# Forward LLVM coverage instrumentation env vars when present
# (allows cargo-llvm-cov to collect profraw data from E2E runs).
-29
View File
@@ -133,32 +133,3 @@ async def wait_for_port_line(process, pattern: str, *, timeout: float = 60) -> i
if match := re.search(pattern, decoded):
return int(match.group(1))
raise TimeoutError(f"Port pattern '{pattern}' not found in stdout after {timeout}s")
# -- API helpers -----------------------------------------------------------
def auth_headers() -> dict[str, str]:
"""Return Authorization header dict for authenticated API calls."""
return {"Authorization": f"Bearer {AUTH_TOKEN}"}
async def api_get(base_url: str, path: str, **kwargs) -> httpx.Response:
"""Make an authenticated GET request to the ironclaw API."""
async with httpx.AsyncClient() as client:
return await client.get(
f"{base_url}{path}",
headers=auth_headers(),
timeout=kwargs.pop("timeout", 10),
**kwargs,
)
async def api_post(base_url: str, path: str, **kwargs) -> httpx.Response:
"""Make an authenticated POST request to the ironclaw API."""
async with httpx.AsyncClient() as client:
return await client.post(
f"{base_url}{path}",
headers=auth_headers(),
timeout=kwargs.pop("timeout", 10),
**kwargs,
)
+59 -184
View File
@@ -1,16 +1,11 @@
"""Mock OpenAI-compatible LLM server for E2E tests.
Serves OpenAI-compatible endpoints for chat completions and model listing.
Supports both streaming and non-streaming responses, plus function calling
via TOOL_CALL_PATTERNS.
"""
"""Mock OpenAI-compatible LLM server for E2E tests."""
import argparse
import asyncio
import json
import re
import time
import uuid
from aiohttp import web
CANNED_RESPONSES = [
@@ -18,207 +13,85 @@ CANNED_RESPONSES = [
(re.compile(r"2\s*\+\s*2|two plus two", re.IGNORECASE), "The answer is 4."),
(re.compile(r"skill|install", re.IGNORECASE), "I can help you with skills management."),
(re.compile(r"html.?test|injection.?test", re.IGNORECASE),
'Here is some content: <script>alert("xss")</script> and <img src=x onerror="alert(1)">'
' and <iframe src="javascript:alert(2)"></iframe> end of content.'),
'Here is some content: <script>alert("xss")</script> and <img src=x onerror="alert(1)"> and <iframe src="javascript:alert(2)"></iframe> end of content.'),
]
DEFAULT_RESPONSE = "I understand your request."
TOOL_CALL_PATTERNS = [
(re.compile(r"echo (.+)", re.IGNORECASE), "echo", lambda m: {"message": m.group(1)}),
(re.compile(r"what time|current time", re.IGNORECASE), "time", lambda _: {"operation": "now"}),
]
def _last_user_content(messages: list[dict]) -> str:
def match_response(messages: list[dict]) -> str:
"""Find canned response for the last user message."""
for msg in reversed(messages):
if msg.get("role") == "user":
content = msg.get("content", "")
# Handle content that may be a list (multi-modal)
if isinstance(content, list):
content = " ".join(
p.get("text", "") for p in content if p.get("type") == "text"
part.get("text", "") for part in content if part.get("type") == "text"
)
return content
return ""
def match_response(messages: list[dict]) -> str:
content = _last_user_content(messages)
for pattern, response in CANNED_RESPONSES:
if pattern.search(content):
return response
for pattern, response in CANNED_RESPONSES:
if pattern.search(content):
return response
return DEFAULT_RESPONSE
return DEFAULT_RESPONSE
def match_tool_call(messages: list[dict], has_tools: bool) -> dict | None:
if not has_tools:
return None
content = _last_user_content(messages)
for pattern, tool_name, args_fn in TOOL_CALL_PATTERNS:
m = pattern.search(content)
if m:
return {"tool_name": tool_name, "arguments": args_fn(m)}
return None
def _extract_tool_name(msg: dict) -> str:
"""Extract tool name from a message, checking both 'name' field and XML content."""
name = msg.get("name")
if name:
return name
# ironclaw wraps tool output as <tool_output name="...">
content = msg.get("content", "")
m = re.search(r'<tool_output\s+name="([^"]+)"', content)
if m:
return m.group(1)
return "unknown"
def _find_tool_result(messages: list[dict]) -> dict | None:
"""Find a pending tool result that appears after the last user message.
Only returns a tool result if it's a fresh result the agent is waiting
for the LLM to summarize (i.e., it follows the most recent user message).
This prevents stale tool results from earlier conversation turns from
being re-processed.
"""
# Find the position of the last user message
last_user_idx = -1
for i in range(len(messages) - 1, -1, -1):
if messages[i].get("role") == "user":
last_user_idx = i
break
# Only look for tool results after the last user message
for i in range(len(messages) - 1, last_user_idx, -1):
if messages[i].get("role") == "tool":
return {"name": _extract_tool_name(messages[i]),
"content": messages[i].get("content", "")}
return None
def _make_base(completion_id: str) -> dict:
return {"id": completion_id, "object": "chat.completion.chunk",
"created": int(time.time()), "model": "mock-model"}
async def _send_sse(resp: web.StreamResponse, data: dict):
await resp.write(f"data: {json.dumps(data)}\n\n".encode())
async def chat_completions(request: web.Request) -> web.StreamResponse:
"""Handle POST /v1/chat/completions and /chat/completions."""
"""Handle POST /v1/chat/completions."""
body = await request.json()
messages = body.get("messages", [])
stream = body.get("stream", False)
has_tools = bool(body.get("tools"))
cid = f"mock-{uuid.uuid4().hex[:8]}"
response_text = match_response(messages)
completion_id = f"mock-{uuid.uuid4().hex[:8]}"
# Tool result in messages -> text summary
tr = _find_tool_result(messages)
if tr:
text = f"The {tr['name']} tool returned: {tr['content']}"
if not stream:
return _text_response(cid, text)
return await _stream_text(request, cid, text)
# Tool-call pattern match
tc = match_tool_call(messages, has_tools)
if tc:
if not stream:
return _tool_call_response(cid, tc)
return await _stream_tool_call(request, cid, tc)
# Default text response
text = match_response(messages)
if not stream:
return _text_response(cid, text)
return await _stream_text(request, cid, text)
return web.json_response({
"id": completion_id,
"object": "chat.completion",
"created": int(time.time()),
"model": "mock-model",
"choices": [{
"index": 0,
"message": {"role": "assistant", "content": response_text},
"finish_reason": "stop",
}],
"usage": {"prompt_tokens": 10, "completion_tokens": len(response_text.split()), "total_tokens": 15},
})
def _text_response(cid: str, text: str) -> web.Response:
return web.json_response({
"id": cid, "object": "chat.completion", "created": int(time.time()),
"model": "mock-model",
"choices": [{"index": 0, "message": {"role": "assistant", "content": text},
"finish_reason": "stop"}],
"usage": {"prompt_tokens": 10, "completion_tokens": len(text.split()), "total_tokens": 15},
})
def _tool_call_response(cid: str, tc: dict) -> web.Response:
return web.json_response({
"id": cid, "object": "chat.completion", "created": int(time.time()),
"model": "mock-model",
"choices": [{"index": 0, "message": {
"role": "assistant", "content": None,
"tool_calls": [{"id": f"call_{uuid.uuid4().hex[:8]}", "type": "function",
"function": {"name": tc["tool_name"],
"arguments": json.dumps(tc["arguments"])}}],
}, "finish_reason": "tool_calls"}],
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
})
async def _stream_text(request: web.Request, cid: str, text: str) -> web.StreamResponse:
resp = web.StreamResponse(status=200, headers={
"Content-Type": "text/event-stream", "Cache-Control": "no-cache"})
# Streaming response: split into word-boundary chunks
resp = web.StreamResponse(
status=200,
headers={"Content-Type": "text/event-stream", "Cache-Control": "no-cache"},
)
await resp.prepare(request)
base = _make_base(cid)
chunk = {**base, "choices": [{"index": 0, "delta": {"role": "assistant", "content": ""},
"finish_reason": None}]}
await _send_sse(resp, chunk)
for i, word in enumerate(text.split(" ")):
chunk["choices"][0]["delta"] = {"content": word if i == 0 else f" {word}"}
await _send_sse(resp, chunk)
# First chunk: role
chunk = {
"id": completion_id,
"object": "chat.completion.chunk",
"created": int(time.time()),
"model": "mock-model",
"choices": [{"index": 0, "delta": {"role": "assistant", "content": ""}, "finish_reason": None}],
}
await resp.write(f"data: {json.dumps(chunk)}\n\n".encode())
# Content chunks: split on spaces
words = response_text.split(" ")
for i, word in enumerate(words):
text = word if i == 0 else f" {word}"
chunk["choices"][0]["delta"] = {"content": text}
await resp.write(f"data: {json.dumps(chunk)}\n\n".encode())
# Final chunk: finish_reason
chunk["choices"][0]["delta"] = {}
chunk["choices"][0]["finish_reason"] = "stop"
await _send_sse(resp, chunk)
await resp.write(f"data: {json.dumps(chunk)}\n\n".encode())
await resp.write(b"data: [DONE]\n\n")
return resp
async def _stream_tool_call(request: web.Request, cid: str, tc: dict) -> web.StreamResponse:
resp = web.StreamResponse(status=200, headers={
"Content-Type": "text/event-stream", "Cache-Control": "no-cache"})
await resp.prepare(request)
call_id = f"call_{uuid.uuid4().hex[:8]}"
base = _make_base(cid)
# First chunk: role + tool call header with empty arguments
chunk = {**base, "choices": [{"index": 0, "delta": {
"role": "assistant",
"tool_calls": [{"index": 0, "id": call_id, "type": "function",
"function": {"name": tc["tool_name"], "arguments": ""}}],
}, "finish_reason": None}]}
await _send_sse(resp, chunk)
# Second chunk: arguments payload
chunk["choices"][0]["delta"] = {
"tool_calls": [{"index": 0, "function": {"arguments": json.dumps(tc["arguments"])}}]}
await _send_sse(resp, chunk)
# Final chunk: finish reason
chunk["choices"][0]["delta"] = {}
chunk["choices"][0]["finish_reason"] = "tool_calls"
await _send_sse(resp, chunk)
await resp.write(b"data: [DONE]\n\n")
return resp
async def oauth_exchange(request: web.Request) -> web.Response:
"""Mock OAuth token exchange proxy for E2E tests.
Accepts form params (code, redirect_uri, code_verifier) and returns
a fake token response. Called by ironclaw's exchange_via_proxy() when
IRONCLAW_OAUTH_EXCHANGE_URL is set.
"""
data = await request.post()
code = data.get("code", "")
return web.json_response({
"access_token": f"mock-token-{code}",
"refresh_token": "mock-refresh-token",
"expires_in": 3600,
})
async def models(_request: web.Request) -> web.Response:
"""Handle GET /v1/models."""
return web.json_response({
"object": "list",
"data": [{"id": "mock-model", "object": "model", "owned_by": "test"}],
@@ -229,21 +102,23 @@ def main():
parser = argparse.ArgumentParser()
parser.add_argument("--port", type=int, default=0)
args = parser.parse_args()
app = web.Application()
# Register both /v1/ and non-/v1/ paths (rig-core omits the /v1/ prefix)
app.router.add_post("/v1/chat/completions", chat_completions)
app.router.add_post("/chat/completions", chat_completions)
app.router.add_get("/v1/models", models)
app.router.add_get("/models", models)
app.router.add_post("/oauth/exchange", oauth_exchange)
# Use aiohttp's runner to get the actual bound port
import asyncio
async def start():
runner = web.AppRunner(app)
await runner.setup()
site = web.TCPSite(runner, "127.0.0.1", args.port)
await site.start()
# Extract the actual port from the bound socket
port = site._server.sockets[0].getsockname()[1]
print(f"MOCK_LLM_PORT={port}", flush=True)
# Block forever
await asyncio.Event().wait()
asyncio.run(start())
-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)
-264
View File
@@ -1,264 +0,0 @@
"""Extension OAuth round-trip e2e tests.
Tests the full internal OAuth callback pipeline: install gmail configure
(get auth_url) simulate OAuth callback verify token stored. Uses gateway
callback mode + mock token exchange (no real Google login).
The conftest sets IRONCLAW_OAUTH_CALLBACK_URL (non-loopback, forces gateway
mode) and IRONCLAW_OAUTH_EXCHANGE_URL (points to mock_llm.py's /oauth/exchange).
"""
from urllib.parse import parse_qs, urlparse
import httpx
import pytest
from helpers import api_get, api_post
# Module-level state
_gmail_installed = False
_auth_url = None
_csrf_state = None
def _extract_state(auth_url: str) -> str:
"""Extract the CSRF state parameter from an OAuth authorization URL."""
parsed = urlparse(auth_url)
qs = parse_qs(parsed.query)
assert "state" in qs, f"auth_url should contain state param: {auth_url}"
state = qs["state"][0]
assert len(state) > 0
return state
async def _get_extension(base_url, name):
"""Get a specific extension from the extensions list, or None."""
r = await api_get(base_url, "/api/extensions")
for ext in r.json().get("extensions", []):
if ext["name"] == name:
return ext
return None
async def _ensure_removed(base_url, name):
"""Remove extension if already installed."""
ext = await _get_extension(base_url, name)
if ext:
await api_post(base_url, f"/api/extensions/{name}/remove", timeout=30)
# ── Section A: Install + OAuth Initiation ────────────────────────────────
async def test_oauth_install_gmail(ironclaw_server):
"""Install gmail from registry for OAuth testing."""
global _gmail_installed
await _ensure_removed(ironclaw_server, "gmail")
r = await api_post(
ironclaw_server,
"/api/extensions/install",
json={"name": "gmail"},
timeout=180,
)
assert r.status_code == 200
data = r.json()
assert data.get("success") is True, f"Install failed: {data.get('message', '')}"
_gmail_installed = True
async def test_oauth_configure_returns_auth_url(ironclaw_server):
"""Configure with empty secrets returns an OAuth auth_url."""
global _auth_url, _csrf_state
if not _gmail_installed:
pytest.skip("gmail not installed")
r = await api_post(
ironclaw_server,
"/api/extensions/gmail/setup",
json={"secrets": {}},
timeout=30,
)
assert r.status_code == 200
data = r.json()
assert data.get("success") is True, f"Configure failed: {data.get('message', '')}"
_auth_url = data.get("auth_url")
assert _auth_url is not None, f"Expected auth_url in response: {data}"
assert "accounts.google.com" in _auth_url, (
f"auth_url should point to Google: {_auth_url}"
)
_csrf_state = _extract_state(_auth_url)
async def test_oauth_activate_returns_auth_url(ironclaw_server):
"""Activate on un-authenticated gmail returns auth_url."""
if not _gmail_installed:
pytest.skip("gmail not installed")
r = await api_post(
ironclaw_server, "/api/extensions/gmail/activate", timeout=30
)
assert r.status_code == 200
data = r.json()
# Activation may fail with auth_url or succeed with auth_url
auth_url = data.get("auth_url")
assert auth_url is not None, f"Expected auth_url in activate response: {data}"
# ── Section B: Internal OAuth Round-Trip ─────────────────────────────────
async def test_oauth_callback_exchanges_token(ironclaw_server):
"""Simulate OAuth callback with mock code — verifies token exchange."""
global _csrf_state
if not _csrf_state:
pytest.skip("No CSRF state from configure step")
# Re-configure to get a fresh pending flow (previous configure may have
# been consumed by the activate test above)
r = await api_post(
ironclaw_server,
"/api/extensions/gmail/setup",
json={"secrets": {}},
timeout=30,
)
data = r.json()
auth_url = data.get("auth_url")
if auth_url:
_csrf_state = _extract_state(auth_url)
# Hit the OAuth callback endpoint directly (public route, no auth header).
# The callback handler looks up the pending flow by state, calls
# exchange_via_proxy() which hits mock_llm.py's /oauth/exchange, and
# stores the returned fake token.
async with httpx.AsyncClient() as client:
r = await client.get(
f"{ironclaw_server}/oauth/callback",
params={"code": "mock_auth_code", "state": _csrf_state},
timeout=30,
follow_redirects=True,
)
assert r.status_code == 200, f"Callback returned {r.status_code}: {r.text[:300]}"
body = r.text.lower()
# The landing page says "<name> Connected" on success, "failed" on error
assert "connected" in body or "success" in body, (
f"Callback HTML should indicate success: {r.text[:500]}"
)
async def test_oauth_callback_replay_rejected(ironclaw_server):
"""Replaying the same callback is rejected (flow consumed on first use)."""
if not _csrf_state:
pytest.skip("No CSRF state")
async with httpx.AsyncClient() as client:
r = await client.get(
f"{ironclaw_server}/oauth/callback",
params={"code": "mock_auth_code", "state": _csrf_state},
timeout=10,
follow_redirects=True,
)
# Should fail — the flow was already consumed
body = r.text.lower()
assert "error" in body or "fail" in body or "expired" in body or r.status_code >= 400, (
f"Replay should be rejected, got status={r.status_code}: {r.text[:500]}"
)
async def test_oauth_callback_invalid_state(ironclaw_server):
"""Callback with bogus state is rejected."""
async with httpx.AsyncClient() as client:
r = await client.get(
f"{ironclaw_server}/oauth/callback",
params={"code": "x", "state": "totally-bogus-state-value"},
timeout=10,
follow_redirects=True,
)
body = r.text.lower()
assert "error" in body or "fail" in body or "expired" in body or r.status_code >= 400, (
f"Invalid state should be rejected, got status={r.status_code}: {r.text[:500]}"
)
async def test_oauth_extension_authenticated(ironclaw_server):
"""After OAuth callback, gmail shows authenticated=True."""
if not _gmail_installed:
pytest.skip("gmail not installed")
ext = await _get_extension(ironclaw_server, "gmail")
assert ext is not None, "gmail not in extensions list"
assert ext["authenticated"] is True, (
f"gmail should be authenticated after OAuth callback: {ext}"
)
async def test_oauth_tools_registered(ironclaw_server):
"""After OAuth authentication, gmail tools appear in tools endpoint."""
if not _gmail_installed:
pytest.skip("gmail not installed")
ext = await _get_extension(ironclaw_server, "gmail")
assert ext is not None
# Check the extension's tools array
tools = ext.get("tools", [])
assert len(tools) > 0, (
f"gmail should have tools registered after auth: {ext}"
)
async def test_remove_during_pending_oauth_invalidates_callback(ironclaw_server):
"""Removing an extension while OAuth is pending invalidates the callback state."""
if not _gmail_installed:
pytest.skip("gmail not installed")
r = await api_post(
ironclaw_server,
"/api/extensions/gmail/setup",
json={"secrets": {}},
timeout=30,
)
assert r.status_code == 200
data = r.json()
auth_url = data.get("auth_url")
assert auth_url is not None, f"Expected auth_url in response: {data}"
callback_state = _extract_state(auth_url)
remove_r = await api_post(
ironclaw_server, "/api/extensions/gmail/remove", timeout=30
)
assert remove_r.status_code == 200
assert remove_r.json().get("success") is True, (
f"Removing gmail during pending OAuth should succeed: {remove_r.text[:300]}"
)
async with httpx.AsyncClient() as client:
callback_r = await client.get(
f"{ironclaw_server}/oauth/callback",
params={"code": "mock_auth_code", "state": callback_state},
timeout=30,
follow_redirects=True,
)
assert callback_r.status_code == 200
body = callback_r.text.lower()
assert "error" in body or "fail" in body or "expired" in body, (
f"Callback after removal should fail: {callback_r.text[:500]}"
)
ext = await _get_extension(ironclaw_server, "gmail")
assert ext is None, "gmail should remain removed after invalidated callback"
# ── Section C: Cleanup ──────────────────────────────────────────────────
async def test_cleanup_gmail(ironclaw_server):
"""Remove gmail (cleanup for other test files)."""
await _ensure_removed(ironclaw_server, "gmail")
ext = await _get_extension(ironclaw_server, "gmail")
assert ext is None, "gmail should be removed"

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