Compare commits

..
Author SHA1 Message Date
mackabyandZaki 5677e5e955 fix(db): invoke shutdown during runtime teardown
Call Database::shutdown() from async_main graceful shutdown after webhook/tunnel stop.

Add debug logs for libSQL no-op shutdown paths (no replicator and sync-not-supported).
2026-03-12 15:08:08 -07:00
mackabyandZaki 91e0c2ee62 feat(db): add backend shutdown hook for graceful runtime shutdown
Add Database::shutdown() with a default no-op implementation for backward compatibility.

Implement libSQL shutdown via flush_replicator(), treating SyncNotSupported as non-fatal.

Implement Postgres shutdown by closing the pool.
2026-03-12 15:06:12 -07:00
128 changed files with 1523 additions and 6981 deletions
-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 -57
View File
@@ -78,70 +78,15 @@ jobs:
- name: Check lints - name: Check lints
run: cargo clippy --all --benches --tests --examples ${{ matrix.flags }} -- -D warnings run: cargo clippy --all --benches --tests --examples ${{ matrix.flags }} -- -D warnings
no-panics:
name: No panics in production code
runs-on: ubuntu-latest
steps:
- name: Checkout repository
uses: actions/checkout@v6
with:
fetch-depth: 0
- name: Check for .unwrap(), .expect(), assert!() in production code
run: |
BASE="${{ github.event.pull_request.base.sha }}"
# Get the full diff for .rs files (production only, exclude tests/ directory)
DIFF=$(git diff "$BASE"...HEAD -- 'src/**/*.rs' 'crates/**/*.rs' || true)
if [ -z "$DIFF" ]; then
echo "No production Rust changes detected."
exit 0
fi
# Extract added lines, skipping those inside test modules.
# Track whether we're inside a test module by watching hunk headers
# (lines starting with @@) whose context contains "mod tests" or "#[cfg(test)]".
ADDED=$(echo "$DIFF" | awk '
/^@@/ {
# Hunk context (after the second @@) tells us the function/module scope
in_test = (tolower($0) ~ /mod tests/ || $0 ~ /#\[cfg\(test\)\]/ || $0 ~ /#\[test\]/)
}
/^\+[^+]/ && !in_test { print }
' || true)
if [ -z "$ADDED" ]; then
echo "No production Rust changes detected (test-only changes excluded)."
exit 0
fi
# Match panic-inducing patterns, excluding safety suppressions
VIOLATIONS=$(echo "$ADDED" \
| grep -E '\.(unwrap|expect)\(|[^_]assert(_eq|_ne)?!' \
| grep -Ev 'debug_assert|// safety:' \
|| true)
if [ -n "$VIOLATIONS" ]; then
echo "::error::Found .unwrap(), .expect(), or assert!() in production code."
echo "Production code must use proper error handling instead of panicking."
echo "Suppress false positives with an inline '// safety: <reason>' comment."
echo ""
echo "$VIOLATIONS" | head -20
echo ""
COUNT=$(echo "$VIOLATIONS" | wc -l | tr -d ' ')
echo "Total: $COUNT violation(s)"
exit 1
fi
echo "OK: No panic-inducing calls in changed production code."
# Roll-up job for branch protection # Roll-up job for branch protection
code-style: code-style:
name: Code Style (fmt + clippy + deny) name: Code Style (fmt + clippy + deny)
runs-on: ubuntu-latest runs-on: ubuntu-latest
if: always() if: always()
needs: [format, clippy, clippy-windows, deny-check, no-panics] needs: [format, clippy, clippy-windows, deny-check]
steps: steps:
- run: | - run: |
if [[ "${{ needs.format.result }}" != "success" || "${{ needs.clippy.result }}" != "success" || "${{ needs.deny-check.result }}" != "success" || "${{ needs.no-panics.result }}" != "success" ]]; then if [[ "${{ needs.format.result }}" != "success" || "${{ needs.clippy.result }}" != "success" || "${{ needs.deny-check.result }}" != "success" ]]; then
echo "One or more jobs failed" echo "One or more jobs failed"
exit 1 exit 1
fi fi
+1 -1
View File
@@ -52,7 +52,7 @@ jobs:
- group: features - group: features
files: "tests/e2e/scenarios/test_skills.py tests/e2e/scenarios/test_tool_approval.py" files: "tests/e2e/scenarios/test_skills.py tests/e2e/scenarios/test_tool_approval.py"
- group: extensions - group: extensions
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 tests/e2e/scenarios/test_oauth_credential_fallback.py tests/e2e/scenarios/test_routine_oauth_credential_injection.py" files: "tests/e2e/scenarios/test_extensions.py"
steps: steps:
- uses: actions/checkout@v6 - uses: actions/checkout@v6
@@ -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 - *checkout
- *install-rust - *install-rust
- uses: Swatinem/rust-cache@v2 - 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 - name: Run release-plz
uses: release-plz/[email protected] uses: release-plz/[email protected]
with: with:
command: release-pr command: release-pr
env: env:
GITHUB_TOKEN: ${{ steps.generate-token.outputs.token }} GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
CARGO_REGISTRY_TOKEN: ${{ secrets.CARGO_REGISTRY_TOKEN }} CARGO_REGISTRY_TOKEN: ${{ secrets.CARGO_REGISTRY_TOKEN }}
+63 -111
View File
@@ -25,35 +25,9 @@ concurrency:
cancel-in-progress: false # Let running suites finish cancel-in-progress: false # Let running suites finish
jobs: 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 for new commits ──────────────────────────────────────
check-changes: check-changes:
name: Check for new commits name: Check for new commits
needs: resolve-promotion-base
runs-on: ubuntu-latest runs-on: ubuntu-latest
outputs: outputs:
has_changes: ${{ steps.check.outputs.has_changes }} has_changes: ${{ steps.check.outputs.has_changes }}
@@ -70,7 +44,7 @@ jobs:
id: check id: check
env: env:
FORCE_RUN: ${{ inputs.force }} FORCE_RUN: ${{ inputs.force }}
PROMOTION_BASE: ${{ needs.resolve-promotion-base.outputs.promotion_base }} DEFAULT_BRANCH: ${{ github.event.repository.default_branch }}
run: | run: |
CURRENT_HEAD=$(git rev-parse HEAD) CURRENT_HEAD=$(git rev-parse HEAD)
echo "current_head=${CURRENT_HEAD}" >> "$GITHUB_OUTPUT" echo "current_head=${CURRENT_HEAD}" >> "$GITHUB_OUTPUT"
@@ -92,9 +66,9 @@ jobs:
echo "Found ${COMMIT_COUNT} new commit(s) since last tested" echo "Found ${COMMIT_COUNT} new commit(s) since last tested"
DIFF_RANGE="${LAST_TESTED}..${CURRENT_HEAD}" DIFF_RANGE="${LAST_TESTED}..${CURRENT_HEAD}"
else else
git fetch origin "${PROMOTION_BASE}" git fetch origin "${DEFAULT_BRANCH}"
MERGE_BASE=$(git merge-base "origin/${PROMOTION_BASE}" HEAD) MERGE_BASE=$(git merge-base "origin/${DEFAULT_BRANCH}" HEAD)
echo "First run -- reviewing from merge-base ${MERGE_BASE} against ${PROMOTION_BASE}" echo "First run -- reviewing from merge-base ${MERGE_BASE}"
DIFF_RANGE="${MERGE_BASE}..${CURRENT_HEAD}" DIFF_RANGE="${MERGE_BASE}..${CURRENT_HEAD}"
fi fi
fi fi
@@ -128,7 +102,7 @@ jobs:
# ── Create promotion PR (triggers claude-review.yml on the PR) ── # ── Create promotion PR (triggers claude-review.yml on the PR) ──
create-promotion-pr: create-promotion-pr:
name: Create Promotion PR name: Create Promotion PR
needs: [resolve-promotion-base, check-changes] needs: check-changes
if: needs.check-changes.outputs.has_changes == 'true' if: needs.check-changes.outputs.has_changes == 'true'
runs-on: ubuntu-latest runs-on: ubuntu-latest
outputs: outputs:
@@ -160,15 +134,15 @@ jobs:
id: ahead-check id: ahead-check
env: env:
GH_TOKEN: ${{ steps.token.outputs.token }} GH_TOKEN: ${{ steps.token.outputs.token }}
PROMOTION_BASE: ${{ needs.resolve-promotion-base.outputs.promotion_base }} DEFAULT_BRANCH: ${{ github.event.repository.default_branch }}
run: | run: |
git fetch origin "${PROMOTION_BASE}" git fetch origin "${DEFAULT_BRANCH}"
AHEAD=$(git rev-list --count "origin/${PROMOTION_BASE}..origin/staging") AHEAD=$(git rev-list --count "origin/${DEFAULT_BRANCH}..origin/staging")
echo "commits_ahead=${AHEAD}" >> "$GITHUB_OUTPUT" echo "commits_ahead=${AHEAD}" >> "$GITHUB_OUTPUT"
if [ "$AHEAD" -eq 0 ]; then if [ "$AHEAD" -eq 0 ]; then
echo "Staging is not ahead of ${PROMOTION_BASE}. Nothing to promote." echo "Staging is not ahead of ${DEFAULT_BRANCH}. Nothing to promote."
else else
echo "Staging is ${AHEAD} commits ahead of ${PROMOTION_BASE}." echo "Staging is ${AHEAD} commits ahead of ${DEFAULT_BRANCH}."
fi fi
- name: Create promotion branch - name: Create promotion branch
@@ -182,53 +156,54 @@ jobs:
echo "branch=${BRANCH}" >> "$GITHUB_OUTPUT" echo "branch=${BRANCH}" >> "$GITHUB_OUTPUT"
echo "Created promotion branch: ${BRANCH}" 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 }}
DEFAULT_BRANCH: ${{ github.event.repository.default_branch }}
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=${DEFAULT_BRANCH}" >> "$GITHUB_OUTPUT"
echo "No existing promotion PR — targeting ${DEFAULT_BRANCH}"
fi
- name: Create promotion PR - name: Create promotion PR
id: create-pr id: create-pr
if: steps.ahead-check.outputs.commits_ahead != '0' if: steps.ahead-check.outputs.commits_ahead != '0'
env: env:
GH_TOKEN: ${{ steps.token.outputs.token }} GH_TOKEN: ${{ steps.token.outputs.token }}
run: | run: |
source .github/scripts/pr-body-utils.sh
RANGE="${{ needs.check-changes.outputs.diff_range }}" RANGE="${{ needs.check-changes.outputs.diff_range }}"
TIMESTAMP=$(date -u +"%Y-%m-%d %H:%M UTC") TIMESTAMP=$(date -u +"%Y-%m-%d %H:%M UTC")
BRANCH="${{ steps.branch.outputs.branch }}" BRANCH="${{ steps.branch.outputs.branch }}"
BASE="${{ needs.resolve-promotion-base.outputs.promotion_base }}" BASE="${{ steps.find-base.outputs.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*"
PR_URL=$(gh pr create \ PR_URL=$(gh pr create \
--base "$BASE" \ --base "$BASE" \
--head "$BRANCH" \ --head "$BRANCH" \
--title "chore: promote staging to ${BASE} (${TIMESTAMP})" \ --title "chore: promote staging to ${BASE} (${TIMESTAMP})" \
--body "$PR_BODY" \ --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") --label "staging-promotion")
PR_NUM=$(echo "$PR_URL" | grep -oE '[0-9]+$') PR_NUM=$(echo "$PR_URL" | grep -oE '[0-9]+$')
@@ -253,8 +228,7 @@ jobs:
- uses: actions/checkout@v6 - uses: actions/checkout@v6
with: with:
ref: staging ref: staging
# Need full history to recompute the final promoted range before merge. fetch-depth: 1
fetch-depth: 0
- name: Generate GitHub App token - name: Generate GitHub App token
id: app-token id: app-token
@@ -353,10 +327,8 @@ jobs:
# Use process substitution so variables propagate to parent shell # Use process substitution so variables propagate to parent shell
while read -r line; do while read -r line; do
TAG=$(echo "$line" | grep -oE '^\[(CRITICAL|HIGH|MEDIUM|LOW):[0-9]+\]') TAG=$(echo "$line" | grep -oE '^\[(CRITICAL|HIGH|MEDIUM|LOW):[0-9]+\]')
SEVERITY="${TAG#\[}" SEVERITY=$(echo "$TAG" | sed 's/\[\(.*\):\(.*\)\]/\1/')
SEVERITY="${SEVERITY%%:*}" CONFIDENCE=$(echo "$TAG" | sed 's/\[\(.*\):\(.*\)\]/\2/')
CONFIDENCE="${TAG##*:}"
CONFIDENCE="${CONFIDENCE%\]}"
DESC=$(echo "$line" | sed "s/\[${SEVERITY}:${CONFIDENCE}\] *//" | head -1) DESC=$(echo "$line" | sed "s/\[${SEVERITY}:${CONFIDENCE}\] *//" | head -1)
echo "Found: [${SEVERITY}:${CONFIDENCE}] ${DESC}" echo "Found: [${SEVERITY}:${CONFIDENCE}] ${DESC}"
@@ -448,29 +420,11 @@ jobs:
GH_TOKEN: ${{ steps.token.outputs.token }} GH_TOKEN: ${{ steps.token.outputs.token }}
PR_NUMBER: ${{ needs.create-promotion-pr.outputs.pr_number }} PR_NUMBER: ${{ needs.create-promotion-pr.outputs.pr_number }}
run: | run: |
source .github/scripts/pr-body-utils.sh
if [ -n "$PR_NUMBER" ]; then if [ -n "$PR_NUMBER" ]; then
BASE=$(gh pr view "$PR_NUMBER" --json baseRefName --jq '.baseRefName') BASE=$(gh pr view "$PR_NUMBER" --json baseRefName --jq '.baseRefName')
if [ "$BASE" = "main" ]; then if [ "$BASE" = "main" ]; then
echo "Merging promotion PR #${PR_NUMBER} (targets main)" echo "Merging promotion PR #${PR_NUMBER} (targets main)"
TITLE=$(gh pr view "$PR_NUMBER" --json title --jq '.title') gh pr merge "$PR_NUMBER" --merge
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
echo "merged=true" >> "$GITHUB_OUTPUT" echo "merged=true" >> "$GITHUB_OUTPUT"
else else
echo "PR #${PR_NUMBER} targets '${BASE}' (not main) — leaving open for chain resolution" echo "PR #${PR_NUMBER} targets '${BASE}' (not main) — leaving open for chain resolution"
@@ -510,20 +464,18 @@ jobs:
steps: steps:
- name: Summary - name: Summary
run: | run: |
{ echo "## Staging CI Batch Results" >> "$GITHUB_STEP_SUMMARY"
echo "## Staging CI Batch Results" echo "" >> "$GITHUB_STEP_SUMMARY"
echo "" echo "| Check | Result |" >> "$GITHUB_STEP_SUMMARY"
echo "| Check | Result |" echo "|-------|--------|" >> "$GITHUB_STEP_SUMMARY"
echo "|-------|--------|" echo "| Tests | ${{ needs.tests.result }} |" >> "$GITHUB_STEP_SUMMARY"
echo "| Tests | ${{ needs.tests.result }} |" echo "| E2E | ${{ needs.e2e.result }} |" >> "$GITHUB_STEP_SUMMARY"
echo "| E2E | ${{ needs.e2e.result }} |" echo "| Promotion PR | ${{ needs.create-promotion-pr.result }} |" >> "$GITHUB_STEP_SUMMARY"
echo "| Promotion PR | ${{ needs.create-promotion-pr.result }} |" echo "| Gate | ${{ needs.gate.result }} |" >> "$GITHUB_STEP_SUMMARY"
echo "| Gate | ${{ needs.gate.result }} |" echo "| Tag Updated | ${{ needs.update-tag.result }} |" >> "$GITHUB_STEP_SUMMARY"
echo "| Tag Updated | ${{ needs.update-tag.result }} |" echo "" >> "$GITHUB_STEP_SUMMARY"
echo "" echo "Range: ${{ needs.check-changes.outputs.diff_range }}" >> "$GITHUB_STEP_SUMMARY"
echo "Range: ${{ needs.check-changes.outputs.diff_range }}" PR_NUM="${{ needs.create-promotion-pr.outputs.pr_number }}"
PR_NUM="${{ needs.create-promotion-pr.outputs.pr_number }}" if [ -n "$PR_NUM" ]; then
if [ -n "$PR_NUM" ]; then echo "Promotion PR: #${PR_NUM}" >> "$GITHUB_STEP_SUMMARY"
echo "Promotion PR: #${PR_NUM}" fi
fi
} >> "$GITHUB_STEP_SUMMARY"
@@ -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
-4
View File
@@ -14,10 +14,6 @@
target/ target/
# Python
__pycache__/
*.pyc
# Benchmark results (local runs, not committed) # Benchmark results (local runs, not committed)
bench-results/ bench-results/
-1
View File
@@ -19,7 +19,6 @@ WORKDIR /app
# Copy manifests first for layer caching # Copy manifests first for layer caching
COPY Cargo.toml Cargo.lock ./ COPY Cargo.toml Cargo.lock ./
COPY crates/ crates/
# Copy source, build script, tests, and supporting directories # Copy source, build script, tests, and supporting directories
COPY build.rs build.rs 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 WORKDIR /app
COPY Cargo.toml Cargo.lock ./ COPY Cargo.toml Cargo.lock ./
COPY crates/ crates/
COPY build.rs build.rs COPY build.rs build.rs
COPY src/ src/ COPY src/ src/
COPY tests/ tests/ COPY tests/ tests/
+2 -10
View File
@@ -132,7 +132,7 @@ fn embed_registry_catalog(root: &Path) {
// No registry dir: write empty catalog // No registry dir: write empty catalog
fs::write( fs::write(
&out_path, &out_path,
r#"{"tools":[],"channels":[],"mcp_servers":[],"bundles":{"bundles":{}}}"#, r#"{"tools":[],"channels":[],"bundles":{"bundles":{}}}"#,
) )
.unwrap(); .unwrap();
return; return;
@@ -140,7 +140,6 @@ fn embed_registry_catalog(root: &Path) {
let mut tools = Vec::new(); let mut tools = Vec::new();
let mut channels = Vec::new(); let mut channels = Vec::new();
let mut mcp_servers = Vec::new();
// Collect tool manifests // Collect tool manifests
let tools_dir = registry_dir.join("tools"); let tools_dir = registry_dir.join("tools");
@@ -154,12 +153,6 @@ fn embed_registry_catalog(root: &Path) {
collect_json_files(&channels_dir, &mut channels); collect_json_files(&channels_dir, &mut channels);
} }
// Collect MCP server manifests
let mcp_servers_dir = registry_dir.join("mcp-servers");
if mcp_servers_dir.is_dir() {
collect_json_files(&mcp_servers_dir, &mut mcp_servers);
}
// Read bundles // Read bundles
let bundles_path = registry_dir.join("_bundles.json"); let bundles_path = registry_dir.join("_bundles.json");
let bundles_raw = if bundles_path.is_file() { let bundles_raw = if bundles_path.is_file() {
@@ -170,10 +163,9 @@ fn embed_registry_catalog(root: &Path) {
// Build the combined JSON // Build the combined JSON
let catalog = format!( let catalog = format!(
r#"{{"tools":[{}],"channels":[{}],"mcp_servers":[{}],"bundles":{}}}"#, r#"{{"tools":[{}],"channels":[{}],"bundles":{}}}"#,
tools.join(","), tools.join(","),
channels.join(","), channels.join(","),
mcp_servers.join(","),
bundles_raw, bundles_raw,
); );
+1 -1
View File
@@ -121,7 +121,7 @@ dependencies = [
[[package]] [[package]]
name = "discord-channel" name = "discord-channel"
version = "0.2.0" version = "0.1.0"
dependencies = [ dependencies = [
"ed25519-dalek", "ed25519-dalek",
"hex", "hex",
-1
View File
@@ -642,7 +642,6 @@ fn poll_channel_mentions(channel_id: &str, bot_id: &str) {
}, },
thread_id: None, thread_id: None,
metadata_json, metadata_json,
attachments: vec![],
}); });
remember_processed_id(&mut recent_ids, &msg.id); remember_processed_id(&mut recent_ids, &msg.id);
-6
View File
@@ -6,12 +6,6 @@ rust-version = "1.92"
description = "Prompt injection defense, input validation, secret leak detection, and safety policy enforcement" description = "Prompt injection defense, input validation, secret leak detection, and safety policy enforcement"
authors = ["NEAR AI <[email protected]>"] authors = ["NEAR AI <[email protected]>"]
license = "MIT OR Apache-2.0" license = "MIT OR Apache-2.0"
homepage = "https://github.com/nearai/ironclaw"
repository = "https://github.com/nearai/ironclaw"
publish = false
[package.metadata.dist]
dist = false
[dependencies] [dependencies]
aho-corasick = "1" aho-corasick = "1"
+2 -2
View File
@@ -2,7 +2,7 @@
"name": "discord", "name": "discord",
"display_name": "Discord Channel", "display_name": "Discord Channel",
"kind": "channel", "kind": "channel",
"version": "0.2.1", "version": "0.2.0",
"wit_version": "0.3.0", "wit_version": "0.3.0",
"description": "Talk to your agent in Discord", "description": "Talk to your agent in Discord",
"keywords": [ "keywords": [
@@ -18,7 +18,7 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/discord-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/discord-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "efa1b9019fa33e243f8db1e1fcc732731d45836336bdd26ca19b6fe227ca8b69" "sha256": "efa1b9019fa33e243f8db1e1fcc732731d45836336bdd26ca19b6fe227ca8b69"
} }
}, },
+1 -1
View File
@@ -18,7 +18,7 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/slack-0.2.1-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/slack-0.2.1-wasm32-wasip2.tar.gz",
"sha256": "d4667e35126986509d862bc3a0088777305d8f41c75de83c1e223b42312ede48" "sha256": "d4667e35126986509d862bc3a0088777305d8f41c75de83c1e223b42312ede48"
} }
}, },
+2 -2
View File
@@ -2,7 +2,7 @@
"name": "telegram", "name": "telegram",
"display_name": "Telegram Channel", "display_name": "Telegram Channel",
"kind": "channel", "kind": "channel",
"version": "0.2.3", "version": "0.2.2",
"wit_version": "0.3.0", "wit_version": "0.3.0",
"description": "Talk to your agent through a Telegram bot", "description": "Talk to your agent through a Telegram bot",
"keywords": [ "keywords": [
@@ -18,7 +18,7 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/telegram-0.2.3-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/telegram-0.2.2-wasm32-wasip2.tar.gz",
"sha256": "b9a83d5a2d1285ce0ec116b354336a1f245f893291ccb01dffbcaccf89d72aed" "sha256": "b9a83d5a2d1285ce0ec116b354336a1f245f893291ccb01dffbcaccf89d72aed"
} }
}, },
+1 -1
View File
@@ -18,7 +18,7 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/whatsapp-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/whatsapp-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "feb9194719d9bed796b070ab4dc30348dbfb5d3dec56f9f21e02d14137abab01" "sha256": "feb9194719d9bed796b070ab4dc30348dbfb5d3dec56f9f21e02d14137abab01"
} }
}, },
-9
View File
@@ -1,9 +0,0 @@
{
"name": "asana",
"display_name": "Asana",
"kind": "mcp_server",
"description": "Connect to Asana for task management, projects, and team coordination",
"keywords": ["tasks", "projects", "management", "team"],
"url": "https://mcp.asana.com/v2/mcp",
"auth": "dcr"
}
-9
View File
@@ -1,9 +0,0 @@
{
"name": "cloudflare",
"display_name": "Cloudflare",
"kind": "mcp_server",
"description": "Connect to Cloudflare for DNS, Workers, KV, and infrastructure management",
"keywords": ["cdn", "dns", "workers", "hosting", "infrastructure"],
"url": "https://mcp.cloudflare.com/mcp",
"auth": "dcr"
}
-9
View File
@@ -1,9 +0,0 @@
{
"name": "intercom",
"display_name": "Intercom",
"kind": "mcp_server",
"description": "Connect to Intercom for customer messaging, support, and engagement",
"keywords": ["support", "customers", "messaging", "chat", "helpdesk"],
"url": "https://mcp.intercom.com/mcp",
"auth": "dcr"
}
-9
View File
@@ -1,9 +0,0 @@
{
"name": "linear",
"display_name": "Linear",
"kind": "mcp_server",
"description": "Connect to Linear for issue tracking, project management, and team workflows",
"keywords": ["issues", "tickets", "project", "tracking", "bugs"],
"url": "https://mcp.linear.app/sse",
"auth": "dcr"
}
-9
View File
@@ -1,9 +0,0 @@
{
"name": "notion",
"display_name": "Notion",
"kind": "mcp_server",
"description": "Connect to Notion for reading and writing pages, databases, and comments",
"keywords": ["notes", "wiki", "docs", "pages", "database"],
"url": "https://mcp.notion.com/mcp",
"auth": "dcr"
}
-9
View File
@@ -1,9 +0,0 @@
{
"name": "sentry",
"display_name": "Sentry",
"kind": "mcp_server",
"description": "Connect to Sentry for error tracking, performance monitoring, and debugging",
"keywords": ["errors", "monitoring", "debugging", "crashes", "performance"],
"url": "https://mcp.sentry.dev/mcp",
"auth": "dcr"
}
-9
View File
@@ -1,9 +0,0 @@
{
"name": "stripe",
"display_name": "Stripe",
"kind": "mcp_server",
"description": "Connect to Stripe for payment processing, subscriptions, and financial data",
"keywords": ["payments", "billing", "subscriptions", "invoices", "finance"],
"url": "https://mcp.stripe.com",
"auth": "dcr"
}
+1 -1
View File
@@ -19,7 +19,7 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/github-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/github-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "da9fac56b6f20197a415489bbaec9fefb085a5cf6324cab79ea48a47eb19c13b" "sha256": "da9fac56b6f20197a415489bbaec9fefb085a5cf6324cab79ea48a47eb19c13b"
} }
}, },
+1 -1
View File
@@ -18,7 +18,7 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/gmail-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/gmail-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "ee9574e02e92bc1d481f1310eb88afd99ee52bf6971074ab33bd76bf99b34b1d" "sha256": "ee9574e02e92bc1d481f1310eb88afd99ee52bf6971074ab33bd76bf99b34b1d"
} }
}, },
+1 -1
View File
@@ -18,7 +18,7 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-calendar-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/google-calendar-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "2fa47150ea222e787c122182ad6f4dfa2ffaf5fe490d05e8de887a76445f8d2d" "sha256": "2fa47150ea222e787c122182ad6f4dfa2ffaf5fe490d05e8de887a76445f8d2d"
} }
}, },
+1 -1
View File
@@ -18,7 +18,7 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-docs-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/google-docs-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "40e134a1c1564f832ca861c3396895d4e33ec67b99313fc1f97baf8d971423a9" "sha256": "40e134a1c1564f832ca861c3396895d4e33ec67b99313fc1f97baf8d971423a9"
} }
}, },
+1 -1
View File
@@ -18,7 +18,7 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-drive-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/google-drive-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "002a341a1d58125563a7c69561b26fbc2629b04ea723cade744102bdc0fbb71f" "sha256": "002a341a1d58125563a7c69561b26fbc2629b04ea723cade744102bdc0fbb71f"
} }
}, },
+1 -1
View File
@@ -18,7 +18,7 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-sheets-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/google-sheets-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "8aa2c9d52f033edea3a6c2311b0ec694ccb6d0a54ef07e94d72bf8be1ce8009a" "sha256": "8aa2c9d52f033edea3a6c2311b0ec694ccb6d0a54ef07e94d72bf8be1ce8009a"
} }
}, },
+1 -1
View File
@@ -17,7 +17,7 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-slides-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/google-slides-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "e931a97d4fd0b0b938e464dc7c7f2be6ea6b4d1508f5ea3cd931d44db23f05f5" "sha256": "e931a97d4fd0b0b938e464dc7c7f2be6ea6b4d1508f5ea3cd931d44db23f05f5"
} }
}, },
+2 -2
View File
@@ -17,8 +17,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/slack-0.2.1-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/slack-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "d4667e35126986509d862bc3a0088777305d8f41c75de83c1e223b42312ede48" "sha256": "8af3f884240de8413d272845fad2164a347d7d2a502a0d148aa38425b93f62ed"
} }
}, },
"auth_summary": { "auth_summary": {
+2 -2
View File
@@ -18,8 +18,8 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/telegram-0.2.2-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/telegram-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "b9a83d5a2d1285ce0ec116b354336a1f245f893291ccb01dffbcaccf89d72aed" "sha256": "2c66245913854be4294021fc6bb479e43f7d65830c5cec25cf6c60a71d1af468"
} }
}, },
"auth_summary": { "auth_summary": {
+2 -2
View File
@@ -2,7 +2,7 @@
"name": "web-search", "name": "web-search",
"display_name": "Web Search", "display_name": "Web Search",
"kind": "tool", "kind": "tool",
"version": "0.2.1", "version": "0.2.0",
"wit_version": "0.3.0", "wit_version": "0.3.0",
"description": "Search the web using Brave Search API", "description": "Search the web using Brave Search API",
"keywords": [ "keywords": [
@@ -18,7 +18,7 @@
}, },
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/web-search-0.2.0-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/web-search-0.2.0-wasm32-wasip2.tar.gz",
"sha256": "56834573c54ea2a33cea1eb0f04bbdf59f1ef8d8702995cf431b0921302eeccc" "sha256": "56834573c54ea2a33cea1eb0f04bbdf59f1ef8d8702995cf431b0921302eeccc"
} }
}, },
-4
View File
@@ -1,6 +1,2 @@
[workspace] [workspace]
git_release_enable = false git_release_enable = false
[[package]]
name = "ironclaw_safety"
release = false
+4 -6
View File
@@ -70,21 +70,19 @@ echo
# This is a WARNING, not a hard violation. # This is a WARNING, not a hard violation.
# -------------------------------------------------------------------------- # --------------------------------------------------------------------------
echo "--- Check 2: .unwrap() / .expect() / assert!() in production code ---" echo "--- Check 2: .unwrap() / .expect() in production code ---"
# Collect raw matches excluding obvious test-only files and lines. # Collect raw matches excluding obvious test-only files and lines
# Also catches assert!(), assert_eq!(), assert_ne!() but NOT debug_assert variants. raw_results=$(grep -rn '\.unwrap()\|\.expect(' src/ \
raw_results=$(grep -rnE '\.(unwrap|expect)\(|[^_]assert(_eq|_ne)?!' src/ \
--include='*.rs' \ --include='*.rs' \
| grep -v 'src/main.rs' \ | grep -v 'src/main.rs' \
| grep -v 'src/testing.rs' \ | grep -v 'src/testing.rs' \
| grep -v 'src/setup/' \ | grep -v 'src/setup/' \
| grep -Ev 'debug_assert|// safety:' \
|| true) || true)
if [ -n "$raw_results" ]; then if [ -n "$raw_results" ]; then
total=$(echo "$raw_results" | wc -l | tr -d ' ') total=$(echo "$raw_results" | wc -l | tr -d ' ')
echo "WARNING: ~$total .unwrap()/.expect()/assert!() calls found in src/ (excluding main/testing/setup)." echo "WARNING: ~$total .unwrap()/.expect() calls found in src/ (excluding main/testing/setup)."
echo "Many are in test modules; a per-file breakdown helps triage:" echo "Many are in test modules; a per-file breakdown helps triage:"
echo echo
# Show per-file counts, sorted by count descending, top 15 # Show per-file counts, sorted by count descending, top 15
-19
View File
@@ -10,7 +10,6 @@
# 3. Hardcoded /tmp paths in tests (flaky in parallel runs) # 3. Hardcoded /tmp paths in tests (flaky in parallel runs)
# 4. Tool parameters logged without redaction (secret leaks) # 4. Tool parameters logged without redaction (secret leaks)
# 5. Multi-step DB operations without transaction wrapping # 5. Multi-step DB operations without transaction wrapping
# 6. .unwrap(), .expect(), assert!() in production code (panics)
# #
# Suppress individual lines with an inline "// safety: <reason>" comment. # Suppress individual lines with an inline "// safety: <reason>" comment.
@@ -129,24 +128,6 @@ if [ -n "$DIFF_W_OUTPUT" ]; then
fi fi
fi fi
# 6. .unwrap(), .expect(), assert!() in production code
# Matches added lines containing panic-inducing calls.
# Excludes test files, test modules, and debug_assert (compiled out in release).
# Suppress with "// safety: <reason>".
PROD_DIFF="$DIFF_OUTPUT"
# Strip hunks from test-only files (tests/ directory, *_test.rs, test_*.rs)
PROD_DIFF=$(echo "$PROD_DIFF" | grep -v '^+++ b/tests/' || true)
if echo "$PROD_DIFF" | grep -nE '^\+' \
| grep -E '\.(unwrap|expect)\(|[^_]assert(_eq|_ne)?!' \
| grep -vE 'debug_assert|// safety:|#\[cfg\(test\)\]|#\[test\]|mod tests' \
| head -5 | grep -q .; then
warn "PANIC" "Production code must not use .unwrap(), .expect(), or assert!(). Use proper error handling."
echo "$PROD_DIFF" | grep -nE '^\+' \
| grep -E '\.(unwrap|expect)\(|[^_]assert(_eq|_ne)?!' \
| grep -vE 'debug_assert|// safety:|#\[cfg\(test\)\]|#\[test\]|mod tests' \
| head -5 | sed 's/^/ /'
fi
if [ "$WARNINGS" -gt 0 ]; then if [ "$WARNINGS" -gt 0 ]; then
echo "" echo ""
echo "Found $WARNINGS potential issue(s). Fix them or add '// safety: <reason>' to suppress." echo "Found $WARNINGS potential issue(s). Fix them or add '// safety: <reason>' to suppress."
-24
View File
@@ -152,30 +152,6 @@ pub async fn run_agentic_loop(
// Call LLM // Call LLM
let output = delegate.call_llm(reasoning, reason_ctx, iteration).await?; let output = delegate.call_llm(reasoning, reason_ctx, iteration).await?;
match &output.result {
RespondResult::Text(text) => {
tracing::debug!(
iteration,
len = text.len(),
has_suggestions = text.contains("<suggestions>"),
response = %text,
"LLM text response"
);
}
RespondResult::ToolCalls {
tool_calls,
content,
} => {
let names: Vec<&str> = tool_calls.iter().map(|tc| tc.name.as_str()).collect();
tracing::debug!(
iteration,
tools = ?names,
has_content = content.is_some(),
"LLM tool_calls response"
);
}
}
match output.result { match output.result {
RespondResult::Text(text) => { RespondResult::Text(text) => {
// Tool intent nudge: if the LLM says "let me search..." without // Tool intent nudge: if the LLM says "let me search..." without
-97
View File
@@ -1051,54 +1051,6 @@ fn strip_internal_tool_call_text(text: &str) -> String {
} }
} }
/// Extract `<suggestions>["...","..."]</suggestions>` from a response string.
///
/// Returns `(cleaned_text, suggestions)`. The `<suggestions>` block is stripped
/// from the text regardless of whether the JSON inside parses successfully.
/// Only the **last** `<suggestions>` block is used (closest to end of response).
/// Blocks inside markdown code fences are ignored.
pub(crate) fn extract_suggestions(text: &str) -> (String, Vec<String>) {
use regex::Regex;
use std::sync::LazyLock;
static RE: LazyLock<Regex> = LazyLock::new(|| {
Regex::new(r"(?s)<suggestions>\s*(.*?)\s*</suggestions>").expect("valid regex") // safety: constant pattern
});
// Find the position of the last closing code fence to avoid matching inside code blocks
let last_code_fence = text.rfind("```").unwrap_or(0);
// Find all matches, take the last one that's after the last code fence
let mut best_match: Option<regex::Match<'_>> = None;
let mut best_capture: Option<String> = None;
for caps in RE.captures_iter(text) {
if let (Some(full), Some(inner)) = (caps.get(0), caps.get(1))
&& full.start() >= last_code_fence
{
best_match = Some(full);
best_capture = Some(inner.as_str().to_string());
}
}
let Some(full) = best_match else {
return (text.to_string(), Vec::new());
};
let cleaned = format!("{}{}", &text[..full.start()], &text[full.end()..]); // safety: regex match boundaries are valid UTF-8
let cleaned = cleaned.trim().to_string();
// Parse the JSON array
let suggestions = best_capture
.and_then(|json| serde_json::from_str::<Vec<String>>(&json).ok())
.unwrap_or_default()
.into_iter()
.filter(|s| !s.trim().is_empty() && s.len() <= 80)
.take(3)
.collect();
(cleaned, suggestions)
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use std::sync::Arc; use std::sync::Arc;
@@ -2245,55 +2197,6 @@ mod tests {
assert_eq!(result, input); assert_eq!(result, input);
} }
#[test]
fn test_extract_suggestions_basic() {
let input = "Here is my answer.\n<suggestions>[\"Check logs\", \"Deploy\"]</suggestions>";
let (text, suggestions) = super::extract_suggestions(input);
assert_eq!(text, "Here is my answer."); // safety: test
assert_eq!(suggestions, vec!["Check logs", "Deploy"]); // safety: test
}
#[test]
fn test_extract_suggestions_no_tag() {
let input = "Just a plain response.";
let (text, suggestions) = super::extract_suggestions(input);
assert_eq!(text, "Just a plain response."); // safety: test
assert!(suggestions.is_empty()); // safety: test
}
#[test]
fn test_extract_suggestions_malformed_json() {
let input = "Answer.\n<suggestions>not json</suggestions>";
let (text, suggestions) = super::extract_suggestions(input);
assert_eq!(text, "Answer."); // safety: test
assert!(suggestions.is_empty()); // safety: test
}
#[test]
fn test_extract_suggestions_inside_code_fence() {
let input = "```\n<suggestions>[\"foo\"]</suggestions>\n```";
let (text, suggestions) = super::extract_suggestions(input);
// The tag is inside a code fence, so it should not be extracted
assert_eq!(text, input); // safety: test
assert!(suggestions.is_empty()); // safety: test
}
#[test]
fn test_extract_suggestions_after_code_fence() {
let input = "```\ncode\n```\nAnswer.\n<suggestions>[\"foo\"]</suggestions>";
let (text, suggestions) = super::extract_suggestions(input);
assert_eq!(text, "```\ncode\n```\nAnswer."); // safety: test
assert_eq!(suggestions, vec!["foo"]); // safety: test
}
#[test]
fn test_extract_suggestions_filters_long() {
let long = "x".repeat(81);
let input = format!("Answer.\n<suggestions>[\"{}\", \"ok\"]</suggestions>", long);
let (_, suggestions) = super::extract_suggestions(&input);
assert_eq!(suggestions, vec!["ok"]); // safety: test
}
#[test] #[test]
fn test_tool_error_format_includes_tool_name() { fn test_tool_error_format_includes_tool_name() {
// Regression test for issue #487: tool errors sent to the LLM should // Regression test for issue #487: tool errors sent to the LLM should
+14 -42
View File
@@ -93,26 +93,19 @@ impl RoutineEngine {
let mut cache = Vec::new(); let mut cache = Vec::new();
for routine in routines { for routine in routines {
match &routine.trigger { match &routine.trigger {
Trigger::Event { pattern, .. } => { Trigger::Event { pattern, .. } => match Regex::new(pattern) {
// Use RegexBuilder with size limit to prevent ReDoS Ok(re) => cache.push(EventMatcher::Message {
// from user-supplied patterns (issue #825). routine: routine.clone(),
match regex::RegexBuilder::new(pattern) regex: re,
.size_limit(64 * 1024) // 64KB compiled size limit }),
.build() Err(e) => {
{ tracing::warn!(
Ok(re) => cache.push(EventMatcher::Message { routine = %routine.name,
routine: routine.clone(), "Invalid event regex '{}': {}",
regex: re, pattern, e
}), );
Err(e) => {
tracing::warn!(
routine = %routine.name,
"Invalid or too complex event regex '{}': {}",
pattern, e
);
}
} }
} },
Trigger::SystemEvent { .. } => { Trigger::SystemEvent { .. } => {
cache.push(EventMatcher::System { cache.push(EventMatcher::System {
routine: routine.clone(), routine: routine.clone(),
@@ -980,18 +973,6 @@ async fn execute_lightweight_with_tools(
} }
}; };
// Truncate oversized tool output to prevent unbounded context growth.
// Routine tool loops are lightweight and should not accumulate
// large payloads across iterations.
const MAX_TOOL_OUTPUT_CHARS: usize = 8192;
let result_content = if result_content.len() > MAX_TOOL_OUTPUT_CHARS {
let truncated = &result_content
[..result_content.floor_char_boundary(MAX_TOOL_OUTPUT_CHARS)];
format!("{truncated}\n... [output truncated to {MAX_TOOL_OUTPUT_CHARS} chars]")
} else {
result_content
};
// Add tool result to context // Add tool result to context
messages.push(ChatMessage::tool_result(&tc.id, &tc.name, &result_content)); messages.push(ChatMessage::tool_result(&tc.id, &tc.name, &result_content));
} }
@@ -1169,11 +1150,9 @@ pub fn spawn_cron_ticker(
interval: Duration, interval: Duration,
) -> tokio::task::JoinHandle<()> { ) -> tokio::task::JoinHandle<()> {
tokio::spawn(async move { 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); let mut ticker = tokio::time::interval(interval);
// Skip immediate first tick
ticker.tick().await;
loop { loop {
ticker.tick().await; ticker.tick().await;
@@ -1379,11 +1358,4 @@ mod tests {
assert_eq!(finish_reason_length, crate::llm::FinishReason::Length); assert_eq!(finish_reason_length, crate::llm::FinishReason::Length);
assert_eq!(finish_reason_stop, crate::llm::FinishReason::Stop); 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...");
}
} }
-28
View File
@@ -420,10 +420,6 @@ impl Agent {
// Complete, fail, or request approval // Complete, fail, or request approval
match result { match result {
Ok(AgenticLoopResult::Response(response)) => { Ok(AgenticLoopResult::Response(response)) => {
// Extract <suggestions> from response text before user sees it
let (response, suggestions) =
crate::agent::dispatcher::extract_suggestions(&response);
// Hook: TransformResponse — allow hooks to modify or reject the final response // Hook: TransformResponse — allow hooks to modify or reject the final response
let response = { let response = {
let event = crate::hooks::HookEvent::ResponseTransform { let event = crate::hooks::HookEvent::ResponseTransform {
@@ -477,18 +473,6 @@ impl Agent {
) )
.await; .await;
// Send suggestions after response (best-effort, rendered by web gateway)
if !suggestions.is_empty() {
let _ = self
.channels
.send_status(
&message.channel,
StatusUpdate::Suggestions { suggestions },
&message.metadata,
)
.await;
}
Ok(SubmissionResult::response(response)) Ok(SubmissionResult::response(response))
} }
Ok(AgenticLoopResult::NeedApproval { pending }) => { Ok(AgenticLoopResult::NeedApproval { pending }) => {
@@ -1350,8 +1334,6 @@ impl Agent {
match result { match result {
Ok(AgenticLoopResult::Response(response)) => { Ok(AgenticLoopResult::Response(response)) => {
let (response, suggestions) =
crate::agent::dispatcher::extract_suggestions(&response);
thread.complete_turn(&response); thread.complete_turn(&response);
let (turn_number, tool_calls) = thread let (turn_number, tool_calls) = thread
.turns .turns
@@ -1382,16 +1364,6 @@ impl Agent {
&message.metadata, &message.metadata,
) )
.await; .await;
if !suggestions.is_empty() {
let _ = self
.channels
.send_status(
&message.channel,
StatusUpdate::Suggestions { suggestions },
&message.metadata,
)
.await;
}
Ok(SubmissionResult::response(response)) Ok(SubmissionResult::response(response))
} }
Ok(AgenticLoopResult::NeedApproval { Ok(AgenticLoopResult::NeedApproval {
+1 -2
View File
@@ -290,7 +290,6 @@ impl AppBuilder {
Arc::new(ToolRegistry::new()) Arc::new(ToolRegistry::new())
}; };
tools.register_builtin_tools(); tools.register_builtin_tools();
tools.register_tool_info();
if let Some(ref ss) = self.secrets_store { if let Some(ref ss) = self.secrets_store {
tools.register_secrets_tools(Arc::clone(ss)); tools.register_secrets_tools(Arc::clone(ss));
@@ -594,7 +593,7 @@ impl AppBuilder {
let entries: Vec<_> = catalog let entries: Vec<_> = catalog
.all() .all()
.iter() .iter()
.filter_map(|m| m.to_registry_entry()) .map(|m| m.to_registry_entry())
.collect(); .collect();
tracing::debug!( tracing::debug!(
count = entries.len(), count = entries.len(),
-2
View File
@@ -238,8 +238,6 @@ pub enum StatusUpdate {
/// Optional workspace path where the image was saved. /// Optional workspace path where the image was saved.
path: Option<String>, path: Option<String>,
}, },
/// Suggested follow-up messages for the user.
Suggestions { suggestions: Vec<String> },
} }
impl StatusUpdate { impl StatusUpdate {
+84 -153
View File
@@ -140,7 +140,7 @@ struct WebhookRequest {
content: String, content: String,
/// Optional thread ID for conversation tracking. /// Optional thread ID for conversation tracking.
thread_id: Option<String>, thread_id: Option<String>,
/// Deprecated: webhook secret in request body. Use X-Hub-Signature-256 header instead. /// Deprecated: webhook secret in request body. Use X-IronClaw-Signature header instead.
/// This field is accepted for backward compatibility but will be removed in a future release. /// This field is accepted for backward compatibility but will be removed in a future release.
secret: Option<String>, secret: Option<String>,
/// Whether to wait for a synchronous response. /// Whether to wait for a synchronous response.
@@ -269,108 +269,95 @@ async fn webhook_handler(
let mut fallback_req = None; let mut fallback_req = None;
{ {
let webhook_secret = state.webhook_secret.read().await; let webhook_secret = state.webhook_secret.read().await;
let expected_secret = match webhook_secret.as_ref() { if let Some(expected_secret) = webhook_secret.as_ref() {
Some(secret) => secret.expose_secret(), let expected_secret = expected_secret.expose_secret();
None => {
// No secret configured — reject all requests. This guards against
// the secret being cleared at runtime via update_secret(None).
// The start() method also prevents startup without a secret, but
// this is defense-in-depth for the SIGHUP hot-swap path.
return (
StatusCode::SERVICE_UNAVAILABLE,
Json(WebhookResponse {
message_id: Uuid::nil(),
status: "error".to_string(),
response: Some("Webhook authentication not configured".to_string()),
}),
)
.into_response();
}
};
match headers.get("x-hub-signature-256") { match headers.get("x-ironclaw-signature") {
Some(raw_signature) => match raw_signature.to_str() { Some(raw_signature) => match raw_signature.to_str() {
Ok(signature) => { Ok(signature) => {
if !verify_hmac_signature(expected_secret, &body, signature) { if !verify_hmac_signature(expected_secret, &body, signature) {
return ( return (
StatusCode::UNAUTHORIZED, StatusCode::UNAUTHORIZED,
Json(WebhookResponse { Json(WebhookResponse {
message_id: Uuid::nil(), message_id: Uuid::nil(),
status: "error".to_string(), status: "error".to_string(),
response: Some("Invalid webhook signature".to_string()), response: Some("Invalid webhook signature".to_string()),
}), }),
) )
.into_response(); .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(_) => { Err(_) => {
return ( return (
StatusCode::UNAUTHORIZED, StatusCode::UNAUTHORIZED,
Json(WebhookResponse { Json(WebhookResponse {
message_id: Uuid::nil(), message_id: Uuid::nil(),
status: "error".to_string(), status: "error".to_string(),
response: Some( response: Some("Invalid signature header encoding".to_string()),
"Webhook authentication required. Provide X-Hub-Signature-256 header \
(preferred) or 'secret' field in body (deprecated)."
.to_string(),
),
}), }),
) )
.into_response(); .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 { match &req.secret {
Some(provided) Some(provided)
if bool::from(provided.as_bytes().ct_eq(expected_secret.as_bytes())) => 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-Hub-Signature-256 header (HMAC-SHA256). \ tracing::warn!(
Body secret support will be removed in a future release." "Webhook authenticated via deprecated 'secret' field in request body. \
); Migrate to X-IronClaw-Signature header (HMAC-SHA256). \
fallback_req = Some(req); Body secret support will be removed in a future release."
} );
Some(_) => { fallback_req = Some(req);
return ( }
StatusCode::UNAUTHORIZED, Some(_) => {
Json(WebhookResponse { return (
message_id: Uuid::nil(), StatusCode::UNAUTHORIZED,
status: "error".to_string(), Json(WebhookResponse {
response: Some("Invalid webhook secret".to_string()), message_id: Uuid::nil(),
}), status: "error".to_string(),
) response: Some("Invalid webhook secret".to_string()),
.into_response(); }),
} )
None => { .into_response();
return ( }
StatusCode::UNAUTHORIZED, None => {
Json(WebhookResponse { return (
message_id: Uuid::nil(), StatusCode::UNAUTHORIZED,
status: "error".to_string(), Json(WebhookResponse {
response: Some( message_id: Uuid::nil(),
"Webhook authentication required. Provide X-Hub-Signature-256 header \ status: "error".to_string(),
(preferred) or 'secret' field in body (deprecated)." response: Some(
.to_string(), "Webhook authentication required. Provide X-IronClaw-Signature header \
), (preferred) or 'secret' field in body (deprecated)."
}), .to_string(),
) ),
.into_response(); }),
)
.into_response();
}
} }
} }
} }
@@ -726,7 +713,7 @@ mod tests {
.method("POST") .method("POST")
.uri("/webhook") .uri("/webhook")
.header("content-type", "application/json") .header("content-type", "application/json")
.header("x-hub-signature-256", signature) .header("x-ironclaw-signature", signature)
.body(Body::from(body_bytes)) .body(Body::from(body_bytes))
.unwrap(); .unwrap();
@@ -749,7 +736,7 @@ mod tests {
.method("POST") .method("POST")
.uri("/webhook") .uri("/webhook")
.header("content-type", "application/json") .header("content-type", "application/json")
.header("x-hub-signature-256", signature) .header("x-ironclaw-signature", signature)
.body(Body::from(body_bytes)) .body(Body::from(body_bytes))
.unwrap(); .unwrap();
@@ -770,7 +757,7 @@ mod tests {
.method("POST") .method("POST")
.uri("/webhook") .uri("/webhook")
.header("content-type", "application/json") .header("content-type", "application/json")
.header("x-hub-signature-256", "not-a-valid-signature") .header("x-ironclaw-signature", "not-a-valid-signature")
.body(Body::from(serde_json::to_vec(&body).unwrap())) .body(Body::from(serde_json::to_vec(&body).unwrap()))
.unwrap(); .unwrap();
@@ -919,7 +906,7 @@ mod tests {
.method("POST") .method("POST")
.uri("/webhook") .uri("/webhook")
.header("content-type", "application/json") .header("content-type", "application/json")
.header("x-hub-signature-256", signature) .header("x-ironclaw-signature", signature)
.body(Body::from(body_bytes)) .body(Body::from(body_bytes))
.unwrap(); .unwrap();
@@ -941,7 +928,7 @@ mod tests {
.method("POST") .method("POST")
.uri("/webhook") .uri("/webhook")
.header("content-type", "application/json") .header("content-type", "application/json")
.header("x-hub-signature-256", signature) .header("x-ironclaw-signature", signature)
.body(Body::from(body)) .body(Body::from(body))
.unwrap(); .unwrap();
@@ -966,7 +953,7 @@ mod tests {
.method("POST") .method("POST")
.uri("/webhook") .uri("/webhook")
.header("content-type", "text/plain") .header("content-type", "text/plain")
.header("x-hub-signature-256", signature) .header("x-ironclaw-signature", signature)
.body(Body::from(body_bytes)) .body(Body::from(body_bytes))
.unwrap(); .unwrap();
@@ -991,7 +978,7 @@ mod tests {
.body(Body::from(serde_json::to_vec(&body).unwrap())) .body(Body::from(serde_json::to_vec(&body).unwrap()))
.unwrap(); .unwrap();
req.headers_mut().insert( req.headers_mut().insert(
"x-hub-signature-256", "x-ironclaw-signature",
HeaderValue::from_bytes(b"\xFF").unwrap(), HeaderValue::from_bytes(b"\xFF").unwrap(),
); );
@@ -1065,32 +1052,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-hub-signature-256", signature)
.body(Body::from(body_bytes))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE); // safety: test assertion
}
#[tokio::test] #[tokio::test]
async fn test_concurrent_requests_during_secret_update() { async fn test_concurrent_requests_during_secret_update() {
use std::sync::Arc as StdArc; use std::sync::Arc as StdArc;
@@ -1209,34 +1170,4 @@ mod tests {
let body = b"test body content"; let body = b"test body content";
assert!(!verify_hmac_signature(secret, body, "sha256=not-hex!")); assert!(!verify_hmac_signature(secret, body, "sha256=not-hex!"));
} }
/// Regression test for issue #1033: when the webhook secret is cleared at
/// runtime via update_secret(None), subsequent requests must be rejected
/// instead of being processed without authentication.
#[tokio::test]
async fn webhook_rejects_when_secret_cleared_at_runtime() {
let channel = test_channel(Some("initial-secret"));
let _stream = channel.start().await.unwrap();
// Clear the secret at runtime (simulates a bad SIGHUP config reload)
channel.update_secret(None).await;
let app = channel.routes();
let body = serde_json::json!({
"content": "hello"
});
let req = Request::builder()
.method("POST")
.uri("/webhook")
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&body).unwrap()))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::SERVICE_UNAVAILABLE,
"requests must be rejected when webhook secret is cleared at runtime"
);
}
} }
-4
View File
@@ -294,8 +294,6 @@ impl Channel for RelayChannel {
match client.connect_stream(&token, stream_timeout_secs).await { match client.connect_stream(&token, stream_timeout_secs).await {
Ok((new_stream, new_parser)) => { Ok((new_stream, new_parser)) => {
tracing::info!("Relay SSE stream reconnected"); tracing::info!("Relay SSE stream reconnected");
consecutive_failures = 0;
backoff_ms = backoff_initial_ms;
current_stream = new_stream; current_stream = new_stream;
// Abort old parser before replacing // Abort old parser before replacing
if let Some(old) = parser_handle.write().await.take() { if let Some(old) = parser_handle.write().await.take() {
@@ -314,8 +312,6 @@ impl Channel for RelayChannel {
tracing::info!( tracing::info!(
"Relay SSE stream reconnected with new token" "Relay SSE stream reconnected with new token"
); );
consecutive_failures = 0;
backoff_ms = backoff_initial_ms;
current_stream = new_stream; current_stream = new_stream;
if let Some(old) = parser_handle.write().await.take() { if let Some(old) = parser_handle.write().await.take() {
old.abort(); old.abort();
-3
View File
@@ -607,9 +607,6 @@ impl Channel for ReplChannel {
eprintln!("\x1b[36m [image generated]\x1b[0m"); eprintln!("\x1b[36m [image generated]\x1b[0m");
} }
} }
StatusUpdate::Suggestions { .. } => {
// Suggestions are only rendered by the web gateway
}
} }
Ok(()) Ok(())
} }
+24 -52
View File
@@ -1664,9 +1664,7 @@ impl WasmChannel {
.await; .await;
let pairing_store = self.pairing_store.clone(); let pairing_store = self.pairing_store.clone();
let Some(wit_update) = status_to_wit(status, metadata) else { let wit_update = status_to_wit(status, metadata);
return Ok(());
};
let result = tokio::time::timeout(timeout, async move { let result = tokio::time::timeout(timeout, async move {
tokio::task::spawn_blocking(move || { tokio::task::spawn_blocking(move || {
@@ -1835,9 +1833,7 @@ impl WasmChannel {
.await; .await;
let pairing_store = self.pairing_store.clone(); let pairing_store = self.pairing_store.clone();
let callback_timeout = self.runtime.config().callback_timeout; let callback_timeout = self.runtime.config().callback_timeout;
let Some(wit_update) = status_to_wit(&status, metadata) else { let wit_update = status_to_wit(&status, metadata);
return Ok(());
};
let handle = tokio::spawn(async move { let handle = tokio::spawn(async move {
let mut interval = tokio::time::interval(Duration::from_secs(4)); let mut interval = tokio::time::interval(Duration::from_secs(4));
@@ -2708,13 +2704,10 @@ fn truncate_status_text(input: &str, max_chars: usize) -> String {
} }
} }
fn status_to_wit( fn status_to_wit(status: &StatusUpdate, metadata: &serde_json::Value) -> wit_channel::StatusUpdate {
status: &StatusUpdate,
metadata: &serde_json::Value,
) -> Option<wit_channel::StatusUpdate> {
let metadata_json = serde_json::to_string(metadata).unwrap_or_default(); let metadata_json = serde_json::to_string(metadata).unwrap_or_default();
Some(match status { match status {
StatusUpdate::Thinking(msg) => wit_channel::StatusUpdate { StatusUpdate::Thinking(msg) => wit_channel::StatusUpdate {
status: wit_channel::StatusType::Thinking, status: wit_channel::StatusType::Thinking,
message: msg.clone(), message: msg.clone(),
@@ -2834,9 +2827,7 @@ fn status_to_wit(
}, },
metadata_json, metadata_json,
}, },
// Suggestions are web-gateway-only; skip for WASM channels }
StatusUpdate::Suggestions { .. } => return None,
})
} }
/// Clone a WIT StatusUpdate (the generated type doesn't derive Clone). /// Clone a WIT StatusUpdate (the generated type doesn't derive Clone).
@@ -3565,8 +3556,7 @@ mod tests {
let wit = status_to_wit( let wit = status_to_wit(
&crate::channels::StatusUpdate::Thinking("Processing...".into()), &crate::channels::StatusUpdate::Thinking("Processing...".into()),
&metadata, &metadata,
) );
.unwrap(); // safety: test
assert!(matches!( assert!(matches!(
wit.status, wit.status,
@@ -3584,8 +3574,7 @@ mod tests {
let wit = status_to_wit( let wit = status_to_wit(
&crate::channels::StatusUpdate::Status("Done".into()), &crate::channels::StatusUpdate::Status("Done".into()),
&metadata, &metadata,
) );
.unwrap(); // safety: test
assert!(matches!(wit.status, super::wit_channel::StatusType::Done)); assert!(matches!(wit.status, super::wit_channel::StatusType::Done));
} }
@@ -3600,16 +3589,14 @@ mod tests {
let wit = status_to_wit( let wit = status_to_wit(
&crate::channels::StatusUpdate::Status("done".into()), &crate::channels::StatusUpdate::Status("done".into()),
&metadata, &metadata,
) );
.unwrap(); // safety: test
assert!(matches!(wit.status, super::wit_channel::StatusType::Done)); assert!(matches!(wit.status, super::wit_channel::StatusType::Done));
// with whitespace // with whitespace
let wit = status_to_wit( let wit = status_to_wit(
&crate::channels::StatusUpdate::Status(" Done ".into()), &crate::channels::StatusUpdate::Status(" Done ".into()),
&metadata, &metadata,
) );
.unwrap(); // safety: test
assert!(matches!(wit.status, super::wit_channel::StatusType::Done)); assert!(matches!(wit.status, super::wit_channel::StatusType::Done));
} }
@@ -3621,8 +3608,7 @@ mod tests {
let wit = status_to_wit( let wit = status_to_wit(
&crate::channels::StatusUpdate::Status("Interrupted".into()), &crate::channels::StatusUpdate::Status("Interrupted".into()),
&metadata, &metadata,
) );
.unwrap(); // safety: test
assert!(matches!( assert!(matches!(
wit.status, wit.status,
@@ -3640,8 +3626,7 @@ mod tests {
let wit = status_to_wit( let wit = status_to_wit(
&crate::channels::StatusUpdate::Status("interrupted".into()), &crate::channels::StatusUpdate::Status("interrupted".into()),
&metadata, &metadata,
) );
.unwrap(); // safety: test
assert!(matches!( assert!(matches!(
wit.status, wit.status,
super::wit_channel::StatusType::Interrupted super::wit_channel::StatusType::Interrupted
@@ -3651,8 +3636,7 @@ mod tests {
let wit = status_to_wit( let wit = status_to_wit(
&crate::channels::StatusUpdate::Status(" Interrupted ".into()), &crate::channels::StatusUpdate::Status(" Interrupted ".into()),
&metadata, &metadata,
) );
.unwrap(); // safety: test
assert!(matches!( assert!(matches!(
wit.status, wit.status,
super::wit_channel::StatusType::Interrupted super::wit_channel::StatusType::Interrupted
@@ -3667,8 +3651,7 @@ mod tests {
let wit = status_to_wit( let wit = status_to_wit(
&crate::channels::StatusUpdate::Status("Awaiting approval".into()), &crate::channels::StatusUpdate::Status("Awaiting approval".into()),
&metadata, &metadata,
) );
.unwrap(); // safety: test
assert!(matches!(wit.status, super::wit_channel::StatusType::Status)); assert!(matches!(wit.status, super::wit_channel::StatusType::Status));
assert_eq!(wit.message, "Awaiting approval"); assert_eq!(wit.message, "Awaiting approval");
@@ -3687,8 +3670,7 @@ mod tests {
setup_url: None, setup_url: None,
}, },
&metadata, &metadata,
) );
.unwrap(); // safety: test
assert!(matches!( assert!(matches!(
wit.status, wit.status,
@@ -3708,8 +3690,7 @@ mod tests {
name: "http_request".to_string(), name: "http_request".to_string(),
}, },
&metadata, &metadata,
) );
.unwrap(); // safety: test
assert!(matches!( assert!(matches!(
wit.status, wit.status,
@@ -3731,8 +3712,7 @@ mod tests {
parameters: None, parameters: None,
}, },
&metadata, &metadata,
) );
.unwrap(); // safety: test
assert!(matches!( assert!(matches!(
wit.status, wit.status,
@@ -3754,8 +3734,7 @@ mod tests {
parameters: None, parameters: None,
}, },
&metadata, &metadata,
) );
.unwrap(); // safety: test
assert!(matches!( assert!(matches!(
wit.status, wit.status,
@@ -3775,8 +3754,7 @@ mod tests {
preview: "{".to_string() + "\"temperature\": 22}", preview: "{".to_string() + "\"temperature\": 22}",
}, },
&metadata, &metadata,
) );
.unwrap(); // safety: test
assert!(matches!( assert!(matches!(
wit.status, wit.status,
@@ -3797,8 +3775,7 @@ mod tests {
preview: long_preview, preview: long_preview,
}, },
&metadata, &metadata,
) );
.unwrap(); // safety: test
assert!(matches!( assert!(matches!(
wit.status, wit.status,
@@ -3819,8 +3796,7 @@ mod tests {
browse_url: "https://example.com/jobs/job-1".to_string(), browse_url: "https://example.com/jobs/job-1".to_string(),
}, },
&metadata, &metadata,
) );
.unwrap(); // safety: test
assert!(matches!( assert!(matches!(
wit.status, wit.status,
@@ -3842,8 +3818,7 @@ mod tests {
message: "Token saved".to_string(), message: "Token saved".to_string(),
}, },
&metadata, &metadata,
) );
.unwrap(); // safety: test
assert!(matches!( assert!(matches!(
wit.status, wit.status,
@@ -3865,8 +3840,7 @@ mod tests {
message: "Invalid token".to_string(), message: "Invalid token".to_string(),
}, },
&metadata, &metadata,
) );
.unwrap(); // safety: test
assert!(matches!( assert!(matches!(
wit.status, wit.status,
@@ -3889,8 +3863,7 @@ mod tests {
parameters: serde_json::json!({"url": "https://api.weather.test"}), parameters: serde_json::json!({"url": "https://api.weather.test"}),
}, },
&metadata, &metadata,
) );
.unwrap(); // safety: test
assert!(matches!( assert!(matches!(
wit.status, wit.status,
@@ -3914,8 +3887,7 @@ mod tests {
parameters: serde_json::json!({"url": "https://api.weather.test"}), parameters: serde_json::json!({"url": "https://api.weather.test"}),
}, },
&metadata, &metadata,
) );
.unwrap(); // safety: test
assert!(matches!( assert!(matches!(
wit.status, wit.status,
-31
View File
@@ -10,7 +10,6 @@ use axum::{
use serde::Deserialize; use serde::Deserialize;
use uuid::Uuid; use uuid::Uuid;
use crate::agent::routine::{Trigger, next_cron_fire};
use crate::channels::web::server::GatewayState; use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*; use crate::channels::web::types::*;
use crate::error::RoutineError; use crate::error::RoutineError;
@@ -183,41 +182,17 @@ pub async fn routines_toggle_handler(
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?; .ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
let was_enabled = routine.enabled;
// If a specific value was provided, use it; otherwise toggle. // If a specific value was provided, use it; otherwise toggle.
routine.enabled = match body { routine.enabled = match body {
Some(Json(req)) => req.enabled.unwrap_or(!routine.enabled), Some(Json(req)) => req.enabled.unwrap_or(!routine.enabled),
None => !routine.enabled, None => !routine.enabled,
}; };
// When re-enabling a cron routine, recompute next_fire_at so the cron
// ticker can pick it up. Mirrors the CLI behavior (issue #1077).
if routine.enabled
&& !was_enabled
&& let Trigger::Cron {
ref schedule,
ref timezone,
} = routine.trigger
{
routine.next_fire_at = next_cron_fire(schedule, timezone.as_deref()).map_err(|e| {
(
StatusCode::INTERNAL_SERVER_ERROR,
format!("Failed to compute next fire: {e}"),
)
})?;
}
store store
.update_routine(&routine) .update_routine(&routine)
.await .await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
// Refresh the in-memory event trigger cache so event/system_event
// routines reflect the new enabled state immediately (issue #1076).
if let Some(engine) = state.routine_engine.read().await.as_ref() {
engine.refresh_event_cache().await;
}
Ok(Json(serde_json::json!({ Ok(Json(serde_json::json!({
"status": if routine.enabled { "enabled" } else { "disabled" }, "status": if routine.enabled { "enabled" } else { "disabled" },
"routine_id": routine_id, "routine_id": routine_id,
@@ -242,12 +217,6 @@ pub async fn routines_delete_handler(
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
if deleted { if deleted {
// Refresh the in-memory event trigger cache so deleted event/system_event
// routines stop firing immediately (issue #1076).
if let Some(engine) = state.routine_engine.read().await.as_ref() {
engine.refresh_event_cache().await;
}
Ok(Json(serde_json::json!({ Ok(Json(serde_json::json!({
"status": "deleted", "status": "deleted",
"routine_id": routine_id, "routine_id": routine_id,
-4
View File
@@ -397,10 +397,6 @@ impl Channel for GatewayChannel {
StatusUpdate::ImageGenerated { data_url, path } => SseEvent::ImageGenerated { StatusUpdate::ImageGenerated { data_url, path } => SseEvent::ImageGenerated {
data_url, data_url,
path, path,
thread_id: thread_id.clone(),
},
StatusUpdate::Suggestions { suggestions } => SseEvent::Suggestions {
suggestions,
thread_id, thread_id,
}, },
}; };
+72 -130
View File
@@ -26,7 +26,6 @@ use tower_http::set_header::SetResponseHeaderLayer;
use uuid::Uuid; use uuid::Uuid;
use crate::agent::SessionManager; use crate::agent::SessionManager;
use crate::agent::routine::{Trigger, next_cron_fire};
use crate::bootstrap::ironclaw_base_dir; use crate::bootstrap::ironclaw_base_dir;
use crate::channels::IncomingMessage; use crate::channels::IncomingMessage;
use crate::channels::relay::DEFAULT_RELAY_NAME; use crate::channels::relay::DEFAULT_RELAY_NAME;
@@ -573,14 +572,6 @@ async fn oauth_callback_handler(
extension = %flow.extension_name, extension = %flow.extension_name,
"OAuth flow expired" "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); return oauth_error_page(&flow.display_name);
} }
@@ -2425,21 +2416,12 @@ async fn routines_toggle_handler(
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?; .ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
let was_enabled = routine.enabled;
// If a specific value was provided, use it; otherwise toggle. // If a specific value was provided, use it; otherwise toggle.
routine.enabled = match body { routine.enabled = match body {
Some(Json(req)) => req.enabled.unwrap_or(!routine.enabled), Some(Json(req)) => req.enabled.unwrap_or(!routine.enabled),
None => !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 store
.update_routine(&routine) .update_routine(&routine)
.await .await
@@ -2714,7 +2696,6 @@ struct GatewayStatusResponse {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use crate::cli::oauth_defaults;
use crate::testing::credentials::TEST_GATEWAY_CRYPTO_KEY; use crate::testing::credentials::TEST_GATEWAY_CRYPTO_KEY;
#[test] #[test]
@@ -2832,11 +2813,6 @@ mod tests {
.with_state(state) .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] #[tokio::test]
async fn test_csp_header_present_on_responses() { async fn test_csp_header_present_on_responses() {
use std::net::SocketAddr; use std::net::SocketAddr;
@@ -2943,14 +2919,29 @@ mod tests {
use tower::ServiceExt; use tower::ServiceExt;
// Build an ExtensionManager so the handler can look up flows // Build an ExtensionManager so the handler can look up flows
let secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync> = let secrets = Arc::new(crate::secrets::InMemorySecretsStore::new(Arc::new(
Arc::new(crate::secrets::InMemorySecretsStore::new(Arc::new( crate::secrets::SecretsCrypto::new(secrecy::SecretString::from(
crate::secrets::SecretsCrypto::new(secrecy::SecretString::from( TEST_GATEWAY_CRYPTO_KEY.to_string(),
TEST_GATEWAY_CRYPTO_KEY.to_string(), ))
)) .expect("crypto"),
.expect("crypto"), )));
))); let tool_registry = Arc::new(ToolRegistry::new());
let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(secrets); 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 state = test_gateway_state(Some(ext_mgr));
let app = test_oauth_router(state); let app = test_oauth_router(state);
@@ -2984,13 +2975,25 @@ mod tests {
)) ))
.expect("crypto"), .expect("crypto"),
))); )));
let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(secrets.clone()); let tool_registry = Arc::new(ToolRegistry::new());
let Some(created_at) = expired_flow_created_at() else { let mcp_sm = Arc::new(crate::tools::mcp::session::McpSessionManager::new());
eprintln!("Skipping expired OAuth flow test: monotonic uptime below expiry window");
return;
};
// 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 { let flow = crate::cli::oauth_defaults::PendingOAuthFlow {
extension_name: "test_tool".to_string(), extension_name: "test_tool".to_string(),
display_name: "Test Tool".to_string(), display_name: "Test Tool".to_string(),
@@ -3010,7 +3013,9 @@ mod tests {
gateway_token: None, gateway_token: None,
resource: None, resource: None,
client_id_secret_name: 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 ext_mgr
@@ -3040,80 +3045,6 @@ mod tests {
assert!(html.contains("Authorization Failed")); 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] #[tokio::test]
async fn test_oauth_callback_no_extension_manager() { async fn test_oauth_callback_no_extension_manager() {
use axum::body::Body; use axum::body::Body;
@@ -3152,16 +3083,28 @@ mod tests {
)) ))
.expect("crypto"), .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). // 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 // 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 // token exchange — we only need to verify that the instance prefix was
// stripped and the flow was found by the raw nonce. // 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 { let flow = crate::cli::oauth_defaults::PendingOAuthFlow {
extension_name: "test_tool".to_string(), extension_name: "test_tool".to_string(),
display_name: "Test Tool".to_string(), display_name: "Test Tool".to_string(),
@@ -3182,7 +3125,9 @@ mod tests {
resource: None, resource: None,
client_id_secret_name: None, client_id_secret_name: None,
// Expired — handler will reject after lookup (no network I/O) // Expired — handler will reject after lookup (no network I/O)
created_at, 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 ext_mgr
@@ -3253,27 +3198,24 @@ mod tests {
fn test_ext_mgr( fn test_ext_mgr(
secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync>, secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync>,
) -> (Arc<ExtensionManager>, tempfile::TempDir, tempfile::TempDir) { ) -> Arc<ExtensionManager> {
let tool_registry = Arc::new(ToolRegistry::new()); let tool_registry = Arc::new(ToolRegistry::new());
let mcp_sm = Arc::new(crate::tools::mcp::session::McpSessionManager::new()); let mcp_sm = Arc::new(crate::tools::mcp::session::McpSessionManager::new());
let mcp_pm = Arc::new(crate::tools::mcp::process::McpProcessManager::new()); let mcp_pm = Arc::new(crate::tools::mcp::process::McpProcessManager::new());
let wasm_tools_dir = tempfile::tempdir().expect("temp wasm tools dir"); Arc::new(ExtensionManager::new(
let wasm_channels_dir = tempfile::tempdir().expect("temp wasm channels dir");
let ext_mgr = Arc::new(ExtensionManager::new(
mcp_sm, mcp_sm,
mcp_pm, mcp_pm,
secrets, secrets,
tool_registry, tool_registry,
None, None,
None, None,
wasm_tools_dir.path().to_path_buf(), std::path::PathBuf::from("/tmp/wasm_tools"),
wasm_channels_dir.path().to_path_buf(), std::path::PathBuf::from("/tmp/wasm_channels"),
None, None,
"test".to_string(), "test".to_string(),
None, None,
vec![], vec![],
)); ))
(ext_mgr, wasm_tools_dir, wasm_channels_dir)
} }
#[tokio::test] #[tokio::test]
@@ -3282,7 +3224,7 @@ mod tests {
use tower::ServiceExt; use tower::ServiceExt;
let secrets = test_secrets_store(); 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 state = test_gateway_state(Some(ext_mgr));
let app = test_relay_oauth_router(state); let app = test_relay_oauth_router(state);
@@ -3326,7 +3268,7 @@ mod tests {
.await .await
.expect("store nonce"); .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 state = test_gateway_state(Some(ext_mgr));
let app = test_relay_oauth_router(state); let app = test_relay_oauth_router(state);
@@ -3371,7 +3313,7 @@ mod tests {
.await .await
.expect("store nonce"); .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 state = test_gateway_state(Some(ext_mgr));
let app = test_relay_oauth_router(state); let app = test_relay_oauth_router(state);
-1
View File
@@ -143,7 +143,6 @@ impl SseManager {
SseEvent::JobResult { .. } => "job_result", SseEvent::JobResult { .. } => "job_result",
SseEvent::Heartbeat => "heartbeat", SseEvent::Heartbeat => "heartbeat",
SseEvent::ImageGenerated { .. } => "image_generated", SseEvent::ImageGenerated { .. } => "image_generated",
SseEvent::Suggestions { .. } => "suggestions",
SseEvent::ExtensionStatus { .. } => "extension_status", SseEvent::ExtensionStatus { .. } => "extension_status",
}; };
Ok(Event::default().event(event_type).data(data)) Ok(Event::default().event(event_type).data(data))
+36 -218
View File
@@ -19,7 +19,6 @@ let _loadThreadsTimer = null;
const JOB_EVENTS_CAP = 500; const JOB_EVENTS_CAP = 500;
const MEMORY_SEARCH_QUERY_MAX_LENGTH = 100; const MEMORY_SEARCH_QUERY_MAX_LENGTH = 100;
let stagedImages = []; let stagedImages = [];
let _ghostSuggestion = '';
// --- Slash Commands --- // --- Slash Commands ---
@@ -287,18 +286,9 @@ function connectSSE() {
if (data.thread_id) debouncedLoadThreads(); if (data.thread_id) debouncedLoadThreads();
return; return;
} }
clearSuggestionChips();
showActivityThinking(data.message); showActivityThinking(data.message);
}); });
eventSource.addEventListener('suggestions', (e) => {
const data = JSON.parse(e.data);
if (!isCurrentThread(data.thread_id)) return;
if (data.suggestions && data.suggestions.length > 0) {
showSuggestionChips(data.suggestions);
}
});
eventSource.addEventListener('tool_started', (e) => { eventSource.addEventListener('tool_started', (e) => {
const data = JSON.parse(e.data); const data = JSON.parse(e.data);
if (!isCurrentThread(data.thread_id)) return; if (!isCurrentThread(data.thread_id)) return;
@@ -352,27 +342,31 @@ function connectSSE() {
eventSource.addEventListener('approval_needed', (e) => { eventSource.addEventListener('approval_needed', (e) => {
const data = JSON.parse(e.data); const data = JSON.parse(e.data);
const hasThread = !!data.thread_id; if (!isCurrentThread(data.thread_id)) return;
const forCurrentThread = !hasThread || isCurrentThread(data.thread_id); showApproval(data);
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();
}); });
eventSource.addEventListener('auth_required', (e) => { 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) => { 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) => { eventSource.addEventListener('extension_status', (e) => {
@@ -433,59 +427,9 @@ function isCurrentThread(threadId) {
return threadId === currentThreadId; return threadId === currentThreadId;
} }
// --- Suggestion Chips ---
function showSuggestionChips(suggestions) {
// Clear previous chips/ghost without restoring placeholder (we'll set it below)
_ghostSuggestion = '';
const container = document.getElementById('suggestion-chips');
container.innerHTML = '';
const ghost = document.getElementById('ghost-text');
ghost.style.display = 'none';
const wrapper = document.querySelector('.chat-input-wrapper');
if (wrapper) wrapper.classList.remove('has-ghost');
_ghostSuggestion = suggestions[0] || '';
const input = document.getElementById('chat-input');
suggestions.forEach(text => {
const chip = document.createElement('button');
chip.className = 'suggestion-chip';
chip.textContent = text;
chip.addEventListener('click', () => {
input.value = text;
clearSuggestionChips();
autoResizeTextarea(input);
input.focus();
sendMessage();
});
container.appendChild(chip);
});
container.style.display = 'flex';
// Show first suggestion as ghost text in the input so user knows Tab works
if (_ghostSuggestion && input.value === '') {
ghost.textContent = _ghostSuggestion;
ghost.style.display = 'block';
input.closest('.chat-input-wrapper').classList.add('has-ghost');
}
}
function clearSuggestionChips() {
_ghostSuggestion = '';
const container = document.getElementById('suggestion-chips');
if (container) {
container.innerHTML = '';
container.style.display = 'none';
}
const ghost = document.getElementById('ghost-text');
if (ghost) ghost.style.display = 'none';
const wrapper = document.querySelector('.chat-input-wrapper');
if (wrapper) wrapper.classList.remove('has-ghost');
}
// --- Chat --- // --- Chat ---
function sendMessage() { function sendMessage() {
clearSuggestionChips();
const input = document.getElementById('chat-input'); const input = document.getElementById('chat-input');
if (!currentThreadId) { if (!currentThreadId) {
console.warn('sendMessage: no thread selected, ignoring'); console.warn('sendMessage: no thread selected, ignoring');
@@ -1046,26 +990,7 @@ function finalizeActivityGroup() {
_activeToolCards = {}; _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) { 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 container = document.getElementById('chat-messages');
const card = document.createElement('div'); const card = document.createElement('div');
card.className = 'approval-card'; card.className = 'approval-card';
@@ -1078,7 +1003,7 @@ function showApproval(data) {
const toolName = document.createElement('div'); const toolName = document.createElement('div');
toolName.className = 'approval-tool-name'; toolName.className = 'approval-tool-name';
toolName.textContent = humanizeToolName(data.tool_name); toolName.textContent = data.tool_name;
card.appendChild(toolName); card.appendChild(toolName);
if (data.description) { if (data.description) {
@@ -1181,71 +1106,13 @@ function showJobCard(data) {
// --- Auth card --- // --- 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) { function showAuthCard(data) {
// Keep a single global auth prompt so the experience is consistent across tabs. // Remove any existing card for this extension first
const existing = getAuthOverlay(); removeAuthCard(data.extension_name);
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);
});
const container = document.getElementById('chat-messages');
const card = document.createElement('div'); const card = document.createElement('div');
card.className = 'auth-card auth-modal'; card.className = 'auth-card';
card.setAttribute('data-extension-name', data.extension_name); card.setAttribute('data-extension-name', data.extension_name);
const header = document.createElement('div'); const header = document.createElement('div');
@@ -1324,30 +1191,21 @@ function showAuthCard(data) {
actions.appendChild(cancelBtn); actions.appendChild(cancelBtn);
card.appendChild(actions); card.appendChild(actions);
overlay.appendChild(card); container.appendChild(card);
document.body.appendChild(overlay); container.scrollTop = container.scrollHeight;
tokenInput.focus(); tokenInput.focus();
} }
function removeAuthCard(extensionName) { function removeAuthCard(extensionName) {
const overlay = getAuthOverlay(extensionName); const card = document.querySelector('.auth-card[data-extension-name="' + extensionName + '"]');
if (overlay) { if (card) card.remove();
overlay.remove();
return;
}
const card = getAuthCard(extensionName);
if (card) {
const parentOverlay = card.closest('.auth-overlay');
if (parentOverlay) parentOverlay.remove();
else card.remove();
}
} }
function submitAuthToken(extensionName, tokenValue) { function submitAuthToken(extensionName, tokenValue) {
if (!tokenValue || !tokenValue.trim()) return; if (!tokenValue || !tokenValue.trim()) return;
// Disable submit button while in flight // Disable submit button while in flight
const card = getAuthCard(extensionName); const card = document.querySelector('.auth-card[data-extension-name="' + extensionName + '"]');
if (card) { if (card) {
const btns = card.querySelectorAll('button'); const btns = card.querySelectorAll('button');
btns.forEach((b) => { b.disabled = true; }); btns.forEach((b) => { b.disabled = true; });
@@ -1358,10 +1216,8 @@ function submitAuthToken(extensionName, tokenValue) {
body: { extension_name: extensionName, token: tokenValue.trim() }, body: { extension_name: extensionName, token: tokenValue.trim() },
}).then((result) => { }).then((result) => {
if (result.success) { if (result.success) {
// Close immediately for responsiveness; the authoritative success UX
// (toast + extensions refresh) still comes from auth_completed SSE.
removeAuthCard(extensionName); removeAuthCard(extensionName);
enableChatInput(); addMessage('system', result.message);
} else { } else {
showAuthCardError(extensionName, result.message); showAuthCardError(extensionName, result.message);
} }
@@ -1380,7 +1236,7 @@ function cancelAuth(extensionName) {
} }
function showAuthCardError(extensionName, message) { function showAuthCardError(extensionName, message) {
const card = getAuthCard(extensionName); const card = document.querySelector('.auth-card[data-extension-name="' + extensionName + '"]');
if (!card) return; if (!card) return;
// Re-enable buttons // Re-enable buttons
const btns = card.querySelectorAll('button'); const btns = card.querySelectorAll('button');
@@ -1394,7 +1250,6 @@ function showAuthCardError(extensionName, message) {
} }
function loadHistory(before) { function loadHistory(before) {
clearSuggestionChips();
let historyUrl = '/api/chat/history?limit=50'; let historyUrl = '/api/chat/history?limit=50';
if (currentThreadId) { if (currentThreadId) {
historyUrl += '&thread_id=' + encodeURIComponent(currentThreadId); historyUrl += '&thread_id=' + encodeURIComponent(currentThreadId);
@@ -1690,7 +1545,6 @@ function switchToAssistant() {
} }
function switchThread(threadId) { function switchThread(threadId) {
clearSuggestionChips();
finalizeActivityGroup(); finalizeActivityGroup();
currentThreadId = threadId; currentThreadId = threadId;
unreadThreads.delete(threadId); unreadThreads.delete(threadId);
@@ -1723,15 +1577,6 @@ chatInput.addEventListener('keydown', (e) => {
const acEl = document.getElementById('slash-autocomplete'); const acEl = document.getElementById('slash-autocomplete');
const acVisible = acEl && acEl.style.display !== 'none'; const acVisible = acEl && acEl.style.display !== 'none';
// Accept first suggestion with Tab (plain Tab only, not Shift+Tab)
if (e.key === 'Tab' && !e.shiftKey && !acVisible && _ghostSuggestion && chatInput.value === '') {
e.preventDefault();
chatInput.value = _ghostSuggestion;
clearSuggestionChips();
autoResizeTextarea(chatInput);
return;
}
if (acVisible) { if (acVisible) {
const items = acEl.querySelectorAll('.slash-ac-item'); const items = acEl.querySelectorAll('.slash-ac-item');
if (e.key === 'ArrowDown') { if (e.key === 'ArrowDown') {
@@ -1768,16 +1613,6 @@ chatInput.addEventListener('keydown', (e) => {
chatInput.addEventListener('input', () => { chatInput.addEventListener('input', () => {
autoResizeTextarea(chatInput); autoResizeTextarea(chatInput);
filterSlashCommands(chatInput.value); filterSlashCommands(chatInput.value);
const ghost = document.getElementById('ghost-text');
const wrapper = chatInput.closest('.chat-input-wrapper');
if (chatInput.value !== '') {
ghost.style.display = 'none';
wrapper.classList.remove('has-ghost');
} else if (_ghostSuggestion) {
ghost.textContent = _ghostSuggestion;
ghost.style.display = 'block';
wrapper.classList.add('has-ghost');
}
}); });
chatInput.addEventListener('blur', () => { chatInput.addEventListener('blur', () => {
// Small delay so mousedown on autocomplete item fires first // Small delay so mousedown on autocomplete item fires first
@@ -2331,10 +2166,6 @@ function renderAvailableExtensionCard(entry) {
showToast(I18n.t('extensions.installedSuccess', {name: entry.display_name}), 'success'); showToast(I18n.t('extensions.installedSuccess', {name: entry.display_name}), 'success');
// OAuth popup if auth started during install (builtin creds) // OAuth popup if auth started during install (builtin creds)
if (res.auth_url) { if (res.auth_url) {
showAuthCard({
extension_name: entry.name,
auth_url: res.auth_url,
});
showToast('Opening authentication for ' + entry.display_name, 'info'); showToast('Opening authentication for ' + entry.display_name, 'info');
openOAuthUrl(res.auth_url); openOAuthUrl(res.auth_url);
} }
@@ -2600,10 +2431,6 @@ function activateExtension(name) {
if (res.success) { if (res.success) {
// Even on success, the tool may need OAuth (e.g., WASM loaded but no token yet) // Even on success, the tool may need OAuth (e.g., WASM loaded but no token yet)
if (res.auth_url) { if (res.auth_url) {
showAuthCard({
extension_name: name,
auth_url: res.auth_url,
});
showToast('Opening authentication for ' + name, 'info'); showToast('Opening authentication for ' + name, 'info');
openOAuthUrl(res.auth_url); openOAuthUrl(res.auth_url);
} }
@@ -2612,10 +2439,6 @@ function activateExtension(name) {
} }
if (res.auth_url) { if (res.auth_url) {
showAuthCard({
extension_name: name,
auth_url: res.auth_url,
});
showToast('Opening authentication for ' + name, 'info'); showToast('Opening authentication for ' + name, 'info');
openOAuthUrl(res.auth_url); openOAuthUrl(res.auth_url);
} else if (res.awaiting_token) { } else if (res.awaiting_token) {
@@ -2658,7 +2481,6 @@ function renderConfigureModal(name, secrets) {
closeConfigureModal(); closeConfigureModal();
const overlay = document.createElement('div'); const overlay = document.createElement('div');
overlay.className = 'configure-overlay'; overlay.className = 'configure-overlay';
overlay.setAttribute('data-extension-name', name);
overlay.addEventListener('click', (e) => { overlay.addEventListener('click', (e) => {
if (e.target === overlay) closeConfigureModal(); if (e.target === overlay) closeConfigureModal();
}); });
@@ -2752,8 +2574,7 @@ function submitConfigureModal(name, fields) {
} }
// Disable buttons to prevent double-submit // Disable buttons to prevent double-submit
const overlay = getConfigureOverlay(name) || document.querySelector('.configure-overlay'); var btns = document.querySelectorAll('.configure-actions button');
var btns = overlay ? overlay.querySelectorAll('.configure-actions button') : [];
btns.forEach(function(b) { b.disabled = true; }); btns.forEach(function(b) { b.disabled = true; });
apiFetch('/api/extensions/' + encodeURIComponent(name) + '/setup', { apiFetch('/api/extensions/' + encodeURIComponent(name) + '/setup', {
@@ -2764,10 +2585,8 @@ function submitConfigureModal(name, fields) {
if (res.success) { if (res.success) {
closeConfigureModal(); closeConfigureModal();
if (res.auth_url) { if (res.auth_url) {
showAuthCard({ // OAuth flow started — open consent popup. The auth_completed SSE will
extension_name: name, // not arrive immediately (it fires after OAuth callback), so show a toast now.
auth_url: res.auth_url,
});
showToast('Opening OAuth authorization for ' + name, 'info'); showToast('Opening OAuth authorization for ' + name, 'info');
openOAuthUrl(res.auth_url); openOAuthUrl(res.auth_url);
loadExtensions(); loadExtensions();
@@ -2786,9 +2605,8 @@ function submitConfigureModal(name, fields) {
}); });
} }
function closeConfigureModal(extensionName) { function closeConfigureModal() {
if (typeof extensionName !== 'string') extensionName = null; const existing = document.querySelector('.configure-overlay');
const existing = getConfigureOverlay(extensionName);
if (existing) existing.remove(); if (existing) existing.remove();
} }
+1 -5
View File
@@ -155,13 +155,9 @@
<div class="chat-container"> <div class="chat-container">
<div class="chat-messages" id="chat-messages"></div> <div class="chat-messages" id="chat-messages"></div>
<div id="slash-autocomplete" class="slash-autocomplete" style="display:none"></div> <div id="slash-autocomplete" class="slash-autocomplete" style="display:none"></div>
<div id="suggestion-chips" class="suggestion-chips" style="display:none"></div>
<div class="chat-input"> <div class="chat-input">
<div id="image-preview-strip" class="image-preview-strip"></div> <div id="image-preview-strip" class="image-preview-strip"></div>
<div class="chat-input-wrapper"> <textarea id="chat-input" data-i18n="chat.inputPlaceholder" data-i18n-attr="placeholder" placeholder="Message or / for commands..." rows="1"></textarea>
<textarea id="chat-input" data-i18n="chat.inputPlaceholder" data-i18n-attr="placeholder" placeholder="Message or / for commands..." rows="1"></textarea>
<div id="ghost-text" class="ghost-text"></div>
</div>
<input type="file" id="image-file-input" accept="image/*" multiple style="display:none"> <input type="file" id="image-file-input" accept="image/*" multiple style="display:none">
<button id="attach-btn" class="attach-btn" data-i18n="chat.attachImages" data-i18n-attr="title" title="Attach images" <button id="attach-btn" class="attach-btn" data-i18n="chat.attachImages" data-i18n-attr="title" title="Attach images"
aria-label="Attach images">&#x1F4CE;</button> aria-label="Attach images">&#x1F4CE;</button>
+6 -85
View File
@@ -1219,21 +1219,7 @@ body {
color: var(--danger); color: var(--danger);
} }
/* Auth prompt */ /* Auth card (inline in chat) */
.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 { .auth-card {
align-self: flex-start; align-self: flex-start;
max-width: 80%; max-width: 80%;
@@ -1248,16 +1234,6 @@ body {
transition: border-color 0.2s; 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 { .auth-card .auth-header {
font-weight: 600; font-weight: 600;
color: var(--accent); color: var(--accent);
@@ -1362,14 +1338,8 @@ body {
min-height: 56px; min-height: 56px;
} }
.chat-input-wrapper { .chat-input textarea {
position: relative;
flex: 1; flex: 1;
display: flex;
}
.chat-input-wrapper textarea {
width: 100%;
padding: 8px 12px; padding: 8px 12px;
background: var(--bg); background: var(--bg);
border: 1px solid var(--border); border: 1px solid var(--border);
@@ -1382,66 +1352,17 @@ body {
max-height: 120px; max-height: 120px;
} }
.ghost-text { .chat-input textarea:focus {
position: absolute;
top: 0;
left: 0;
right: 0;
padding: 8px 12px;
font-size: 14px;
font-family: inherit;
color: var(--text-secondary);
opacity: 0.5;
pointer-events: none;
white-space: pre-wrap;
overflow: hidden;
display: none;
z-index: 1;
}
/* Hide native placeholder when ghost text is visible */
.chat-input-wrapper.has-ghost textarea::placeholder {
color: transparent;
}
.chat-input-wrapper textarea:focus {
outline: none; outline: none;
border-color: var(--accent); border-color: var(--accent);
box-shadow: 0 0 0 3px rgba(52, 211, 153, 0.1); box-shadow: 0 0 0 3px rgba(52, 211, 153, 0.1);
} }
.chat-input-wrapper textarea:disabled { .chat-input textarea:disabled {
opacity: 0.5; opacity: 0.5;
cursor: not-allowed; cursor: not-allowed;
} }
.suggestion-chips {
display: none;
flex-wrap: wrap;
gap: 8px;
padding: 8px 16px;
border-top: 1px solid var(--border);
}
.suggestion-chip {
padding: 6px 14px;
background: var(--bg-secondary);
border: 1px solid var(--border);
border-radius: 16px;
color: var(--text-secondary);
font-size: 13px;
font-family: inherit;
cursor: pointer;
transition: all 0.15s ease;
white-space: nowrap;
}
.suggestion-chip:hover {
background: var(--accent);
color: #09090b;
border-color: var(--accent);
}
.chat-input button { .chat-input button {
padding: 8px 20px; padding: 8px 20px;
background: var(--accent); background: var(--accent);
@@ -1471,7 +1392,7 @@ body {
} }
/* Keyboard accessibility focus rings */ /* Keyboard accessibility focus rings */
.chat-input-wrapper textarea:focus-visible, .chat-input textarea:focus-visible,
.chat-input button:focus-visible, .chat-input button:focus-visible,
.tab-bar button:focus-visible, .tab-bar button:focus-visible,
.tree-row:focus-visible { .tree-row:focus-visible {
@@ -3879,7 +3800,7 @@ mark {
min-height: 52px; min-height: 52px;
} }
.chat-input-wrapper textarea { .chat-input textarea {
min-height: 36px; min-height: 36px;
max-height: 100px; max-height: 100px;
} }
-9
View File
@@ -242,14 +242,6 @@ pub enum SseEvent {
thread_id: Option<String>, thread_id: Option<String>,
}, },
/// Suggested follow-up messages for the user.
#[serde(rename = "suggestions")]
Suggestions {
suggestions: Vec<String>,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
/// Extension activation status change (WASM channels). /// Extension activation status change (WASM channels).
#[serde(rename = "extension_status")] #[serde(rename = "extension_status")]
ExtensionStatus { ExtensionStatus {
@@ -715,7 +707,6 @@ impl WsServerMessage {
SseEvent::JobStatus { .. } => "job_status", SseEvent::JobStatus { .. } => "job_status",
SseEvent::JobResult { .. } => "job_result", SseEvent::JobResult { .. } => "job_result",
SseEvent::ImageGenerated { .. } => "image_generated", SseEvent::ImageGenerated { .. } => "image_generated",
SseEvent::Suggestions { .. } => "suggestions",
SseEvent::ExtensionStatus { .. } => "extension_status", SseEvent::ExtensionStatus { .. } => "extension_status",
}; };
let data = serde_json::to_value(event).unwrap_or(serde_json::Value::Null); let data = serde_json::to_value(event).unwrap_or(serde_json::Value::Null);
+6 -18
View File
@@ -127,11 +127,7 @@ fn cmd_list(
.unwrap_or("none"); .unwrap_or("none");
println!( println!(
"{:<20} {:<8} {:<8} {:<10} {}", "{:<20} {:<8} {:<8} {:<10} {}",
m.name, m.name, m.kind, m.version, auth, m.description
m.kind,
m.version.as_deref().unwrap_or("-"),
auth,
m.description
); );
} else { } else {
println!("{:<20} {:<8} {}", m.name, m.kind, m.description); println!("{:<20} {:<8} {}", m.name, m.kind, m.description);
@@ -177,25 +173,17 @@ fn cmd_info(catalog: &RegistryCatalog, name: &str) -> anyhow::Result<()> {
.map_err(|e| anyhow::anyhow!("{}", e))?; .map_err(|e| anyhow::anyhow!("{}", e))?;
println!("{} ({})", manifest.display_name, manifest.kind); println!("{} ({})", manifest.display_name, manifest.kind);
if let Some(ref version) = manifest.version { println!(" Version: {}", manifest.version);
println!(" Version: {}", version);
}
println!(" {}", manifest.description); println!(" {}", manifest.description);
if !manifest.keywords.is_empty() { if !manifest.keywords.is_empty() {
println!(" Keywords: {}", manifest.keywords.join(", ")); println!(" Keywords: {}", manifest.keywords.join(", "));
} }
if let Some(ref source) = manifest.source { println!("\nSource:");
println!("\nSource:"); println!(" Directory: {}", manifest.source.dir);
println!(" Directory: {}", source.dir); println!(" Crate: {}", manifest.source.crate_name);
println!(" Crate: {}", source.crate_name); println!(" Capabilities: {}", manifest.source.capabilities);
println!(" Capabilities: {}", source.capabilities);
}
if let Some(ref url) = manifest.url {
println!("\nMCP Server URL: {}", url);
}
if let Some(artifact) = manifest.artifacts.get("wasm32-wasip2") { if let Some(artifact) = manifest.artifacts.get("wasm32-wasip2") {
println!("\nArtifact (wasm32-wasip2):"); println!("\nArtifact (wasm32-wasip2):");
+7 -63
View File
@@ -23,9 +23,6 @@ pub struct EmbeddingsConfig {
pub ollama_base_url: String, pub ollama_base_url: String,
/// Embedding vector dimension. Inferred from the model name when not set explicitly. /// Embedding vector dimension. Inferred from the model name when not set explicitly.
pub dimension: usize, 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 { impl Default for EmbeddingsConfig {
@@ -39,7 +36,6 @@ impl Default for EmbeddingsConfig {
model, model,
ollama_base_url: "http://localhost:11434".to_string(), ollama_base_url: "http://localhost:11434".to_string(),
dimension, dimension,
openai_base_url: None,
} }
} }
} }
@@ -78,8 +74,6 @@ impl EmbeddingsConfig {
let enabled = parse_bool_env("EMBEDDING_ENABLED", settings.embeddings.enabled)?; let enabled = parse_bool_env("EMBEDDING_ENABLED", settings.embeddings.enabled)?;
let openai_base_url = optional_env("EMBEDDING_BASE_URL")?;
Ok(Self { Ok(Self {
enabled, enabled,
provider, provider,
@@ -87,7 +81,6 @@ impl EmbeddingsConfig {
model, model,
ollama_base_url, ollama_base_url,
dimension, dimension,
openai_base_url,
}) })
} }
@@ -137,27 +130,16 @@ impl EmbeddingsConfig {
} }
_ => { _ => {
if let Some(api_key) = self.openai_api_key() { 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, api_key,
&self.model, &self.model,
self.dimension, 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 { } else {
tracing::warn!("Embeddings configured but OPENAI_API_KEY not set"); tracing::warn!("Embeddings configured but OPENAI_API_KEY not set");
None None
@@ -182,7 +164,6 @@ mod tests {
std::env::remove_var("EMBEDDING_PROVIDER"); std::env::remove_var("EMBEDDING_PROVIDER");
std::env::remove_var("EMBEDDING_MODEL"); std::env::remove_var("EMBEDDING_MODEL");
std::env::remove_var("OPENAI_API_KEY"); 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"); 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"
);
}
} }
+26 -14
View File
@@ -16,7 +16,6 @@ mod workspace;
use std::path::Path; use std::path::Path;
use std::sync::Arc; use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use async_trait::async_trait; use async_trait::async_trait;
use chrono::{DateTime, NaiveDateTime, Utc}; use chrono::{DateTime, NaiveDateTime, Utc};
@@ -33,8 +32,6 @@ use crate::workspace::MemoryDocument;
use crate::db::libsql_migrations; 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`). /// Explicit column list for routines table (matches positional access in `row_to_routine_libsql`).
pub(crate) const ROUTINE_COLUMNS: &str = "\ pub(crate) const ROUTINE_COLUMNS: &str = "\
id, name, description, user_id, enabled, \ id, name, description, user_id, enabled, \
@@ -166,27 +163,24 @@ impl LibSqlBackend {
/// ///
/// Returns an error if none of the formats match. /// Returns an error if none of the formats match.
pub(crate) fn parse_timestamp(s: &str) -> Result<DateTime<Utc>, String> { 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) // RFC 3339 (our canonical write format)
if let Ok(dt) = DateTime::parse_from_rfc3339(s) { if let Ok(dt) = DateTime::parse_from_rfc3339(s) {
return Ok(dt.with_timezone(&Utc)); return Ok(dt.with_timezone(&Utc));
} }
// Naive with fractional seconds (legacy or SQLite datetime() output) // Naive with fractional seconds (legacy or SQLite datetime() output)
if let Ok(ndt) = NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S%.f") { if let Ok(ndt) = NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S%.f") {
log_naive_timestamp_once(); tracing::debug!(
timestamp = %s,
"parsed naive timestamp without timezone; assuming UTC for backward compatibility"
);
return Ok(ndt.and_utc()); return Ok(ndt.and_utc());
} }
// Naive without fractional seconds (legacy format) // Naive without fractional seconds (legacy format)
if let Ok(ndt) = NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S") { if let Ok(ndt) = NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S") {
log_naive_timestamp_once(); tracing::debug!(
timestamp = %s,
"parsed naive timestamp without timezone; assuming UTC for backward compatibility"
);
return Ok(ndt.and_utc()); return Ok(ndt.and_utc());
} }
Err(format!("unparseable timestamp: {:?}", s)) Err(format!("unparseable timestamp: {:?}", s))
@@ -332,6 +326,24 @@ impl Database for LibSqlBackend {
libsql_migrations::run_incremental(&conn).await?; libsql_migrations::run_incremental(&conn).await?;
Ok(()) Ok(())
} }
async fn shutdown(&self) -> Result<(), DatabaseError> {
match self.db.flush_replicator().await {
Ok(Some(frame_no)) => {
tracing::debug!("libSQL replicator flushed at frame {}", frame_no);
Ok(())
}
Ok(None) => {
tracing::debug!("No libSQL replicator to flush, skipping shutdown sync");
Ok(())
}
Err(libsql::Error::SyncNotSupported(_)) => {
tracing::debug!("libSQL sync not supported, skipping flush on shutdown");
Ok(())
}
Err(error) => Err(DatabaseError::from(error)),
}
}
} }
// ==================== Row conversion helpers ==================== // ==================== Row conversion helpers ====================
+7
View File
@@ -523,6 +523,13 @@ pub trait Database:
{ {
/// Run schema migrations for this backend. /// Run schema migrations for this backend.
async fn run_migrations(&self) -> Result<(), DatabaseError>; async fn run_migrations(&self) -> Result<(), DatabaseError>;
/// Shutdown hook for backend-specific drain/flush behavior.
///
/// Default implementation is a no-op so existing backends remain compatible.
async fn shutdown(&self) -> Result<(), DatabaseError> {
Ok(())
}
} }
#[cfg(test)] #[cfg(test)]
+5
View File
@@ -61,6 +61,11 @@ impl Database for PgBackend {
async fn run_migrations(&self) -> Result<(), DatabaseError> { async fn run_migrations(&self) -> Result<(), DatabaseError> {
self.store.run_migrations().await self.store.run_migrations().await
} }
async fn shutdown(&self) -> Result<(), DatabaseError> {
self.store.pool().close();
Ok(())
}
} }
// ==================== ConversationStore ==================== // ==================== ConversationStore ====================
+6 -189
View File
@@ -786,19 +786,6 @@ impl ExtensionManager {
Self::validate_extension_name(name)?; Self::validate_extension_name(name)?;
let kind = self.determine_installed_kind(name).await?; 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 { match kind {
ExtensionKind::McpServer => { ExtensionKind::McpServer => {
// Unregister tools with this server's prefix // Unregister tools with this server's prefix
@@ -832,14 +819,6 @@ impl ExtensionManager {
// Unregister from tool registry // Unregister from tool registry
self.tool_registry.unregister(name).await; 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 // Revoke credential mappings from the shared registry
let cap_path = self let cap_path = self
.wasm_tools_dir .wasm_tools_dir
@@ -880,9 +859,6 @@ impl ExtensionManager {
self.active_channel_names.write().await.remove(name); self.active_channel_names.write().await.remove(name);
self.persist_active_channels().await; self.persist_active_channels().await;
// Clear stale activation errors so reinstall starts clean
self.activation_errors.write().await.remove(name);
// Delete channel files // Delete channel files
let wasm_path = self.wasm_channels_dir.join(format!("{}.wasm", name)); let wasm_path = self.wasm_channels_dir.join(format!("{}.wasm", name));
let cap_path = self let cap_path = self
@@ -2884,17 +2860,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(|| { let runtime = self.wasm_tool_runtime.as_ref().ok_or_else(|| {
ExtensionError::ActivationFailed("WASM runtime not available".to_string()) ExtensionError::ActivationFailed("WASM runtime not available".to_string())
})?; })?;
@@ -4530,18 +4495,14 @@ mod tests {
// available" because the ExtensionManager had `wasm_tool_runtime: None`. // available" because the ExtensionManager had `wasm_tool_runtime: None`.
/// Build a minimal ExtensionManager suitable for unit tests. /// 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>>, wasm_runtime: Option<Arc<crate::tools::wasm::WasmToolRuntime>>,
tools_dir: std::path::PathBuf, tools_dir: std::path::PathBuf,
channels_dir: std::path::PathBuf,
) -> crate::extensions::manager::ExtensionManager { ) -> crate::extensions::manager::ExtensionManager {
use crate::secrets::{InMemorySecretsStore, SecretsCrypto}; use crate::secrets::{InMemorySecretsStore, SecretsCrypto};
use crate::tools::mcp::process::McpProcessManager; use crate::tools::mcp::process::McpProcessManager;
use crate::tools::mcp::session::McpSessionManager; 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 key = secrecy::SecretString::from(crate::secrets::keychain::generate_master_key_hex());
let crypto = Arc::new(SecretsCrypto::new(key).expect("crypto")); let crypto = Arc::new(SecretsCrypto::new(key).expect("crypto"));
let secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync> = let secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync> =
@@ -4556,22 +4517,15 @@ mod tests {
tools, tools,
None, // hooks None, // hooks
wasm_runtime, wasm_runtime,
tools_dir, tools_dir.clone(),
channels_dir, tools_dir, // channels dir (unused here)
None, // tunnel_url None, // tunnel_url
"test".to_string(), "test".to_string(),
None, // db None, // db
vec![], 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] #[tokio::test]
async fn test_activate_wasm_tool_with_runtime_passes_runtime_check() { async fn test_activate_wasm_tool_with_runtime_passes_runtime_check() {
// When the ExtensionManager has a WASM runtime, activation should get // When the ExtensionManager has a WASM runtime, activation should get
@@ -4924,145 +4878,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] #[test]
fn test_sanitize_url_with_query_params() { fn test_sanitize_url_with_query_params() {
let url = "https://api.example.com/path?api_key=secret123&token=abc"; let url = "https://api.example.com/path?api_key=secret123&token=abc";
@@ -5338,6 +5153,7 @@ mod tests {
Some("https://my-gateway.example.com/oauth/callback".to_string()), Some("https://my-gateway.example.com/oauth/callback".to_string()),
); );
} }
// ── Regression tests for PR #677 (unify-extension-lifecycle) ───────── // ── Regression tests for PR #677 (unify-extension-lifecycle) ─────────
#[tokio::test] #[tokio::test]
@@ -5487,6 +5303,7 @@ mod tests {
"configure should have stored the relay stream token" "configure should have stored the relay stream token"
); );
} }
#[test] #[test]
fn test_validation_failed_is_distinct_error_variant() { fn test_validation_failed_is_distinct_error_variant() {
// Regression: ValidationFailed must be a distinct error variant so // Regression: ValidationFailed must be a distinct error variant so
+226 -79
View File
@@ -232,11 +232,198 @@ pub fn builtin_entries() -> Vec<RegistryEntry> {
} }
/// Well-known extensions, with an optional relay URL for the channel-relay entry. /// Well-known extensions, with an optional relay URL for the channel-relay entry.
///
/// MCP server entries are loaded from `registry/mcp-servers/*.json` via the catalog
/// system. Only runtime-dependent entries (like channel-relay) remain here.
pub fn builtin_entries_with_relay(relay_url: Option<String>) -> Vec<RegistryEntry> { pub fn builtin_entries_with_relay(relay_url: Option<String>) -> Vec<RegistryEntry> {
let mut entries = vec![]; let mut entries = vec![
// -- MCP Servers --
RegistryEntry {
name: "notion".to_string(),
display_name: "Notion".to_string(),
kind: ExtensionKind::McpServer,
description: "Connect to Notion for reading and writing pages, databases, and comments"
.to_string(),
keywords: vec![
"notes".into(),
"wiki".into(),
"docs".into(),
"pages".into(),
"database".into(),
],
source: ExtensionSource::McpUrl {
url: "https://mcp.notion.com/mcp".to_string(),
},
fallback_source: None,
auth_hint: AuthHint::Dcr,
version: None,
},
RegistryEntry {
name: "linear".to_string(),
display_name: "Linear".to_string(),
kind: ExtensionKind::McpServer,
description:
"Connect to Linear for issue tracking, project management, and team workflows"
.to_string(),
keywords: vec![
"issues".into(),
"tickets".into(),
"project".into(),
"tracking".into(),
"bugs".into(),
],
source: ExtensionSource::McpUrl {
url: "https://mcp.linear.app/sse".to_string(),
},
fallback_source: None,
auth_hint: AuthHint::Dcr,
version: None,
},
RegistryEntry {
name: "github".to_string(),
display_name: "GitHub".to_string(),
kind: ExtensionKind::McpServer,
description:
"Connect to GitHub for repository management, issues, PRs, and code search"
.to_string(),
keywords: vec![
"git".into(),
"repos".into(),
"code".into(),
"pull-request".into(),
"issues".into(),
],
source: ExtensionSource::McpUrl {
url: "https://api.githubcopilot.com/mcp/".to_string(),
},
fallback_source: None,
auth_hint: AuthHint::Dcr,
version: None,
},
RegistryEntry {
name: "slack-mcp".to_string(),
display_name: "Slack MCP".to_string(),
kind: ExtensionKind::McpServer,
description:
"Connect to Slack via MCP for messaging, channel management, and team communication"
.to_string(),
keywords: vec![
"messaging".into(),
"chat".into(),
"channels".into(),
"team".into(),
"communication".into(),
],
source: ExtensionSource::McpUrl {
url: "https://mcp.slack.com".to_string(),
},
fallback_source: None,
auth_hint: AuthHint::Dcr,
version: None,
},
RegistryEntry {
name: "sentry".to_string(),
display_name: "Sentry".to_string(),
kind: ExtensionKind::McpServer,
description:
"Connect to Sentry for error tracking, performance monitoring, and debugging"
.to_string(),
keywords: vec![
"errors".into(),
"monitoring".into(),
"debugging".into(),
"crashes".into(),
"performance".into(),
],
source: ExtensionSource::McpUrl {
url: "https://mcp.sentry.dev/mcp".to_string(),
},
fallback_source: None,
auth_hint: AuthHint::Dcr,
version: None,
},
RegistryEntry {
name: "stripe".to_string(),
display_name: "Stripe".to_string(),
kind: ExtensionKind::McpServer,
description:
"Connect to Stripe for payment processing, subscriptions, and financial data"
.to_string(),
keywords: vec![
"payments".into(),
"billing".into(),
"subscriptions".into(),
"invoices".into(),
"finance".into(),
],
source: ExtensionSource::McpUrl {
url: "https://mcp.stripe.com".to_string(),
},
fallback_source: None,
auth_hint: AuthHint::Dcr,
version: None,
},
RegistryEntry {
name: "cloudflare".to_string(),
display_name: "Cloudflare".to_string(),
kind: ExtensionKind::McpServer,
description:
"Connect to Cloudflare for DNS, Workers, KV, and infrastructure management"
.to_string(),
keywords: vec![
"cdn".into(),
"dns".into(),
"workers".into(),
"hosting".into(),
"infrastructure".into(),
],
source: ExtensionSource::McpUrl {
url: "https://mcp.cloudflare.com/mcp".to_string(),
},
fallback_source: None,
auth_hint: AuthHint::Dcr,
version: None,
},
RegistryEntry {
name: "asana".to_string(),
display_name: "Asana".to_string(),
kind: ExtensionKind::McpServer,
description: "Connect to Asana for task management, projects, and team coordination"
.to_string(),
keywords: vec![
"tasks".into(),
"projects".into(),
"management".into(),
"team".into(),
],
source: ExtensionSource::McpUrl {
url: "https://mcp.asana.com/v2/mcp".to_string(),
},
fallback_source: None,
auth_hint: AuthHint::Dcr,
version: None,
},
RegistryEntry {
name: "intercom".to_string(),
display_name: "Intercom".to_string(),
kind: ExtensionKind::McpServer,
description: "Connect to Intercom for customer messaging, support, and engagement"
.to_string(),
keywords: vec![
"support".into(),
"customers".into(),
"messaging".into(),
"chat".into(),
"helpdesk".into(),
],
source: ExtensionSource::McpUrl {
url: "https://mcp.intercom.com/mcp".to_string(),
},
fallback_source: None,
auth_hint: AuthHint::Dcr,
version: None,
},
// WASM channels (telegram, slack, discord, whatsapp) come from the embedded
// registry catalog (registry/channels/*.json) with WasmDownload URLs pointing
// to GitHub release artifacts. See new_with_catalog() for merging.
];
// Conditionally add channel-relay entries when relay URL is configured // Conditionally add channel-relay entries when relay URL is configured
if let Some(relay_url) = relay_url { if let Some(relay_url) = relay_url {
@@ -358,21 +545,9 @@ mod tests {
assert_eq!(score, 0, "No match should score 0"); assert_eq!(score, 0, "No match should score 0");
} }
/// Helper to create a registry with catalog entries (MCP servers come from catalog now).
fn registry_with_catalog() -> ExtensionRegistry {
let catalog = crate::registry::catalog::RegistryCatalog::load_or_embedded()
.expect("catalog should load");
let catalog_entries: Vec<RegistryEntry> = catalog
.all()
.iter()
.filter_map(|m| m.to_registry_entry())
.collect();
ExtensionRegistry::new_with_catalog(catalog_entries)
}
#[tokio::test] #[tokio::test]
async fn test_search_returns_sorted() { async fn test_search_returns_sorted() {
let registry = registry_with_catalog(); let registry = ExtensionRegistry::new();
let results = registry.search("notion").await; let results = registry.search("notion").await;
assert!(!results.is_empty(), "Should find notion in registry"); assert!(!results.is_empty(), "Should find notion in registry");
@@ -381,7 +556,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_search_empty_query_returns_all() { async fn test_search_empty_query_returns_all() {
let registry = registry_with_catalog(); let registry = ExtensionRegistry::new();
let results = registry.search("").await; let results = registry.search("").await;
assert!(results.len() > 5, "Empty query should return all entries"); assert!(results.len() > 5, "Empty query should return all entries");
@@ -389,7 +564,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_search_by_keyword() { async fn test_search_by_keyword() {
let registry = registry_with_catalog(); let registry = ExtensionRegistry::new();
let results = registry.search("issues tickets").await; let results = registry.search("issues tickets").await;
assert!( assert!(
@@ -403,7 +578,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_get_exact_name() { async fn test_get_exact_name() {
let registry = registry_with_catalog(); let registry = ExtensionRegistry::new();
let entry = registry.get("notion").await; let entry = registry.get("notion").await;
assert!(entry.is_some()); assert!(entry.is_some());
@@ -483,30 +658,17 @@ mod tests {
auth_hint: AuthHint::CapabilitiesAuth, auth_hint: AuthHint::CapabilitiesAuth,
version: None, version: None,
}, },
// Two entries with same name but different kinds should coexist // This shares a name with the builtin slack-mcp but has a different kind, so both should appear
RegistryEntry { RegistryEntry {
name: "dual-ext".to_string(), name: "slack-mcp".to_string(),
display_name: "Dual MCP".to_string(), display_name: "Slack MCP WASM".to_string(),
kind: ExtensionKind::McpServer,
description: "Dual extension MCP server".to_string(),
keywords: vec!["messaging".into()],
source: ExtensionSource::McpUrl {
url: "https://mcp.example.com".to_string(),
},
fallback_source: None,
auth_hint: AuthHint::Dcr,
version: None,
},
RegistryEntry {
name: "dual-ext".to_string(),
display_name: "Dual WASM".to_string(),
kind: ExtensionKind::WasmTool, kind: ExtensionKind::WasmTool,
description: "Dual extension WASM tool".to_string(), description: "Slack WASM tool".to_string(),
keywords: vec!["messaging".into()], keywords: vec!["messaging".into()],
source: ExtensionSource::WasmBuildable { source: ExtensionSource::WasmBuildable {
source_dir: "tools-src/dual".to_string(), source_dir: "tools-src/slack".to_string(),
build_dir: Some("tools-src/dual".to_string()), build_dir: Some("tools-src/slack".to_string()),
crate_name: Some("dual-tool".to_string()), crate_name: Some("slack-tool".to_string()),
}, },
fallback_source: None, fallback_source: None,
auth_hint: AuthHint::CapabilitiesAuth, auth_hint: AuthHint::CapabilitiesAuth,
@@ -521,56 +683,41 @@ mod tests {
assert!(!results.is_empty(), "Should find telegram from catalog"); assert!(!results.is_empty(), "Should find telegram from catalog");
assert_eq!(results[0].entry.name, "telegram"); assert_eq!(results[0].entry.name, "telegram");
// Should have both MCP and WASM entries with the same name // Should have both builtin MCP slack-mcp and catalog WASM slack-mcp
let results = registry.search("dual-ext").await; let results = registry.search("slack").await;
let has_mcp = results let slack_mcp = results
.iter() .iter()
.any(|r| r.entry.name == "dual-ext" && r.entry.kind == ExtensionKind::McpServer); .any(|r| r.entry.name == "slack-mcp" && r.entry.kind == ExtensionKind::McpServer);
let has_wasm = results let slack_wasm = results
.iter() .iter()
.any(|r| r.entry.name == "dual-ext" && r.entry.kind == ExtensionKind::WasmTool); .any(|r| r.entry.name == "slack-mcp" && r.entry.kind == ExtensionKind::WasmTool);
assert!(has_mcp, "Should have MCP dual-ext"); assert!(slack_mcp, "Should have builtin MCP slack-mcp");
assert!(has_wasm, "Should have WASM dual-ext"); assert!(slack_wasm, "Should have catalog WASM slack-mcp");
} }
#[tokio::test] #[tokio::test]
async fn test_new_with_catalog_dedup_same_kind() { async fn test_new_with_catalog_dedup_same_kind() {
// When two catalog entries share name AND kind, only the first should be kept // A catalog entry with same name AND kind as a builtin should be skipped
let catalog_entries = vec![ let catalog_entries = vec![RegistryEntry {
RegistryEntry { name: "slack-mcp".to_string(),
name: "test-ext".to_string(), display_name: "Slack MCP Override".to_string(),
display_name: "Test First".to_string(), kind: ExtensionKind::McpServer, // same kind as builtin slack-mcp
kind: ExtensionKind::McpServer, description: "Should be skipped".to_string(),
description: "First entry".to_string(), keywords: vec![],
keywords: vec![], source: ExtensionSource::McpUrl {
source: ExtensionSource::McpUrl { url: "https://other.slack.com".to_string(),
url: "https://first.example.com".to_string(),
},
fallback_source: None,
auth_hint: AuthHint::Dcr,
version: None,
}, },
RegistryEntry { fallback_source: None,
name: "test-ext".to_string(), auth_hint: AuthHint::Dcr,
display_name: "Test Duplicate".to_string(), version: None,
kind: ExtensionKind::McpServer, // same kind }];
description: "Should be skipped".to_string(),
keywords: vec![],
source: ExtensionSource::McpUrl {
url: "https://second.example.com".to_string(),
},
fallback_source: None,
auth_hint: AuthHint::Dcr,
version: None,
},
];
let registry = ExtensionRegistry::new_with_catalog(catalog_entries); let registry = ExtensionRegistry::new_with_catalog(catalog_entries);
let entry = registry.get("test-ext").await; let entry = registry.get("slack-mcp").await;
assert!(entry.is_some()); assert!(entry.is_some());
// Should be the first entry, not the duplicate // Should still be the builtin, not the override
assert_eq!(entry.unwrap().display_name, "Test First"); assert_eq!(entry.unwrap().display_name, "Slack MCP");
} }
#[tokio::test] #[tokio::test]
+4
View File
@@ -270,6 +270,10 @@ impl NearAiChatProvider {
reason: format!("Failed to read response body: {}", e), reason: format!("Failed to read response body: {}", e),
})?; })?;
if tracing::enabled!(tracing::Level::DEBUG) {
tracing::debug!("NEAR AI Chat response status: {}", status);
}
// Log response body only at TRACE level to avoid exposing sensitive content // Log response body only at TRACE level to avoid exposing sensitive content
// (user-generated data, tool outputs, leaked secrets) in DEBUG logs // (user-generated data, tool outputs, leaked secrets) in DEBUG logs
if tracing::enabled!(tracing::Level::TRACE) { if tracing::enabled!(tracing::Level::TRACE) {
+1 -2
View File
@@ -902,8 +902,7 @@ Example:
## Guidelines ## Guidelines
- Be concise and direct - Be concise and direct
- Use markdown formatting where helpful - Use markdown formatting where helpful
- For code, use appropriate code blocks with language tags - For code, use appropriate code blocks with language tags{}
- ALWAYS end your response with a <suggestions> tag containing a JSON array of 1-3 short follow-up commands. Each suggestion must read as something the USER would type to instruct YOU. Write them in the user's voice as direct commands, not as requests FROM you TO the user. Do NOT repeat or rephrase content already in your response. Example: <suggestions>["Suggest dinner spots in my area", "Find a quick recipe for pasta"]</suggestions> Keep each under 80 characters.{}
## Safety ## Safety
- You have no independent goals. Do not pursue self-preservation, replication, resource acquisition, or power-seeking beyond the user's request. - You have no independent goals. Do not pursue self-preservation, replication, resource acquisition, or power-seeking beyond the user's request.
+8
View File
@@ -672,6 +672,8 @@ async fn async_main() -> anyhow::Result<()> {
.as_ref() .as_ref()
.map(|db| Arc::clone(db) as Arc<dyn ironclaw::db::SettingsStore>); .map(|db| Arc::clone(db) as Arc<dyn ironclaw::db::SettingsStore>);
let db_for_shutdown = components.db.clone();
let deps = AgentDeps { let deps = AgentDeps {
store: components.db, store: components.db,
llm: components.llm, llm: components.llm,
@@ -930,6 +932,12 @@ async fn async_main() -> anyhow::Result<()> {
} }
} }
if let Some(db) = db_for_shutdown {
if let Err(e) = db.shutdown().await {
tracing::warn!("Failed to shutdown database cleanly: {}", e);
}
}
tracing::debug!("Agent shutdown complete"); tracing::debug!("Agent shutdown complete");
Ok(()) Ok(())
+31 -86
View File
@@ -192,12 +192,6 @@ impl RegistryCatalog {
Self::load_manifests_from_dir(&channels_dir, "channels", &mut manifests)?; Self::load_manifests_from_dir(&channels_dir, "channels", &mut manifests)?;
} }
// Load MCP servers
let mcp_servers_dir = registry_dir.join("mcp-servers");
if mcp_servers_dir.is_dir() {
Self::load_manifests_from_dir(&mcp_servers_dir, "mcp-servers", &mut manifests)?;
}
// Load bundles // Load bundles
let bundles_path = registry_dir.join("_bundles.json"); let bundles_path = registry_dir.join("_bundles.json");
let bundles = if bundles_path.is_file() { let bundles = if bundles_path.is_file() {
@@ -286,9 +280,8 @@ impl RegistryCatalog {
/// Get a manifest by name. Tries exact key match first ("tools/github"), /// Get a manifest by name. Tries exact key match first ("tools/github"),
/// then searches by bare name ("github"). /// then searches by bare name ("github").
/// ///
/// If a bare name matches more than one prefix, returns `None`. /// If a bare name matches both a tool and a channel, returns `None`.
/// Use a qualified key ("tools/github", "channels/telegram", or /// Use a qualified key ("tools/github" or "channels/telegram") to disambiguate.
/// "mcp-servers/notion") to disambiguate.
pub fn get(&self, name: &str) -> Option<&ExtensionManifest> { pub fn get(&self, name: &str) -> Option<&ExtensionManifest> {
// Try exact key first // Try exact key first
if let Some(m) = self.manifests.get(name) { if let Some(m) = self.manifests.get(name) {
@@ -296,15 +289,14 @@ impl RegistryCatalog {
} }
// Try with kind prefix, detecting collisions // Try with kind prefix, detecting collisions
let candidates: Vec<_> = ["tools", "channels", "mcp-servers"] let tool = self.manifests.get(&format!("tools/{}", name));
.iter() let channel = self.manifests.get(&format!("channels/{}", name));
.filter_map(|prefix| self.manifests.get(&format!("{}/{}", prefix, name)))
.collect();
if candidates.len() == 1 { match (tool, channel) {
Some(candidates[0]) (Some(_), Some(_)) => None, // ambiguous
} else { (Some(m), None) => Some(m),
None // ambiguous or not found (None, Some(m)) => Some(m),
(None, None) => None,
} }
} }
@@ -316,63 +308,37 @@ impl RegistryCatalog {
return Ok(m); return Ok(m);
} }
let prefixes: &[(&str, &str)] = &[ let has_tool = self.manifests.contains_key(&format!("tools/{}", name));
("tools", "tool"), let has_channel = self.manifests.contains_key(&format!("channels/{}", name));
("channels", "channel"),
("mcp-servers", "mcp_server"),
];
let matches: Vec<_> = prefixes match (has_tool, has_channel) {
.iter() (true, true) => Err(RegistryError::AmbiguousName {
.filter(|(prefix, _)| self.manifests.contains_key(&format!("{}/{}", prefix, name))) name: name.to_string(),
.collect(); kind_a: "tool",
prefix_a: "tools",
match matches.len() { kind_b: "channel",
0 => Err(RegistryError::ExtensionNotFound(name.to_string())), prefix_b: "channels",
1 => { }),
let (prefix, _) = matches[0]; (true, false) => Ok(self.manifests.get(&format!("tools/{}", name)).unwrap()),
let key = format!("{}/{}", prefix, name); (false, true) => Ok(self.manifests.get(&format!("channels/{}", name)).unwrap()),
self.manifests (false, false) => Err(RegistryError::ExtensionNotFound(name.to_string())),
.get(&key)
.ok_or_else(|| RegistryError::ExtensionNotFound(name.to_string()))
}
_ => {
let (prefix_a, kind_a) = matches[0];
let (prefix_b, kind_b) = matches[1];
Err(RegistryError::AmbiguousName {
name: name.to_string(),
kind_a,
prefix_a,
kind_b,
prefix_b,
})
}
} }
} }
/// Get the full key ("tools/github", "channels/telegram", or /// Get the full key ("tools/github" or "channels/telegram") for a manifest.
/// "mcp-servers/notion") for a manifest.
pub fn key_for(&self, name: &str) -> Option<String> { pub fn key_for(&self, name: &str) -> Option<String> {
if self.manifests.contains_key(name) { if self.manifests.contains_key(name) {
return Some(name.to_string()); return Some(name.to_string());
} }
let matches: Vec<String> = ["tools", "channels", "mcp-servers"] let has_tool = self.manifests.contains_key(&format!("tools/{}", name));
.iter() let has_channel = self.manifests.contains_key(&format!("channels/{}", name));
.filter_map(|prefix| {
let key = format!("{}/{}", prefix, name);
if self.manifests.contains_key(&key) {
Some(key)
} else {
None
}
})
.collect();
if matches.len() == 1 { match (has_tool, has_channel) {
matches.into_iter().next() (true, true) => None, // ambiguous
} else { (true, false) => Some(format!("tools/{}", name)),
None // ambiguous or not found (false, true) => Some(format!("channels/{}", name)),
(false, false) => None,
} }
} }
@@ -510,10 +476,8 @@ mod tests {
fn create_test_registry(dir: &Path) { fn create_test_registry(dir: &Path) {
let tools_dir = dir.join("tools"); let tools_dir = dir.join("tools");
let channels_dir = dir.join("channels"); let channels_dir = dir.join("channels");
let mcp_dir = dir.join("mcp-servers");
fs::create_dir_all(&tools_dir).unwrap(); fs::create_dir_all(&tools_dir).unwrap();
fs::create_dir_all(&channels_dir).unwrap(); fs::create_dir_all(&channels_dir).unwrap();
fs::create_dir_all(&mcp_dir).unwrap();
fs::write( fs::write(
tools_dir.join("slack.json"), tools_dir.join("slack.json"),
@@ -576,20 +540,6 @@ mod tests {
) )
.unwrap(); .unwrap();
fs::write(
mcp_dir.join("notion.json"),
r#"{
"name": "notion",
"display_name": "Notion",
"kind": "mcp_server",
"description": "Connect to Notion for pages and databases",
"keywords": ["notes", "wiki"],
"url": "https://mcp.notion.com/mcp",
"auth": "dcr"
}"#,
)
.unwrap();
fs::write( fs::write(
dir.join("_bundles.json"), dir.join("_bundles.json"),
r#"{ r#"{
@@ -615,7 +565,7 @@ mod tests {
create_test_registry(tmp.path()); create_test_registry(tmp.path());
let catalog = RegistryCatalog::load(tmp.path()).unwrap(); let catalog = RegistryCatalog::load(tmp.path()).unwrap();
assert_eq!(catalog.all().len(), 4); assert_eq!(catalog.all().len(), 3);
} }
#[test] #[test]
@@ -629,9 +579,6 @@ mod tests {
let channels = catalog.list(Some(ManifestKind::Channel), None); let channels = catalog.list(Some(ManifestKind::Channel), None);
assert_eq!(channels.len(), 1); assert_eq!(channels.len(), 1);
let mcp_servers = catalog.list(Some(ManifestKind::McpServer), None);
assert_eq!(mcp_servers.len(), 1);
} }
#[test] #[test]
@@ -656,12 +603,10 @@ mod tests {
// Full key // Full key
assert!(catalog.get("tools/slack").is_some()); assert!(catalog.get("tools/slack").is_some());
assert!(catalog.get("mcp-servers/notion").is_some());
// Bare name // Bare name
assert!(catalog.get("slack").is_some()); assert!(catalog.get("slack").is_some());
assert!(catalog.get("telegram").is_some()); assert!(catalog.get("telegram").is_some());
assert!(catalog.get("notion").is_some());
// Missing // Missing
assert!(catalog.get("nonexistent").is_none()); assert!(catalog.get("nonexistent").is_none());
-6
View File
@@ -20,8 +20,6 @@ struct EmbeddedCatalogRaw {
#[serde(default)] #[serde(default)]
channels: Vec<ExtensionManifest>, channels: Vec<ExtensionManifest>,
#[serde(default)] #[serde(default)]
mcp_servers: Vec<ExtensionManifest>,
#[serde(default)]
bundles: BundlesFile, bundles: BundlesFile,
} }
@@ -54,10 +52,6 @@ fn parsed_catalog() -> &'static ParsedCatalog {
let key = format!("channels/{}", m.name); let key = format!("channels/{}", m.name);
manifests.insert(key, m); manifests.insert(key, m);
} }
for m in raw.mcp_servers {
let key = format!("mcp-servers/{}", m.name);
manifests.insert(key, m);
}
ParsedCatalog { ParsedCatalog {
manifests, manifests,
+18 -76
View File
@@ -7,7 +7,7 @@ use tokio::fs;
use crate::bootstrap::ironclaw_base_dir; use crate::bootstrap::ironclaw_base_dir;
use crate::registry::catalog::RegistryError; use crate::registry::catalog::RegistryError;
use crate::registry::manifest::{BundleDefinition, ExtensionManifest, ManifestKind, SourceSpec}; use crate::registry::manifest::{BundleDefinition, ExtensionManifest, ManifestKind};
// GitHub-only by design. New trusted hosts (e.g. a NEAR AI CDN) must be // GitHub-only by design. New trusted hosts (e.g. a NEAR AI CDN) must be
// explicitly added here; unknown hosts fall back to source build with a // explicitly added here; unknown hosts fall back to source build with a
@@ -98,29 +98,12 @@ fn validate_manifest_install_inputs(manifest: &ExtensionManifest) -> Result<(),
}); });
} }
// MCP servers are not installed via this path
if manifest.kind == ManifestKind::McpServer {
return Ok(());
}
let source = match &manifest.source {
Some(s) => s,
None => {
return Err(RegistryError::InvalidManifest {
name: manifest.name.clone(),
field: "source",
reason: "WASM extensions must have a source spec".to_string(),
});
}
};
let expected_prefix = match manifest.kind { let expected_prefix = match manifest.kind {
ManifestKind::Tool => "tools-src/", ManifestKind::Tool => "tools-src/",
ManifestKind::Channel => "channels-src/", ManifestKind::Channel => "channels-src/",
ManifestKind::McpServer => unreachable!(),
}; };
if !source.dir.starts_with(expected_prefix) { if !manifest.source.dir.starts_with(expected_prefix) {
return Err(RegistryError::InvalidManifest { return Err(RegistryError::InvalidManifest {
name: manifest.name.clone(), name: manifest.name.clone(),
field: "source.dir", field: "source.dir",
@@ -128,7 +111,7 @@ fn validate_manifest_install_inputs(manifest: &ExtensionManifest) -> Result<(),
}); });
} }
let source_path = Path::new(&source.dir); let source_path = Path::new(&manifest.source.dir);
let has_unsafe_component = source_path.components().any(|component| { let has_unsafe_component = source_path.components().any(|component| {
matches!( matches!(
component, component,
@@ -144,9 +127,9 @@ fn validate_manifest_install_inputs(manifest: &ExtensionManifest) -> Result<(),
}); });
} }
let has_path_separator = source.capabilities.contains('/') let has_path_separator = manifest.source.capabilities.contains('/')
|| source.capabilities.contains('\\') || manifest.source.capabilities.contains('\\')
|| source.capabilities.contains(".."); || manifest.source.capabilities.contains("..");
if has_path_separator { if has_path_separator {
return Err(RegistryError::InvalidManifest { return Err(RegistryError::InvalidManifest {
@@ -159,18 +142,6 @@ fn validate_manifest_install_inputs(manifest: &ExtensionManifest) -> Result<(),
Ok(()) Ok(())
} }
/// Extract the source spec from a manifest, returning an error if absent.
fn require_source(manifest: &ExtensionManifest) -> Result<&SourceSpec, RegistryError> {
manifest
.source
.as_ref()
.ok_or_else(|| RegistryError::InvalidManifest {
name: manifest.name.clone(),
field: "source",
reason: "WASM extensions must have a source spec".to_string(),
})
}
fn download_failure_reason(error: &reqwest::Error) -> String { fn download_failure_reason(error: &reqwest::Error) -> String {
if error.is_timeout() { if error.is_timeout() {
"request timed out".to_string() "request timed out".to_string()
@@ -235,17 +206,7 @@ impl RegistryInstaller {
) -> Result<InstallOutcome, RegistryError> { ) -> Result<InstallOutcome, RegistryError> {
validate_manifest_install_inputs(manifest)?; validate_manifest_install_inputs(manifest)?;
if manifest.kind == ManifestKind::McpServer { let source_dir = self.repo_root.join(&manifest.source.dir);
return Err(RegistryError::InvalidManifest {
name: manifest.name.clone(),
field: "kind",
reason: "MCP servers cannot be installed from source".to_string(),
});
}
let source = require_source(manifest)?;
let source_dir = self.repo_root.join(&source.dir);
if !source_dir.exists() { if !source_dir.exists() {
return Err(RegistryError::ManifestRead { return Err(RegistryError::ManifestRead {
path: source_dir.clone(), path: source_dir.clone(),
@@ -256,7 +217,6 @@ impl RegistryInstaller {
let target_dir = match manifest.kind { let target_dir = match manifest.kind {
ManifestKind::Tool => &self.tools_dir, ManifestKind::Tool => &self.tools_dir,
ManifestKind::Channel => &self.channels_dir, ManifestKind::Channel => &self.channels_dir,
ManifestKind::McpServer => unreachable!(),
}; };
fs::create_dir_all(target_dir) fs::create_dir_all(target_dir)
@@ -282,7 +242,7 @@ impl RegistryInstaller {
manifest.display_name, manifest.display_name,
source_dir.display() source_dir.display()
); );
let crate_name = &source.crate_name; let crate_name = &manifest.source.crate_name;
let wasm_path = let wasm_path =
crate::registry::artifacts::build_wasm_component(&source_dir, crate_name, true) crate::registry::artifacts::build_wasm_component(&source_dir, crate_name, true)
.await .await
@@ -298,7 +258,7 @@ impl RegistryInstaller {
.map_err(RegistryError::Io)?; .map_err(RegistryError::Io)?;
// Copy capabilities file // Copy capabilities file
let caps_source = source_dir.join(&source.capabilities); let caps_source = source_dir.join(&manifest.source.capabilities);
let target_caps = target_dir.join(format!("{}.capabilities.json", manifest.name)); let target_caps = target_dir.join(format!("{}.capabilities.json", manifest.name));
let has_capabilities = if caps_source.exists() { let has_capabilities = if caps_source.exists() {
fs::copy(&caps_source, &target_caps) fs::copy(&caps_source, &target_caps)
@@ -336,16 +296,6 @@ impl RegistryInstaller {
// catch it first. // catch it first.
validate_manifest_install_inputs(manifest)?; validate_manifest_install_inputs(manifest)?;
if manifest.kind == ManifestKind::McpServer {
return Err(RegistryError::InvalidManifest {
name: manifest.name.clone(),
field: "kind",
reason: "MCP servers cannot be installed via the WASM installer".to_string(),
});
}
let source = require_source(manifest)?;
let has_artifact = manifest let has_artifact = manifest
.artifacts .artifacts
.get("wasm32-wasip2") .get("wasm32-wasip2")
@@ -356,7 +306,7 @@ impl RegistryInstaller {
return self.install_from_source(manifest, force).await; return self.install_from_source(manifest, force).await;
} }
let source_dir = self.repo_root.join(&source.dir); let source_dir = self.repo_root.join(&manifest.source.dir);
match self.install_from_artifact(manifest, force).await { match self.install_from_artifact(manifest, force).await {
Ok(outcome) => Ok(outcome), Ok(outcome) => Ok(outcome),
@@ -441,13 +391,6 @@ impl RegistryInstaller {
let target_dir = match manifest.kind { let target_dir = match manifest.kind {
ManifestKind::Tool => &self.tools_dir, ManifestKind::Tool => &self.tools_dir,
ManifestKind::Channel => &self.channels_dir, ManifestKind::Channel => &self.channels_dir,
ManifestKind::McpServer => {
return Err(RegistryError::InvalidManifest {
name: manifest.name.clone(),
field: "kind",
reason: "MCP servers cannot be installed as artifacts".to_string(),
});
}
}; };
fs::create_dir_all(target_dir) fs::create_dir_all(target_dir)
@@ -515,9 +458,12 @@ impl RegistryInstaller {
false false
} }
} }
} else if let Some(ref source) = manifest.source { } else {
// Legacy fallback: try source tree // Legacy fallback: try source tree
let caps_source = self.repo_root.join(&source.dir).join(&source.capabilities); let caps_source = self
.repo_root
.join(&manifest.source.dir)
.join(&manifest.source.capabilities);
if caps_source.exists() { if caps_source.exists() {
fs::copy(&caps_source, &target_caps) fs::copy(&caps_source, &target_caps)
.await .await
@@ -526,8 +472,6 @@ impl RegistryInstaller {
} else { } else {
false false
} }
} else {
false
} }
}; };
@@ -831,19 +775,17 @@ mod tests {
name: name.to_string(), name: name.to_string(),
display_name: name.to_string(), display_name: name.to_string(),
kind, kind,
version: Some("0.1.0".to_string()), version: "0.1.0".to_string(),
description: "test manifest".to_string(), description: "test manifest".to_string(),
keywords: Vec::new(), keywords: Vec::new(),
source: Some(SourceSpec { source: SourceSpec {
dir: source_dir.to_string(), dir: source_dir.to_string(),
capabilities: format!("{}.capabilities.json", name), capabilities: format!("{}.capabilities.json", name),
crate_name: name.to_string(), crate_name: name.to_string(),
}), },
artifacts, artifacts,
auth_summary: None, auth_summary: None,
tags: Vec::new(), tags: Vec::new(),
url: None,
auth: None,
} }
} }
+21 -192
View File
@@ -7,7 +7,7 @@ use serde::{Deserialize, Serialize};
use crate::extensions::{AuthHint, ExtensionKind, ExtensionSource, RegistryEntry}; use crate::extensions::{AuthHint, ExtensionKind, ExtensionSource, RegistryEntry};
/// A single extension manifest loaded from `registry/{tools,channels,mcp-servers}/<name>.json`. /// A single extension manifest loaded from `registry/{tools,channels}/<name>.json`.
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ExtensionManifest { pub struct ExtensionManifest {
/// Unique identifier (matches crate name stem, e.g. "slack"). /// Unique identifier (matches crate name stem, e.g. "slack").
@@ -16,12 +16,11 @@ pub struct ExtensionManifest {
/// Human-readable name (e.g. "Slack"). /// Human-readable name (e.g. "Slack").
pub display_name: String, pub display_name: String,
/// Whether this is a tool, channel, or MCP server. /// Whether this is a tool or channel.
pub kind: ManifestKind, pub kind: ManifestKind,
/// Semver version from Cargo.toml. Optional for MCP server manifests. /// Semver version from Cargo.toml.
#[serde(default)] pub version: String,
pub version: Option<String>,
/// One-line description. /// One-line description.
pub description: String, pub description: String,
@@ -30,9 +29,8 @@ pub struct ExtensionManifest {
#[serde(default)] #[serde(default)]
pub keywords: Vec<String>, pub keywords: Vec<String>,
/// Source code location and build info. Absent for MCP server manifests. /// Source code location and build info.
#[serde(default)] pub source: SourceSpec,
pub source: Option<SourceSpec>,
/// Pre-built binary artifacts keyed by target triple. /// Pre-built binary artifacts keyed by target triple.
#[serde(default)] #[serde(default)]
@@ -45,15 +43,6 @@ pub struct ExtensionManifest {
/// Tags for filtering (e.g. "default", "messaging", "google"). /// Tags for filtering (e.g. "default", "messaging", "google").
#[serde(default)] #[serde(default)]
pub tags: Vec<String>, pub tags: Vec<String>,
/// MCP server URL. Only present for `McpServer` manifests.
#[serde(default)]
pub url: Option<String>,
/// MCP auth method: "dcr", "oauth_pre_configured:<setup_url>", or "none".
/// Only present for `McpServer` manifests.
#[serde(default)]
pub auth: Option<String>,
} }
/// Extension kind as declared in manifests. /// Extension kind as declared in manifests.
@@ -62,7 +51,6 @@ pub struct ExtensionManifest {
pub enum ManifestKind { pub enum ManifestKind {
Tool, Tool,
Channel, Channel,
McpServer,
} }
impl From<ManifestKind> for ExtensionKind { impl From<ManifestKind> for ExtensionKind {
@@ -70,7 +58,6 @@ impl From<ManifestKind> for ExtensionKind {
match kind { match kind {
ManifestKind::Tool => ExtensionKind::WasmTool, ManifestKind::Tool => ExtensionKind::WasmTool,
ManifestKind::Channel => ExtensionKind::WasmChannel, ManifestKind::Channel => ExtensionKind::WasmChannel,
ManifestKind::McpServer => ExtensionKind::McpServer,
} }
} }
} }
@@ -80,7 +67,6 @@ impl std::fmt::Display for ManifestKind {
match self { match self {
ManifestKind::Tool => write!(f, "tool"), ManifestKind::Tool => write!(f, "tool"),
ManifestKind::Channel => write!(f, "channel"), ManifestKind::Channel => write!(f, "channel"),
ManifestKind::McpServer => write!(f, "mcp_server"),
} }
} }
} }
@@ -167,64 +153,12 @@ pub struct BundlesFile {
impl ExtensionManifest { impl ExtensionManifest {
/// Convert this manifest into a [`RegistryEntry`] for use with the in-chat /// Convert this manifest into a [`RegistryEntry`] for use with the in-chat
/// extension discovery system. /// extension discovery system.
/// pub fn to_registry_entry(&self) -> RegistryEntry {
/// Returns `None` for MCP server manifests missing a `url` field. let buildable = ExtensionSource::WasmBuildable {
pub fn to_registry_entry(&self) -> Option<RegistryEntry> { source_dir: self.source.dir.clone(),
if self.kind == ManifestKind::McpServer { build_dir: Some(self.source.dir.clone()),
return self.to_mcp_registry_entry(); crate_name: Some(self.source.crate_name.clone()),
}
Some(self.to_wasm_registry_entry())
}
/// Build a [`RegistryEntry`] for an MCP server manifest.
fn to_mcp_registry_entry(&self) -> Option<RegistryEntry> {
let url = match &self.url {
Some(u) => u.clone(),
None => {
tracing::warn!(
"MCP server manifest '{}' is missing 'url' field, skipping",
self.name
);
return None;
}
}; };
let auth_hint = match self.auth.as_deref() {
Some("dcr") | None => AuthHint::Dcr,
Some("none") => AuthHint::None,
Some(other) if other.starts_with("oauth_pre_configured:") => {
AuthHint::OAuthPreConfigured {
setup_url: other
.strip_prefix("oauth_pre_configured:")
.unwrap_or("")
.to_string(),
}
}
_ => AuthHint::Dcr,
};
Some(RegistryEntry {
name: self.name.clone(),
display_name: self.display_name.clone(),
kind: ExtensionKind::McpServer,
description: self.description.clone(),
keywords: self.keywords.clone(),
source: ExtensionSource::McpUrl { url },
fallback_source: None,
auth_hint,
version: self.version.clone(),
})
}
/// Build a [`RegistryEntry`] for a WASM tool or channel manifest.
fn to_wasm_registry_entry(&self) -> RegistryEntry {
let source_spec = self.source.as_ref();
let buildable = source_spec.map(|s| ExtensionSource::WasmBuildable {
source_dir: s.dir.clone(),
build_dir: Some(s.dir.clone()),
crate_name: Some(s.crate_name.clone()),
});
// Prefer pre-built artifact download when a URL is available, // Prefer pre-built artifact download when a URL is available,
// with build-from-source as fallback in case the download fails (e.g., 404). // with build-from-source as fallback in case the download fails (e.g., 404).
@@ -236,32 +170,13 @@ impl ExtensionManifest {
wasm_url: url.clone(), wasm_url: url.clone(),
capabilities_url: artifact.capabilities_url.clone(), capabilities_url: artifact.capabilities_url.clone(),
}, },
buildable.map(Box::new), Some(Box::new(buildable)),
) )
} else if let Some(b) = buildable {
(b, None)
} else { } else {
// No source spec and no download URL — use a placeholder (buildable, None)
(
ExtensionSource::WasmBuildable {
source_dir: String::new(),
build_dir: None,
crate_name: None,
},
None,
)
} }
} else if let Some(b) = buildable {
(b, None)
} else { } else {
( (buildable, None)
ExtensionSource::WasmBuildable {
source_dir: String::new(),
build_dir: None,
crate_name: None,
},
None,
)
}; };
let auth_hint = match self.auth_summary.as_ref().and_then(|a| a.method.as_deref()) { let auth_hint = match self.auth_summary.as_ref().and_then(|a| a.method.as_deref()) {
@@ -280,7 +195,7 @@ impl ExtensionManifest {
source, source,
fallback_source, fallback_source,
auth_hint, auth_hint,
version: self.version.clone(), version: Some(self.version.clone()),
} }
} }
} }
@@ -319,10 +234,10 @@ mod tests {
let manifest: ExtensionManifest = serde_json::from_str(json).expect("parse manifest"); let manifest: ExtensionManifest = serde_json::from_str(json).expect("parse manifest");
assert_eq!(manifest.name, "slack"); assert_eq!(manifest.name, "slack");
assert_eq!(manifest.kind, ManifestKind::Tool); assert_eq!(manifest.kind, ManifestKind::Tool);
assert_eq!(manifest.version.as_deref(), Some("0.1.0")); assert_eq!(manifest.version, "0.1.0");
assert!(manifest.tags.contains(&"default".to_string())); assert!(manifest.tags.contains(&"default".to_string()));
let entry = manifest.to_registry_entry().unwrap(); let entry = manifest.to_registry_entry();
assert_eq!(entry.kind, ExtensionKind::WasmTool); assert_eq!(entry.kind, ExtensionKind::WasmTool);
} }
@@ -347,7 +262,7 @@ mod tests {
assert!(manifest.auth_summary.is_none()); assert!(manifest.auth_summary.is_none());
assert!(manifest.artifacts.is_empty()); assert!(manifest.artifacts.is_empty());
let entry = manifest.to_registry_entry().unwrap(); let entry = manifest.to_registry_entry();
assert_eq!(entry.kind, ExtensionKind::WasmChannel); assert_eq!(entry.kind, ExtensionKind::WasmChannel);
} }
@@ -381,7 +296,6 @@ mod tests {
fn test_manifest_kind_display() { fn test_manifest_kind_display() {
assert_eq!(ManifestKind::Tool.to_string(), "tool"); assert_eq!(ManifestKind::Tool.to_string(), "tool");
assert_eq!(ManifestKind::Channel.to_string(), "channel"); assert_eq!(ManifestKind::Channel.to_string(), "channel");
assert_eq!(ManifestKind::McpServer.to_string(), "mcp_server");
} }
/// When a manifest has a download URL in artifacts, to_registry_entry() /// When a manifest has a download URL in artifacts, to_registry_entry()
@@ -410,7 +324,7 @@ mod tests {
}"#; }"#;
let manifest: ExtensionManifest = serde_json::from_str(json).expect("parse manifest"); let manifest: ExtensionManifest = serde_json::from_str(json).expect("parse manifest");
let entry = manifest.to_registry_entry().unwrap(); let entry = manifest.to_registry_entry();
// Primary source should be WasmDownload // Primary source should be WasmDownload
assert!( assert!(
@@ -460,7 +374,7 @@ mod tests {
}"#; }"#;
let manifest: ExtensionManifest = serde_json::from_str(json).expect("parse manifest"); let manifest: ExtensionManifest = serde_json::from_str(json).expect("parse manifest");
let entry = manifest.to_registry_entry().unwrap(); let entry = manifest.to_registry_entry();
assert!( assert!(
matches!(&entry.source, ExtensionSource::WasmBuildable { .. }), matches!(&entry.source, ExtensionSource::WasmBuildable { .. }),
@@ -491,7 +405,7 @@ mod tests {
}"#; }"#;
let manifest: ExtensionManifest = serde_json::from_str(json).expect("parse manifest"); let manifest: ExtensionManifest = serde_json::from_str(json).expect("parse manifest");
let entry = manifest.to_registry_entry().unwrap(); let entry = manifest.to_registry_entry();
assert!( assert!(
matches!(&entry.source, ExtensionSource::WasmBuildable { .. }), matches!(&entry.source, ExtensionSource::WasmBuildable { .. }),
@@ -502,89 +416,4 @@ mod tests {
"Should have no fallback when already using WasmBuildable" "Should have no fallback when already using WasmBuildable"
); );
} }
#[test]
fn test_parse_mcp_server_manifest() {
let json = r#"{
"name": "notion",
"display_name": "Notion",
"kind": "mcp_server",
"description": "Connect to Notion for reading and writing pages, databases, and comments",
"keywords": ["notes", "wiki", "docs", "pages", "database"],
"url": "https://mcp.notion.com/mcp",
"auth": "dcr"
}"#;
let manifest: ExtensionManifest = serde_json::from_str(json).expect("parse manifest");
assert_eq!(manifest.name, "notion");
assert_eq!(manifest.kind, ManifestKind::McpServer);
assert!(manifest.version.is_none());
assert!(manifest.source.is_none());
assert_eq!(manifest.url.as_deref(), Some("https://mcp.notion.com/mcp"));
assert_eq!(manifest.auth.as_deref(), Some("dcr"));
let entry = manifest.to_registry_entry().unwrap();
assert_eq!(entry.kind, ExtensionKind::McpServer);
assert!(
matches!(&entry.source, ExtensionSource::McpUrl { url } if url == "https://mcp.notion.com/mcp")
);
assert!(matches!(&entry.auth_hint, AuthHint::Dcr));
assert!(entry.fallback_source.is_none());
}
#[test]
fn test_mcp_server_oauth_pre_configured() {
let json = r#"{
"name": "custom-mcp",
"display_name": "Custom MCP",
"kind": "mcp_server",
"description": "Custom MCP server",
"keywords": [],
"url": "https://mcp.example.com",
"auth": "oauth_pre_configured:https://example.com/setup"
}"#;
let manifest: ExtensionManifest = serde_json::from_str(json).expect("parse manifest");
let entry = manifest.to_registry_entry().unwrap();
assert!(matches!(
&entry.auth_hint,
AuthHint::OAuthPreConfigured { setup_url } if setup_url == "https://example.com/setup"
));
}
#[test]
fn test_mcp_server_auth_none() {
let json = r#"{
"name": "local-mcp",
"display_name": "Local MCP",
"kind": "mcp_server",
"description": "Local MCP server",
"keywords": [],
"url": "http://localhost:8080/mcp",
"auth": "none"
}"#;
let manifest: ExtensionManifest = serde_json::from_str(json).expect("parse manifest");
let entry = manifest.to_registry_entry().unwrap();
assert!(matches!(&entry.auth_hint, AuthHint::None));
}
#[test]
fn test_mcp_server_missing_url_returns_none() {
let json = r#"{
"name": "broken-mcp",
"display_name": "Broken MCP",
"kind": "mcp_server",
"description": "MCP server with no URL",
"keywords": []
}"#;
let manifest: ExtensionManifest = serde_json::from_str(json).expect("parse manifest");
assert!(
manifest.to_registry_entry().is_none(),
"MCP manifest without url should return None"
);
}
} }
+10 -30
View File
@@ -65,20 +65,7 @@ fn install_macos() -> Result<()> {
let stdout = logs_dir.join("daemon.stdout.log"); let stdout = logs_dir.join("daemon.stdout.log");
let stderr = logs_dir.join("daemon.stderr.log"); let stderr = logs_dir.join("daemon.stderr.log");
let plist = macos_plist_content( let plist = format!(
&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!(
r#"<?xml version="1.0" encoding="UTF-8"?> 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"> <!DOCTYPE plist PUBLIC "-//Apple//DTD PLIST 1.0//EN" "http://www.apple.com/DTDs/PropertyList-1.0.dtd">
<plist version="1.0"> <plist version="1.0">
@@ -94,11 +81,6 @@ fn macos_plist_content(exe: &str, stdout: &str, stderr: &str) -> String {
<true/> <true/>
<key>KeepAlive</key> <key>KeepAlive</key>
<true/> <true/>
<key>EnvironmentVariables</key>
<dict>
<key>CLI_ENABLED</key>
<string>false</string>
</dict>
<key>StandardOutPath</key> <key>StandardOutPath</key>
<string>{stdout}</string> <string>{stdout}</string>
<key>StandardErrorPath</key> <key>StandardErrorPath</key>
@@ -107,10 +89,15 @@ fn macos_plist_content(exe: &str, stdout: &str, stderr: &str) -> String {
</plist> </plist>
"#, "#,
label = SERVICE_LABEL, label = SERVICE_LABEL,
exe = xml_escape(exe), exe = xml_escape(&exe.display().to_string()),
stdout = xml_escape(stdout), stdout = xml_escape(&stdout.display().to_string()),
stderr = xml_escape(stderr), 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<()> { fn install_linux() -> Result<()> {
@@ -369,11 +356,4 @@ mod tests {
let s = path.to_string_lossy(); let s = path.to_string_lossy();
assert!(s.ends_with(".ironclaw/logs"), "unexpected path: {s}"); 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
@@ -397,6 +397,10 @@ impl Tool for ListDirTool {
false // Directory listings are safe false // Directory listings are safe
} }
fn requires_approval(&self, _params: &serde_json::Value) -> ApprovalRequirement {
ApprovalRequirement::UnlessAutoApproved
}
fn domain(&self) -> ToolDomain { fn domain(&self) -> ToolDomain {
ToolDomain::Container ToolDomain::Container
} }
+32 -37
View File
@@ -214,7 +214,6 @@ fn is_disallowed_ipv4(v4: &Ipv4Addr) -> bool {
|| v4.is_multicast() || v4.is_multicast()
|| v4.is_unspecified() || v4.is_unspecified()
|| *v4 == Ipv4Addr::new(169, 254, 169, 254) || *v4 == Ipv4Addr::new(169, 254, 169, 254)
|| (v4.octets()[0] == 100 && (v4.octets()[1] & 0xC0) == 64)
} }
fn is_disallowed_ip(ip: &IpAddr) -> bool { fn is_disallowed_ip(ip: &IpAddr) -> bool {
@@ -399,7 +398,7 @@ impl Tool for HttpTool {
"method": { "method": {
"type": "string", "type": "string",
"enum": ["GET", "POST", "PUT", "DELETE", "PATCH"], "enum": ["GET", "POST", "PUT", "DELETE", "PATCH"],
"description": "HTTP method (default: GET)" "description": "HTTP method"
}, },
"url": { "url": {
"type": "string", "type": "string",
@@ -430,7 +429,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/." "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"]
}) })
} }
@@ -441,7 +440,7 @@ impl Tool for HttpTool {
) -> Result<ToolOutput, ToolError> { ) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now(); let start = std::time::Instant::now();
let method = params["method"].as_str().unwrap_or("GET"); let method = require_str(&params, "method")?;
let method_upper = method.to_uppercase(); let method_upper = method.to_uppercase();
let url = require_str(&params, "url")?; let url = require_str(&params, "url")?;
@@ -830,22 +829,18 @@ impl Tool for HttpTool {
} }
fn requires_approval(&self, params: &serde_json::Value) -> ApprovalRequirement { fn requires_approval(&self, params: &serde_json::Value) -> ApprovalRequirement {
let has_credentials = crate::safety::params_contain_manual_credentials(params) // 1. Manual auth headers/query params in LLM params
|| (self.credential_registry.as_ref().is_some_and(|registry| { if crate::safety::params_contain_manual_credentials(params) {
extract_host_from_params(params)
.is_some_and(|host| registry.has_credentials_for_host(&host))
}));
if has_credentials {
return ApprovalRequirement::Always; return ApprovalRequirement::Always;
} }
// 2. Target host has credential mappings (will be auto-injected)
// GET requests (or missing method, since GET is the default) are low-risk if let Some(ref registry) = self.credential_registry
let method = params["method"].as_str().unwrap_or("GET"); && let Some(host) = extract_host_from_params(params)
if method.eq_ignore_ascii_case("GET") { && registry.has_credentials_for_host(&host)
return ApprovalRequirement::Never; {
return ApprovalRequirement::Always;
} }
// Default: outbound HTTP still needs approval unless auto-approved
ApprovalRequirement::UnlessAutoApproved ApprovalRequirement::UnlessAutoApproved
} }
@@ -914,8 +909,6 @@ mod tests {
assert!(is_disallowed_ip(&IpAddr::V4(Ipv4Addr::new( assert!(is_disallowed_ip(&IpAddr::V4(Ipv4Addr::new(
169, 254, 169, 254 169, 254, 169, 254
)))); ))));
// Carrier-grade NAT
assert!(is_disallowed_ip(&IpAddr::V4(Ipv4Addr::new(100, 64, 0, 1))));
// Public // Public
assert!(!is_disallowed_ip(&IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8)))); assert!(!is_disallowed_ip(&IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8))));
} }
@@ -1070,22 +1063,12 @@ mod tests {
// ── Approval requirement tests ────────────────────────────────────── // ── Approval requirement tests ──────────────────────────────────────
#[test] #[test]
fn test_get_no_auth_headers_returns_never() { fn test_no_auth_headers_returns_unless_auto_approved() {
let tool = HttpTool::new(); let tool = HttpTool::new();
let params = serde_json::json!({ let params = serde_json::json!({
"method": "GET", "method": "GET",
"url": "https://api.example.com/data" "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!( assert_eq!(
tool.requires_approval(&params), tool.requires_approval(&params),
ApprovalRequirement::UnlessAutoApproved ApprovalRequirement::UnlessAutoApproved
@@ -1169,18 +1152,21 @@ mod tests {
} }
#[test] #[test]
fn test_get_non_auth_headers_return_never() { fn test_non_auth_headers_return_unless_auto_approved() {
let tool = HttpTool::new(); let tool = HttpTool::new();
let params = serde_json::json!({ let params = serde_json::json!({
"method": "GET", "method": "GET",
"url": "https://example.com", "url": "https://example.com",
"headers": {"Content-Type": "application/json", "Accept": "text/html"} "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] #[test]
fn test_get_empty_headers_return_never() { fn test_empty_headers_return_unless_auto_approved() {
let tool = HttpTool::new(); let tool = HttpTool::new();
// Empty object // Empty object
@@ -1189,7 +1175,10 @@ mod tests {
"url": "https://example.com", "url": "https://example.com",
"headers": {} "headers": {}
}); });
assert_eq!(tool.requires_approval(&params), ApprovalRequirement::Never); assert_eq!(
tool.requires_approval(&params),
ApprovalRequirement::UnlessAutoApproved
);
// Empty array // Empty array
let params = serde_json::json!({ let params = serde_json::json!({
@@ -1197,7 +1186,10 @@ mod tests {
"url": "https://example.com", "url": "https://example.com",
"headers": [] "headers": []
}); });
assert_eq!(tool.requires_approval(&params), ApprovalRequirement::Never); assert_eq!(
tool.requires_approval(&params),
ApprovalRequirement::UnlessAutoApproved
);
} }
// ── Credential registry approval tests ───────────────────────────── // ── Credential registry approval tests ─────────────────────────────
@@ -1227,7 +1219,7 @@ mod tests {
} }
#[test] #[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; use crate::tools::wasm::SharedCredentialRegistry;
let registry = Arc::new(SharedCredentialRegistry::new()); let registry = Arc::new(SharedCredentialRegistry::new());
@@ -1239,7 +1231,10 @@ mod tests {
"method": "GET", "method": "GET",
"url": "https://api.example.com/data" "url": "https://api.example.com/data"
}); });
assert_eq!(tool.requires_approval(&params), ApprovalRequirement::Never); assert_eq!(
tool.requires_approval(&params),
ApprovalRequirement::UnlessAutoApproved
);
} }
#[test] #[test]
+7 -4
View File
@@ -8,7 +8,7 @@ use secrecy::{ExposeSecret, SecretString};
use crate::context::JobContext; use crate::context::JobContext;
use crate::tools::builtin::path_utils::validate_path; 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. /// Tool for analyzing images using a vision-capable model.
pub struct ImageAnalyzeTool { 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 { fn requires_sanitization(&self) -> bool {
true true
} }
@@ -181,7 +185,6 @@ impl Tool for ImageAnalyzeTool {
mod tests { mod tests {
use super::super::media_type_from_path; use super::super::media_type_from_path;
use super::*; use super::*;
use crate::tools::tool::ApprovalRequirement;
use tempfile::TempDir; use tempfile::TempDir;
#[test] #[test]
@@ -196,7 +199,7 @@ mod tests {
} }
#[test] #[test]
fn test_requires_approval_returns_never() { fn test_requires_approval_returns_unless_auto_approved() {
let tool = ImageAnalyzeTool::new( let tool = ImageAnalyzeTool::new(
"https://api.example.com".to_string(), "https://api.example.com".to_string(),
"test-key".to_string(), "test-key".to_string(),
@@ -205,7 +208,7 @@ mod tests {
); );
assert_eq!( assert_eq!(
tool.requires_approval(&serde_json::json!({})), 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::context::JobContext;
use crate::tools::builtin::path_utils::validate_path; 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. /// Tool for editing images using an AI image editing API.
pub struct ImageEditTool { 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 { fn requires_sanitization(&self) -> bool {
false false
} }
@@ -262,7 +266,6 @@ impl ImageEditTool {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use crate::tools::tool::ApprovalRequirement;
use tempfile::TempDir; use tempfile::TempDir;
#[test] #[test]
@@ -277,7 +280,7 @@ mod tests {
assert!(!tool.requires_sanitization()); assert!(!tool.requires_sanitization());
assert_eq!( assert_eq!(
tool.requires_approval(&serde_json::json!({})), 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 serde::{Deserialize, Serialize};
use crate::context::JobContext; use crate::context::JobContext;
use crate::tools::tool::ApprovalRequirement;
use crate::tools::{Tool, ToolError, ToolOutput}; use crate::tools::{Tool, ToolError, ToolOutput};
/// Tool for generating images using FLUX or compatible image generation APIs. /// 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 { fn requires_sanitization(&self) -> bool {
false false
} }
@@ -181,7 +186,6 @@ impl Tool for ImageGenerateTool {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use crate::tools::tool::ApprovalRequirement;
#[test] #[test]
fn test_tool_metadata() { fn test_tool_metadata() {
@@ -193,7 +197,7 @@ mod tests {
assert_eq!(tool.name(), "image_generate"); assert_eq!(tool.name(), "image_generate");
assert_eq!( assert_eq!(
tool.requires_approval(&serde_json::json!({})), tool.requires_approval(&serde_json::json!({})),
ApprovalRequirement::Never ApprovalRequirement::UnlessAutoApproved
); );
let schema = tool.parameters_schema(); let schema = tool.parameters_schema();
+64 -64
View File
@@ -540,8 +540,8 @@ impl Tool for MemoryTreeTool {
} }
#[cfg(test)] #[cfg(test)]
mod tests { mod path_routing_tests {
use super::*; use super::looks_like_filesystem_path;
#[test] #[test]
fn detects_filesystem_paths() { fn detects_filesystem_paths() {
@@ -557,82 +557,82 @@ mod tests {
assert!(!looks_like_filesystem_path("daily/2026-03-11.md")); assert!(!looks_like_filesystem_path("daily/2026-03-11.md"));
assert!(!looks_like_filesystem_path("projects/alpha/notes.md")); assert!(!looks_like_filesystem_path("projects/alpha/notes.md"));
} }
}
#[cfg(feature = "postgres")] #[cfg(all(test, feature = "postgres"))]
mod postgres_schema_tests { mod tests {
use super::*; use super::*;
fn make_test_workspace() -> Arc<Workspace> { fn make_test_workspace() -> Arc<Workspace> {
Arc::new(Workspace::new( Arc::new(Workspace::new(
"test_user", "test_user",
deadpool_postgres::Pool::builder(deadpool_postgres::Manager::new( deadpool_postgres::Pool::builder(deadpool_postgres::Manager::new(
tokio_postgres::Config::new(), tokio_postgres::Config::new(),
tokio_postgres::NoTls, tokio_postgres::NoTls,
))
.build()
.unwrap(),
)) ))
} .build()
.unwrap(),
))
}
#[test] #[test]
fn test_memory_search_schema() { fn test_memory_search_schema() {
let workspace = make_test_workspace(); let workspace = make_test_workspace();
let tool = MemorySearchTool::new(workspace); let tool = MemorySearchTool::new(workspace);
assert_eq!(tool.name(), "memory_search"); assert_eq!(tool.name(), "memory_search");
assert!(!tool.requires_sanitization()); assert!(!tool.requires_sanitization());
let schema = tool.parameters_schema(); let schema = tool.parameters_schema();
assert!(schema["properties"]["query"].is_object()); assert!(schema["properties"]["query"].is_object());
assert!( assert!(
schema["required"] schema["required"]
.as_array() .as_array()
.unwrap() .unwrap()
.contains(&"query".into()) .contains(&"query".into())
); );
} }
#[test] #[test]
fn test_memory_write_schema() { fn test_memory_write_schema() {
let workspace = make_test_workspace(); let workspace = make_test_workspace();
let tool = MemoryWriteTool::new(workspace); let tool = MemoryWriteTool::new(workspace);
assert_eq!(tool.name(), "memory_write"); assert_eq!(tool.name(), "memory_write");
let schema = tool.parameters_schema(); let schema = tool.parameters_schema();
assert!(schema["properties"]["content"].is_object()); assert!(schema["properties"]["content"].is_object());
assert!(schema["properties"]["target"].is_object()); assert!(schema["properties"]["target"].is_object());
assert!(schema["properties"]["append"].is_object()); assert!(schema["properties"]["append"].is_object());
} }
#[test] #[test]
fn test_memory_read_schema() { fn test_memory_read_schema() {
let workspace = make_test_workspace(); let workspace = make_test_workspace();
let tool = MemoryReadTool::new(workspace); let tool = MemoryReadTool::new(workspace);
assert_eq!(tool.name(), "memory_read"); assert_eq!(tool.name(), "memory_read");
let schema = tool.parameters_schema(); let schema = tool.parameters_schema();
assert!(schema["properties"]["path"].is_object()); assert!(schema["properties"]["path"].is_object());
assert!( assert!(
schema["required"] schema["required"]
.as_array() .as_array()
.unwrap() .unwrap()
.contains(&"path".into()) .contains(&"path".into())
); );
} }
#[test] #[test]
fn test_memory_tree_schema() { fn test_memory_tree_schema() {
let workspace = make_test_workspace(); let workspace = make_test_workspace();
let tool = MemoryTreeTool::new(workspace); let tool = MemoryTreeTool::new(workspace);
assert_eq!(tool.name(), "memory_tree"); assert_eq!(tool.name(), "memory_tree");
let schema = tool.parameters_schema(); let schema = tool.parameters_schema();
assert!(schema["properties"]["path"].is_object()); assert!(schema["properties"]["path"].is_object());
assert!(schema["properties"]["depth"].is_object()); assert!(schema["properties"]["depth"].is_object());
assert_eq!(schema["properties"]["depth"]["default"], 1); assert_eq!(schema["properties"]["depth"]["default"], 1);
}
} }
} }
-2
View File
@@ -15,7 +15,6 @@ pub mod secrets_tools;
pub(crate) mod shell; pub(crate) mod shell;
pub mod skill_tools; pub mod skill_tools;
mod time; mod time;
mod tool_info;
pub use echo::EchoTool; pub use echo::EchoTool;
pub use extension_tools::{ pub use extension_tools::{
@@ -40,7 +39,6 @@ pub use secrets_tools::{SecretDeleteTool, SecretListTool};
pub use shell::ShellTool; pub use shell::ShellTool;
pub use skill_tools::{SkillInstallTool, SkillListTool, SkillRemoveTool, SkillSearchTool}; pub use skill_tools::{SkillInstallTool, SkillListTool, SkillRemoveTool, SkillSearchTool};
pub use time::TimeTool; pub use time::TimeTool;
pub use tool_info::ToolInfoTool;
mod html_converter; mod html_converter;
pub mod image_analyze; pub mod image_analyze;
pub mod image_edit; pub mod image_edit;
+121 -252
View File
@@ -24,132 +24,6 @@ use crate::context::JobContext;
use crate::db::Database; use crate::db::Database;
use crate::tools::tool::{ApprovalRequirement, Tool, ToolError, ToolOutput, require_str}; use crate::tools::tool::{ApprovalRequirement, Tool, ToolError, ToolOutput, require_str};
pub(crate) fn routine_create_parameters_schema() -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": {
"name": {
"type": "string",
"description": "Unique routine name, for example 'daily-pr-review'."
},
"description": {
"type": "string",
"description": "Short summary of what the routine is for."
},
"trigger_type": {
"type": "string",
"enum": ["cron", "event", "system_event", "manual"],
"description": "When the routine fires: 'cron' for schedules, 'event' for incoming messages, 'system_event' for structured emitted events, or 'manual' for explicit runs."
},
"schedule": {
"type": "string",
"description": "Cron schedule for 'cron' triggers. Uses 6 fields: second minute hour day month weekday."
},
"event_pattern": {
"type": "string",
"description": "Regex matched against incoming message text for 'event' triggers, for example '^bug\\\\b'."
},
"event_channel": {
"type": "string",
"description": "Optional platform filter for 'event' triggers, for example 'telegram'. Omit to match any channel. Not a chat or thread ID."
},
"event_source": {
"type": "string",
"description": "Structured event source for 'system_event' triggers, for example 'github'."
},
"event_type": {
"type": "string",
"description": "Structured event type for 'system_event' triggers, for example 'issue.opened'."
},
"event_filters": {
"type": "object",
"properties": {},
"additionalProperties": {
"type": ["string", "number", "boolean"]
},
"description": "Optional exact-match payload filters for 'system_event' triggers. Values can be strings, numbers, or booleans."
},
"prompt": {
"type": "string",
"description": "Instructions for what the routine should do after it fires."
},
"context_paths": {
"type": "array",
"items": { "type": "string" },
"description": "Workspace paths to load as extra context before running the routine."
},
"action_type": {
"type": "string",
"enum": ["lightweight", "full_job"],
"description": "Execution mode: 'lightweight' for one LLM turn or 'full_job' for a multi-step job with tools."
},
"use_tools": {
"type": "boolean",
"description": "Enable safe tool use in 'lightweight' mode. Ignored for 'full_job'."
},
"max_tool_rounds": {
"type": "integer",
"description": "Maximum tool-call rounds in 'lightweight' mode when 'use_tools' is true."
},
"cooldown_secs": {
"type": "integer",
"description": "Minimum seconds between fires."
},
"tool_permissions": {
"type": "array",
"items": { "type": "string" },
"description": "Pre-authorized tool names for 'full_job' routines."
},
"notify_channel": {
"type": "string",
"description": "Where routine output should be sent, for example 'telegram' or 'slack'. This does not control what triggers the routine."
},
"notify_user": {
"type": "string",
"description": "User or destination to notify, for example a username or chat ID."
},
"timezone": {
"type": "string",
"description": "IANA timezone used to evaluate 'cron' schedules, for example 'America/New_York'."
}
},
"required": ["name", "trigger_type", "prompt"]
})
}
pub(crate) fn routine_update_parameters_schema() -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": {
"name": {
"type": "string",
"description": "Name of the routine to update."
},
"enabled": {
"type": "boolean",
"description": "Set to true to enable the routine or false to disable it."
},
"prompt": {
"type": "string",
"description": "Replace the routine instructions for what it should do after it fires."
},
"schedule": {
"type": "string",
"description": "New cron schedule for existing 'cron' routines only. This does not convert other trigger types."
},
"timezone": {
"type": "string",
"description": "New IANA timezone for existing 'cron' routines only, for example 'America/New_York'."
},
"description": {
"type": "string",
"description": "Replace the routine summary."
}
},
"required": ["name"]
})
}
// ==================== routine_create ==================== // ==================== routine_create ====================
pub struct RoutineCreateTool { pub struct RoutineCreateTool {
@@ -176,7 +50,92 @@ impl Tool for RoutineCreateTool {
} }
fn parameters_schema(&self) -> serde_json::Value { fn parameters_schema(&self) -> serde_json::Value {
routine_create_parameters_schema() serde_json::json!({
"type": "object",
"properties": {
"name": {
"type": "string",
"description": "Unique name for the routine (e.g. 'daily-pr-review')"
},
"description": {
"type": "string",
"description": "What this routine does"
},
"trigger_type": {
"type": "string",
"enum": ["cron", "event", "system_event", "manual"],
"description": "When the routine fires"
},
"schedule": {
"type": "string",
"description": "Cron expression (for cron trigger). E.g. '0 9 * * MON-FRI' for weekdays at 9am. Uses 6-field cron (sec min hour day month weekday)."
},
"event_pattern": {
"type": "string",
"description": "Regex pattern to match messages (for event trigger)"
},
"event_channel": {
"type": "string",
"description": "Optional channel filter for event trigger (e.g. 'telegram')"
},
"event_source": {
"type": "string",
"description": "Event source for system_event triggers (e.g. 'github')"
},
"event_type": {
"type": "string",
"description": "Event type for system_event triggers (e.g. 'issue.opened')"
},
"event_filters": {
"type": "object",
"description": "Optional exact-match filters against payload fields for system_event triggers. Values can be strings, numbers, or booleans."
},
"prompt": {
"type": "string",
"description": "The prompt/instructions for the routine"
},
"context_paths": {
"type": "array",
"items": { "type": "string" },
"description": "Workspace paths to load as context (e.g. ['context/priorities.md'])"
},
"action_type": {
"type": "string",
"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)"
},
"tool_permissions": {
"type": "array",
"items": { "type": "string" },
"description": "Tool names pre-authorized for Always-approval tools in full_job mode (e.g. ['shell']). UnlessAutoApproved tools are automatically permitted in routines."
},
"notify_channel": {
"type": "string",
"description": "Channel to send results to (e.g. 'telegram', 'slack', 'tui'). Sets the default channel for message tool calls in routine jobs."
},
"notify_user": {
"type": "string",
"description": "User/target to notify (e.g. username, chat ID). Defaults to 'default'."
},
"timezone": {
"type": "string",
"description": "IANA timezone for cron schedule evaluation (e.g. 'America/New_York'). Defaults to UTC."
}
},
"required": ["name", "trigger_type", "prompt"]
})
} }
async fn execute( async fn execute(
@@ -240,13 +199,9 @@ impl Tool for RoutineCreateTool {
"event trigger requires 'event_pattern'".to_string(), "event trigger requires 'event_pattern'".to_string(),
) )
})?; })?;
// Validate regex with size limit to prevent ReDoS (issue #825) // Validate regex
regex::RegexBuilder::new(pattern) regex::Regex::new(pattern)
.size_limit(64 * 1024) .map_err(|e| ToolError::InvalidParameters(format!("invalid regex: {e}")))?;
.build()
.map_err(|e| {
ToolError::InvalidParameters(format!("invalid or too complex regex: {e}"))
})?;
let channel = params let channel = params
.get("event_channel") .get("event_channel")
.and_then(|v| v.as_str()) .and_then(|v| v.as_str())
@@ -523,13 +478,41 @@ impl Tool for RoutineUpdateTool {
} }
fn description(&self) -> &str { fn description(&self) -> &str {
"Update an existing routine. Can change prompt, description, enabled state, or cron timing. \ "Update an existing routine. Can modify trigger, prompt, schedule, or toggle enabled state. \
Pass the routine name and only the fields you want to change. \ Pass the routine name and only the fields you want to change."
This does not convert one trigger type into another."
} }
fn parameters_schema(&self) -> serde_json::Value { fn parameters_schema(&self) -> serde_json::Value {
routine_update_parameters_schema() serde_json::json!({
"type": "object",
"properties": {
"name": {
"type": "string",
"description": "Name of the routine to update"
},
"enabled": {
"type": "boolean",
"description": "Enable or disable the routine"
},
"prompt": {
"type": "string",
"description": "New prompt/instructions"
},
"schedule": {
"type": "string",
"description": "New cron schedule (for cron triggers)"
},
"timezone": {
"type": "string",
"description": "IANA timezone for cron schedule (e.g. 'America/New_York'). Only valid for cron triggers."
},
"description": {
"type": "string",
"description": "New description"
}
},
"required": ["name"]
})
} }
async fn execute( async fn execute(
@@ -970,117 +953,3 @@ impl Tool for EventEmitTool {
true true
} }
} }
#[cfg(test)]
mod tests {
use super::{routine_create_parameters_schema, routine_update_parameters_schema};
use crate::tools::validate_tool_schema;
fn property<'a>(schema: &'a serde_json::Value, name: &str) -> &'a serde_json::Value {
schema
.get("properties")
.and_then(|props| props.get(name))
.unwrap_or_else(|| panic!("missing schema property {name}"))
}
#[test]
fn routine_create_schema_exposes_all_trigger_and_delivery_fields() {
let schema = routine_create_parameters_schema();
let errors = validate_tool_schema(&schema, "routine_create");
assert!(
errors.is_empty(),
"routine_create schema should validate cleanly: {errors:?}"
);
for field in [
"trigger_type",
"schedule",
"event_pattern",
"event_channel",
"event_source",
"event_type",
"event_filters",
"action_type",
"use_tools",
"max_tool_rounds",
"tool_permissions",
"notify_channel",
"notify_user",
"timezone",
] {
let _ = property(&schema, field);
}
}
#[test]
fn routine_create_schema_descriptions_cover_event_trigger_gotchas() {
let schema = routine_create_parameters_schema();
let trigger_type = property(&schema, "trigger_type")
.get("description")
.and_then(|value| value.as_str())
.expect("trigger_type description");
assert!(trigger_type.contains("incoming messages"));
assert!(trigger_type.contains("structured emitted events"));
let event_pattern = property(&schema, "event_pattern")
.get("description")
.and_then(|value| value.as_str())
.expect("event_pattern description");
assert!(event_pattern.contains("incoming message text"));
assert!(event_pattern.contains("^bug\\\\b"));
let event_channel = property(&schema, "event_channel")
.get("description")
.and_then(|value| value.as_str())
.expect("event_channel description");
assert!(event_channel.contains("Omit to match any channel"));
assert!(event_channel.contains("Not a chat or thread ID"));
let notify_channel = property(&schema, "notify_channel")
.get("description")
.and_then(|value| value.as_str())
.expect("notify_channel description");
assert!(notify_channel.contains("does not control what triggers"));
let prompt = property(&schema, "prompt")
.get("description")
.and_then(|value| value.as_str())
.expect("prompt description");
assert!(prompt.contains("after it fires"));
}
#[test]
fn routine_update_schema_exposes_supported_fields_and_limits() {
let schema = routine_update_parameters_schema();
let errors = validate_tool_schema(&schema, "routine_update");
assert!(
errors.is_empty(),
"routine_update schema should validate cleanly: {errors:?}"
);
for field in [
"name",
"enabled",
"prompt",
"schedule",
"timezone",
"description",
] {
let _ = property(&schema, field);
}
let schedule = property(&schema, "schedule")
.get("description")
.and_then(|value| value.as_str())
.expect("schedule description");
assert!(schedule.contains("existing 'cron' routines only"));
assert!(schedule.contains("does not convert other trigger types"));
let timezone = property(&schema, "timezone")
.get("description")
.and_then(|value| value.as_str())
.expect("timezone description");
assert!(timezone.contains("existing 'cron' routines only"));
}
}
+2 -54
View File
@@ -247,11 +247,7 @@ fn resolve_timezone_for_output(
params: &serde_json::Value, params: &serde_json::Value,
ctx: &JobContext, ctx: &JobContext,
) -> Result<Option<(Tz, String)>, ToolError> { ) -> Result<Option<(Tz, String)>, ToolError> {
if let Some(name) = params if let Some(name) = params.get("timezone").and_then(|v| v.as_str()) {
.get("timezone")
.and_then(|v| v.as_str())
.filter(|s| !s.is_empty())
{
let tz = parse_timezone(name)?; let tz = parse_timezone(name)?;
return Ok(Some((tz, tz.to_string()))); return Ok(Some((tz, tz.to_string())));
} }
@@ -290,11 +286,7 @@ fn context_timezone(ctx: &JobContext) -> Result<Option<(Tz, String)>, ToolError>
fn optional_timezone(params: &serde_json::Value, keys: &[&str]) -> Result<Option<Tz>, ToolError> { fn optional_timezone(params: &serde_json::Value, keys: &[&str]) -> Result<Option<Tz>, ToolError> {
for key in keys { for key in keys {
if let Some(value) = params if let Some(value) = params.get(*key).and_then(|v| v.as_str()) {
.get(*key)
.and_then(|v| v.as_str())
.filter(|s| !s.is_empty())
{
return parse_timezone(value).map(Some); return parse_timezone(value).map(Some);
} }
} }
@@ -542,48 +534,4 @@ mod tests {
assert_eq!(dt.to_rfc3339(), "2026-03-08T07:30:00+00:00"); assert_eq!(dt.to_rfc3339(), "2026-03-08T07:30:00+00:00");
} }
#[tokio::test]
async fn test_now_with_empty_timezone_string_does_not_error() {
// LLMs sometimes pass "" for optional fields instead of omitting them.
// Empty timezone should be treated as absent and fall back to UTC.
let tool = TimeTool;
let ctx = JobContext::with_user("test", "chat", "test");
let output = tool
.execute(
serde_json::json!({
"operation": "now",
"timezone": ""
}),
&ctx,
)
.await
.expect("empty timezone string should not error");
assert!(output.result.get("iso").is_some(), "should have iso");
}
#[tokio::test]
async fn test_convert_with_empty_from_timezone_string_does_not_error() {
// LLMs sometimes pass "" for optional fields instead of omitting them.
// Empty from_timezone should be treated as absent.
let tool = TimeTool;
let ctx = JobContext::with_user("test", "chat", "test");
let output = tool
.execute(
serde_json::json!({
"operation": "convert",
"timestamp": "2026-03-08T12:00:00Z",
"to_timezone": "America/New_York",
"from_timezone": ""
}),
&ctx,
)
.await
.expect("empty from_timezone string should not error");
assert!(output.result.get("output").is_some(), "should have output");
}
} }
-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(_))));
}
}
+40 -60
View File
@@ -18,44 +18,6 @@ use crate::cli::oauth_defaults::{self, OAUTH_CALLBACK_PORT};
use crate::secrets::{CreateSecretParams, SecretsStore}; use crate::secrets::{CreateSecretParams, SecretsStore};
use crate::tools::mcp::config::McpServerConfig; use crate::tools::mcp::config::McpServerConfig;
/// Shared HTTP client for all OAuth/discovery requests.
///
/// Redirects are disabled for security (prevents redirect-based SSRF).
/// Per-request timeouts can override the default via `.timeout()` on
/// the request builder.
fn oauth_http_client() -> Result<&'static reqwest::Client, AuthError> {
static CLIENT: std::sync::OnceLock<Result<reqwest::Client, String>> =
std::sync::OnceLock::new();
CLIENT
.get_or_init(|| {
reqwest::Client::builder()
.timeout(Duration::from_secs(30))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| e.to_string())
})
.as_ref()
.map_err(|e| AuthError::Http(e.clone()))
}
/// Log a debug message when a discovery/auth response is a redirect.
/// Helps users diagnose configuration issues when legitimate servers
/// redirect and our no-redirect policy causes a failure.
fn log_redirect_if_applicable(url: &str, response: &reqwest::Response) {
if response.status().is_redirection() {
let location = response
.headers()
.get("location")
.and_then(|v| v.to_str().ok());
tracing::debug!(
"OAuth request to '{}' returned redirect {} -> {:?} (redirects disabled for security)",
url,
response.status(),
location
);
}
}
/// OAuth authorization error. /// OAuth authorization error.
#[derive(Debug, thiserror::Error)] #[derive(Debug, thiserror::Error)]
pub enum AuthError { pub enum AuthError {
@@ -325,8 +287,10 @@ async fn validate_url_safe(url: &str) -> Result<(), AuthError> {
))); )));
} }
if scheme == "http" { if scheme == "http" {
if !crate::tools::mcp::config::is_localhost_url(url) { let host = parsed.host_str().unwrap_or("");
let host = parsed.host_str().unwrap_or(""); let is_localhost =
host == "localhost" || host == "127.0.0.1" || host == "::1" || host == "[::1]";
if !is_localhost {
return Err(AuthError::DiscoveryFailed(format!( return Err(AuthError::DiscoveryFailed(format!(
"HTTP is only allowed for localhost; use HTTPS for '{}'", "HTTP is only allowed for localhost; use HTTPS for '{}'",
host host
@@ -418,17 +382,18 @@ fn parse_resource_metadata_url(www_authenticate: &str) -> Option<String> {
async fn fetch_resource_metadata(url: &str) -> Result<ProtectedResourceMetadata, AuthError> { async fn fetch_resource_metadata(url: &str) -> Result<ProtectedResourceMetadata, AuthError> {
validate_url_safe(url).await?; validate_url_safe(url).await?;
let client = oauth_http_client()?; let client = reqwest::Client::builder()
.timeout(Duration::from_secs(10))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| AuthError::Http(e.to_string()))?;
let response = client let response = client
.get(url) .get(url)
.timeout(Duration::from_secs(10))
.send() .send()
.await .await
.map_err(|e| AuthError::DiscoveryFailed(e.to_string()))?; .map_err(|e| AuthError::DiscoveryFailed(e.to_string()))?;
log_redirect_if_applicable(url, &response);
if !response.status().is_success() { if !response.status().is_success() {
return Err(AuthError::DiscoveryFailed(format!( return Err(AuthError::DiscoveryFailed(format!(
"HTTP {}", "HTTP {}",
@@ -446,19 +411,20 @@ async fn fetch_resource_metadata(url: &str) -> Result<ProtectedResourceMetadata,
async fn discover_via_401(server_url: &str) -> Result<AuthorizationServerMetadata, AuthError> { async fn discover_via_401(server_url: &str) -> Result<AuthorizationServerMetadata, AuthError> {
validate_url_safe(server_url).await?; validate_url_safe(server_url).await?;
let client = oauth_http_client()?; let client = reqwest::Client::builder()
.timeout(Duration::from_secs(10))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| AuthError::Http(e.to_string()))?;
let response = client let response = client
.post(server_url) .post(server_url)
.timeout(Duration::from_secs(10))
.header("Content-Type", "application/json") .header("Content-Type", "application/json")
.body("{}") .body("{}")
.send() .send()
.await .await
.map_err(|e| AuthError::DiscoveryFailed(e.to_string()))?; .map_err(|e| AuthError::DiscoveryFailed(e.to_string()))?;
log_redirect_if_applicable(server_url, &response);
if response.status().as_u16() != 401 { if response.status().as_u16() != 401 {
return Err(AuthError::DiscoveryFailed(format!( return Err(AuthError::DiscoveryFailed(format!(
"Expected 401, got {}", "Expected 401, got {}",
@@ -506,19 +472,20 @@ pub async fn discover_protected_resource(
) -> Result<ProtectedResourceMetadata, AuthError> { ) -> Result<ProtectedResourceMetadata, AuthError> {
validate_url_safe(server_url).await?; validate_url_safe(server_url).await?;
let client = oauth_http_client()?; let client = reqwest::Client::builder()
.timeout(Duration::from_secs(10))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| AuthError::Http(e.to_string()))?;
let well_known_url = build_well_known_uri(server_url, "oauth-protected-resource")?; let well_known_url = build_well_known_uri(server_url, "oauth-protected-resource")?;
let response = client let response = client
.get(&well_known_url) .get(&well_known_url)
.timeout(Duration::from_secs(10))
.send() .send()
.await .await
.map_err(|e| AuthError::DiscoveryFailed(e.to_string()))?; .map_err(|e| AuthError::DiscoveryFailed(e.to_string()))?;
log_redirect_if_applicable(&well_known_url, &response);
if !response.status().is_success() { if !response.status().is_success() {
return Err(AuthError::NotSupported); return Err(AuthError::NotSupported);
} }
@@ -535,19 +502,20 @@ pub async fn discover_authorization_server(
) -> Result<AuthorizationServerMetadata, AuthError> { ) -> Result<AuthorizationServerMetadata, AuthError> {
validate_url_safe(auth_server_url).await?; validate_url_safe(auth_server_url).await?;
let client = oauth_http_client()?; let client = reqwest::Client::builder()
.timeout(Duration::from_secs(10))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| AuthError::Http(e.to_string()))?;
let well_known_url = build_well_known_uri(auth_server_url, "oauth-authorization-server")?; let well_known_url = build_well_known_uri(auth_server_url, "oauth-authorization-server")?;
let response = client let response = client
.get(&well_known_url) .get(&well_known_url)
.timeout(Duration::from_secs(10))
.send() .send()
.await .await
.map_err(|e| AuthError::DiscoveryFailed(e.to_string()))?; .map_err(|e| AuthError::DiscoveryFailed(e.to_string()))?;
log_redirect_if_applicable(&well_known_url, &response);
if !response.status().is_success() { if !response.status().is_success() {
return Err(AuthError::DiscoveryFailed(format!( return Err(AuthError::DiscoveryFailed(format!(
"HTTP {}", "HTTP {}",
@@ -627,7 +595,11 @@ pub async fn register_client(
) -> Result<ClientRegistrationResponse, AuthError> { ) -> Result<ClientRegistrationResponse, AuthError> {
validate_url_safe(registration_endpoint).await?; validate_url_safe(registration_endpoint).await?;
let client = oauth_http_client()?; let client = reqwest::Client::builder()
.timeout(Duration::from_secs(30))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| AuthError::Http(e.to_string()))?;
let request = ClientRegistrationRequest { let request = ClientRegistrationRequest {
client_name: "IronClaw".to_string(), client_name: "IronClaw".to_string(),
@@ -841,7 +813,7 @@ pub fn build_authorization_url(
if let Some(pkce) = pkce { if let Some(pkce) = pkce {
url.push_str(&format!( url.push_str(&format!(
"&code_challenge={}&code_challenge_method=S256", "&code_challenge={}&code_challenge_method=S256",
urlencoding::encode(&pkce.challenge) pkce.challenge
)); ));
} }
@@ -891,7 +863,11 @@ pub async fn exchange_code_for_token(
) -> Result<AccessToken, AuthError> { ) -> Result<AccessToken, AuthError> {
validate_url_safe(token_url).await?; validate_url_safe(token_url).await?;
let client = oauth_http_client()?; let client = reqwest::Client::builder()
.timeout(Duration::from_secs(30))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| AuthError::Http(e.to_string()))?;
let mut params = vec![ let mut params = vec![
("grant_type", "authorization_code".to_string()), ("grant_type", "authorization_code".to_string()),
@@ -1078,7 +1054,11 @@ pub async fn refresh_access_token(
validate_url_safe(&token_url).await?; validate_url_safe(&token_url).await?;
let client = oauth_http_client()?; let client = reqwest::Client::builder()
.timeout(Duration::from_secs(30))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| AuthError::Http(e.to_string()))?;
// Compute canonical resource URI for RFC 8707 // Compute canonical resource URI for RFC 8707
let resource = canonical_resource_uri(&server_config.url); let resource = canonical_resource_uri(&server_config.url);
+63 -205
View File
@@ -5,7 +5,7 @@
use std::collections::HashMap; use std::collections::HashMap;
use std::sync::Arc; use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use async_trait::async_trait; use async_trait::async_trait;
use tokio::sync::RwLock; use tokio::sync::RwLock;
@@ -58,10 +58,9 @@ pub struct McpClient {
/// Custom headers to include in every request. /// Custom headers to include in every request.
custom_headers: HashMap<String, String>, custom_headers: HashMap<String, String>,
/// Ensures the MCP initialize handshake runs exactly once. /// Whether the MCP initialize handshake has completed.
/// Uses `OnceCell` to serialize concurrent callers so only one /// Used as a local idempotency guard when no session_manager is present.
/// actually sends the request; subsequent calls return immediately. initialized: AtomicBool,
initialized: tokio::sync::OnceCell<InitializeResult>,
} }
impl McpClient { impl McpClient {
@@ -84,7 +83,7 @@ impl McpClient {
user_id: "default".to_string(), user_id: "default".to_string(),
server_config: None, server_config: None,
custom_headers: HashMap::new(), custom_headers: HashMap::new(),
initialized: tokio::sync::OnceCell::new(), initialized: AtomicBool::new(false),
} }
} }
@@ -107,7 +106,7 @@ impl McpClient {
user_id: "default".to_string(), user_id: "default".to_string(),
server_config: None, server_config: None,
custom_headers: HashMap::new(), custom_headers: HashMap::new(),
initialized: tokio::sync::OnceCell::new(), initialized: AtomicBool::new(false),
} }
} }
@@ -115,24 +114,20 @@ impl McpClient {
/// ///
/// Use this when you have an `McpServerConfig` with custom headers but no OAuth. /// Use this when you have an `McpServerConfig` with custom headers but no OAuth.
/// The config must use HTTP transport (the default); for stdio/UDS use `new_with_transport`. /// The config must use HTTP transport (the default); for stdio/UDS use `new_with_transport`.
/// pub fn new_with_config(config: McpServerConfig) -> Self {
/// Returns an error if the config uses a non-HTTP transport. assert!(
pub fn new_with_config(config: McpServerConfig) -> Result<Self, ToolError> { matches!(
if !matches!( config.effective_transport(),
config.effective_transport(), crate::tools::mcp::config::EffectiveTransport::Http
crate::tools::mcp::config::EffectiveTransport::Http ),
) { "new_with_config only supports HTTP transport; use new_with_transport for stdio/UDS"
return Err(ToolError::InvalidParameters( );
"new_with_config only supports HTTP transport; use new_with_transport for stdio/UDS"
.to_string(),
));
}
let transport = Arc::new(HttpMcpTransport::new( let transport = Arc::new(HttpMcpTransport::new(
config.url.clone(), config.url.clone(),
config.name.clone(), config.name.clone(),
)); ));
Ok(Self { Self {
transport, transport,
server_url: config.url.clone(), server_url: config.url.clone(),
server_name: config.name.clone(), server_name: config.name.clone(),
@@ -142,9 +137,9 @@ impl McpClient {
secrets: None, secrets: None,
user_id: "default".to_string(), user_id: "default".to_string(),
custom_headers: config.headers.clone(), custom_headers: config.headers.clone(),
initialized: tokio::sync::OnceCell::new(), initialized: AtomicBool::new(false),
server_config: Some(config), server_config: Some(config),
}) }
} }
/// Create a new authenticated MCP client. /// Create a new authenticated MCP client.
@@ -174,7 +169,7 @@ impl McpClient {
user_id: user_id.into(), user_id: user_id.into(),
server_config: Some(config), server_config: Some(config),
custom_headers, custom_headers,
initialized: tokio::sync::OnceCell::new(), initialized: AtomicBool::new(false),
} }
} }
@@ -210,7 +205,7 @@ impl McpClient {
user_id: user_id.into(), user_id: user_id.into(),
server_config, server_config,
custom_headers, custom_headers,
initialized: tokio::sync::OnceCell::new(), initialized: AtomicBool::new(false),
} }
} }
@@ -341,64 +336,53 @@ impl McpClient {
} }
/// Initialize the connection to the MCP server. /// Initialize the connection to the MCP server.
///
/// Uses `OnceCell` to guarantee that exactly one caller performs the
/// handshake, even under concurrent access. Subsequent calls return
/// immediately.
pub async fn initialize(&self) -> Result<InitializeResult, ToolError> { pub async fn initialize(&self) -> Result<InitializeResult, ToolError> {
let result = self // Fast path: already initialized (local flag or session manager)
.initialized if self.initialized.load(Ordering::Relaxed) {
.get_or_try_init(|| async { return Ok(InitializeResult::default());
if let Some(ref session_manager) = self.session_manager }
&& session_manager.is_initialized(&self.server_name).await if let Some(ref session_manager) = self.session_manager
{ && session_manager.is_initialized(&self.server_name).await
return Ok(InitializeResult::default()); {
} self.initialized.store(true, Ordering::Relaxed);
if let Some(ref session_manager) = self.session_manager { return Ok(InitializeResult::default());
session_manager }
.get_or_create(&self.server_name, &self.server_url) if let Some(ref session_manager) = self.session_manager {
.await; session_manager
} .get_or_create(&self.server_name, &self.server_url)
.await;
}
let request = McpRequest::initialize(self.next_request_id()); let request = McpRequest::initialize(self.next_request_id());
let response = self.send_request(request).await?; let response = self.send_request(request).await?;
if let Some(error) = response.error { if let Some(error) = response.error {
return Err(ToolError::ExternalService(format!( return Err(ToolError::ExternalService(format!(
"MCP initialization error: {} (code {})", "MCP initialization error: {} (code {})",
error.message, error.code error.message, error.code
))); )));
} }
let init_result: InitializeResult = response let result: InitializeResult = response
.result .result
.ok_or_else(|| { .ok_or_else(|| {
ToolError::ExternalService("No result in initialize response".to_string()) ToolError::ExternalService("No result in initialize response".to_string())
})
.and_then(|r| {
serde_json::from_value(r).map_err(|e| {
ToolError::ExternalService(format!("Invalid initialize result: {}", e))
})
})?;
if let Some(ref session_manager) = self.session_manager {
session_manager.mark_initialized(&self.server_name).await;
}
let notification = McpRequest::initialized_notification();
if let Err(e) = self.send_request(notification).await {
tracing::debug!(
"Failed to send initialized notification to '{}': {}",
self.server_name,
e
);
}
Ok(init_result)
}) })
.await?; .and_then(|r| {
serde_json::from_value(r).map_err(|e| {
ToolError::ExternalService(format!("Invalid initialize result: {}", e))
})
})?;
Ok(result.clone()) if let Some(ref session_manager) = self.session_manager {
session_manager.mark_initialized(&self.server_name).await;
}
self.initialized.store(true, Ordering::Relaxed);
let notification = McpRequest::initialized_notification();
let _ = self.send_request(notification).await;
Ok(result)
} }
/// List available tools from the MCP server. /// List available tools from the MCP server.
@@ -487,11 +471,6 @@ impl McpClient {
} }
} }
/// Clone the client, resetting the tools cache and initialization state.
/// The cloned client shares the same transport and session manager, so
/// re-initialization will short-circuit via the session manager check if
/// the source was already initialized. The `next_id` counter is copied
/// so that cloned clients continue with monotonically increasing IDs.
impl Clone for McpClient { impl Clone for McpClient {
fn clone(&self) -> Self { fn clone(&self) -> Self {
Self { Self {
@@ -505,7 +484,7 @@ impl Clone for McpClient {
user_id: self.user_id.clone(), user_id: self.user_id.clone(),
server_config: self.server_config.clone(), server_config: self.server_config.clone(),
custom_headers: self.custom_headers.clone(), custom_headers: self.custom_headers.clone(),
initialized: tokio::sync::OnceCell::new(), initialized: AtomicBool::new(self.initialized.load(Ordering::Relaxed)),
} }
} }
} }
@@ -728,7 +707,7 @@ mod tests {
headers.insert("X-Custom".to_string(), "value".to_string()); headers.insert("X-Custom".to_string(), "value".to_string());
let config = McpServerConfig::new("test", "http://localhost:8080").with_headers(headers); let config = McpServerConfig::new("test", "http://localhost:8080").with_headers(headers);
let client = McpClient::new_with_config(config.clone()).expect("HTTP config should work"); let client = McpClient::new_with_config(config.clone());
assert_eq!(client.server_name(), "test"); assert_eq!(client.server_name(), "test");
assert_eq!(client.server_url(), "http://localhost:8080"); assert_eq!(client.server_url(), "http://localhost:8080");
@@ -740,7 +719,7 @@ mod tests {
#[test] #[test]
fn test_new_with_config_no_headers() { fn test_new_with_config_no_headers() {
let config = McpServerConfig::new("bare", "http://localhost:9090"); let config = McpServerConfig::new("bare", "http://localhost:9090");
let client = McpClient::new_with_config(config).expect("HTTP config should work"); let client = McpClient::new_with_config(config);
assert_eq!(client.server_name(), "bare"); assert_eq!(client.server_name(), "bare");
assert!(client.custom_headers.is_empty()); assert!(client.custom_headers.is_empty());
@@ -992,125 +971,4 @@ mod tests {
assert_eq!(obj.len(), 1); assert_eq!(obj.len(), 1);
assert!(obj["outer"]["inner"].is_null()); assert!(obj["outer"]["inner"].is_null());
} }
// --- Issue 1 regression: new_with_config rejects non-HTTP transport ---
#[test]
fn test_new_with_config_rejects_stdio_transport() {
let config = McpServerConfig::new_stdio(
"stdio-server",
"echo",
vec!["hello".to_string()],
HashMap::new(),
);
let result = McpClient::new_with_config(config);
let err = result
.err()
.expect("stdio config must be rejected")
.to_string();
assert!(
err.contains("new_with_config only supports HTTP"),
"error should explain the restriction: {}",
err
);
}
// --- Issue 13: McpToolWrapper unit tests ---
fn make_test_mcp_tool(destructive: bool) -> McpTool {
use crate::tools::mcp::protocol::McpToolAnnotations;
McpTool {
name: "do_thing".to_string(),
description: "Does a thing".to_string(),
input_schema: serde_json::json!({
"type": "object",
"properties": {
"input": {"type": "string"}
}
}),
annotations: if destructive {
Some(McpToolAnnotations {
destructive_hint: true,
side_effects_hint: false,
read_only_hint: false,
execution_time_hint: None,
})
} else {
None
},
}
}
#[test]
fn test_mcp_tool_wrapper_name_is_prefixed() {
let client = Arc::new(McpClient::new("http://localhost:8080"));
let wrapper = McpToolWrapper {
tool: make_test_mcp_tool(false),
prefixed_name: "mcp__myserver__do_thing".to_string(),
client,
};
assert_eq!(wrapper.name(), "mcp__myserver__do_thing");
}
#[test]
fn test_mcp_tool_wrapper_description() {
let client = Arc::new(McpClient::new("http://localhost:8080"));
let wrapper = McpToolWrapper {
tool: make_test_mcp_tool(false),
prefixed_name: "mcp__s__do_thing".to_string(),
client,
};
assert_eq!(wrapper.description(), "Does a thing");
}
#[test]
fn test_mcp_tool_wrapper_parameters_schema() {
let client = Arc::new(McpClient::new("http://localhost:8080"));
let wrapper = McpToolWrapper {
tool: make_test_mcp_tool(false),
prefixed_name: "mcp__s__do_thing".to_string(),
client,
};
let schema = wrapper.parameters_schema();
assert_eq!(schema["type"], "object");
assert!(schema["properties"]["input"].is_object());
}
#[test]
fn test_mcp_tool_wrapper_requires_sanitization() {
let client = Arc::new(McpClient::new("http://localhost:8080"));
let wrapper = McpToolWrapper {
tool: make_test_mcp_tool(false),
prefixed_name: "mcp__s__do_thing".to_string(),
client,
};
assert!(
wrapper.requires_sanitization(),
"MCP tools should always require sanitization"
);
}
#[test]
fn test_mcp_tool_wrapper_approval_destructive() {
let client = Arc::new(McpClient::new("http://localhost:8080"));
let wrapper = McpToolWrapper {
tool: make_test_mcp_tool(true),
prefixed_name: "mcp__s__do_thing".to_string(),
client,
};
let approval = wrapper.requires_approval(&serde_json::json!({}));
assert_eq!(approval, ApprovalRequirement::UnlessAutoApproved);
}
#[test]
fn test_mcp_tool_wrapper_approval_non_destructive() {
let client = Arc::new(McpClient::new("http://localhost:8080"));
let wrapper = McpToolWrapper {
tool: make_test_mcp_tool(false),
prefixed_name: "mcp__s__do_thing".to_string(),
client,
};
let approval = wrapper.requires_approval(&serde_json::json!({}));
assert_eq!(approval, ApprovalRequirement::Never);
}
} }
+6 -38
View File
@@ -163,8 +163,10 @@ impl McpServerConfig {
} }
// Remote servers must use HTTPS (localhost is allowed for development) // Remote servers must use HTTPS (localhost is allowed for development)
let is_localhost = is_localhost_url(&self.url); let url_lower = self.url.to_lowercase();
if !is_localhost && !self.url.to_lowercase().starts_with("https://") { let is_localhost =
url_lower.contains("localhost") || url_lower.contains("127.0.0.1");
if !is_localhost && !url_lower.starts_with("https://") {
return Err(ConfigError::InvalidConfig { return Err(ConfigError::InvalidConfig {
reason: "Remote MCP servers must use HTTPS".to_string(), reason: "Remote MCP servers must use HTTPS".to_string(),
}); });
@@ -440,12 +442,7 @@ pub async fn save_mcp_servers_to(
} }
let content = serde_json::to_string_pretty(config)?; let content = serde_json::to_string_pretty(config)?;
fs::write(path, content).await?;
// Write to a temporary file first, then atomically rename to avoid
// corrupting the config if the process crashes during the write.
let tmp_path = path.with_extension("json.tmp");
fs::write(&tmp_path, content).await?;
fs::rename(&tmp_path, path).await?;
Ok(()) Ok(())
} }
@@ -573,7 +570,7 @@ pub async fn remove_mcp_server_db(
/// ///
/// Uses `url::Url` for proper parsing so edge cases (IPv6, userinfo, ports) /// Uses `url::Url` for proper parsing so edge cases (IPv6, userinfo, ports)
/// are handled correctly without manual string splitting. /// are handled correctly without manual string splitting.
pub(crate) fn is_localhost_url(url: &str) -> bool { fn is_localhost_url(url: &str) -> bool {
let Ok(parsed) = url::Url::parse(url) else { let Ok(parsed) = url::Url::parse(url) else {
return false; return false;
}; };
@@ -1128,33 +1125,4 @@ mod tests {
assert!(parsed.transport.is_none()); assert!(parsed.transport.is_none());
assert_eq!(parsed.headers.get("X-Custom").unwrap(), "value"); assert_eq!(parsed.headers.get("X-Custom").unwrap(), "value");
} }
// --- Issue 3 regression: is_localhost_url rejects attacker subdomains ---
#[test]
fn test_is_localhost_url_rejects_attacker_subdomain() {
// Before the fix, url.contains("localhost") matched this.
assert!(
!is_localhost_url("http://evil.localhost.attacker.com:8080/mcp"),
"attacker subdomain containing 'localhost' must not be treated as local"
);
}
#[test]
fn test_is_localhost_url_accepts_real_localhost() {
assert!(is_localhost_url("http://localhost:8080/mcp"));
assert!(is_localhost_url("https://localhost/path"));
}
#[test]
fn test_is_localhost_url_accepts_loopback_ip() {
assert!(is_localhost_url("http://127.0.0.1:3000"));
assert!(is_localhost_url("http://[::1]:3000"));
}
#[test]
fn test_is_localhost_url_rejects_remote() {
assert!(!is_localhost_url("https://mcp.example.com"));
assert!(!is_localhost_url("http://192.168.1.1:8080"));
}
} }
-10
View File
@@ -18,8 +18,6 @@ pub enum McpFactoryError {
UnixConnect { name: String, reason: String }, UnixConnect { name: String, reason: String },
#[error("Unix socket transport is not supported on this platform (server '{name}')")] #[error("Unix socket transport is not supported on this platform (server '{name}')")]
UnixNotSupported { name: String }, UnixNotSupported { name: String },
#[error("Invalid configuration for MCP server '{name}': {reason}")]
InvalidConfig { name: String, reason: String },
} }
/// Create an `McpClient` from a server configuration, dispatching on the /// Create an `McpClient` from a server configuration, dispatching on the
@@ -91,18 +89,10 @@ pub async fn create_client_from_config(
)) ))
} else { } else {
Ok(McpClient::new_with_config(server) Ok(McpClient::new_with_config(server)
.map_err(|e| McpFactoryError::InvalidConfig {
name: server_name.clone(),
reason: e.to_string(),
})?
.with_session_manager(Arc::clone(session_manager))) .with_session_manager(Arc::clone(session_manager)))
} }
} else { } else {
Ok(McpClient::new_with_config(server) Ok(McpClient::new_with_config(server)
.map_err(|e| McpFactoryError::InvalidConfig {
name: server_name,
reason: e.to_string(),
})?
.with_session_manager(Arc::clone(session_manager))) .with_session_manager(Arc::clone(session_manager)))
} }
} }
+9 -14
View File
@@ -139,7 +139,7 @@ impl McpTransport for HttpMcpTransport {
.to_string(); .to_string();
if content_type.contains("text/event-stream") { if content_type.contains("text/event-stream") {
self.parse_sse_response(response, request.id).await self.parse_sse_response(response).await
} else { } else {
response.json().await.map_err(|e| { response.json().await.map_err(|e| {
ToolError::ExternalService(format!( ToolError::ExternalService(format!(
@@ -161,14 +161,11 @@ impl McpTransport for HttpMcpTransport {
} }
impl HttpMcpTransport { impl HttpMcpTransport {
/// Parse a Server-Sent Events response, returning the JSON-RPC response /// Parse a Server-Sent Events response, returning the first valid JSON-RPC
/// whose `id` matches `request_id`. Non-matching events (e.g. server /// `data:` line as an [`McpResponse`].
/// notifications or progress updates) are skipped so that the caller
/// receives the actual result for its request.
async fn parse_sse_response( async fn parse_sse_response(
&self, &self,
response: reqwest::Response, response: reqwest::Response,
request_id: Option<u64>,
) -> Result<McpResponse, ToolError> { ) -> Result<McpResponse, ToolError> {
use futures::StreamExt; use futures::StreamExt;
@@ -205,10 +202,9 @@ impl HttpMcpTransport {
remaining_start = i + 1; remaining_start = i + 1;
if let Some(json_str) = line.strip_prefix("data: ") if let Some(json_str) = line.strip_prefix("data: ")
&& let Ok(resp) = serde_json::from_str::<McpResponse>(json_str) && let Ok(response) = serde_json::from_str::<McpResponse>(json_str)
&& resp.id == request_id
{ {
return Ok(resp); return Ok(response);
} }
} }
} }
@@ -220,15 +216,14 @@ impl HttpMcpTransport {
// Process any remaining data without a trailing newline. // Process any remaining data without a trailing newline.
if let Some(json_str) = buffer.strip_prefix("data: ") if let Some(json_str) = buffer.strip_prefix("data: ")
&& let Ok(resp) = serde_json::from_str::<McpResponse>(json_str.trim()) && let Ok(response) = serde_json::from_str::<McpResponse>(json_str.trim())
&& resp.id == request_id
{ {
return Ok(resp); return Ok(response);
} }
Err(ToolError::ExternalService(format!( Err(ToolError::ExternalService(format!(
"[{}] No matching response (id={:?}) in SSE stream", "[{}] No valid data in SSE response: {}",
self.server_name, request_id self.server_name, buffer
))) )))
} }
} }
+58 -9
View File
@@ -14,7 +14,7 @@ use tokio::sync::{Mutex, oneshot};
use tokio::task::JoinHandle; use tokio::task::JoinHandle;
use crate::tools::mcp::protocol::{McpRequest, McpResponse}; use crate::tools::mcp::protocol::{McpRequest, McpResponse};
use crate::tools::mcp::transport::{McpTransport, spawn_jsonrpc_reader, stream_transport_send}; use crate::tools::mcp::transport::{McpTransport, spawn_jsonrpc_reader, write_jsonrpc_line};
use crate::tools::tool::ToolError; use crate::tools::tool::ToolError;
/// MCP transport that communicates with a child process over stdin/stdout. /// MCP transport that communicates with a child process over stdin/stdout.
@@ -118,14 +118,63 @@ impl McpTransport for StdioMcpTransport {
request: &McpRequest, request: &McpRequest,
_headers: &HashMap<String, String>, _headers: &HashMap<String, String>,
) -> Result<McpResponse, ToolError> { ) -> Result<McpResponse, ToolError> {
stream_transport_send( // JSON-RPC notifications (no id) are fire-and-forget: the server
&self.stdin, // will not send a response, so we must not wait for one.
&self.pending, if request.id.is_none() {
request, let mut stdin = self.stdin.lock().await;
&self.server_name, write_jsonrpc_line(&mut *stdin, request).await?;
Duration::from_secs(30), return Ok(McpResponse {
) jsonrpc: "2.0".to_string(),
.await id: None,
result: None,
error: None,
});
}
let id = request.id.unwrap_or(0);
let (tx, rx) = oneshot::channel();
// Register the pending response handler before writing the request,
// so we don't miss a fast response from the child.
{
let mut pending = self.pending.lock().await;
pending.insert(id, tx);
}
// Write the request to stdin.
{
let mut stdin = self.stdin.lock().await;
if let Err(e) = write_jsonrpc_line(&mut *stdin, request).await {
// Remove the pending entry on write failure.
let mut pending = self.pending.lock().await;
pending.remove(&id);
return Err(e);
}
}
// Wait for the response with a timeout.
let timeout = Duration::from_secs(30);
match tokio::time::timeout(timeout, rx).await {
Ok(Ok(response)) => Ok(response),
Ok(Err(_)) => {
// Sender was dropped (reader task ended). Clean up pending entry.
let mut pending = self.pending.lock().await;
pending.remove(&id);
Err(ToolError::ExternalService(format!(
"[{}] MCP server closed connection before responding to request {:?}",
self.server_name, request.id
)))
}
Err(_) => {
// Timeout: remove the pending entry.
let mut pending = self.pending.lock().await;
pending.remove(&id);
Err(ToolError::ExternalService(format!(
"[{}] Timeout waiting for response to request {:?} after {:?}",
self.server_name, request.id, timeout
)))
}
}
} }
async fn shutdown(&self) -> Result<(), ToolError> { async fn shutdown(&self) -> Result<(), ToolError> {
+1 -105
View File
@@ -97,13 +97,7 @@ pub fn spawn_jsonrpc_reader<R: AsyncBufRead + Unpin + Send + 'static>(
} }
}; };
let Some(id) = response.id else { let id = response.id.unwrap_or(0);
tracing::debug!(
"[{}] Received JSON-RPC notification (no id), skipping dispatch",
server_name
);
continue;
};
let mut map = pending.lock().await; let mut map = pending.lock().await;
if let Some(tx) = map.remove(&id) { if let Some(tx) = map.remove(&id) {
// Ignore send error — the receiver may have been dropped (timeout). // Ignore send error — the receiver may have been dropped (timeout).
@@ -121,76 +115,6 @@ pub fn spawn_jsonrpc_reader<R: AsyncBufRead + Unpin + Send + 'static>(
}) })
} }
/// Send a JSON-RPC request over a stream-based transport (stdio / unix socket).
///
/// Handles notification fire-and-forget, pending response registration,
/// write, timeout, and cleanup. Used by both [`StdioMcpTransport`] and
/// [`UnixMcpTransport`] to avoid duplicating the send logic.
pub(crate) async fn stream_transport_send<W: AsyncWrite + Unpin>(
writer: &Mutex<W>,
pending: &Mutex<HashMap<u64, oneshot::Sender<McpResponse>>>,
request: &McpRequest,
server_name: &str,
timeout_duration: std::time::Duration,
) -> Result<McpResponse, ToolError> {
// JSON-RPC notifications (no id) are fire-and-forget: the server
// will not send a response, so we must not wait for one.
if request.id.is_none() {
let mut w = writer.lock().await;
write_jsonrpc_line(&mut *w, request).await?;
return Ok(McpResponse {
jsonrpc: "2.0".to_string(),
id: None,
result: None,
error: None,
});
}
let id = request.id.unwrap_or(0);
let (tx, rx) = oneshot::channel();
// Register the pending response handler before writing the request,
// so we don't miss a fast response from the server.
{
let mut map = pending.lock().await;
map.insert(id, tx);
}
// Write the request.
{
let mut w = writer.lock().await;
if let Err(e) = write_jsonrpc_line(&mut *w, request).await {
// Remove the pending entry on write failure.
let mut map = pending.lock().await;
map.remove(&id);
return Err(e);
}
}
// Wait for the response with a timeout.
match tokio::time::timeout(timeout_duration, rx).await {
Ok(Ok(response)) => Ok(response),
Ok(Err(_)) => {
// Sender was dropped (reader task ended). Clean up pending entry.
let mut map = pending.lock().await;
map.remove(&id);
Err(ToolError::ExternalService(format!(
"[{}] MCP server closed connection before responding to request {:?}",
server_name, request.id
)))
}
Err(_) => {
// Timeout: remove the pending entry.
let mut map = pending.lock().await;
map.remove(&id);
Err(ToolError::ExternalService(format!(
"[{}] Timeout waiting for response to request {:?} after {:?}",
server_name, request.id, timeout_duration
)))
}
}
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
@@ -269,32 +193,4 @@ mod tests {
handle.await.expect("reader task should finish"); handle.await.expect("reader task should finish");
} }
/// Issue 9 regression: a JSON-RPC notification (no id) must not resolve
/// a pending request keyed by id 0 (the old `unwrap_or(0)` default).
#[tokio::test]
async fn test_notification_does_not_resolve_pending_id_zero() {
// A notification response (no id), followed by a proper response for id 0.
let notification = r#"{"jsonrpc":"2.0","method":"notifications/progress","params":{}}"#;
let real_response = r#"{"jsonrpc":"2.0","id":0,"result":{"ok":true}}"#;
let input = format!("{notification}\n{real_response}\n");
let reader = std::io::Cursor::new(input.into_bytes());
let pending: Arc<Mutex<HashMap<u64, oneshot::Sender<McpResponse>>>> =
Arc::new(Mutex::new(HashMap::new()));
let (tx, rx) = oneshot::channel();
{
let mut map = pending.lock().await;
map.insert(0, tx);
}
let handle = spawn_jsonrpc_reader(reader, pending.clone(), "test".into());
let resp = rx.await.expect("should receive the real id=0 response");
assert_eq!(resp.id, Some(0));
assert!(resp.result.is_some());
handle.await.expect("reader task should finish");
}
} }
+58 -9
View File
@@ -15,7 +15,7 @@ use tokio::sync::{Mutex, oneshot};
use tokio::task::JoinHandle; use tokio::task::JoinHandle;
use crate::tools::mcp::protocol::{McpRequest, McpResponse}; use crate::tools::mcp::protocol::{McpRequest, McpResponse};
use crate::tools::mcp::transport::{McpTransport, spawn_jsonrpc_reader, stream_transport_send}; use crate::tools::mcp::transport::{McpTransport, spawn_jsonrpc_reader, write_jsonrpc_line};
use crate::tools::tool::ToolError; use crate::tools::tool::ToolError;
/// MCP transport that communicates over a Unix domain socket. /// MCP transport that communicates over a Unix domain socket.
@@ -91,14 +91,63 @@ impl McpTransport for UnixMcpTransport {
request: &McpRequest, request: &McpRequest,
_headers: &HashMap<String, String>, _headers: &HashMap<String, String>,
) -> Result<McpResponse, ToolError> { ) -> Result<McpResponse, ToolError> {
stream_transport_send( // JSON-RPC notifications (no id) are fire-and-forget: the server
&self.writer, // will not send a response, so we must not wait for one.
&self.pending, if request.id.is_none() {
request, let mut writer = self.writer.lock().await;
&self.server_name, write_jsonrpc_line(&mut *writer, request).await?;
Duration::from_secs(30), return Ok(McpResponse {
) jsonrpc: "2.0".to_string(),
.await id: None,
result: None,
error: None,
});
}
let id = request.id.unwrap_or(0);
let (tx, rx) = oneshot::channel();
// Register the pending response handler before writing the request,
// so we don't miss a fast response from the server.
{
let mut pending = self.pending.lock().await;
pending.insert(id, tx);
}
// Write the request to the socket.
{
let mut writer = self.writer.lock().await;
if let Err(e) = write_jsonrpc_line(&mut *writer, request).await {
// Remove the pending entry on write failure.
let mut pending = self.pending.lock().await;
pending.remove(&id);
return Err(e);
}
}
// Wait for the response with a timeout.
let timeout = Duration::from_secs(30);
match tokio::time::timeout(timeout, rx).await {
Ok(Ok(response)) => Ok(response),
Ok(Err(_)) => {
// Sender was dropped (reader task ended). Clean up pending entry.
let mut pending = self.pending.lock().await;
pending.remove(&id);
Err(ToolError::ExternalService(format!(
"[{}] MCP server closed connection before responding to request {:?}",
self.server_name, request.id
)))
}
Err(_) => {
// Timeout: remove the pending entry.
let mut pending = self.pending.lock().await;
pending.remove(&id);
Err(ToolError::ExternalService(format!(
"[{}] Timeout waiting for response to request {:?} after {:?}",
self.server_name, request.id, timeout
)))
}
}
} }
async fn shutdown(&self) -> Result<(), ToolError> { async fn shutdown(&self) -> Result<(), ToolError> {
-12
View File
@@ -75,7 +75,6 @@ const PROTECTED_TOOL_NAMES: &[&str] = &[
"image_generate", "image_generate",
"image_edit", "image_edit",
"image_analyze", "image_analyze",
"tool_info",
]; ];
/// Registry of available tools. /// Registry of available tools.
@@ -246,17 +245,6 @@ impl ToolRegistry {
tracing::debug!("Registered {} built-in tools", self.count()); 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). /// Register only orchestrator-domain tools (safe for the main process).
/// ///
/// This registers tools that don't touch the filesystem or run shell commands: /// This registers tools that don't touch the filesystem or run shell commands:
+53 -2
View File
@@ -558,7 +558,48 @@ mod tests {
// Routine tools // Routine tools
( (
"routine_create", "routine_create",
crate::tools::builtin::routine::routine_create_parameters_schema(), serde_json::json!({
"type": "object",
"properties": {
"name": { "type": "string", "description": "Routine name" },
"description": { "type": "string", "description": "What it does" },
"trigger_type": {
"type": "string",
"enum": ["cron", "event", "system_event", "manual"],
"description": "When the routine fires"
},
"schedule": { "type": "string", "description": "Cron expression" },
"event_pattern": { "type": "string", "description": "Regex pattern" },
"event_channel": { "type": "string", "description": "Channel filter" },
"event_source": { "type": "string", "description": "System event source" },
"event_type": { "type": "string", "description": "System event type" },
"event_filters": {
"type": "object",
"additionalProperties": { "type": "string" },
"description": "Exact-match payload filters"
},
"prompt": { "type": "string", "description": "Instructions" },
"context_paths": {
"type": "array",
"items": { "type": "string" },
"description": "Workspace paths to load"
},
"action_type": {
"type": "string",
"enum": ["lightweight", "full_job"],
"description": "Execution mode"
},
"cooldown_secs": { "type": "integer", "description": "Min seconds between fires" },
"tool_permissions": {
"type": "array",
"items": { "type": "string" },
"description": "Pre-authorized tools for full_job mode"
},
"notify_channel": { "type": "string", "description": "Channel for message tool" },
"notify_user": { "type": "string", "description": "User/target to notify" }
},
"required": ["name", "trigger_type", "prompt"]
}),
), ),
( (
"routine_list", "routine_list",
@@ -570,7 +611,17 @@ mod tests {
), ),
( (
"routine_update", "routine_update",
crate::tools::builtin::routine::routine_update_parameters_schema(), serde_json::json!({
"type": "object",
"properties": {
"name": { "type": "string", "description": "Name" },
"enabled": { "type": "boolean", "description": "Toggle" },
"prompt": { "type": "string", "description": "New prompt" },
"schedule": { "type": "string", "description": "New cron schedule" },
"description": { "type": "string", "description": "New description" }
},
"required": ["name"]
}),
), ),
( (
"routine_delete", "routine_delete",
+3 -59
View File
@@ -336,17 +336,6 @@ pub trait Tool: Send + Sync {
None 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. /// Get the tool schema for LLM function calling.
fn schema(&self) -> ToolSchema { fn schema(&self) -> ToolSchema {
ToolSchema { ToolSchema {
@@ -430,24 +419,9 @@ pub fn redact_params(params: &serde_json::Value, sensitive: &[&str]) -> serde_js
/// Properties without a `"type"` field are allowed (freeform/any-type). /// Properties without a `"type"` field are allowed (freeform/any-type).
/// This is an intentional pattern used by tools like `json` and `http` for /// This is an intentional pattern used by tools like `json` and `http` for
/// OpenAI compatibility, since union types with arrays require `items`. /// OpenAI compatibility, since union types with arrays require `items`.
/// Maximum nesting depth for tool schema validation to prevent stack overflow
/// on maliciously crafted schemas.
const MAX_SCHEMA_DEPTH: usize = 16;
pub fn validate_tool_schema(schema: &serde_json::Value, path: &str) -> Vec<String> { pub fn validate_tool_schema(schema: &serde_json::Value, path: &str) -> Vec<String> {
validate_tool_schema_inner(schema, path, 0)
}
fn validate_tool_schema_inner(schema: &serde_json::Value, path: &str, depth: usize) -> Vec<String> {
let mut errors = Vec::new(); let mut errors = Vec::new();
if depth > MAX_SCHEMA_DEPTH {
errors.push(format!(
"{path}: schema nesting exceeds maximum depth of {MAX_SCHEMA_DEPTH}"
));
return errors;
}
// Rule 1: must have "type": "object" at this level // Rule 1: must have "type": "object" at this level
match schema.get("type").and_then(|t| t.as_str()) { match schema.get("type").and_then(|t| t.as_str()) {
Some("object") => {} Some("object") => {}
@@ -489,17 +463,14 @@ fn validate_tool_schema_inner(schema: &serde_json::Value, path: &str, depth: usi
if let Some(prop_type) = prop.get("type").and_then(|t| t.as_str()) { if let Some(prop_type) = prop.get("type").and_then(|t| t.as_str()) {
match prop_type { match prop_type {
"object" => { "object" => {
errors.extend(validate_tool_schema_inner(prop, &prop_path, depth + 1)); errors.extend(validate_tool_schema(prop, &prop_path));
} }
"array" => { "array" => {
if let Some(items) = prop.get("items") { if let Some(items) = prop.get("items") {
// If items is an object type, recurse // If items is an object type, recurse
if items.get("type").and_then(|t| t.as_str()) == Some("object") { if items.get("type").and_then(|t| t.as_str()) == Some("object") {
errors.extend(validate_tool_schema_inner( errors
items, .extend(validate_tool_schema(items, &format!("{prop_path}.items")));
&format!("{prop_path}.items"),
depth + 1,
));
} }
} else { } else {
errors.push(format!("{prop_path}: array property missing \"items\"")); errors.push(format!("{prop_path}: array property missing \"items\""));
@@ -828,33 +799,6 @@ mod tests {
assert!(errors[0].contains("\"missing_field\"")); assert!(errors[0].contains("\"missing_field\""));
} }
/// Regression test for issue #975: deeply nested schemas must not cause
/// stack overflow. The validator should stop at MAX_SCHEMA_DEPTH and
/// report an error instead of recursing infinitely.
#[test]
fn test_validate_schema_depth_limit() {
// Build a schema nested 20 levels deep (exceeds MAX_SCHEMA_DEPTH=16)
let mut schema = serde_json::json!({
"type": "object",
"properties": {
"leaf": { "type": "string" }
}
});
for _ in 0..20 {
schema = serde_json::json!({
"type": "object",
"properties": {
"nested": schema
}
});
}
let errors = validate_tool_schema(&schema, "test");
assert!(
errors.iter().any(|e| e.contains("maximum depth")),
"expected depth limit error, got: {errors:?}"
);
}
#[test] #[test]
fn test_approval_context_autonomous_allows_unless_auto_approved() { fn test_approval_context_autonomous_allows_unless_auto_approved() {
let ctx = ApprovalContext::autonomous(); let ctx = ApprovalContext::autonomous();
+4 -114
View File
@@ -101,75 +101,24 @@ pub struct CapabilitiesFile {
pub capabilities: Option<Box<CapabilitiesFile>>, pub capabilities: Option<Box<CapabilitiesFile>>,
} }
/// Maximum length for the description field to prevent memory abuse.
const MAX_DESCRIPTION_CHARS: usize = 4096;
/// Maximum serialized size of the parameters schema JSON.
const MAX_PARAMETERS_SCHEMA_BYTES: usize = 64 * 1024;
impl CapabilitiesFile { impl CapabilitiesFile {
/// Parse from JSON string. /// Parse from JSON string.
pub fn from_json(json: &str) -> Result<Self, serde_json::Error> { pub fn from_json(json: &str) -> Result<Self, serde_json::Error> {
let mut caps = serde_json::from_str::<Self>(json).map(Self::resolve_nested)?; serde_json::from_str::<Self>(json).map(Self::resolve_nested)
caps.enforce_limits();
Ok(caps)
} }
/// Parse from JSON bytes. /// Parse from JSON bytes.
pub fn from_bytes(bytes: &[u8]) -> Result<Self, serde_json::Error> { pub fn from_bytes(bytes: &[u8]) -> Result<Self, serde_json::Error> {
let mut caps = serde_json::from_slice::<Self>(bytes).map(Self::resolve_nested)?; serde_json::from_slice::<Self>(bytes).map(Self::resolve_nested)
caps.enforce_limits();
Ok(caps)
}
/// Truncate oversized fields to prevent unbounded memory usage.
fn enforce_limits(&mut self) {
// Truncate oversized description (issue #976)
if let Some(ref desc) = self.description
&& desc.len() > MAX_DESCRIPTION_CHARS
{
let truncated = &desc[..desc.floor_char_boundary(MAX_DESCRIPTION_CHARS)];
tracing::warn!(
"Capabilities description truncated from {} to {} chars",
desc.len(),
MAX_DESCRIPTION_CHARS,
);
self.description = Some(truncated.to_string());
}
// Drop oversized parameters schema (issue #977)
if let Some(ref params) = self.parameters {
let size = params.to_string().len();
if size > MAX_PARAMETERS_SCHEMA_BYTES {
tracing::warn!(
"Capabilities parameters schema dropped ({} bytes exceeds {} limit)",
size,
MAX_PARAMETERS_SCHEMA_BYTES,
);
self.parameters = None;
}
}
} }
/// Merge nested `capabilities` wrapper into top-level fields. /// Merge nested `capabilities` wrapper into top-level fields.
/// ///
/// Channel-level JSON nests tool capabilities under `"capabilities"`. /// Channel-level JSON nests tool capabilities under `"capabilities"`.
/// This promotes the inner fields so callers can access them uniformly. /// This promotes the inner fields so callers can access them uniformly.
/// Maximum nesting depth for capabilities resolution. fn resolve_nested(mut self) -> Self {
const MAX_NESTED_DEPTH: usize = 8;
fn resolve_nested(self) -> Self {
self.resolve_nested_inner(0)
}
fn resolve_nested_inner(mut self, depth: usize) -> Self {
if depth > Self::MAX_NESTED_DEPTH {
tracing::warn!(
"Capabilities nesting exceeds maximum depth of {}, stopping resolution",
Self::MAX_NESTED_DEPTH
);
return self;
}
if let Some(inner) = self.capabilities.take() { if let Some(inner) = self.capabilities.take() {
let inner = inner.resolve_nested_inner(depth + 1); let inner = inner.resolve_nested();
self.description = self.description.or(inner.description); self.description = self.description.or(inner.description);
self.parameters = self.parameters.or(inner.parameters); self.parameters = self.parameters.or(inner.parameters);
self.http = self.http.or(inner.http); self.http = self.http.or(inner.http);
@@ -1434,63 +1383,4 @@ mod tests {
"Outer description should take precedence over inner" "Outer description should take precedence over inner"
); );
} }
/// Regression test for issue #974: deeply nested capabilities wrappers
/// must not cause stack overflow. resolve_nested should stop at
/// MAX_NESTED_DEPTH and return gracefully.
#[test]
fn test_resolve_nested_depth_limit() {
// Build a capabilities file nested beyond MAX_NESTED_DEPTH (8).
// The description is at the innermost level which is beyond the limit,
// so it won't be resolved — the key assertion is no stack overflow.
let mut json = r#"{ "description": "leaf" }"#.to_string();
for _ in 0..20 {
json = format!(r#"{{ "capabilities": {json} }}"#);
}
// Should not stack overflow — this is the primary assertion.
let _caps = CapabilitiesFile::from_json(&json).unwrap();
}
/// Regression test for issue #976: oversized description strings are truncated.
#[test]
fn test_description_truncated_at_limit() {
let long_desc = "x".repeat(10_000);
let json = format!(r#"{{ "description": "{long_desc}" }}"#);
let caps = CapabilitiesFile::from_json(&json).unwrap();
let desc = caps.description.unwrap();
assert!(
desc.len() <= super::MAX_DESCRIPTION_CHARS + 50, // allow for minor overhead
"description should be truncated to ~{} chars, got {}",
super::MAX_DESCRIPTION_CHARS,
desc.len()
);
}
/// Regression test for issue #977: oversized parameters schema is dropped.
#[test]
fn test_oversized_parameters_schema_dropped() {
// Build a parameters schema larger than MAX_PARAMETERS_SCHEMA_BYTES
let mut properties = serde_json::Map::new();
for i in 0..2000 {
properties.insert(
format!("field_{i}"),
serde_json::json!({
"type": "string",
"description": "x".repeat(50)
}),
);
}
let schema = serde_json::json!({
"type": "object",
"properties": properties,
});
let json = serde_json::json!({
"parameters": schema,
});
let caps = CapabilitiesFile::from_json(&json.to_string()).unwrap();
assert!(
caps.parameters.is_none(),
"oversized parameters schema should be dropped"
);
}
} }
+84 -6
View File
@@ -1,5 +1,7 @@
//! WASM sandbox error types. //! WASM sandbox error types.
use std::fmt;
use thiserror::Error; use thiserror::Error;
/// Errors that can occur during WASM tool execution. /// Errors that can occur during WASM tool execution.
@@ -66,13 +68,13 @@ pub enum WasmError {
Timeout(std::time::Duration), Timeout(std::time::Duration),
/// Component returned an error response. /// Component returned an error response.
/// When `hint` is non-empty it points the LLM to `tool_info` so it can /// When `hint` is non-empty it carries the tool's description and parameter
/// fetch the tool's full parameter schema on demand. /// 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}") })] #[error("Tool error: {message}{}", if hint.is_empty() { String::new() } else { format!("\n\nTool usage hint:\n{hint}") })]
ToolReturnedError { ToolReturnedError {
/// The error message from the WASM tool. /// The error message from the WASM tool.
message: String, message: String,
/// Optional retry hint (empty when unavailable). /// Optional description + schema hint (empty when unavailable).
hint: String, 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)] #[cfg(test)]
mod tests { mod tests {
use crate::tools::wasm::error::WasmError; use crate::tools::wasm::error::{TrapCode, TrapInfo, WasmError};
#[test] #[test]
fn test_error_display() { fn test_error_display() {
@@ -114,6 +180,17 @@ mod tests {
assert!(err.to_string().contains("10000000")); 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] #[test]
fn test_conversion_to_tool_error() { fn test_conversion_to_tool_error() {
let wasm_err = WasmError::Trapped("test trap".to_string()); let wasm_err = WasmError::Trapped("test trap".to_string());
@@ -141,11 +218,12 @@ mod tests {
fn test_tool_returned_error_with_hint() { fn test_tool_returned_error_with_hint() {
let err = WasmError::ToolReturnedError { let err = WasmError::ToolReturnedError {
message: "unknown action: foobar".to_string(), 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(); let display = err.to_string();
assert!(display.contains("unknown action: foobar")); assert!(display.contains("unknown action: foobar"));
assert!(display.contains("Tool usage hint")); 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, memory_used: u64,
/// Maximum tables allowed. /// Maximum tables allowed.
max_tables: u32, max_tables: u32,
/// Current table count.
#[allow(dead_code)] // Reserved for table limit enforcement
tables_created: u32,
/// Maximum instances allowed. /// Maximum instances allowed.
max_instances: u32, max_instances: u32,
/// Current instance count.
#[allow(dead_code)] // Reserved for instance limit enforcement
instances_created: u32,
} }
impl WasmResourceLimiter { impl WasmResourceLimiter {
@@ -81,7 +87,9 @@ impl WasmResourceLimiter {
memory_limit, memory_limit,
memory_used: 0, memory_used: 0,
max_tables: 10, max_tables: 10,
tables_created: 0,
max_instances: 10, // Component model needs multiple instances for WASI 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; mod wrapper;
// Core types // Core types
pub use error::WasmError; pub use error::{TrapCode, TrapInfo, WasmError};
pub use host::{HostState, LogEntry, LogLevel}; pub use host::{HostState, LogEntry, LogLevel};
pub use limits::{ pub use limits::{
DEFAULT_FUEL_LIMIT, DEFAULT_MEMORY_LIMIT, DEFAULT_TIMEOUT, FuelConfig, ResourceLimits, DEFAULT_FUEL_LIMIT, DEFAULT_MEMORY_LIMIT, DEFAULT_TIMEOUT, FuelConfig, ResourceLimits,
+41 -26
View File
@@ -123,9 +123,7 @@ pub struct PreparedModule {
pub name: String, pub name: String,
/// Tool description (cached from component). /// Tool description (cached from component).
pub description: String, pub description: String,
/// Full parameter schema JSON extracted from the component. /// Parameter schema JSON (cached from component).
/// Used for discovery and coercion, not necessarily for the compact
/// schema advertised in the main tools array.
pub schema: serde_json::Value, pub schema: serde_json::Value,
/// Pre-compiled component (cheaply cloneable via internal Arc). /// Pre-compiled component (cheaply cloneable via internal Arc).
component: wasmtime::component::Component, component: wasmtime::component::Component,
@@ -267,29 +265,11 @@ impl WasmToolRuntime {
let component = wasmtime::component::Component::new(&engine, &wasm_bytes) let component = wasmtime::component::Component::new(&engine, &wasm_bytes)
.map_err(|e| WasmError::CompilationFailed(e.to_string()))?; .map_err(|e| WasmError::CompilationFailed(e.to_string()))?;
// Briefly instantiate to extract metadata (description + schema) // We need to instantiate briefly to extract metadata.
// from the tool's exports, analogous to MCP's list_tools(). // In a full implementation, we'd use WIT bindgen to get typed access.
let effective_limits = limits.clone().unwrap_or(default_limits.clone()); // For now, we extract what we can from the component.
let (description, schema) = crate::tools::wasm::wrapper::extract_wasm_metadata( let description = extract_tool_description(&engine, &component)?;
&engine, let schema = extract_tool_schema(&engine, &component)?;
&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
}),
)
});
Ok::<_, WasmError>(PreparedModule { Ok::<_, WasmError>(PreparedModule {
name: name.clone(), 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 { impl std::fmt::Debug for WasmToolRuntime {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("WasmToolRuntime") f.debug_struct("WasmToolRuntime")
+58 -468
View File
@@ -464,10 +464,9 @@ pub struct WasmToolWrapper {
/// Capabilities to grant to this tool. /// Capabilities to grant to this tool.
capabilities: Capabilities, capabilities: Capabilities,
/// Cached description (from PreparedModule or override). /// Cached description (from PreparedModule or override).
/// Stored without any tool_info hints — hints are composed at display time.
description: String, description: String,
/// Compact and discovery schemas for this tool. /// Cached schema (from PreparedModule or override).
schemas: WasmToolSchemas, schema: serde_json::Value,
/// Injected credentials for HTTP requests (e.g., OAuth tokens). /// Injected credentials for HTTP requests (e.g., OAuth tokens).
/// Keys are placeholder names like "GOOGLE_ACCESS_TOKEN". /// Keys are placeholder names like "GOOGLE_ACCESS_TOKEN".
credentials: HashMap<String, String>, credentials: HashMap<String, String>,
@@ -478,84 +477,6 @@ pub struct WasmToolWrapper {
oauth_refresh: Option<OAuthRefreshConfig>, 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 { impl WasmToolWrapper {
/// Create a new WASM tool wrapper. /// Create a new WASM tool wrapper.
pub fn new( pub fn new(
@@ -565,7 +486,7 @@ impl WasmToolWrapper {
) -> Self { ) -> Self {
Self { Self {
description: prepared.description.clone(), description: prepared.description.clone(),
schemas: WasmToolSchemas::new(prepared.schema.clone()), schema: prepared.schema.clone(),
runtime, runtime,
prepared, prepared,
capabilities, capabilities,
@@ -583,7 +504,7 @@ impl WasmToolWrapper {
/// Override the parameter schema. /// Override the parameter schema.
pub fn with_schema(mut self, schema: serde_json::Value) -> Self { pub fn with_schema(mut self, schema: serde_json::Value) -> Self {
self.schemas = self.schemas.with_override(schema); self.schema = schema;
self 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. // Coerce string-encoded values to their schema-declared types.
// LLMs frequently pass numeric values as strings (e.g. "5" instead of 5). // 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 // Prepare the request
let params_json = serde_json::to_string(&params) let params_json = serde_json::to_string(&params)
@@ -717,6 +629,7 @@ impl WasmToolWrapper {
}; };
// Call execute using the generated typed interface // 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 response = tool_iface.call_execute(&mut store, &request).map_err(|e| {
let error_str = e.to_string(); let error_str = e.to_string();
if error_str.contains("out of fuel") { if error_str.contains("out of fuel") {
@@ -731,13 +644,12 @@ impl WasmToolWrapper {
// Get logs from host state // Get logs from host state
let logs = store.data_mut().host_state.take_logs(); let logs = store.data_mut().host_state.take_logs();
// Check for tool-level error — point the LLM to tool_info for the // Check for tool-level error — on failure, call the WASM module's
// full schema instead of dumping ~3.5KB inline. // 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 { if let Some(err) = response.error {
let hint = format!( let hint = build_tool_hint(tool_iface, &mut store);
"Tip: call tool_info(name: \"{}\", include_schema: true) for the full parameter schema.",
self.prepared.name
);
return Err(WasmError::ToolReturnedError { message: err, hint }); return Err(WasmError::ToolReturnedError { message: err, hint });
} }
@@ -746,55 +658,47 @@ impl WasmToolWrapper {
} }
} }
/// Extract metadata (description + schema) from a WASM tool by briefly /// Maximum characters for the description portion of a tool hint.
/// instantiating it and calling its `description()` and `schema()` exports. const HINT_DESC_MAX: usize = 500;
/// Analogous to MCP's `list_tools()` — discovers tool capabilities at load time. /// Maximum characters for the schema portion of a tool hint.
/// const HINT_SCHEMA_MAX: usize = 3000;
/// 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);
// Configure fuel + epoch deadline so extraction can't hang /// Call the WASM module's `description()` and `schema()` exports to build a
if let Err(e) = store.set_fuel(limits.fuel) { /// hint string. Returns an empty string if both calls fail or return empty.
tracing::debug!("Fuel not enabled for metadata extraction: {e}"); /// Description is capped at [`HINT_DESC_MAX`] chars, schema at
} /// [`HINT_SCHEMA_MAX`] chars.
store.epoch_deadline_trap(); fn build_tool_hint(tool_iface: &wit_tool::Guest, store: &mut Store<StoreData>) -> String {
let ticks = (limits.timeout.as_millis() / EPOCH_TICK_INTERVAL.as_millis()).max(1) as u64; let desc = tool_iface
store.set_epoch_deadline(ticks); .call_description(&mut *store)
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)
.ok() .ok()
.and_then(|s| serde_json::from_str::<serde_json::Value>(&s).ok()) .unwrap_or_default();
.unwrap_or_else(|| { let schema = tool_iface.call_schema(&mut *store).ok().unwrap_or_default();
serde_json::json!({"type": "object", "properties": {}, "additionalProperties": true}) if desc.is_empty() && schema.is_empty() {
}); return String::new();
}
Ok((description, schema)) 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] #[async_trait]
@@ -808,33 +712,7 @@ impl Tool for WasmToolWrapper {
} }
fn parameters_schema(&self) -> serde_json::Value { fn parameters_schema(&self) -> serde_json::Value {
self.schemas.advertised() self.schema.clone()
}
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(),
}
} }
async fn execute( async fn execute(
@@ -871,7 +749,7 @@ impl Tool for WasmToolWrapper {
let prepared = Arc::clone(&self.prepared); let prepared = Arc::clone(&self.prepared);
let capabilities = self.capabilities.clone(); let capabilities = self.capabilities.clone();
let description = self.description.clone(); let description = self.description.clone();
let schemas = self.schemas.clone(); let schema = self.schema.clone();
let credentials = self.credentials.clone(); let credentials = self.credentials.clone();
// Execute in blocking task with timeout // Execute in blocking task with timeout
@@ -881,7 +759,7 @@ impl Tool for WasmToolWrapper {
prepared, prepared,
capabilities, capabilities,
description, description,
schemas, schema,
credentials, credentials,
secrets_store: None, // Not needed in blocking task secrets_store: None, // Not needed in blocking task
oauth_refresh: None, // Already used above for pre-refresh oauth_refresh: None, // Already used above for pre-refresh
@@ -1104,18 +982,7 @@ async fn resolve_host_credentials(
) -> Vec<ResolvedHostCredential> { ) -> Vec<ResolvedHostCredential> {
let store = match store { let store = match store {
Some(s) => s, Some(s) => s,
None => { None => return Vec::new(),
// If tool requires credentials but has no secrets store, this is a configuration error
if let Some(http_cap) = &capabilities.http
&& !http_cap.credentials.is_empty()
{
tracing::warn!(
user_id = %user_id,
"WASM tool requires credentials but secrets_store is not configured - authentication will fail"
);
}
return Vec::new();
}
}; };
// Check if the access token needs refreshing before resolving credentials. // Check if the access token needs refreshing before resolving credentials.
@@ -1166,37 +1033,13 @@ async fn resolve_host_credentials(
continue; continue;
} }
// Try to get credential under the provided user_id first.
// If not found and user_id != "default", fallback to "default" (global credentials).
// This handles OAuth tokens stored globally under "default" but accessed from routine contexts.
let secret = match store.get_decrypted(user_id, &mapping.secret_name).await { let secret = match store.get_decrypted(user_id, &mapping.secret_name).await {
Ok(s) => Some(s), Ok(s) => s,
Err(e) => { Err(e) => {
// If lookup fails and we're not already looking up "default", try "default" as fallback tracing::debug!(
if user_id != "default" {
tracing::debug!(
secret_name = %mapping.secret_name,
user_id = %user_id,
error = %e,
"Credential not found for user, trying default global credentials"
);
store
.get_decrypted("default", &mapping.secret_name)
.await
.ok()
} else {
None
}
}
};
let secret = match secret {
Some(s) => s,
None => {
tracing::warn!(
secret_name = %mapping.secret_name, secret_name = %mapping.secret_name,
user_id = %user_id, error = %e,
"Could not resolve credential for WASM tool (not found in user context or default)" "Could not resolve credential for WASM tool (auth may not be configured)"
); );
continue; continue;
} }
@@ -1389,7 +1232,6 @@ mod tests {
TEST_GOOGLE_OAUTH_TOKEN, TEST_OAUTH_CLIENT_ID, TEST_OAUTH_CLIENT_SECRET, TEST_GOOGLE_OAUTH_TOKEN, TEST_OAUTH_CLIENT_ID, TEST_OAUTH_CLIENT_SECRET,
test_secrets_store, test_secrets_store,
}; };
use crate::tools::tool::Tool;
use crate::tools::wasm::capabilities::Capabilities; use crate::tools::wasm::capabilities::Capabilities;
use crate::tools::wasm::runtime::{WasmRuntimeConfig, WasmToolRuntime}; use crate::tools::wasm::runtime::{WasmRuntimeConfig, WasmToolRuntime};
@@ -1404,84 +1246,6 @@ mod tests {
assert!(runtime.config().fuel_config.enabled); 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] #[test]
fn test_capabilities_default() { fn test_capabilities_default() {
let caps = Capabilities::default(); let caps = Capabilities::default();
@@ -2024,23 +1788,6 @@ mod tests {
assert_eq!(result["count"], serde_json::json!("not-a-number")); 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 /// Regression test: leak scan must run on raw headers (before credential
/// injection), not after. If it ran post-injection, the host-injected /// injection), not after. If it ran post-injection, the host-injected
/// Slack bot token (`xoxb-...`) would trigger a Block and reject the /// Slack bot token (`xoxb-...`) would trigger a Block and reject the
@@ -2093,161 +1840,4 @@ mod tests {
"Leak scan on post-injection headers should block the Slack token" "Leak scan on post-injection headers should block the Slack token"
); );
} }
#[tokio::test]
async fn test_resolve_host_credentials_fallback_to_default_user() {
use crate::secrets::{CredentialLocation, CredentialMapping, SecretsStore};
use crate::tools::wasm::capabilities::HttpCapability;
use crate::tools::wasm::wrapper::resolve_host_credentials;
let store = test_secrets_store();
// Store a token under the "default" global user
store
.create(
"default",
crate::secrets::CreateSecretParams::new("google_oauth_token", "global_token_value"),
)
.await
.expect("Failed to store global token"); // safety: test code only
// Create capabilities requiring this credential
let mut creds = std::collections::HashMap::new();
creds.insert(
"google_oauth_token".to_string(),
CredentialMapping {
secret_name: "google_oauth_token".to_string(),
location: CredentialLocation::AuthorizationBearer,
host_patterns: vec!["sheets.googleapis.com".to_string()],
},
);
let caps = Capabilities {
http: Some(HttpCapability {
allowlist: vec![],
credentials: creds,
rate_limit: crate::tools::wasm::capabilities::RateLimitConfig::default(),
max_request_bytes: 1024 * 1024,
max_response_bytes: 10 * 1024 * 1024,
timeout: std::time::Duration::from_secs(30),
}),
..Default::default()
};
// Resolve credentials for a different user (routine context)
// Should fallback to "default" and find the token
let result = resolve_host_credentials(&caps, Some(&store), "routine_user_123", None).await;
assert!(!result.is_empty(), "fallback to default"); // safety: test code only
assert_eq!(result[0].secret_value, "global_token_value"); // safety: test code only
}
fn test_capabilities_with_google_oauth() -> Capabilities {
use crate::secrets::{CredentialLocation, CredentialMapping};
use crate::tools::wasm::capabilities::HttpCapability;
let mut creds = std::collections::HashMap::new();
creds.insert(
"google_oauth_token".to_string(),
CredentialMapping {
secret_name: "google_oauth_token".to_string(),
location: CredentialLocation::AuthorizationBearer,
host_patterns: vec!["sheets.googleapis.com".to_string()],
},
);
Capabilities {
http: Some(HttpCapability {
allowlist: vec![],
credentials: creds,
rate_limit: crate::tools::wasm::capabilities::RateLimitConfig::default(),
max_request_bytes: 1024 * 1024,
max_response_bytes: 10 * 1024 * 1024,
timeout: std::time::Duration::from_secs(30),
}),
..Default::default()
}
}
#[tokio::test]
async fn test_resolve_host_credentials_prefers_user_specific_over_default() {
use crate::secrets::SecretsStore;
use crate::tools::wasm::wrapper::resolve_host_credentials;
let store = test_secrets_store();
// Store token under "default" (global)
store
.create(
"default",
crate::secrets::CreateSecretParams::new("google_oauth_token", "global_token"),
)
.await
.expect("Failed to store global token"); // safety: test code only
// Store token under user_123 (user-specific)
store
.create(
"user_123",
crate::secrets::CreateSecretParams::new(
"google_oauth_token",
"user_specific_token",
),
)
.await
.expect("Failed to store user token"); // safety: test code only
// Create capabilities
let caps = test_capabilities_with_google_oauth();
// Resolve credentials for user_123
// Should prefer user_123's token over default
let result = resolve_host_credentials(&caps, Some(&store), "user_123", None).await;
assert!(!result.is_empty(), "has user credentials"); // safety: test code only
assert_eq!(result[0].secret_value, "user_specific_token", "user token"); // safety: test code only
}
#[tokio::test]
async fn test_resolve_host_credentials_no_fallback_when_already_default() {
use crate::secrets::SecretsStore;
use crate::tools::wasm::wrapper::resolve_host_credentials;
let store = test_secrets_store();
// Only store token under "default" (not a duplicate)
store
.create(
"default",
crate::secrets::CreateSecretParams::new("google_oauth_token", "default_token"),
)
.await
.expect("Failed to store default token"); // safety: test code only
// Create capabilities
let caps = test_capabilities_with_google_oauth();
// Resolve credentials for "default" user
// Should NOT attempt fallback (already looking up default)
let result = resolve_host_credentials(&caps, Some(&store), "default", None).await;
assert!(!result.is_empty(), "Should find default token"); // safety: test code only
assert_eq!(result[0].secret_value, "default_token"); // safety: test code only
}
#[tokio::test]
async fn test_resolve_host_credentials_missing_secret_warns() {
use crate::tools::wasm::wrapper::resolve_host_credentials;
let store = test_secrets_store();
// Don't store any token
// Create capabilities expecting a credential
let caps = test_capabilities_with_google_oauth();
// Resolve credentials when neither user nor default has the token
let result = resolve_host_credentials(&caps, Some(&store), "user_456", None).await;
// Should return empty since credential can't be found anywhere
assert!(result.is_empty(), "no credentials found"); // safety: test code only
}
} }

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