mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-25 14:53:34 +00:00
Compare commits
41
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ece6f23957 | ||
|
|
bc58e07ec5 | ||
|
|
c2b876c333 | ||
|
|
1e00b1fed5 | ||
|
+12 |
5e7758598f | ||
|
|
c47237b9c7 | ||
|
|
a89cf37993 | ||
|
|
3c619b6272 | ||
|
|
15c5d3e2e2 | ||
|
|
cd1245afc0 | ||
|
|
9fbdd42988 | ||
|
|
8a60fa2d37 | ||
|
|
c7dec64b2d | ||
|
|
c54f739354 | ||
|
|
1ba6a83ca4 | ||
|
|
d8bcfe15cf | ||
|
|
6f00490900 | ||
|
|
e522d33a53 | ||
|
|
7a9cbb3b50 | ||
|
|
442a42d996 | ||
|
|
d5828b271d | ||
|
|
e1691a8d42 | ||
|
|
8ac24e775b | ||
|
|
bcda73c2e0 | ||
|
|
5dfa666691 | ||
|
|
fd574b2859 | ||
|
|
c592c50dad | ||
|
|
0b122cb28f | ||
|
|
c94ecf19db | ||
|
|
8df51c04ae | ||
|
|
6bbf87ba3a | ||
|
|
006c15e79c | ||
|
|
d420abfa6a | ||
|
|
863702a87a | ||
|
+7 |
f776d96395 | ||
|
|
4faf81ab61 | ||
|
|
8a26cfae73 | ||
|
|
0b81342b5c | ||
|
|
c26f116a98 | ||
|
|
ef34943c14 | ||
|
|
c937dfa315 |
@@ -70,6 +70,12 @@ NEARAI_AUTH_URL=https://private.near.ai
|
||||
# LLM_BASE_URL=https://api.fireworks.ai/inference/v1
|
||||
# LLM_API_KEY=fw_...
|
||||
|
||||
# === MiniMax ===
|
||||
# LLM_BACKEND=minimax
|
||||
# MINIMAX_API_KEY=...
|
||||
# MINIMAX_MODEL=MiniMax-M2.5
|
||||
# MINIMAX_BASE_URL=https://api.minimax.io/v1 # default (global); use https://api.minimaxi.com/v1 for China
|
||||
|
||||
# === Anthropic Direct ===
|
||||
# LLM_BACKEND=anthropic
|
||||
# ANTHROPIC_MODEL=claude-sonnet-4-6
|
||||
|
||||
Executable
+23
@@ -0,0 +1,23 @@
|
||||
#!/usr/bin/env bash
|
||||
set -euo pipefail
|
||||
|
||||
# Pre-push hook: run clippy and tests before pushing.
|
||||
# Install: git config core.hooksPath .githooks
|
||||
|
||||
echo "pre-push: running clippy..."
|
||||
if ! cargo clippy --all --benches --tests --examples --all-features -- -D warnings; then
|
||||
echo ""
|
||||
echo "Push blocked: clippy warnings found."
|
||||
echo "To bypass: git push --no-verify"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "pre-push: running tests..."
|
||||
if ! cargo test; then
|
||||
echo ""
|
||||
echo "Push blocked: tests failed."
|
||||
echo "To bypass: git push --no-verify"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "pre-push: all checks passed."
|
||||
@@ -0,0 +1,59 @@
|
||||
#!/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
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
#!/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
|
||||
@@ -0,0 +1,53 @@
|
||||
#!/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
|
||||
@@ -48,11 +48,11 @@ jobs:
|
||||
matrix:
|
||||
include:
|
||||
- group: core
|
||||
files: "tests/e2e/scenarios/test_connection.py tests/e2e/scenarios/test_chat.py tests/e2e/scenarios/test_sse_reconnect.py tests/e2e/scenarios/test_html_injection.py"
|
||||
files: "tests/e2e/scenarios/test_connection.py tests/e2e/scenarios/test_chat.py tests/e2e/scenarios/test_sse_reconnect.py tests/e2e/scenarios/test_html_injection.py tests/e2e/scenarios/test_csp.py"
|
||||
- group: features
|
||||
files: "tests/e2e/scenarios/test_skills.py tests/e2e/scenarios/test_tool_approval.py"
|
||||
- group: extensions
|
||||
files: "tests/e2e/scenarios/test_extensions.py"
|
||||
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"
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
|
||||
|
||||
@@ -13,6 +13,11 @@ jobs:
|
||||
with:
|
||||
fetch-depth: 0
|
||||
|
||||
- name: Fetch PR head and base
|
||||
run: |
|
||||
git fetch origin ${{ github.event.pull_request.base.ref }}
|
||||
git fetch origin pull/${{ github.event.pull_request.number }}/head:pr-head
|
||||
|
||||
- name: Check for regression tests
|
||||
env:
|
||||
PR_TITLE: ${{ github.event.pull_request.title }}
|
||||
@@ -21,6 +26,8 @@ jobs:
|
||||
set -euo pipefail
|
||||
|
||||
BASE_REF="origin/${{ github.event.pull_request.base.ref }}"
|
||||
# Use the actual PR head, not the merge commit that actions/checkout checks out
|
||||
HEAD_REF="pr-head"
|
||||
|
||||
# --- 1. Is this a fix PR? Check title first, then commit messages ---
|
||||
IS_FIX=false
|
||||
@@ -30,7 +37,7 @@ jobs:
|
||||
fi
|
||||
|
||||
if [ "$IS_FIX" = false ]; then
|
||||
COMMITS=$(git log --format='%s' "${BASE_REF}..HEAD")
|
||||
COMMITS=$(git log --format='%s' "${BASE_REF}..${HEAD_REF}")
|
||||
if grep -qiE '^(fix(\(.*\))?|hotfix|bugfix):' <<< "$COMMITS"; then
|
||||
IS_FIX=true
|
||||
fi
|
||||
@@ -49,14 +56,14 @@ jobs:
|
||||
exit 0
|
||||
fi
|
||||
|
||||
COMMIT_BODIES=$(git log --format='%B' "${BASE_REF}..HEAD")
|
||||
COMMIT_BODIES=$(git log --format='%B' "${BASE_REF}..${HEAD_REF}")
|
||||
if grep -qF '[skip-regression-check]' <<< "$COMMIT_BODIES"; then
|
||||
echo "[skip-regression-check] found in commit message — skipping."
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# --- 3. Exempt static-only / docs-only changes ---
|
||||
CHANGED_FILES=$(git diff --name-only "${BASE_REF}...HEAD")
|
||||
CHANGED_FILES=$(git diff --name-only "${BASE_REF}...${HEAD_REF}")
|
||||
|
||||
if [ -z "$CHANGED_FILES" ]; then
|
||||
echo "No changed files — skipping."
|
||||
@@ -80,13 +87,13 @@ jobs:
|
||||
# --- 4. Look for test changes ---
|
||||
|
||||
# Fast path: new test attributes or test modules in added lines.
|
||||
if git diff "${BASE_REF}...HEAD" -U0 -- '*.rs' | grep -qE '^\+.*(#\[test\]|#\[tokio::test\]|#\[cfg\(test\)\]|mod tests)'; then
|
||||
if git diff "${BASE_REF}...${HEAD_REF}" -U0 -- '*.rs' | grep -qE '^\+.*(#\[test\]|#\[tokio::test\]|#\[cfg\(test\)\]|mod tests)'; then
|
||||
echo "Test changes found in .rs files."
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# Whole-function context: detect edits inside existing test functions.
|
||||
if git diff "${BASE_REF}...HEAD" -W -- '*.rs' | awk '
|
||||
if git diff "${BASE_REF}...${HEAD_REF}" -W -- '*.rs' | awk '
|
||||
/^@@/ { if (has_test && has_add) { found=1; exit } has_test=0; has_add=0 }
|
||||
/^ .*#\[test\]/ || /^ .*#\[tokio::test\]/ || /^ .*#\[cfg\(test\)\]/ || /^ .*mod tests/ { has_test=1 }
|
||||
/^\+.*#\[test\]/ || /^\+.*#\[tokio::test\]/ || /^\+.*#\[cfg\(test\)\]/ || /^\+.*mod tests/ { has_test=1 }
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
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
|
||||
@@ -58,10 +58,16 @@ jobs:
|
||||
- *checkout
|
||||
- *install-rust
|
||||
- uses: Swatinem/rust-cache@v2
|
||||
- name: Generate GitHub token
|
||||
uses: actions/create-github-app-token@v2
|
||||
id: generate-token
|
||||
with:
|
||||
app-id: ${{ secrets.GH_RELEASES_MANAGER_APP_ID }}
|
||||
private-key: ${{ secrets.GH_RELEASES_MANAGER_APP_PRIVATE_KEY }}
|
||||
- name: Run release-plz
|
||||
uses: release-plz/[email protected]
|
||||
with:
|
||||
command: release-pr
|
||||
env:
|
||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
GITHUB_TOKEN: ${{ steps.generate-token.outputs.token }}
|
||||
CARGO_REGISTRY_TOKEN: ${{ secrets.CARGO_REGISTRY_TOKEN }}
|
||||
|
||||
@@ -25,9 +25,35 @@ concurrency:
|
||||
cancel-in-progress: false # Let running suites finish
|
||||
|
||||
jobs:
|
||||
# ── Resolve promotion base branch ───────────────────────────────
|
||||
resolve-promotion-base:
|
||||
name: Resolve promotion base
|
||||
runs-on: ubuntu-latest
|
||||
outputs:
|
||||
promotion_base: ${{ steps.resolve.outputs.promotion_base }}
|
||||
steps:
|
||||
- name: Resolve promotion base
|
||||
id: resolve
|
||||
env:
|
||||
GH_TOKEN: ${{ github.token }}
|
||||
FALLBACK_BRANCH: main
|
||||
REPO: ${{ github.repository }}
|
||||
run: |
|
||||
LATEST=$(gh pr list --repo "${REPO}" --label staging-promotion --state open \
|
||||
--json headRefName,createdAt \
|
||||
--jq '[.[] | select(.headRefName | startswith("staging-promote/"))] | sort_by(.createdAt) | last | .headRefName // empty')
|
||||
if [ -n "$LATEST" ]; then
|
||||
echo "promotion_base=${LATEST}" >> "$GITHUB_OUTPUT"
|
||||
echo "Using open promotion branch as base: ${LATEST}"
|
||||
else
|
||||
echo "promotion_base=${FALLBACK_BRANCH}" >> "$GITHUB_OUTPUT"
|
||||
echo "No open promotion branch found. Using ${FALLBACK_BRANCH}."
|
||||
fi
|
||||
|
||||
# ── Check for new commits ──────────────────────────────────────
|
||||
check-changes:
|
||||
name: Check for new commits
|
||||
needs: resolve-promotion-base
|
||||
runs-on: ubuntu-latest
|
||||
outputs:
|
||||
has_changes: ${{ steps.check.outputs.has_changes }}
|
||||
@@ -44,6 +70,7 @@ jobs:
|
||||
id: check
|
||||
env:
|
||||
FORCE_RUN: ${{ inputs.force }}
|
||||
PROMOTION_BASE: ${{ needs.resolve-promotion-base.outputs.promotion_base }}
|
||||
run: |
|
||||
CURRENT_HEAD=$(git rev-parse HEAD)
|
||||
echo "current_head=${CURRENT_HEAD}" >> "$GITHUB_OUTPUT"
|
||||
@@ -65,9 +92,9 @@ jobs:
|
||||
echo "Found ${COMMIT_COUNT} new commit(s) since last tested"
|
||||
DIFF_RANGE="${LAST_TESTED}..${CURRENT_HEAD}"
|
||||
else
|
||||
git fetch origin main
|
||||
MERGE_BASE=$(git merge-base origin/main HEAD)
|
||||
echo "First run -- reviewing from merge-base ${MERGE_BASE}"
|
||||
git fetch origin "${PROMOTION_BASE}"
|
||||
MERGE_BASE=$(git merge-base "origin/${PROMOTION_BASE}" HEAD)
|
||||
echo "First run -- reviewing from merge-base ${MERGE_BASE} against ${PROMOTION_BASE}"
|
||||
DIFF_RANGE="${MERGE_BASE}..${CURRENT_HEAD}"
|
||||
fi
|
||||
fi
|
||||
@@ -101,7 +128,7 @@ jobs:
|
||||
# ── Create promotion PR (triggers claude-review.yml on the PR) ──
|
||||
create-promotion-pr:
|
||||
name: Create Promotion PR
|
||||
needs: check-changes
|
||||
needs: [resolve-promotion-base, check-changes]
|
||||
if: needs.check-changes.outputs.has_changes == 'true'
|
||||
runs-on: ubuntu-latest
|
||||
outputs:
|
||||
@@ -129,18 +156,19 @@ jobs:
|
||||
echo "token=${{ github.token }}" >> "$GITHUB_OUTPUT"
|
||||
fi
|
||||
|
||||
- name: Check if staging is ahead of main
|
||||
- name: Check if staging is ahead of target branch
|
||||
id: ahead-check
|
||||
env:
|
||||
GH_TOKEN: ${{ steps.token.outputs.token }}
|
||||
PROMOTION_BASE: ${{ needs.resolve-promotion-base.outputs.promotion_base }}
|
||||
run: |
|
||||
git fetch origin main
|
||||
AHEAD=$(git rev-list --count origin/main..origin/staging)
|
||||
git fetch origin "${PROMOTION_BASE}"
|
||||
AHEAD=$(git rev-list --count "origin/${PROMOTION_BASE}..origin/staging")
|
||||
echo "commits_ahead=${AHEAD}" >> "$GITHUB_OUTPUT"
|
||||
if [ "$AHEAD" -eq 0 ]; then
|
||||
echo "Staging is not ahead of main. Nothing to promote."
|
||||
echo "Staging is not ahead of ${PROMOTION_BASE}. Nothing to promote."
|
||||
else
|
||||
echo "Staging is ${AHEAD} commits ahead of main."
|
||||
echo "Staging is ${AHEAD} commits ahead of ${PROMOTION_BASE}."
|
||||
fi
|
||||
|
||||
- name: Create promotion branch
|
||||
@@ -154,53 +182,53 @@ jobs:
|
||||
echo "branch=${BRANCH}" >> "$GITHUB_OUTPUT"
|
||||
echo "Created promotion branch: ${BRANCH}"
|
||||
|
||||
- name: Find base branch
|
||||
id: find-base
|
||||
if: steps.ahead-check.outputs.commits_ahead != '0'
|
||||
env:
|
||||
GH_TOKEN: ${{ steps.token.outputs.token }}
|
||||
run: |
|
||||
# Find the newest open promotion PR with a staging-promote/* head branch
|
||||
LATEST=$(gh pr list --label staging-promotion --state open \
|
||||
--json headRefName,createdAt \
|
||||
--jq '[.[] | select(.headRefName | startswith("staging-promote/"))] | sort_by(.createdAt) | last | .headRefName // empty')
|
||||
if [ -n "$LATEST" ]; then
|
||||
echo "base=${LATEST}" >> "$GITHUB_OUTPUT"
|
||||
echo "Chaining onto existing promotion branch: ${LATEST}"
|
||||
else
|
||||
echo "base=main" >> "$GITHUB_OUTPUT"
|
||||
echo "No existing promotion PR — targeting main"
|
||||
fi
|
||||
|
||||
- name: Create promotion PR
|
||||
id: create-pr
|
||||
if: steps.ahead-check.outputs.commits_ahead != '0'
|
||||
env:
|
||||
GH_TOKEN: ${{ steps.token.outputs.token }}
|
||||
run: |
|
||||
source .github/scripts/pr-body-utils.sh
|
||||
RANGE="${{ needs.check-changes.outputs.diff_range }}"
|
||||
TIMESTAMP=$(date -u +"%Y-%m-%d %H:%M UTC")
|
||||
BRANCH="${{ steps.branch.outputs.branch }}"
|
||||
BASE="${{ steps.find-base.outputs.base }}"
|
||||
BASE="${{ needs.resolve-promotion-base.outputs.promotion_base }}"
|
||||
|
||||
MAX_COMMITS=50
|
||||
load_commit_summary "${RANGE}" "${MAX_COMMITS}"
|
||||
|
||||
# Build PR body via concatenation to avoid heredoc shell expansion
|
||||
# (commit messages in COMMIT_MD may contain $, backticks, or backslashes)
|
||||
PR_BODY="## Auto-promotion from staging CI"
|
||||
PR_BODY+=$'\n\n'"**Batch range:** \`${RANGE}\`"
|
||||
PR_BODY+=$'\n'"**Promotion branch:** \`${BRANCH}\`"
|
||||
PR_BODY+=$'\n'"**Base:** \`${BASE}\`"
|
||||
PR_BODY+=$'\n'"**Triggered by:** Staging CI batch at ${TIMESTAMP}"
|
||||
PR_BODY+=$'\n\n'"### Commits in this batch (${COMMIT_COUNT}):"
|
||||
PR_BODY+=$'\n'"${COMMIT_MD}"
|
||||
PR_BODY+=$'\n\n'"<!-- staging-ci-current:start -->"
|
||||
PR_BODY+=$'\n'"### Current commits in this promotion (${COMMIT_COUNT})"
|
||||
PR_BODY+=$'\n'
|
||||
PR_BODY+=$'\n'"**Current base:** \`${BASE}\`"
|
||||
PR_BODY+=$'\n'"**Current head:** \`${BRANCH}\`"
|
||||
PR_BODY+=$'\n'"**Current range:** \`origin/${BASE}..origin/${BRANCH}\`"
|
||||
PR_BODY+=$'\n'
|
||||
PR_BODY+=$'\n'"${COMMIT_MD}"
|
||||
PR_BODY+=$'\n'
|
||||
PR_BODY+=$'\n'"*Auto-updated by staging promotion metadata workflow*"
|
||||
PR_BODY+=$'\n'"<!-- staging-ci-current:end -->"
|
||||
PR_BODY+=$'\n\n'"Waiting for gates:"
|
||||
PR_BODY+=$'\n'"- Tests: pending"
|
||||
PR_BODY+=$'\n'"- E2E: pending"
|
||||
PR_BODY+=$'\n'"- Claude Code review: pending (will post comments on this PR)"
|
||||
PR_BODY+=$'\n\n'"---"
|
||||
PR_BODY+=$'\n'"*Auto-created by staging-ci workflow*"
|
||||
|
||||
PR_URL=$(gh pr create \
|
||||
--base "$BASE" \
|
||||
--head "$BRANCH" \
|
||||
--title "chore: promote staging to main (${TIMESTAMP})" \
|
||||
--body "## Auto-promotion from staging CI
|
||||
|
||||
**Batch range:** \`${RANGE}\`
|
||||
**Promotion branch:** \`${BRANCH}\`
|
||||
**Base:** \`${BASE}\`
|
||||
**Triggered by:** Staging CI batch at ${TIMESTAMP}
|
||||
|
||||
Waiting for gates:
|
||||
- Tests: pending
|
||||
- E2E: pending
|
||||
- Claude Code review: pending (will post comments on this PR)
|
||||
|
||||
---
|
||||
*Auto-created by staging-ci workflow*" \
|
||||
--title "chore: promote staging to ${BASE} (${TIMESTAMP})" \
|
||||
--body "$PR_BODY" \
|
||||
--label "staging-promotion")
|
||||
|
||||
PR_NUM=$(echo "$PR_URL" | grep -oE '[0-9]+$')
|
||||
@@ -225,7 +253,8 @@ jobs:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: staging
|
||||
fetch-depth: 1
|
||||
# Need full history to recompute the final promoted range before merge.
|
||||
fetch-depth: 0
|
||||
|
||||
- name: Generate GitHub App token
|
||||
id: app-token
|
||||
@@ -324,8 +353,10 @@ jobs:
|
||||
# Use process substitution so variables propagate to parent shell
|
||||
while read -r line; do
|
||||
TAG=$(echo "$line" | grep -oE '^\[(CRITICAL|HIGH|MEDIUM|LOW):[0-9]+\]')
|
||||
SEVERITY=$(echo "$TAG" | sed 's/\[\(.*\):\(.*\)\]/\1/')
|
||||
CONFIDENCE=$(echo "$TAG" | sed 's/\[\(.*\):\(.*\)\]/\2/')
|
||||
SEVERITY="${TAG#\[}"
|
||||
SEVERITY="${SEVERITY%%:*}"
|
||||
CONFIDENCE="${TAG##*:}"
|
||||
CONFIDENCE="${CONFIDENCE%\]}"
|
||||
DESC=$(echo "$line" | sed "s/\[${SEVERITY}:${CONFIDENCE}\] *//" | head -1)
|
||||
|
||||
echo "Found: [${SEVERITY}:${CONFIDENCE}] ${DESC}"
|
||||
@@ -417,11 +448,29 @@ jobs:
|
||||
GH_TOKEN: ${{ steps.token.outputs.token }}
|
||||
PR_NUMBER: ${{ needs.create-promotion-pr.outputs.pr_number }}
|
||||
run: |
|
||||
source .github/scripts/pr-body-utils.sh
|
||||
if [ -n "$PR_NUMBER" ]; then
|
||||
BASE=$(gh pr view "$PR_NUMBER" --json baseRefName --jq '.baseRefName')
|
||||
if [ "$BASE" = "main" ]; then
|
||||
echo "Merging promotion PR #${PR_NUMBER} (targets main)"
|
||||
gh pr merge "$PR_NUMBER" --merge
|
||||
TITLE=$(gh pr view "$PR_NUMBER" --json title --jq '.title')
|
||||
HEAD_BRANCH=$(gh pr view "$PR_NUMBER" --json headRefName --jq '.headRefName')
|
||||
git fetch origin "${BASE}" "${HEAD_BRANCH}"
|
||||
CURRENT_RANGE="origin/${BASE}..origin/${HEAD_BRANCH}"
|
||||
MAX_COMMITS=50
|
||||
load_commit_summary "${CURRENT_RANGE}" "${MAX_COMMITS}"
|
||||
{
|
||||
echo "staging-promotion-summary-v1"
|
||||
echo "promotion-pr: #${PR_NUMBER}"
|
||||
echo "base: ${BASE}"
|
||||
echo "head: ${HEAD_BRANCH}"
|
||||
echo "current-range: ${CURRENT_RANGE}"
|
||||
echo "current-commit-count: ${COMMIT_COUNT}"
|
||||
echo ""
|
||||
echo "Current commits in this promotion (${COMMIT_COUNT}):"
|
||||
echo "${COMMIT_MD}"
|
||||
} > /tmp/staging-promotion-merge-body.md
|
||||
gh pr merge "$PR_NUMBER" --merge --subject "#${PR_NUMBER} $TITLE" --body-file /tmp/staging-promotion-merge-body.md
|
||||
echo "merged=true" >> "$GITHUB_OUTPUT"
|
||||
else
|
||||
echo "PR #${PR_NUMBER} targets '${BASE}' (not main) — leaving open for chain resolution"
|
||||
@@ -461,18 +510,20 @@ jobs:
|
||||
steps:
|
||||
- name: Summary
|
||||
run: |
|
||||
echo "## Staging CI Batch Results" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "| Check | Result |" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "|-------|--------|" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "| Tests | ${{ needs.tests.result }} |" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "| E2E | ${{ needs.e2e.result }} |" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "| Promotion PR | ${{ needs.create-promotion-pr.result }} |" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "| Gate | ${{ needs.gate.result }} |" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "| Tag Updated | ${{ needs.update-tag.result }} |" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "Range: ${{ needs.check-changes.outputs.diff_range }}" >> "$GITHUB_STEP_SUMMARY"
|
||||
PR_NUM="${{ needs.create-promotion-pr.outputs.pr_number }}"
|
||||
if [ -n "$PR_NUM" ]; then
|
||||
echo "Promotion PR: #${PR_NUM}" >> "$GITHUB_STEP_SUMMARY"
|
||||
fi
|
||||
{
|
||||
echo "## Staging CI Batch Results"
|
||||
echo ""
|
||||
echo "| Check | Result |"
|
||||
echo "|-------|--------|"
|
||||
echo "| Tests | ${{ needs.tests.result }} |"
|
||||
echo "| E2E | ${{ needs.e2e.result }} |"
|
||||
echo "| Promotion PR | ${{ needs.create-promotion-pr.result }} |"
|
||||
echo "| Gate | ${{ needs.gate.result }} |"
|
||||
echo "| Tag Updated | ${{ needs.update-tag.result }} |"
|
||||
echo ""
|
||||
echo "Range: ${{ needs.check-changes.outputs.diff_range }}"
|
||||
PR_NUM="${{ needs.create-promotion-pr.outputs.pr_number }}"
|
||||
if [ -n "$PR_NUM" ]; then
|
||||
echo "Promotion PR: #${PR_NUM}"
|
||||
fi
|
||||
} >> "$GITHUB_STEP_SUMMARY"
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
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
|
||||
@@ -7,6 +7,15 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
||||
|
||||
## [Unreleased]
|
||||
|
||||
## [0.18.0](https://github.com/nearai/ironclaw/compare/v0.17.0...v0.18.0) - 2026-03-11
|
||||
|
||||
### Other
|
||||
|
||||
- Merge pull request #907 from nearai/staging-promote/b0214fef-22930316561
|
||||
- promote staging to main (2026-03-10 15:19 UTC) ([#865](https://github.com/nearai/ironclaw/pull/865))
|
||||
- Merge pull request #830 from nearai/staging-promote/3a2989d0-22888378864
|
||||
- update WASM artifact SHA256 checksums [skip ci] ([#876](https://github.com/nearai/ironclaw/pull/876))
|
||||
|
||||
## [0.17.0](https://github.com/nearai/ironclaw/compare/v0.16.1...v0.17.0) - 2026-03-10
|
||||
|
||||
### Added
|
||||
|
||||
Generated
+1
-1
@@ -3350,7 +3350,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "ironclaw"
|
||||
version = "0.17.0"
|
||||
version = "0.18.0"
|
||||
dependencies = [
|
||||
"aes-gcm",
|
||||
"aho-corasick",
|
||||
|
||||
+1
-1
@@ -20,7 +20,7 @@ exclude = [
|
||||
|
||||
[package]
|
||||
name = "ironclaw"
|
||||
version = "0.17.0"
|
||||
version = "0.18.0"
|
||||
edition = "2024"
|
||||
rust-version = "1.92"
|
||||
description = "Secure personal AI assistant that protects your data and expands its capabilities on the fly"
|
||||
|
||||
@@ -19,6 +19,7 @@ WORKDIR /app
|
||||
|
||||
# Copy manifests first for layer caching
|
||||
COPY Cargo.toml Cargo.lock ./
|
||||
COPY crates/ crates/
|
||||
|
||||
# Copy source, build script, tests, and supporting directories
|
||||
COPY build.rs build.rs
|
||||
|
||||
@@ -20,6 +20,7 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
WORKDIR /app
|
||||
|
||||
COPY Cargo.toml Cargo.lock ./
|
||||
COPY crates/ crates/
|
||||
COPY build.rs build.rs
|
||||
COPY src/ src/
|
||||
COPY tests/ tests/
|
||||
|
||||
+1
-1
@@ -170,7 +170,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
||||
| `nodes` | ✅ | ❌ | P3 | Device management, remove/clear flows |
|
||||
| `plugins` | ✅ | ❌ | P3 | Plugin management |
|
||||
| `hooks` | ✅ | ✅ | P2 | Lifecycle hooks |
|
||||
| `cron` | ✅ | ❌ | P2 | Scheduled jobs (model/thinking fields in edit) |
|
||||
| `cron` | ✅ | 🚧 | P2 | list/create/edit/enable/disable/delete/history; TODO: `cron run`, model/thinking fields |
|
||||
| `webhooks` | ✅ | ❌ | P3 | Webhook config |
|
||||
| `message send` | ✅ | ❌ | P2 | Send to channels |
|
||||
| `browser` | ✅ | ❌ | P3 | Browser automation |
|
||||
|
||||
Generated
+207
-2
@@ -20,12 +20,27 @@ version = "1.0.102"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7f202df86484c868dbad7eaa557ef785d5c66295e41b460ef922eca0723b842c"
|
||||
|
||||
[[package]]
|
||||
name = "base64ct"
|
||||
version = "1.8.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "2af50177e190e07a26ab74f8b1efbfe2ef87da2116221318cb1c2e82baf7de06"
|
||||
|
||||
[[package]]
|
||||
name = "bitflags"
|
||||
version = "2.11.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "843867be96c8daad0d758b57df9392b6d8d271134fce549de6ce169ff98a92af"
|
||||
|
||||
[[package]]
|
||||
name = "block-buffer"
|
||||
version = "0.10.4"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "3078c7629b62d3f0439517fa394996acacc5cbc91c5a20d8c658e77abd503a71"
|
||||
dependencies = [
|
||||
"generic-array",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "cfg-if"
|
||||
version = "1.0.4"
|
||||
@@ -33,20 +48,134 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801"
|
||||
|
||||
[[package]]
|
||||
name = "discord-channel"
|
||||
version = "0.1.0"
|
||||
name = "const-oid"
|
||||
version = "0.9.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c2459377285ad874054d797f3ccebf984978aa39129f6eafde5cdc8315b612f8"
|
||||
|
||||
[[package]]
|
||||
name = "cpufeatures"
|
||||
version = "0.2.17"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "59ed5838eebb26a2bb2e58f6d5b5316989ae9d08bab10e0e6d103e656d1b0280"
|
||||
dependencies = [
|
||||
"libc",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "crypto-common"
|
||||
version = "0.1.7"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a"
|
||||
dependencies = [
|
||||
"generic-array",
|
||||
"typenum",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "curve25519-dalek"
|
||||
version = "4.1.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "97fb8b7c4503de7d6ae7b42ab72a5a59857b4c937ec27a3d4539dba95b5ab2be"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"cpufeatures",
|
||||
"curve25519-dalek-derive",
|
||||
"digest",
|
||||
"fiat-crypto",
|
||||
"rustc_version",
|
||||
"subtle",
|
||||
"zeroize",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "curve25519-dalek-derive"
|
||||
version = "0.1.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f46882e17999c6cc590af592290432be3bce0428cb0d5f8b6715e4dc7b383eb3"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "der"
|
||||
version = "0.7.10"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e7c1832837b905bbfb5101e07cc24c8deddf52f93225eee6ead5f4d63d53ddcb"
|
||||
dependencies = [
|
||||
"const-oid",
|
||||
"zeroize",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "digest"
|
||||
version = "0.10.7"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292"
|
||||
dependencies = [
|
||||
"block-buffer",
|
||||
"crypto-common",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "discord-channel"
|
||||
version = "0.2.0"
|
||||
dependencies = [
|
||||
"ed25519-dalek",
|
||||
"hex",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"wit-bindgen",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ed25519"
|
||||
version = "2.2.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "115531babc129696a58c64a4fef0a8bf9e9698629fb97e9e40767d235cfbcd53"
|
||||
dependencies = [
|
||||
"pkcs8",
|
||||
"signature",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "ed25519-dalek"
|
||||
version = "2.2.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "70e796c081cee67dc755e1a36a0a172b897fab85fc3f6bc48307991f64e4eca9"
|
||||
dependencies = [
|
||||
"curve25519-dalek",
|
||||
"ed25519",
|
||||
"serde",
|
||||
"sha2",
|
||||
"subtle",
|
||||
"zeroize",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "equivalent"
|
||||
version = "1.0.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f"
|
||||
|
||||
[[package]]
|
||||
name = "fiat-crypto"
|
||||
version = "0.2.9"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "28dea519a9695b9977216879a3ebfddf92f1c08c05d984f8996aecd6ecdc811d"
|
||||
|
||||
[[package]]
|
||||
name = "generic-array"
|
||||
version = "0.14.7"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a"
|
||||
dependencies = [
|
||||
"typenum",
|
||||
"version_check",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "hashbrown"
|
||||
version = "0.14.5"
|
||||
@@ -68,6 +197,12 @@ version = "0.5.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea"
|
||||
|
||||
[[package]]
|
||||
name = "hex"
|
||||
version = "0.4.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70"
|
||||
|
||||
[[package]]
|
||||
name = "id-arena"
|
||||
version = "2.3.0"
|
||||
@@ -98,6 +233,12 @@ version = "0.2.5"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "884e2677b40cc8c339eaefcb701c32ef1fd2493d71118dc0ca4b6a736c93bd67"
|
||||
|
||||
[[package]]
|
||||
name = "libc"
|
||||
version = "0.2.182"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "6800badb6cb2082ffd7b6a67e6125bb39f18782f793520caee8cb8846be06112"
|
||||
|
||||
[[package]]
|
||||
name = "log"
|
||||
version = "0.4.29"
|
||||
@@ -116,6 +257,16 @@ version = "1.21.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "42f5e15c9953c5e4ccceeb2e7382a716482c34515315f7b03532b8b4e8393d2d"
|
||||
|
||||
[[package]]
|
||||
name = "pkcs8"
|
||||
version = "0.10.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f950b2377845cebe5cf8b5165cb3cc1a5e0fa5cfa3e1f7f55707d8fd82e0a7b7"
|
||||
dependencies = [
|
||||
"der",
|
||||
"spki",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "prettyplease"
|
||||
version = "0.2.37"
|
||||
@@ -144,6 +295,15 @@ dependencies = [
|
||||
"proc-macro2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rustc_version"
|
||||
version = "0.4.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "cfcb3a22ef46e85b45de6ee7e79d063319ebb6594faafcf1c225ea92ab6e9b92"
|
||||
dependencies = [
|
||||
"semver",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "semver"
|
||||
version = "1.0.27"
|
||||
@@ -193,6 +353,23 @@ dependencies = [
|
||||
"zmij",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "sha2"
|
||||
version = "0.10.9"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"cpufeatures",
|
||||
"digest",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "signature"
|
||||
version = "2.2.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "77549399552de45a898a580c1b41d445bf730df867cc44e6c0233bbc4b8329de"
|
||||
|
||||
[[package]]
|
||||
name = "smallvec"
|
||||
version = "1.15.1"
|
||||
@@ -208,6 +385,22 @@ dependencies = [
|
||||
"smallvec",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "spki"
|
||||
version = "0.7.3"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d91ed6c858b01f942cd56b37a94b3e0a1798290327d1236e4d9cf4eaca44d29d"
|
||||
dependencies = [
|
||||
"base64ct",
|
||||
"der",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "subtle"
|
||||
version = "2.6.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292"
|
||||
|
||||
[[package]]
|
||||
name = "syn"
|
||||
version = "2.0.117"
|
||||
@@ -219,6 +412,12 @@ dependencies = [
|
||||
"unicode-ident",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "typenum"
|
||||
version = "1.19.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "562d481066bde0658276a35467c4af00bdc6ee726305698a55b86e61d7ad82bb"
|
||||
|
||||
[[package]]
|
||||
name = "unicode-ident"
|
||||
version = "1.0.24"
|
||||
@@ -394,6 +593,12 @@ dependencies = [
|
||||
"syn",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "zeroize"
|
||||
version = "1.8.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b97154e67e32c85465826e8bcc1c59429aaaf107c1e4a9e53c8d8ccd5eff88d0"
|
||||
|
||||
[[package]]
|
||||
name = "zmij"
|
||||
version = "1.0.21"
|
||||
|
||||
@@ -10,6 +10,8 @@ publish = false
|
||||
serde = { version = "1.0", features = ["derive"] }
|
||||
serde_json = "1.0"
|
||||
wit-bindgen = "0.36"
|
||||
ed25519-dalek = { version = "2", default-features = false, features = ["alloc", "fast", "zeroize"] }
|
||||
hex = "0.4"
|
||||
|
||||
[lib]
|
||||
crate-type = ["cdylib"]
|
||||
|
||||
@@ -21,11 +21,10 @@ WASM channel for Discord integration - handle slash commands and button interact
|
||||
ironclaw secret set discord_bot_token YOUR_BOT_TOKEN
|
||||
```
|
||||
|
||||
**Note:** The `discord_bot_token` secret is the only value read directly by this
|
||||
Discord channel WASM component. The `discord_app_id` and `discord_public_key`
|
||||
secrets are used by the IronClaw host (for example, to verify Discord
|
||||
interaction signatures and manage slash command registration) and are not
|
||||
accessed from the WASM module itself.
|
||||
**Note:** The `discord_bot_token` secret is used for Discord REST API calls.
|
||||
Interaction signature verification is performed inside the Discord channel
|
||||
module and uses the channel config field `webhook_secret` (set this to your
|
||||
Discord app public key hex).
|
||||
|
||||
## Discord Configuration
|
||||
|
||||
@@ -87,6 +86,30 @@ If an internal error occurs (e.g., metadata serialization failure), the tool att
|
||||
Check the host logs for detailed error information.
|
||||
|
||||
## Advanced Usage
|
||||
### Mention Polling
|
||||
|
||||
The Discord channel can also poll configured channels for `@bot` mentions.
|
||||
|
||||
Example channel config:
|
||||
|
||||
```json
|
||||
{
|
||||
"require_signature_verification": true,
|
||||
"webhook_secret": "YOUR_DISCORD_PUBLIC_KEY_HEX",
|
||||
"polling_enabled": true,
|
||||
"poll_interval_ms": 30000,
|
||||
"mention_channel_ids": ["123456789012345678"],
|
||||
"owner_id": null,
|
||||
"dm_policy": "pairing",
|
||||
"allow_from": []
|
||||
}
|
||||
```
|
||||
|
||||
### Access Control
|
||||
|
||||
- `owner_id`: when set, only that Discord user can interact with the bot.
|
||||
- `dm_policy`: `open` allows all DMs; `pairing` requires approval.
|
||||
- `allow_from`: allowlist entries for DM pairing checks (`*`, user id, or username).
|
||||
|
||||
### Embeds
|
||||
|
||||
@@ -96,8 +119,11 @@ To send embeds, include an `embeds` array in the `metadata_json` field of the ag
|
||||
|
||||
### "Invalid Signature"
|
||||
|
||||
- Check that `discord_public_key` is set correctly in IronClaw secrets.
|
||||
- This validation happens on the host before reaching the WASM.
|
||||
- Check that `webhook_secret` is set to your Discord app public key hex in the
|
||||
Discord channel config.
|
||||
- Validation happens inside the Discord WASM channel.
|
||||
- If `require_signature_verification` is `true` and `webhook_secret` is empty,
|
||||
the channel returns HTTP `500` with a configuration error.
|
||||
|
||||
### "401 Unauthorized"
|
||||
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
"wit_version": "0.3.0",
|
||||
"type": "channel",
|
||||
"name": "discord",
|
||||
"description": "Discord Gateway/Webhook channel for handling slash commands, buttons, and messages",
|
||||
"description": "Discord webhook channel for slash commands, components, and optional mention polling",
|
||||
"setup": {
|
||||
"required_secrets": [
|
||||
{
|
||||
@@ -41,7 +41,7 @@
|
||||
},
|
||||
"channel": {
|
||||
"allowed_paths": ["/webhook/discord"],
|
||||
"allow_polling": false,
|
||||
"allow_polling": true,
|
||||
"callback_timeout_secs": 45,
|
||||
"workspace_prefix": "channels/discord/",
|
||||
"emit_rate_limit": {
|
||||
@@ -55,8 +55,12 @@
|
||||
},
|
||||
"config": {
|
||||
"require_signature_verification": true,
|
||||
"webhook_secret": null,
|
||||
"polling_enabled": false,
|
||||
"poll_interval_ms": 30000,
|
||||
"mention_channel_ids": [],
|
||||
"owner_id": null,
|
||||
"dm_policy": "pairing",
|
||||
"allow_from": []
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+976
-102
File diff suppressed because it is too large
Load Diff
@@ -1,5 +1,10 @@
|
||||
# WARNING: Replace all CHANGE_ME values before deploying.
|
||||
# Do not use placeholder passwords in production.
|
||||
|
||||
# Pin the Docker image version for deterministic deployments.
|
||||
# Update this value when deploying a new release.
|
||||
# IRONCLAW_VERSION=v1.0.0
|
||||
|
||||
DATABASE_URL=postgres://ironclaw:CHANGE_ME@localhost:5432/ironclaw
|
||||
|
||||
# NEAR AI Cloud (API key auth, Chat Completions API)
|
||||
|
||||
@@ -5,13 +5,17 @@ Requires=cloud-sql-proxy.service
|
||||
|
||||
[Service]
|
||||
Type=simple
|
||||
ExecStartPre=/usr/bin/docker pull us-central1-docker.pkg.dev/ironclaw-prod/ironclaw/agent:latest
|
||||
ExecStart=/usr/bin/docker run --rm \
|
||||
EnvironmentFile=/opt/ironclaw/.env
|
||||
# Pin to a specific version tag or digest instead of :latest to prevent
|
||||
# uncontrolled deployments. Update IRONCLAW_VERSION in /opt/ironclaw/.env
|
||||
# or replace the tag below when deploying a new release.
|
||||
ExecStartPre=/bin/bash -c 'docker pull us-central1-docker.pkg.dev/ironclaw-prod/ironclaw/agent:${IRONCLAW_VERSION:-latest}'
|
||||
ExecStart=/bin/bash -c 'docker run --rm \
|
||||
--name ironclaw \
|
||||
--env-file /opt/ironclaw/.env \
|
||||
--network=host \
|
||||
us-central1-docker.pkg.dev/ironclaw-prod/ironclaw/agent:latest \
|
||||
--no-onboard
|
||||
-p 3000:3000 \
|
||||
us-central1-docker.pkg.dev/ironclaw-prod/ironclaw/agent:${IRONCLAW_VERSION:-latest} \
|
||||
--no-onboard'
|
||||
ExecStop=/usr/bin/docker stop ironclaw
|
||||
Restart=always
|
||||
RestartSec=10
|
||||
|
||||
+8
-1
@@ -24,8 +24,15 @@ systemctl enable docker
|
||||
systemctl start docker
|
||||
|
||||
echo "==> Installing Cloud SQL Auth Proxy"
|
||||
CLOUD_SQL_PROXY_VERSION="v2.14.3"
|
||||
CLOUD_SQL_PROXY_SHA256="75e7cc1f158ab6f97b7810e9d8419c55735cff40bc56d4f19673adfdf2406a59"
|
||||
curl -fsSL -o /usr/local/bin/cloud-sql-proxy \
|
||||
https://storage.googleapis.com/cloud-sql-connectors/cloud-sql-proxy/v2.14.3/cloud-sql-proxy.linux.amd64
|
||||
"https://storage.googleapis.com/cloud-sql-connectors/cloud-sql-proxy/${CLOUD_SQL_PROXY_VERSION}/cloud-sql-proxy.linux.amd64"
|
||||
echo "${CLOUD_SQL_PROXY_SHA256} /usr/local/bin/cloud-sql-proxy" | sha256sum -c - || {
|
||||
echo "ERROR: Cloud SQL Auth Proxy checksum verification failed -- aborting"
|
||||
rm -f /usr/local/bin/cloud-sql-proxy
|
||||
exit 1
|
||||
}
|
||||
chmod +x /usr/local/bin/cloud-sql-proxy
|
||||
|
||||
echo "==> Installing systemd services"
|
||||
|
||||
@@ -15,6 +15,7 @@ configurations.
|
||||
| io.net | `ionet` | `IONET_API_KEY` | Intelligence API |
|
||||
| Mistral | `mistral` | `MISTRAL_API_KEY` | Mistral models |
|
||||
| Yandex AI Studio | `yandex` | `YANDEX_API_KEY` | YandexGPT models |
|
||||
| MiniMax | `minimax` | `MINIMAX_API_KEY` | MiniMax-M2.5 models |
|
||||
| Cloudflare Workers AI | `cloudflare` | `CLOUDFLARE_API_KEY` | Access to Workers AI |
|
||||
| Ollama | `ollama` | No | Local inference |
|
||||
| AWS Bedrock | `bedrock` | AWS credentials | Native Converse API |
|
||||
@@ -74,6 +75,25 @@ Pull a model first: `ollama pull llama3.2`
|
||||
|
||||
---
|
||||
|
||||
## MiniMax
|
||||
|
||||
[MiniMax](https://platform.minimax.io) provides high-performance language models with 204,800 token context windows.
|
||||
|
||||
```env
|
||||
LLM_BACKEND=minimax
|
||||
MINIMAX_API_KEY=...
|
||||
```
|
||||
|
||||
Available models: `MiniMax-M2.5` (default), `MiniMax-M2.5-highspeed`
|
||||
|
||||
To use the China mainland endpoint, set:
|
||||
|
||||
```env
|
||||
MINIMAX_BASE_URL=https://api.minimaxi.com/v1
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## AWS Bedrock (requires `--features bedrock`)
|
||||
|
||||
Uses the native AWS Converse API via `aws-sdk-bedrockruntime`. Supports standard AWS
|
||||
|
||||
@@ -382,6 +382,27 @@
|
||||
"can_list_models": true
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": "minimax",
|
||||
"aliases": [
|
||||
"mini_max"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://api.minimax.io/v1",
|
||||
"api_key_env": "MINIMAX_API_KEY",
|
||||
"api_key_required": true,
|
||||
"base_url_env": "MINIMAX_BASE_URL",
|
||||
"model_env": "MINIMAX_MODEL",
|
||||
"default_model": "MiniMax-M2.5",
|
||||
"description": "MiniMax API (MiniMax-M2.5 and MiniMax-M2.5-highspeed models)",
|
||||
"setup": {
|
||||
"kind": "api_key",
|
||||
"secret_name": "llm_minimax_api_key",
|
||||
"key_url": "https://platform.minimax.io",
|
||||
"display_name": "MiniMax",
|
||||
"can_list_models": false
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": "cloudflare",
|
||||
"aliases": [
|
||||
|
||||
@@ -18,8 +18,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/discord-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/discord-0.2.0-wasm32-wasip2.tar.gz",
|
||||
"sha256": "efa1b9019fa33e243f8db1e1fcc732731d45836336bdd26ca19b6fe227ca8b69"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -18,8 +18,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/slack-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/slack-0.2.1-wasm32-wasip2.tar.gz",
|
||||
"sha256": "d4667e35126986509d862bc3a0088777305d8f41c75de83c1e223b42312ede48"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "telegram",
|
||||
"display_name": "Telegram Channel",
|
||||
"kind": "channel",
|
||||
"version": "0.2.2",
|
||||
"version": "0.2.3",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Talk to your agent through a Telegram bot",
|
||||
"keywords": [
|
||||
@@ -18,8 +18,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/telegram-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/telegram-0.2.3-wasm32-wasip2.tar.gz",
|
||||
"sha256": "b9a83d5a2d1285ce0ec116b354336a1f245f893291ccb01dffbcaccf89d72aed"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -18,8 +18,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/whatsapp-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/whatsapp-0.2.0-wasm32-wasip2.tar.gz",
|
||||
"sha256": "feb9194719d9bed796b070ab4dc30348dbfb5d3dec56f9f21e02d14137abab01"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "github",
|
||||
"display_name": "GitHub",
|
||||
"kind": "tool",
|
||||
"version": "0.2.1",
|
||||
"version": "0.2.0",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "GitHub integration for issues, PRs, repos, and code search",
|
||||
"keywords": [
|
||||
@@ -19,8 +19,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/github-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/github-0.2.0-wasm32-wasip2.tar.gz",
|
||||
"sha256": "da9fac56b6f20197a415489bbaec9fefb085a5cf6324cab79ea48a47eb19c13b"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -18,8 +18,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/gmail-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/gmail-0.2.0-wasm32-wasip2.tar.gz",
|
||||
"sha256": "ee9574e02e92bc1d481f1310eb88afd99ee52bf6971074ab33bd76bf99b34b1d"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -18,8 +18,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-calendar-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-calendar-0.2.0-wasm32-wasip2.tar.gz",
|
||||
"sha256": "2fa47150ea222e787c122182ad6f4dfa2ffaf5fe490d05e8de887a76445f8d2d"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -18,8 +18,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-docs-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-docs-0.2.0-wasm32-wasip2.tar.gz",
|
||||
"sha256": "40e134a1c1564f832ca861c3396895d4e33ec67b99313fc1f97baf8d971423a9"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -18,8 +18,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-drive-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-drive-0.2.0-wasm32-wasip2.tar.gz",
|
||||
"sha256": "002a341a1d58125563a7c69561b26fbc2629b04ea723cade744102bdc0fbb71f"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -18,8 +18,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-sheets-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-sheets-0.2.0-wasm32-wasip2.tar.gz",
|
||||
"sha256": "8aa2c9d52f033edea3a6c2311b0ec694ccb6d0a54ef07e94d72bf8be1ce8009a"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -17,8 +17,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-slides-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/google-slides-0.2.0-wasm32-wasip2.tar.gz",
|
||||
"sha256": "e931a97d4fd0b0b938e464dc7c7f2be6ea6b4d1508f5ea3cd931d44db23f05f5"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
{
|
||||
"name": "llm-context",
|
||||
"display_name": "LLM Context",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Fetch pre-extracted web content from Brave Search for grounding LLM answers (RAG, fact-checking)",
|
||||
"keywords": [
|
||||
"search",
|
||||
"web",
|
||||
"brave",
|
||||
"rag",
|
||||
"grounding",
|
||||
"llm",
|
||||
"context"
|
||||
],
|
||||
"source": {
|
||||
"dir": "tools-src/llm-context",
|
||||
"capabilities": "llm-context-tool.capabilities.json",
|
||||
"crate_name": "llm-context-tool"
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/llm-context-wasm32-wasip2.tar.gz",
|
||||
"sha256": "581cc5867ef3b75116b7ddc8161e63dd92befe2b53e6ad8213c007639aa243c3"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
"method": "manual",
|
||||
"provider": "Brave",
|
||||
"secrets": [
|
||||
"brave_api_key"
|
||||
],
|
||||
"shared_auth": "Same API key as Web Search tool (brave_api_key)",
|
||||
"setup_url": "https://brave.com/search/api/"
|
||||
},
|
||||
"tags": [
|
||||
"default",
|
||||
"search"
|
||||
]
|
||||
}
|
||||
@@ -17,8 +17,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/slack-tool-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/slack-0.2.1-wasm32-wasip2.tar.gz",
|
||||
"sha256": "d4667e35126986509d862bc3a0088777305d8f41c75de83c1e223b42312ede48"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -18,8 +18,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/telegram-mtproto-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/telegram-0.2.2-wasm32-wasip2.tar.gz",
|
||||
"sha256": "b9a83d5a2d1285ce0ec116b354336a1f245f893291ccb01dffbcaccf89d72aed"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -18,8 +18,8 @@
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/web-search-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"url": "https://github.com/nearai/ironclaw/releases/download/v0.18.0/web-search-0.2.0-wasm32-wasip2.tar.gz",
|
||||
"sha256": "56834573c54ea2a33cea1eb0f04bbdf59f1ef8d8702995cf431b0921302eeccc"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
+5
-24
@@ -18,7 +18,7 @@ use crate::agent::self_repair::{DefaultSelfRepair, RepairResult, SelfRepair};
|
||||
use crate::agent::session_manager::SessionManager;
|
||||
use crate::agent::submission::{Submission, SubmissionParser, SubmissionResult};
|
||||
use crate::agent::{HeartbeatConfig as AgentHeartbeatConfig, Router, Scheduler};
|
||||
use crate::channels::{ChannelManager, IncomingMessage, OutgoingResponse, StatusUpdate};
|
||||
use crate::channels::{ChannelManager, IncomingMessage, OutgoingResponse};
|
||||
use crate::config::{AgentConfig, HeartbeatConfig, RoutineConfig, SkillsConfig};
|
||||
use crate::context::ContextManager;
|
||||
use crate::db::Database;
|
||||
@@ -936,29 +936,10 @@ impl Agent {
|
||||
SubmissionResult::Ok { message } => Ok(message),
|
||||
SubmissionResult::Error { message } => Ok(Some(format!("Error: {}", message))),
|
||||
SubmissionResult::Interrupted => Ok(Some("Interrupted.".into())),
|
||||
SubmissionResult::NeedApproval {
|
||||
request_id,
|
||||
tool_name,
|
||||
description,
|
||||
parameters,
|
||||
} => {
|
||||
// Each channel renders the approval prompt via send_status.
|
||||
// Web gateway shows an inline card, REPL prints a formatted prompt, etc.
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::ApprovalNeeded {
|
||||
request_id: request_id.to_string(),
|
||||
tool_name,
|
||||
description,
|
||||
parameters,
|
||||
},
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
|
||||
// Empty string signals the caller to skip respond() (no duplicate text)
|
||||
SubmissionResult::NeedApproval { .. } => {
|
||||
// ApprovalNeeded status was already sent by thread_ops.rs before
|
||||
// returning this result. Empty string signals the caller to skip
|
||||
// respond() (no duplicate text).
|
||||
Ok(Some(String::new()))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -554,6 +554,31 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
|
||||
};
|
||||
|
||||
if needs_approval {
|
||||
// In non-DM relay channels, auto-deny approval-
|
||||
// requiring tools to prevent stuck AwaitingApproval
|
||||
// state and prompt injection from other users.
|
||||
let is_relay = self.message.channel.ends_with("-relay");
|
||||
let is_dm = self
|
||||
.message
|
||||
.metadata
|
||||
.get("event_type")
|
||||
.and_then(|v| v.as_str())
|
||||
== Some("direct_message");
|
||||
if is_relay && !is_dm {
|
||||
tracing::info!(
|
||||
tool = %tc.name,
|
||||
channel = %self.message.channel,
|
||||
"Auto-denying approval-requiring tool in non-DM relay channel"
|
||||
);
|
||||
let reject_msg = format!(
|
||||
"Tool '{}' requires approval and cannot run in shared channels. \
|
||||
Ask the user to message me directly (DM) to use this tool.",
|
||||
tc.name
|
||||
);
|
||||
preflight.push((tc, PreflightOutcome::Rejected(reject_msg)));
|
||||
continue;
|
||||
}
|
||||
|
||||
approval_needed = Some((idx, tc, tool));
|
||||
break;
|
||||
}
|
||||
@@ -2235,4 +2260,51 @@ mod tests {
|
||||
"Present 'data' field should produce non-empty string"
|
||||
);
|
||||
}
|
||||
|
||||
/// Test the relay channel auto-deny decision logic:
|
||||
/// approval-requiring tools in non-DM relay channels must be rejected.
|
||||
#[test]
|
||||
fn test_relay_non_dm_auto_deny_decision() {
|
||||
use crate::channels::IncomingMessage;
|
||||
|
||||
// Case 1: relay channel + non-DM → should auto-deny
|
||||
let msg = IncomingMessage::new("slack-relay", "u1", "hello")
|
||||
.with_metadata(serde_json::json!({ "event_type": "message" }));
|
||||
let is_relay = msg.channel.ends_with("-relay");
|
||||
let is_dm =
|
||||
msg.metadata.get("event_type").and_then(|v| v.as_str()) == Some("direct_message");
|
||||
assert!(is_relay && !is_dm, "Should auto-deny in relay non-DM");
|
||||
|
||||
// Case 2: relay channel + DM → should NOT auto-deny
|
||||
let msg_dm = IncomingMessage::new("slack-relay", "u1", "hello")
|
||||
.with_metadata(serde_json::json!({ "event_type": "direct_message" }));
|
||||
let is_dm_2 =
|
||||
msg_dm.metadata.get("event_type").and_then(|v| v.as_str()) == Some("direct_message");
|
||||
assert!(
|
||||
!msg_dm.channel.ends_with("-relay") || is_dm_2,
|
||||
"Should NOT auto-deny in relay DM"
|
||||
);
|
||||
|
||||
// Case 3: non-relay channel → should NOT auto-deny
|
||||
let msg_web = IncomingMessage::new("web", "u1", "hello")
|
||||
.with_metadata(serde_json::json!({ "event_type": "message" }));
|
||||
assert!(
|
||||
!msg_web.channel.ends_with("-relay"),
|
||||
"Non-relay channel should not trigger auto-deny"
|
||||
);
|
||||
}
|
||||
|
||||
/// Test that the auto-deny produces a PreflightOutcome::Rejected-style message.
|
||||
#[test]
|
||||
fn test_relay_auto_deny_message_format() {
|
||||
let tool_name = "shell";
|
||||
let result_msg = format!(
|
||||
"Tool '{}' requires approval and cannot run in shared channels. \
|
||||
Ask the user to message me directly (DM) to use this tool.",
|
||||
tool_name
|
||||
);
|
||||
assert!(result_msg.contains("shell"));
|
||||
assert!(result_msg.contains("approval"));
|
||||
assert!(result_msg.contains("DM"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -32,7 +32,6 @@ pub mod task;
|
||||
mod thread_ops;
|
||||
pub mod undo;
|
||||
|
||||
pub use crate::worker::{Worker, WorkerDeps};
|
||||
pub(crate) use agent_loop::truncate_for_preview;
|
||||
pub use agent_loop::{Agent, AgentDeps};
|
||||
pub use compaction::{CompactionResult, ContextCompactor};
|
||||
|
||||
+116
-3
@@ -207,7 +207,7 @@ impl Trigger {
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum RoutineAction {
|
||||
/// Single LLM call, no tools. Cheap and fast.
|
||||
/// Single LLM call (optionally with tools). Cheap and fast.
|
||||
Lightweight {
|
||||
/// The prompt sent to the LLM.
|
||||
prompt: String,
|
||||
@@ -217,6 +217,14 @@ pub enum RoutineAction {
|
||||
/// Max output tokens (default: 4096).
|
||||
#[serde(default = "default_max_tokens")]
|
||||
max_tokens: u32,
|
||||
/// Enable tool access (default: false for backward compatibility).
|
||||
/// When true, the LLM can call tools during execution.
|
||||
/// Tools requiring approval are automatically filtered out.
|
||||
#[serde(default)]
|
||||
use_tools: bool,
|
||||
/// Max tool call rounds (default: 3). Only used when use_tools is true.
|
||||
#[serde(default = "default_max_tool_rounds")]
|
||||
max_tool_rounds: u32,
|
||||
},
|
||||
/// Full multi-turn worker job with tool access.
|
||||
FullJob {
|
||||
@@ -243,6 +251,19 @@ fn default_max_iterations() -> u32 {
|
||||
10
|
||||
}
|
||||
|
||||
fn default_max_tool_rounds() -> u32 {
|
||||
3
|
||||
}
|
||||
|
||||
/// Hard upper bound for max_tool_rounds to prevent runaway loops and cost explosion.
|
||||
pub(crate) const MAX_TOOL_ROUNDS_LIMIT: u32 = 20;
|
||||
|
||||
/// Clamp max_tool_rounds to [1, MAX_TOOL_ROUNDS_LIMIT].
|
||||
/// Accepts u64 to avoid truncation before clamping.
|
||||
fn clamp_max_tool_rounds(value: u64) -> u32 {
|
||||
value.clamp(1, MAX_TOOL_ROUNDS_LIMIT as u64) as u32
|
||||
}
|
||||
|
||||
/// Parse a `tool_permissions` JSON array into a `Vec<String>`.
|
||||
pub fn parse_tool_permissions(value: &serde_json::Value) -> Vec<String> {
|
||||
value
|
||||
@@ -290,10 +311,22 @@ impl RoutineAction {
|
||||
.get("max_tokens")
|
||||
.and_then(|v| v.as_u64())
|
||||
.unwrap_or(default_max_tokens() as u64) as u32;
|
||||
let use_tools = config
|
||||
.get("use_tools")
|
||||
.and_then(|v| v.as_bool())
|
||||
.unwrap_or(false);
|
||||
let max_tool_rounds = clamp_max_tool_rounds(
|
||||
config
|
||||
.get("max_tool_rounds")
|
||||
.and_then(|v| v.as_u64())
|
||||
.unwrap_or(default_max_tool_rounds() as u64),
|
||||
);
|
||||
Ok(RoutineAction::Lightweight {
|
||||
prompt,
|
||||
context_paths,
|
||||
max_tokens,
|
||||
use_tools,
|
||||
max_tool_rounds,
|
||||
})
|
||||
}
|
||||
"full_job" => {
|
||||
@@ -339,10 +372,14 @@ impl RoutineAction {
|
||||
prompt,
|
||||
context_paths,
|
||||
max_tokens,
|
||||
use_tools,
|
||||
max_tool_rounds,
|
||||
} => serde_json::json!({
|
||||
"prompt": prompt,
|
||||
"context_paths": context_paths,
|
||||
"max_tokens": max_tokens,
|
||||
"use_tools": use_tools,
|
||||
"max_tool_rounds": max_tool_rounds,
|
||||
}),
|
||||
RoutineAction::FullJob {
|
||||
title,
|
||||
@@ -504,7 +541,8 @@ pub fn next_cron_fire(
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::agent::routine::{
|
||||
RoutineAction, RoutineGuardrails, RunStatus, Trigger, content_hash, next_cron_fire,
|
||||
MAX_TOOL_ROUNDS_LIMIT, RoutineAction, RoutineGuardrails, RunStatus, Trigger, content_hash,
|
||||
next_cron_fire,
|
||||
};
|
||||
|
||||
#[test]
|
||||
@@ -554,11 +592,13 @@ mod tests {
|
||||
prompt: "Check PRs".to_string(),
|
||||
context_paths: vec!["context/priorities.md".to_string()],
|
||||
max_tokens: 2048,
|
||||
use_tools: false,
|
||||
max_tool_rounds: 3,
|
||||
};
|
||||
let json = action.to_config_json();
|
||||
let parsed = RoutineAction::from_db("lightweight", json).expect("parse lightweight");
|
||||
assert!(
|
||||
matches!(parsed, RoutineAction::Lightweight { prompt, context_paths, max_tokens }
|
||||
matches!(parsed, RoutineAction::Lightweight { prompt, context_paths, max_tokens, .. }
|
||||
if prompt == "Check PRs" && context_paths.len() == 1 && max_tokens == 2048)
|
||||
);
|
||||
}
|
||||
@@ -695,4 +735,77 @@ mod tests {
|
||||
);
|
||||
assert_eq!(Trigger::Manual.type_tag(), "manual");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_action_lightweight_backward_compat_no_use_tools() {
|
||||
// Simulate old DB record without use_tools field
|
||||
let json = serde_json::json!({
|
||||
"prompt": "old routine",
|
||||
"context_paths": [],
|
||||
"max_tokens": 4096
|
||||
});
|
||||
let parsed = RoutineAction::from_db("lightweight", json).expect("parse lightweight");
|
||||
assert!(
|
||||
matches!(parsed, RoutineAction::Lightweight { use_tools, max_tool_rounds, .. }
|
||||
if !use_tools && max_tool_rounds == 3),
|
||||
"missing use_tools should default to false, max_tool_rounds to 3"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_max_tool_rounds_clamped_to_upper_bound() {
|
||||
let json = serde_json::json!({
|
||||
"prompt": "test",
|
||||
"use_tools": true,
|
||||
"max_tool_rounds": 9999
|
||||
});
|
||||
let parsed = RoutineAction::from_db("lightweight", json).expect("parse");
|
||||
match parsed {
|
||||
RoutineAction::Lightweight {
|
||||
max_tool_rounds, ..
|
||||
} => {
|
||||
assert_eq!(
|
||||
max_tool_rounds, MAX_TOOL_ROUNDS_LIMIT,
|
||||
"should clamp to MAX_TOOL_ROUNDS_LIMIT"
|
||||
);
|
||||
}
|
||||
_ => panic!("expected Lightweight"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_max_tool_rounds_clamped_to_lower_bound() {
|
||||
let json = serde_json::json!({
|
||||
"prompt": "test",
|
||||
"use_tools": true,
|
||||
"max_tool_rounds": 0
|
||||
});
|
||||
let parsed = RoutineAction::from_db("lightweight", json).expect("parse");
|
||||
match parsed {
|
||||
RoutineAction::Lightweight {
|
||||
max_tool_rounds, ..
|
||||
} => {
|
||||
assert_eq!(max_tool_rounds, 1, "should clamp 0 to 1");
|
||||
}
|
||||
_ => panic!("expected Lightweight"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_max_tool_rounds_normal_value_passes_through() {
|
||||
let json = serde_json::json!({
|
||||
"prompt": "test",
|
||||
"use_tools": true,
|
||||
"max_tool_rounds": 10
|
||||
});
|
||||
let parsed = RoutineAction::from_db("lightweight", json).expect("parse");
|
||||
match parsed {
|
||||
RoutineAction::Lightweight {
|
||||
max_tool_rounds, ..
|
||||
} => {
|
||||
assert_eq!(max_tool_rounds, 10, "normal value should pass through");
|
||||
}
|
||||
_ => panic!("expected Lightweight"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+95
-25
@@ -459,7 +459,20 @@ async fn execute_routine(ctx: EngineContext, routine: Routine, run: RoutineRun)
|
||||
prompt,
|
||||
context_paths,
|
||||
max_tokens,
|
||||
} => execute_lightweight(&ctx, &routine, prompt, context_paths, *max_tokens).await,
|
||||
use_tools,
|
||||
max_tool_rounds,
|
||||
} => {
|
||||
execute_lightweight(
|
||||
&ctx,
|
||||
&routine,
|
||||
prompt,
|
||||
context_paths,
|
||||
*max_tokens,
|
||||
*use_tools,
|
||||
*max_tool_rounds,
|
||||
)
|
||||
.await
|
||||
}
|
||||
RoutineAction::FullJob {
|
||||
title,
|
||||
description,
|
||||
@@ -670,6 +683,8 @@ async fn execute_lightweight(
|
||||
prompt: &str,
|
||||
context_paths: &[String],
|
||||
max_tokens: u32,
|
||||
use_tools: bool,
|
||||
max_tool_rounds: u32,
|
||||
) -> Result<(RunStatus, Option<String>, Option<i32>), RoutineError> {
|
||||
// Load context from workspace
|
||||
let mut context_parts = Vec::new();
|
||||
@@ -732,14 +747,15 @@ async fn execute_lightweight(
|
||||
Err(_) => max_tokens,
|
||||
};
|
||||
|
||||
// If tools are enabled, use the tool execution loop; otherwise, single LLM call
|
||||
if ctx.config.lightweight_tools_enabled {
|
||||
// If tools are enabled (both globally and per-routine), use the tool execution loop
|
||||
if use_tools && ctx.config.lightweight_tools_enabled {
|
||||
execute_lightweight_with_tools(
|
||||
ctx,
|
||||
routine,
|
||||
&system_prompt,
|
||||
&full_prompt,
|
||||
effective_max_tokens,
|
||||
max_tool_rounds,
|
||||
)
|
||||
.await
|
||||
} else {
|
||||
@@ -783,24 +799,12 @@ async fn execute_lightweight_no_tools(
|
||||
reason: e.to_string(),
|
||||
})?;
|
||||
|
||||
let content = response.content.trim();
|
||||
let tokens_used = Some((response.input_tokens + response.output_tokens) as i32);
|
||||
|
||||
// Empty content guard
|
||||
if content.is_empty() {
|
||||
return if response.finish_reason == FinishReason::Length {
|
||||
Err(RoutineError::TruncatedResponse)
|
||||
} else {
|
||||
Err(RoutineError::EmptyResponse)
|
||||
};
|
||||
}
|
||||
|
||||
// Check for the "nothing to do" sentinel
|
||||
if content == "ROUTINE_OK" || content.contains("ROUTINE_OK") {
|
||||
return Ok((RunStatus::Ok, None, tokens_used));
|
||||
}
|
||||
|
||||
Ok((RunStatus::Attention, Some(content.to_string()), tokens_used))
|
||||
handle_text_response(
|
||||
&response.content,
|
||||
response.finish_reason,
|
||||
response.input_tokens,
|
||||
response.output_tokens,
|
||||
)
|
||||
}
|
||||
|
||||
/// Handle a text-only LLM response in lightweight routine execution.
|
||||
@@ -850,6 +854,7 @@ async fn execute_lightweight_with_tools(
|
||||
system_prompt: &str,
|
||||
full_prompt: &str,
|
||||
effective_max_tokens: u32,
|
||||
max_tool_rounds: u32,
|
||||
) -> Result<(RunStatus, Option<String>, Option<i32>), RoutineError> {
|
||||
let mut messages = if system_prompt.is_empty() {
|
||||
vec![ChatMessage::user(full_prompt)]
|
||||
@@ -860,7 +865,9 @@ async fn execute_lightweight_with_tools(
|
||||
]
|
||||
};
|
||||
|
||||
let max_iterations = ctx.config.lightweight_max_iterations.min(5);
|
||||
let max_iterations = max_tool_rounds
|
||||
.min(ctx.config.lightweight_max_iterations)
|
||||
.min(5);
|
||||
let mut iteration = 0;
|
||||
let mut total_input_tokens = 0;
|
||||
let mut total_output_tokens = 0;
|
||||
@@ -906,7 +913,10 @@ async fn execute_lightweight_with_tools(
|
||||
);
|
||||
} else {
|
||||
// Tool-enabled iteration
|
||||
let tool_defs = ctx.tools.tool_definitions().await;
|
||||
let tool_defs = ctx
|
||||
.tools
|
||||
.tool_definitions_excluding(ROUTINE_TOOL_DENYLIST)
|
||||
.await;
|
||||
|
||||
let request = ToolCompletionRequest::new(messages.clone(), tool_defs)
|
||||
.with_max_tokens(effective_max_tokens)
|
||||
@@ -972,12 +982,33 @@ async fn execute_lightweight_with_tools(
|
||||
}
|
||||
}
|
||||
|
||||
/// Tools that must never be callable from lightweight routines.
|
||||
///
|
||||
/// These tools pose autonomy-escalation risks: a routine could self-replicate,
|
||||
/// modify its own triggers/prompts, delete other routines, or restart the agent.
|
||||
const ROUTINE_TOOL_DENYLIST: &[&str] = &[
|
||||
"routine_create",
|
||||
"routine_update",
|
||||
"routine_delete",
|
||||
"routine_fire",
|
||||
"restart",
|
||||
];
|
||||
|
||||
/// Execute a single tool for a lightweight routine.
|
||||
async fn execute_routine_tool(
|
||||
ctx: &EngineContext,
|
||||
job_ctx: &JobContext,
|
||||
tc: &ToolCall,
|
||||
) -> Result<String, Box<dyn std::error::Error + Send + Sync>> {
|
||||
// Block tools that pose autonomy-escalation risks
|
||||
if ROUTINE_TOOL_DENYLIST.contains(&tc.name.as_str()) {
|
||||
return Err(format!(
|
||||
"Tool '{}' is not available in lightweight routines",
|
||||
tc.name
|
||||
)
|
||||
.into());
|
||||
}
|
||||
|
||||
// Check if tool exists
|
||||
let tool = ctx
|
||||
.tools
|
||||
@@ -1119,9 +1150,11 @@ pub fn spawn_cron_ticker(
|
||||
interval: Duration,
|
||||
) -> tokio::task::JoinHandle<()> {
|
||||
tokio::spawn(async move {
|
||||
// Run one check immediately so routines due at startup don't wait
|
||||
// an extra full polling interval.
|
||||
engine.check_cron_triggers().await;
|
||||
|
||||
let mut ticker = tokio::time::interval(interval);
|
||||
// Skip immediate first tick
|
||||
ticker.tick().await;
|
||||
|
||||
loop {
|
||||
ticker.tick().await;
|
||||
@@ -1283,6 +1316,36 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_routine_tool_denylist_blocks_self_management_tools() {
|
||||
let denylisted = vec![
|
||||
"routine_create",
|
||||
"routine_update",
|
||||
"routine_delete",
|
||||
"routine_fire",
|
||||
"restart",
|
||||
];
|
||||
for tool in &denylisted {
|
||||
assert!(
|
||||
super::ROUTINE_TOOL_DENYLIST.contains(tool),
|
||||
"Tool '{}' should be in ROUTINE_TOOL_DENYLIST",
|
||||
tool
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_routine_tool_denylist_allows_safe_tools() {
|
||||
let allowed = vec!["echo", "time", "json", "http", "memory_search", "shell"];
|
||||
for tool in &allowed {
|
||||
assert!(
|
||||
!super::ROUTINE_TOOL_DENYLIST.contains(tool),
|
||||
"Tool '{}' should NOT be in ROUTINE_TOOL_DENYLIST",
|
||||
tool
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_empty_response_handling() {
|
||||
// Simulate the empty content guard logic
|
||||
@@ -1297,4 +1360,11 @@ mod tests {
|
||||
assert_eq!(finish_reason_length, crate::llm::FinishReason::Length);
|
||||
assert_eq!(finish_reason_stop, crate::llm::FinishReason::Stop);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_truncate_adds_ellipsis_when_over_limit() {
|
||||
let input = "abcdefghijk";
|
||||
let out = super::truncate(input, 5);
|
||||
assert_eq!(out, "abcde...");
|
||||
}
|
||||
}
|
||||
|
||||
+18
-3
@@ -486,7 +486,12 @@ impl Agent {
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::Status("Awaiting approval".into()),
|
||||
StatusUpdate::ApprovalNeeded {
|
||||
request_id: request_id.to_string(),
|
||||
tool_name: tool_name.clone(),
|
||||
description: description.clone(),
|
||||
parameters: parameters.clone(),
|
||||
},
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
@@ -1297,7 +1302,12 @@ impl Agent {
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::Status("Awaiting approval".into()),
|
||||
StatusUpdate::ApprovalNeeded {
|
||||
request_id: request_id.to_string(),
|
||||
tool_name: tool_name.clone(),
|
||||
description: description.clone(),
|
||||
parameters: parameters.clone(),
|
||||
},
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
@@ -1368,7 +1378,12 @@ impl Agent {
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::Status("Awaiting approval".into()),
|
||||
StatusUpdate::ApprovalNeeded {
|
||||
request_id: request_id.to_string(),
|
||||
tool_name: tool_name.clone(),
|
||||
description: description.clone(),
|
||||
parameters: parameters.clone(),
|
||||
},
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
|
||||
+75
-1
@@ -9,6 +9,7 @@
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::agent::SessionManager as AgentSessionManager;
|
||||
use crate::channels::web::log_layer::LogBroadcaster;
|
||||
use crate::config::Config;
|
||||
use crate::context::ContextManager;
|
||||
@@ -46,6 +47,8 @@ pub struct AppComponents {
|
||||
pub log_broadcaster: Arc<LogBroadcaster>,
|
||||
pub context_manager: Arc<ContextManager>,
|
||||
pub hooks: Arc<HookRegistry>,
|
||||
/// Shared thread/session manager used by the standard agent runtime.
|
||||
pub agent_session_manager: Arc<AgentSessionManager>,
|
||||
pub skill_registry: Option<Arc<std::sync::RwLock<SkillRegistry>>>,
|
||||
pub skill_catalog: Option<Arc<SkillCatalog>>,
|
||||
pub cost_guard: Arc<crate::agent::cost_guard::CostGuard>,
|
||||
@@ -287,6 +290,7 @@ impl AppBuilder {
|
||||
Arc::new(ToolRegistry::new())
|
||||
};
|
||||
tools.register_builtin_tools();
|
||||
tools.register_tool_info();
|
||||
|
||||
if let Some(ref ss) = self.secrets_store {
|
||||
tools.register_secrets_tools(Arc::clone(ss));
|
||||
@@ -300,7 +304,8 @@ impl AppBuilder {
|
||||
|
||||
// Register memory tools if database is available
|
||||
let workspace = if let Some(ref db) = self.db {
|
||||
let mut ws = Workspace::new_with_db("default", db.clone());
|
||||
let mut ws = Workspace::new_with_db("default", db.clone())
|
||||
.with_search_config(&self.config.search);
|
||||
if let Some(ref emb) = embeddings {
|
||||
ws = ws.with_embeddings(emb.clone());
|
||||
}
|
||||
@@ -689,6 +694,8 @@ impl AppBuilder {
|
||||
|
||||
// Create hook registry early so runtime extension activation can register hooks.
|
||||
let hooks = Arc::new(HookRegistry::new());
|
||||
let agent_session_manager =
|
||||
Arc::new(AgentSessionManager::new().with_hooks(Arc::clone(&hooks)));
|
||||
|
||||
let (
|
||||
mcp_session_manager,
|
||||
@@ -795,6 +802,7 @@ impl AppBuilder {
|
||||
log_broadcaster: self.log_broadcaster,
|
||||
context_manager,
|
||||
hooks,
|
||||
agent_session_manager,
|
||||
skill_registry,
|
||||
skill_catalog,
|
||||
cost_guard,
|
||||
@@ -805,3 +813,69 @@ impl AppBuilder {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
use crate::agent::SessionManager as AgentSessionManager;
|
||||
use crate::hooks::{
|
||||
Hook, HookContext, HookError, HookEvent, HookOutcome, HookPoint, HookRegistry,
|
||||
};
|
||||
|
||||
struct SessionStartHook {
|
||||
tx: mpsc::UnboundedSender<(String, String)>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Hook for SessionStartHook {
|
||||
fn name(&self) -> &str {
|
||||
"session-start-test"
|
||||
}
|
||||
|
||||
fn hook_points(&self) -> &[HookPoint] {
|
||||
&[HookPoint::OnSessionStart]
|
||||
}
|
||||
|
||||
async fn execute(
|
||||
&self,
|
||||
event: &HookEvent,
|
||||
_ctx: &HookContext,
|
||||
) -> Result<HookOutcome, HookError> {
|
||||
if let HookEvent::SessionStart {
|
||||
user_id,
|
||||
session_id,
|
||||
} = event
|
||||
{
|
||||
self.tx
|
||||
.send((user_id.clone(), session_id.clone()))
|
||||
.expect("test channel receiver should be alive");
|
||||
} else {
|
||||
panic!("SessionStartHook received an unexpected event: {event:?}");
|
||||
}
|
||||
Ok(HookOutcome::ok())
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn agent_session_manager_runs_session_start_hooks() {
|
||||
let hooks = Arc::new(HookRegistry::new());
|
||||
let (tx, mut rx) = mpsc::unbounded_channel();
|
||||
hooks.register(Arc::new(SessionStartHook { tx })).await;
|
||||
|
||||
let manager = AgentSessionManager::new().with_hooks(Arc::clone(&hooks));
|
||||
manager.get_or_create_session("user-123").await;
|
||||
|
||||
let (user_id, session_id) =
|
||||
tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv())
|
||||
.await
|
||||
.expect("session start hook should fire")
|
||||
.expect("session start payload should be present");
|
||||
|
||||
assert_eq!(user_id, "user-123");
|
||||
assert!(!session_id.is_empty());
|
||||
}
|
||||
}
|
||||
|
||||
+173
-76
@@ -269,95 +269,105 @@ async fn webhook_handler(
|
||||
let mut fallback_req = None;
|
||||
{
|
||||
let webhook_secret = state.webhook_secret.read().await;
|
||||
if let Some(expected_secret) = webhook_secret.as_ref() {
|
||||
let expected_secret = expected_secret.expose_secret();
|
||||
let Some(expected_secret) = webhook_secret.as_ref() else {
|
||||
return (
|
||||
StatusCode::UNAUTHORIZED,
|
||||
Json(WebhookResponse {
|
||||
message_id: Uuid::nil(),
|
||||
status: "error".to_string(),
|
||||
response: Some(
|
||||
"Webhook authentication required: HTTP webhook secret is not configured."
|
||||
.to_string(),
|
||||
),
|
||||
}),
|
||||
)
|
||||
.into_response();
|
||||
};
|
||||
let expected_secret = expected_secret.expose_secret();
|
||||
|
||||
match headers.get("x-ironclaw-signature") {
|
||||
Some(raw_signature) => match raw_signature.to_str() {
|
||||
Ok(signature) => {
|
||||
if !verify_hmac_signature(expected_secret, &body, signature) {
|
||||
return (
|
||||
StatusCode::UNAUTHORIZED,
|
||||
Json(WebhookResponse {
|
||||
message_id: Uuid::nil(),
|
||||
status: "error".to_string(),
|
||||
response: Some("Invalid webhook signature".to_string()),
|
||||
}),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
match headers.get("x-ironclaw-signature") {
|
||||
Some(raw_signature) => match raw_signature.to_str() {
|
||||
Ok(signature) => {
|
||||
if !verify_hmac_signature(expected_secret, &body, signature) {
|
||||
return (
|
||||
StatusCode::UNAUTHORIZED,
|
||||
Json(WebhookResponse {
|
||||
message_id: Uuid::nil(),
|
||||
status: "error".to_string(),
|
||||
response: Some("Invalid webhook signature".to_string()),
|
||||
}),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
}
|
||||
Err(_) => {
|
||||
return (
|
||||
StatusCode::UNAUTHORIZED,
|
||||
Json(WebhookResponse {
|
||||
message_id: Uuid::nil(),
|
||||
status: "error".to_string(),
|
||||
response: Some("Invalid signature header encoding".to_string()),
|
||||
}),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
},
|
||||
None => {
|
||||
let req: WebhookRequest = match serde_json::from_slice(&body) {
|
||||
Ok(req) => req,
|
||||
Err(_) => {
|
||||
return (
|
||||
StatusCode::UNAUTHORIZED,
|
||||
Json(WebhookResponse {
|
||||
message_id: Uuid::nil(),
|
||||
status: "error".to_string(),
|
||||
response: Some("Invalid signature header encoding".to_string()),
|
||||
response: Some(
|
||||
"Webhook authentication required. Provide X-IronClaw-Signature header \
|
||||
(preferred) or 'secret' field in body (deprecated)."
|
||||
.to_string(),
|
||||
),
|
||||
}),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
},
|
||||
None => {
|
||||
let req: WebhookRequest = match serde_json::from_slice(&body) {
|
||||
Ok(req) => req,
|
||||
Err(_) => {
|
||||
return (
|
||||
StatusCode::UNAUTHORIZED,
|
||||
Json(WebhookResponse {
|
||||
message_id: Uuid::nil(),
|
||||
status: "error".to_string(),
|
||||
response: Some(
|
||||
"Webhook authentication required. Provide X-IronClaw-Signature header \
|
||||
(preferred) or 'secret' field in body (deprecated)."
|
||||
.to_string(),
|
||||
),
|
||||
}),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
};
|
||||
};
|
||||
|
||||
match &req.secret {
|
||||
Some(provided)
|
||||
if bool::from(
|
||||
provided.as_bytes().ct_eq(expected_secret.as_bytes()),
|
||||
) =>
|
||||
{
|
||||
tracing::warn!(
|
||||
"Webhook authenticated via deprecated 'secret' field in request body. \
|
||||
Migrate to X-IronClaw-Signature header (HMAC-SHA256). \
|
||||
Body secret support will be removed in a future release."
|
||||
);
|
||||
fallback_req = Some(req);
|
||||
}
|
||||
Some(_) => {
|
||||
return (
|
||||
StatusCode::UNAUTHORIZED,
|
||||
Json(WebhookResponse {
|
||||
message_id: Uuid::nil(),
|
||||
status: "error".to_string(),
|
||||
response: Some("Invalid webhook secret".to_string()),
|
||||
}),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
None => {
|
||||
return (
|
||||
StatusCode::UNAUTHORIZED,
|
||||
Json(WebhookResponse {
|
||||
message_id: Uuid::nil(),
|
||||
status: "error".to_string(),
|
||||
response: Some(
|
||||
"Webhook authentication required. Provide X-IronClaw-Signature header \
|
||||
(preferred) or 'secret' field in body (deprecated)."
|
||||
.to_string(),
|
||||
),
|
||||
}),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
match &req.secret {
|
||||
Some(provided)
|
||||
if bool::from(provided.as_bytes().ct_eq(expected_secret.as_bytes())) =>
|
||||
{
|
||||
tracing::warn!(
|
||||
"Webhook authenticated via deprecated 'secret' field in request body. \
|
||||
Migrate to X-IronClaw-Signature header (HMAC-SHA256). \
|
||||
Body secret support will be removed in a future release."
|
||||
);
|
||||
fallback_req = Some(req);
|
||||
}
|
||||
Some(_) => {
|
||||
return (
|
||||
StatusCode::UNAUTHORIZED,
|
||||
Json(WebhookResponse {
|
||||
message_id: Uuid::nil(),
|
||||
status: "error".to_string(),
|
||||
response: Some("Invalid webhook secret".to_string()),
|
||||
}),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
None => {
|
||||
return (
|
||||
StatusCode::UNAUTHORIZED,
|
||||
Json(WebhookResponse {
|
||||
message_id: Uuid::nil(),
|
||||
status: "error".to_string(),
|
||||
response: Some(
|
||||
"Webhook authentication required. Provide X-IronClaw-Signature header \
|
||||
(preferred) or 'secret' field in body (deprecated)."
|
||||
.to_string(),
|
||||
),
|
||||
}),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -807,6 +817,67 @@ mod tests {
|
||||
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
|
||||
}
|
||||
|
||||
/// Regression test for issue #869: RwLock read guard was held across
|
||||
/// tx.send(msg).await in `process_message()`, blocking shutdown() from
|
||||
/// acquiring the write lock when the channel buffer was full.
|
||||
///
|
||||
/// This test exercises the actual production code path (`process_message`)
|
||||
/// with a full channel buffer, then verifies shutdown() can still complete.
|
||||
#[tokio::test]
|
||||
async fn shutdown_completes_while_process_message_blocked() {
|
||||
let channel = Arc::new(test_channel(Some("secret")));
|
||||
let stream = channel.start().await.unwrap();
|
||||
|
||||
// Fill all 256 slots in the channel buffer
|
||||
{
|
||||
let tx = {
|
||||
let guard = channel.state.tx.read().await;
|
||||
guard.as_ref().unwrap().clone()
|
||||
};
|
||||
for i in 0..256 {
|
||||
let msg = IncomingMessage::new("http", "user", format!("fill-{}", i));
|
||||
tx.send(msg).await.unwrap();
|
||||
}
|
||||
}
|
||||
|
||||
// Signal so we know the spawned task has started and is about to
|
||||
// call process_message (which will block on the full channel).
|
||||
let started = Arc::new(tokio::sync::Notify::new());
|
||||
let started_clone = started.clone();
|
||||
|
||||
// Spawn a task that calls the actual production code path.
|
||||
// process_message() internally acquires the RwLock read guard and
|
||||
// sends on the channel. With the fix, the guard is released before
|
||||
// send().await; without the fix, shutdown() would deadlock.
|
||||
let state = channel.state.clone();
|
||||
let blocked_send = tokio::spawn(async move {
|
||||
started_clone.notify_one();
|
||||
let msg = IncomingMessage::new("http", "user", "blocked-257th");
|
||||
let _ = process_message(state, msg, false).await;
|
||||
});
|
||||
|
||||
// Wait for the spawned task to start, then give it time to reach
|
||||
// the send().await and verify that it is still pending (i.e., blocked).
|
||||
started.notified().await;
|
||||
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
|
||||
assert!(
|
||||
!blocked_send.is_finished(),
|
||||
"process_message task should still be pending before shutdown()"
|
||||
);
|
||||
|
||||
// shutdown() must complete even though process_message is blocked on
|
||||
// send(). Before the fix, the read guard held across send().await
|
||||
// would prevent shutdown() from acquiring the write lock.
|
||||
let result =
|
||||
tokio::time::timeout(std::time::Duration::from_secs(2), channel.shutdown()).await;
|
||||
assert!(result.is_ok(), "shutdown() must not deadlock");
|
||||
assert!(result.unwrap().is_ok());
|
||||
|
||||
// Drop the stream (receiver) so the blocked send task can complete
|
||||
drop(stream);
|
||||
let _ = blocked_send.await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn webhook_missing_all_auth_returns_unauthorized() {
|
||||
let channel = test_channel(Some("correct-secret"));
|
||||
@@ -991,6 +1062,32 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn webhook_rejects_requests_after_secret_is_cleared() {
|
||||
let secret = "test-secret-123";
|
||||
let channel = test_channel(Some(secret));
|
||||
let _stream = channel.start().await.unwrap();
|
||||
let app = channel.routes();
|
||||
|
||||
channel.update_secret(None).await;
|
||||
|
||||
let body = serde_json::json!({
|
||||
"content": "hello"
|
||||
});
|
||||
let body_bytes = serde_json::to_vec(&body).unwrap();
|
||||
let signature = compute_signature(secret, &body_bytes);
|
||||
let req = Request::builder()
|
||||
.method("POST")
|
||||
.uri("/webhook")
|
||||
.header("content-type", "application/json")
|
||||
.header("x-ironclaw-signature", signature)
|
||||
.body(Body::from(body_bytes))
|
||||
.unwrap();
|
||||
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_concurrent_requests_during_secret_update() {
|
||||
use std::sync::Arc as StdArc;
|
||||
|
||||
@@ -408,12 +408,120 @@ impl Channel for RelayChannel {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Status updates are not forwarded to messaging providers to avoid noise.
|
||||
async fn send_status(
|
||||
&self,
|
||||
_status: StatusUpdate,
|
||||
_metadata: &serde_json::Value,
|
||||
status: StatusUpdate,
|
||||
metadata: &serde_json::Value,
|
||||
) -> Result<(), ChannelError> {
|
||||
// Only handle ApprovalNeeded — all other variants are no-ops
|
||||
let StatusUpdate::ApprovalNeeded {
|
||||
request_id,
|
||||
tool_name,
|
||||
description,
|
||||
parameters,
|
||||
} = status
|
||||
else {
|
||||
return Ok(());
|
||||
};
|
||||
|
||||
// Only send buttons in DMs (dispatcher gates upstream, but guard here too)
|
||||
let event_type = metadata
|
||||
.get("event_type")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("");
|
||||
if event_type != "direct_message" {
|
||||
tracing::warn!(
|
||||
tool = %tool_name,
|
||||
event_type,
|
||||
"Approval requested in non-DM, skipping buttons"
|
||||
);
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
// Extract required metadata — error if missing
|
||||
let channel_id = metadata
|
||||
.get("channel_id")
|
||||
.and_then(|v| v.as_str())
|
||||
.ok_or_else(|| ChannelError::SendFailed {
|
||||
name: self.name().to_string(),
|
||||
reason: "Missing channel_id for approval buttons".into(),
|
||||
})?;
|
||||
let sender_id = metadata
|
||||
.get("sender_id")
|
||||
.and_then(|v| v.as_str())
|
||||
.ok_or_else(|| ChannelError::SendFailed {
|
||||
name: self.name().to_string(),
|
||||
reason: "Missing sender_id for approval buttons".into(),
|
||||
})?;
|
||||
let thread_id = metadata.get("thread_id").and_then(|v| v.as_str());
|
||||
let team_id = metadata
|
||||
.get("team_id")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or(&self.team_id);
|
||||
|
||||
// Button value payload (Slack limits button values to 2000 chars;
|
||||
// safe with typical UUIDs but documented here as a constraint)
|
||||
let value_payload = serde_json::json!({
|
||||
"instance_id": self.instance_id,
|
||||
"team_id": team_id,
|
||||
"channel_id": channel_id,
|
||||
"thread_ts": thread_id,
|
||||
"request_id": request_id,
|
||||
"sender_id": sender_id,
|
||||
});
|
||||
let value_str = value_payload.to_string();
|
||||
|
||||
// Parameters are already redacted via redact_params() in dispatcher.rs
|
||||
let params_display =
|
||||
serde_json::to_string_pretty(¶meters).unwrap_or_else(|_| parameters.to_string());
|
||||
|
||||
let blocks = serde_json::json!([
|
||||
{
|
||||
"type": "section",
|
||||
"text": {
|
||||
"type": "mrkdwn",
|
||||
"text": format!(
|
||||
"*Tool approval required*\n`{tool_name}`: {description}\n```{params_display}```"
|
||||
)
|
||||
}
|
||||
},
|
||||
{
|
||||
"type": "actions",
|
||||
"elements": [
|
||||
{
|
||||
"type": "button",
|
||||
"text": { "type": "plain_text", "text": "Approve" },
|
||||
"style": "primary",
|
||||
"action_id": "approve_tool",
|
||||
"value": value_str,
|
||||
},
|
||||
{
|
||||
"type": "button",
|
||||
"text": { "type": "plain_text", "text": "Deny" },
|
||||
"style": "danger",
|
||||
"action_id": "deny_tool",
|
||||
"value": value_str,
|
||||
}
|
||||
]
|
||||
}
|
||||
]);
|
||||
|
||||
let mut body = serde_json::json!({
|
||||
"channel": channel_id,
|
||||
"text": format!("Tool approval required: {tool_name} - {description}"),
|
||||
"blocks": blocks,
|
||||
});
|
||||
if let Some(tid) = thread_id {
|
||||
body["thread_ts"] = serde_json::Value::String(tid.to_string());
|
||||
}
|
||||
|
||||
self.proxy_send(team_id, "chat.postMessage", body)
|
||||
.await
|
||||
.map_err(|e| ChannelError::SendFailed {
|
||||
name: self.name().to_string(),
|
||||
reason: e.to_string(),
|
||||
})?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -639,4 +747,118 @@ mod tests {
|
||||
// The reconnect loop now skips team validation when team_id is empty,
|
||||
// so the channel remains alive.
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_send_status_non_approval_is_noop() {
|
||||
let channel = RelayChannel::new(
|
||||
test_client(),
|
||||
"token".into(),
|
||||
"T123".into(),
|
||||
"inst1".into(),
|
||||
"user1".into(),
|
||||
);
|
||||
let metadata = serde_json::json!({});
|
||||
let result = channel
|
||||
.send_status(
|
||||
StatusUpdate::ToolStarted {
|
||||
name: "echo".into(),
|
||||
},
|
||||
&metadata,
|
||||
)
|
||||
.await;
|
||||
assert!(result.is_ok());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_send_status_approval_non_dm_skips() {
|
||||
let channel = RelayChannel::new(
|
||||
test_client(),
|
||||
"token".into(),
|
||||
"T123".into(),
|
||||
"inst1".into(),
|
||||
"user1".into(),
|
||||
);
|
||||
let metadata = serde_json::json!({
|
||||
"event_type": "message",
|
||||
"channel_id": "C456",
|
||||
"sender_id": "U789",
|
||||
});
|
||||
let result = channel
|
||||
.send_status(
|
||||
StatusUpdate::ApprovalNeeded {
|
||||
request_id: "req1".into(),
|
||||
tool_name: "shell".into(),
|
||||
description: "run command".into(),
|
||||
parameters: serde_json::json!({}),
|
||||
},
|
||||
&metadata,
|
||||
)
|
||||
.await;
|
||||
// Non-DM approval requests are silently skipped (no HTTP call)
|
||||
assert!(result.is_ok());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_send_status_approval_dm_missing_channel_id_errors() {
|
||||
let channel = RelayChannel::new(
|
||||
test_client(),
|
||||
"token".into(),
|
||||
"T123".into(),
|
||||
"inst1".into(),
|
||||
"user1".into(),
|
||||
);
|
||||
let metadata = serde_json::json!({
|
||||
"event_type": "direct_message",
|
||||
"sender_id": "U789",
|
||||
});
|
||||
let result = channel
|
||||
.send_status(
|
||||
StatusUpdate::ApprovalNeeded {
|
||||
request_id: "req1".into(),
|
||||
tool_name: "shell".into(),
|
||||
description: "run command".into(),
|
||||
parameters: serde_json::json!({}),
|
||||
},
|
||||
&metadata,
|
||||
)
|
||||
.await;
|
||||
assert!(result.is_err());
|
||||
let err = result.unwrap_err().to_string();
|
||||
assert!(
|
||||
err.contains("channel_id"),
|
||||
"expected channel_id error, got: {err}"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_send_status_approval_dm_missing_sender_id_errors() {
|
||||
let channel = RelayChannel::new(
|
||||
test_client(),
|
||||
"token".into(),
|
||||
"T123".into(),
|
||||
"inst1".into(),
|
||||
"user1".into(),
|
||||
);
|
||||
let metadata = serde_json::json!({
|
||||
"event_type": "direct_message",
|
||||
"channel_id": "C456",
|
||||
});
|
||||
let result = channel
|
||||
.send_status(
|
||||
StatusUpdate::ApprovalNeeded {
|
||||
request_id: "req1".into(),
|
||||
tool_name: "shell".into(),
|
||||
description: "run command".into(),
|
||||
parameters: serde_json::json!({}),
|
||||
},
|
||||
&metadata,
|
||||
)
|
||||
.await;
|
||||
assert!(result.is_err());
|
||||
let err = result.unwrap_err().to_string();
|
||||
assert!(
|
||||
err.contains("sender_id"),
|
||||
"expected sender_id error, got: {err}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -63,7 +63,11 @@ const ALLOWED_MIME_PREFIXES: &[&str] = &[
|
||||
"application/x-tar",
|
||||
"application/octet-stream",
|
||||
];
|
||||
|
||||
/// Truncate a string to at most `max_bytes` without splitting UTF-8 code points.
|
||||
fn truncate_utf8(s: &str, max_bytes: usize) -> &str {
|
||||
let end = crate::util::floor_char_boundary(s, max_bytes);
|
||||
&s[..end]
|
||||
}
|
||||
/// A message emitted by a WASM channel to be sent to the agent.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct EmittedMessage {
|
||||
@@ -264,7 +268,7 @@ impl ChannelHostState {
|
||||
max = MAX_MESSAGE_CONTENT_SIZE,
|
||||
"Message content too large, truncating"
|
||||
);
|
||||
let mut truncated = msg.content[..MAX_MESSAGE_CONTENT_SIZE].to_string();
|
||||
let mut truncated = truncate_utf8(&msg.content, MAX_MESSAGE_CONTENT_SIZE).to_string();
|
||||
truncated.push_str("... (truncated)");
|
||||
let msg = EmittedMessage {
|
||||
content: truncated,
|
||||
@@ -631,6 +635,7 @@ mod tests {
|
||||
use crate::channels::wasm::host::{
|
||||
Attachment, ChannelEmitRateLimiter, ChannelHostState, EmittedMessage,
|
||||
MAX_ATTACHMENT_TOTAL_SIZE, MAX_ATTACHMENTS_PER_MESSAGE, MAX_EMITS_PER_EXECUTION,
|
||||
MAX_MESSAGE_CONTENT_SIZE,
|
||||
};
|
||||
|
||||
#[test]
|
||||
@@ -689,6 +694,25 @@ mod tests {
|
||||
assert_eq!(state.emits_dropped(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_emit_message_truncates_utf8_safely() {
|
||||
let caps = ChannelCapabilities::for_channel("test");
|
||||
let mut state = ChannelHostState::new("test", caps);
|
||||
|
||||
let prefix = "a".repeat(MAX_MESSAGE_CONTENT_SIZE - 1);
|
||||
let content = format!("{}🙂suffix", prefix);
|
||||
let msg = EmittedMessage::new("user123", content);
|
||||
|
||||
state.emit_message(msg).unwrap();
|
||||
let messages = state.take_emitted_messages();
|
||||
assert_eq!(messages.len(), 1);
|
||||
|
||||
let emitted = &messages[0].content;
|
||||
assert!(emitted.starts_with(&prefix));
|
||||
assert!(emitted.ends_with("... (truncated)"));
|
||||
assert!(!emitted.contains("🙂"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_workspace_write_prefixing() {
|
||||
let caps = ChannelCapabilities::for_channel("slack");
|
||||
|
||||
@@ -1994,28 +1994,33 @@ impl WasmChannel {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let tx_guard = self.message_tx.read().await;
|
||||
let Some(tx) = tx_guard.as_ref() else {
|
||||
tracing::error!(
|
||||
channel = %self.name,
|
||||
count = messages.len(),
|
||||
"Messages emitted but no sender available - channel may not be started!"
|
||||
);
|
||||
return Ok(());
|
||||
// Clone sender to avoid holding RwLock read guard across send().await in the loop
|
||||
let tx = {
|
||||
let tx_guard = self.message_tx.read().await;
|
||||
let Some(tx) = tx_guard.as_ref() else {
|
||||
tracing::error!(
|
||||
channel = %self.name,
|
||||
count = messages.len(),
|
||||
"Messages emitted but no sender available - channel may not be started!"
|
||||
);
|
||||
return Ok(());
|
||||
};
|
||||
tx.clone()
|
||||
};
|
||||
|
||||
let mut rate_limiter = self.rate_limiter.write().await;
|
||||
|
||||
for emitted in messages {
|
||||
// Check rate limit
|
||||
if !rate_limiter.check_and_record() {
|
||||
tracing::warn!(
|
||||
channel = %self.name,
|
||||
"Message emission rate limited"
|
||||
);
|
||||
return Err(WasmChannelError::EmitRateLimited {
|
||||
name: self.name.clone(),
|
||||
});
|
||||
// Check rate limit — acquire and release the write lock before send().await
|
||||
{
|
||||
let mut rate_limiter = self.rate_limiter.write().await;
|
||||
if !rate_limiter.check_and_record() {
|
||||
tracing::warn!(
|
||||
channel = %self.name,
|
||||
"Message emission rate limited"
|
||||
);
|
||||
return Err(WasmChannelError::EmitRateLimited {
|
||||
name: self.name.clone(),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
// Convert to IncomingMessage
|
||||
@@ -2057,7 +2062,7 @@ impl WasmChannel {
|
||||
self.update_broadcast_metadata(&emitted.metadata_json).await;
|
||||
}
|
||||
|
||||
// Send to stream
|
||||
// Send to stream — no locks held across this await
|
||||
tracing::info!(
|
||||
channel = %self.name,
|
||||
user_id = %emitted.user_id,
|
||||
@@ -2281,28 +2286,33 @@ impl WasmChannel {
|
||||
"Processing emitted messages from polling callback"
|
||||
);
|
||||
|
||||
let tx_guard = message_tx.read().await;
|
||||
let Some(tx) = tx_guard.as_ref() else {
|
||||
tracing::error!(
|
||||
channel = %channel_name,
|
||||
count = messages.len(),
|
||||
"Messages emitted but no sender available - channel may not be started!"
|
||||
);
|
||||
return Ok(());
|
||||
// Clone sender to avoid holding RwLock read guard across send().await in the loop
|
||||
let tx = {
|
||||
let tx_guard = message_tx.read().await;
|
||||
let Some(tx) = tx_guard.as_ref() else {
|
||||
tracing::error!(
|
||||
channel = %channel_name,
|
||||
count = messages.len(),
|
||||
"Messages emitted but no sender available - channel may not be started!"
|
||||
);
|
||||
return Ok(());
|
||||
};
|
||||
tx.clone()
|
||||
};
|
||||
|
||||
let mut limiter = rate_limiter.write().await;
|
||||
|
||||
for emitted in messages {
|
||||
// Check rate limit
|
||||
if !limiter.check_and_record() {
|
||||
tracing::warn!(
|
||||
channel = %channel_name,
|
||||
"Message emission rate limited"
|
||||
);
|
||||
return Err(WasmChannelError::EmitRateLimited {
|
||||
name: channel_name.to_string(),
|
||||
});
|
||||
// Check rate limit — acquire and release the write lock before send().await
|
||||
{
|
||||
let mut limiter = rate_limiter.write().await;
|
||||
if !limiter.check_and_record() {
|
||||
tracing::warn!(
|
||||
channel = %channel_name,
|
||||
"Message emission rate limited"
|
||||
);
|
||||
return Err(WasmChannelError::EmitRateLimited {
|
||||
name: channel_name.to_string(),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
// Convert to IncomingMessage
|
||||
@@ -2350,7 +2360,7 @@ impl WasmChannel {
|
||||
.await;
|
||||
}
|
||||
|
||||
// Send to stream
|
||||
// Send to stream — no locks held across this await
|
||||
tracing::info!(
|
||||
channel = %channel_name,
|
||||
user_id = %emitted.user_id,
|
||||
|
||||
@@ -37,11 +37,17 @@ pub async fn chat_send_handler(
|
||||
let msg_id = msg.id;
|
||||
let thread_id = msg.thread_id.clone();
|
||||
|
||||
let tx_guard = state.msg_tx.read().await;
|
||||
let tx = tx_guard.as_ref().ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Channel not started".to_string(),
|
||||
))?;
|
||||
// Clone sender to avoid holding RwLock read guard across send().await
|
||||
let tx = {
|
||||
let tx_guard = state.msg_tx.read().await;
|
||||
tx_guard
|
||||
.as_ref()
|
||||
.ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Channel not started".to_string(),
|
||||
))?
|
||||
.clone()
|
||||
};
|
||||
|
||||
tx.send(msg).await.map_err(|_| {
|
||||
(
|
||||
@@ -111,11 +117,17 @@ pub async fn chat_approval_handler(
|
||||
|
||||
let msg_id = msg.id;
|
||||
|
||||
let tx_guard = state.msg_tx.read().await;
|
||||
let tx = tx_guard.as_ref().ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Channel not started".to_string(),
|
||||
))?;
|
||||
// Clone sender to avoid holding RwLock read guard across send().await
|
||||
let tx = {
|
||||
let tx_guard = state.msg_tx.read().await;
|
||||
tx_guard
|
||||
.as_ref()
|
||||
.ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Channel not started".to_string(),
|
||||
))?
|
||||
.clone()
|
||||
};
|
||||
|
||||
tx.send(msg).await.map_err(|_| {
|
||||
(
|
||||
|
||||
@@ -10,6 +10,7 @@ use axum::{
|
||||
use serde::Deserialize;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::agent::routine::{Trigger, next_cron_fire};
|
||||
use crate::channels::web::server::GatewayState;
|
||||
use crate::channels::web::types::*;
|
||||
use crate::error::RoutineError;
|
||||
@@ -182,12 +183,21 @@ pub async fn routines_toggle_handler(
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
|
||||
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
|
||||
|
||||
let was_enabled = routine.enabled;
|
||||
// If a specific value was provided, use it; otherwise toggle.
|
||||
routine.enabled = match body {
|
||||
Some(Json(req)) => req.enabled.unwrap_or(!routine.enabled),
|
||||
None => !routine.enabled,
|
||||
};
|
||||
|
||||
if routine.enabled
|
||||
&& !was_enabled
|
||||
&& let Trigger::Cron { schedule, timezone } = &routine.trigger
|
||||
{
|
||||
routine.next_fire_at = next_cron_fire(schedule, timezone.as_deref())
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
}
|
||||
|
||||
store
|
||||
.update_routine(&routine)
|
||||
.await
|
||||
|
||||
+204
-85
@@ -26,6 +26,7 @@ use tower_http::set_header::SetResponseHeaderLayer;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::agent::SessionManager;
|
||||
use crate::agent::routine::{Trigger, next_cron_fire};
|
||||
use crate::bootstrap::ironclaw_base_dir;
|
||||
use crate::channels::IncomingMessage;
|
||||
use crate::channels::relay::DEFAULT_RELAY_NAME;
|
||||
@@ -572,6 +573,14 @@ async fn oauth_callback_handler(
|
||||
extension = %flow.extension_name,
|
||||
"OAuth flow expired"
|
||||
);
|
||||
// Notify UI so auth card can show error instead of staying stuck
|
||||
if let Some(ref sender) = flow.sse_sender {
|
||||
let _ = sender.send(SseEvent::AuthCompleted {
|
||||
extension_name: flow.extension_name.clone(),
|
||||
success: false,
|
||||
message: "OAuth flow expired. Please try again.".to_string(),
|
||||
});
|
||||
}
|
||||
return oauth_error_page(&flow.display_name);
|
||||
}
|
||||
|
||||
@@ -581,7 +590,12 @@ async fn oauth_callback_handler(
|
||||
let exchange_proxy_url = std::env::var("IRONCLAW_OAUTH_EXCHANGE_URL").ok();
|
||||
|
||||
let result: Result<(), String> = async {
|
||||
let token_response = if let Some(ref proxy_url) = exchange_proxy_url {
|
||||
let token_response = if let (Some(proxy_url), None) = (&exchange_proxy_url, &flow.resource)
|
||||
{
|
||||
// Use the platform exchange proxy when configured and no resource
|
||||
// parameter is needed. The proxy holds client_secret server-side so
|
||||
// the container never sees it. MCP flows (resource.is_some()) bypass
|
||||
// the proxy because it doesn't forward the RFC 8707 resource param.
|
||||
let gateway_token = flow.gateway_token.as_deref().unwrap_or_default();
|
||||
oauth_defaults::exchange_via_proxy(
|
||||
proxy_url,
|
||||
@@ -594,7 +608,10 @@ async fn oauth_callback_handler(
|
||||
.await
|
||||
.map_err(|e| e.to_string())?
|
||||
} else {
|
||||
oauth_defaults::exchange_oauth_code(
|
||||
// Direct token exchange: uses exchange_oauth_code_with_resource so MCP
|
||||
// flows can include the RFC 8707 `resource` parameter to scope the
|
||||
// issued token to the specific MCP server.
|
||||
oauth_defaults::exchange_oauth_code_with_resource(
|
||||
&flow.token_url,
|
||||
&flow.client_id,
|
||||
flow.client_secret.as_deref(),
|
||||
@@ -602,6 +619,7 @@ async fn oauth_callback_handler(
|
||||
&flow.redirect_uri,
|
||||
flow.code_verifier.as_deref(),
|
||||
&flow.access_token_field,
|
||||
flow.resource.as_deref(),
|
||||
)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?
|
||||
@@ -628,6 +646,19 @@ async fn oauth_callback_handler(
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
// For MCP OAuth flows (identified by resource field), persist the
|
||||
// client_id so token refresh works without re-authentication.
|
||||
// The CLI flow stores this in authorize_mcp_server(); the gateway
|
||||
// callback must do the same.
|
||||
if let Some(ref client_id_secret) = flow.client_id_secret_name {
|
||||
let params = crate::secrets::CreateSecretParams::new(client_id_secret, &flow.client_id)
|
||||
.with_provider(flow.provider.as_ref().cloned().unwrap_or_default());
|
||||
flow.secrets
|
||||
.create(&flow.user_id, params)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
.await;
|
||||
@@ -659,12 +690,35 @@ async fn oauth_callback_handler(
|
||||
}
|
||||
}
|
||||
|
||||
// After successful OAuth, auto-activate the extension so it moves
|
||||
// from "Installed (Authenticate)" → "Active" without a second click.
|
||||
// OAuth success is independent of activation — tokens are already stored.
|
||||
// Report auth as successful and attempt activation as a bonus step.
|
||||
let final_message = if success {
|
||||
match ext_mgr.activate(&flow.extension_name).await {
|
||||
Ok(result) => result.message,
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
extension = %flow.extension_name,
|
||||
error = %e,
|
||||
"Auto-activation after OAuth failed"
|
||||
);
|
||||
format!(
|
||||
"{} authenticated successfully. Activation failed: {}. Try activating manually.",
|
||||
flow.display_name, e
|
||||
)
|
||||
}
|
||||
}
|
||||
} else {
|
||||
message
|
||||
};
|
||||
|
||||
// Broadcast SSE event to notify the web UI
|
||||
if let Some(ref sender) = flow.sse_sender {
|
||||
let _ = sender.send(SseEvent::AuthCompleted {
|
||||
extension_name: flow.extension_name,
|
||||
success,
|
||||
message,
|
||||
message: final_message.clone(),
|
||||
});
|
||||
}
|
||||
|
||||
@@ -973,11 +1027,17 @@ async fn chat_send_handler(
|
||||
req.images.len()
|
||||
);
|
||||
|
||||
let tx_guard = state.msg_tx.read().await;
|
||||
let tx = tx_guard.as_ref().ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Channel not started".to_string(),
|
||||
))?;
|
||||
// Clone sender to avoid holding RwLock read guard across send().await
|
||||
let tx = {
|
||||
let tx_guard = state.msg_tx.read().await;
|
||||
tx_guard
|
||||
.as_ref()
|
||||
.ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Channel not started".to_string(),
|
||||
))?
|
||||
.clone()
|
||||
};
|
||||
|
||||
tracing::debug!("[chat_send_handler] Sending message through channel");
|
||||
tx.send(msg).await.map_err(|_| {
|
||||
@@ -1043,11 +1103,17 @@ async fn chat_approval_handler(
|
||||
|
||||
let msg_id = msg.id;
|
||||
|
||||
let tx_guard = state.msg_tx.read().await;
|
||||
let tx = tx_guard.as_ref().ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Channel not started".to_string(),
|
||||
))?;
|
||||
// Clone sender to avoid holding RwLock read guard across send().await
|
||||
let tx = {
|
||||
let tx_guard = state.msg_tx.read().await;
|
||||
tx_guard
|
||||
.as_ref()
|
||||
.ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Channel not started".to_string(),
|
||||
))?
|
||||
.clone()
|
||||
};
|
||||
|
||||
tx.send(msg).await.map_err(|_| {
|
||||
(
|
||||
@@ -2359,12 +2425,21 @@ async fn routines_toggle_handler(
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
|
||||
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
|
||||
|
||||
let was_enabled = routine.enabled;
|
||||
// If a specific value was provided, use it; otherwise toggle.
|
||||
routine.enabled = match body {
|
||||
Some(Json(req)) => req.enabled.unwrap_or(!routine.enabled),
|
||||
None => !routine.enabled,
|
||||
};
|
||||
|
||||
if routine.enabled
|
||||
&& !was_enabled
|
||||
&& let Trigger::Cron { schedule, timezone } = &routine.trigger
|
||||
{
|
||||
routine.next_fire_at = next_cron_fire(schedule, timezone.as_deref())
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
}
|
||||
|
||||
store
|
||||
.update_routine(&routine)
|
||||
.await
|
||||
@@ -2639,6 +2714,7 @@ struct GatewayStatusResponse {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::cli::oauth_defaults;
|
||||
use crate::testing::credentials::TEST_GATEWAY_CRYPTO_KEY;
|
||||
|
||||
#[test]
|
||||
@@ -2756,6 +2832,11 @@ mod tests {
|
||||
.with_state(state)
|
||||
}
|
||||
|
||||
fn expired_flow_created_at() -> Option<std::time::Instant> {
|
||||
std::time::Instant::now()
|
||||
.checked_sub(oauth_defaults::OAUTH_FLOW_EXPIRY + std::time::Duration::from_secs(1))
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_csp_header_present_on_responses() {
|
||||
use std::net::SocketAddr;
|
||||
@@ -2862,29 +2943,14 @@ mod tests {
|
||||
use tower::ServiceExt;
|
||||
|
||||
// Build an ExtensionManager so the handler can look up flows
|
||||
let secrets = Arc::new(crate::secrets::InMemorySecretsStore::new(Arc::new(
|
||||
crate::secrets::SecretsCrypto::new(secrecy::SecretString::from(
|
||||
TEST_GATEWAY_CRYPTO_KEY.to_string(),
|
||||
))
|
||||
.expect("crypto"),
|
||||
)));
|
||||
let tool_registry = Arc::new(ToolRegistry::new());
|
||||
let mcp_sm = Arc::new(crate::tools::mcp::session::McpSessionManager::new());
|
||||
|
||||
let ext_mgr = Arc::new(ExtensionManager::new(
|
||||
mcp_sm,
|
||||
Arc::new(crate::tools::mcp::process::McpProcessManager::new()),
|
||||
secrets,
|
||||
tool_registry,
|
||||
None,
|
||||
None,
|
||||
std::path::PathBuf::from("/tmp/wasm_tools"),
|
||||
std::path::PathBuf::from("/tmp/wasm_channels"),
|
||||
None,
|
||||
"test".to_string(),
|
||||
None,
|
||||
vec![],
|
||||
));
|
||||
let secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync> =
|
||||
Arc::new(crate::secrets::InMemorySecretsStore::new(Arc::new(
|
||||
crate::secrets::SecretsCrypto::new(secrecy::SecretString::from(
|
||||
TEST_GATEWAY_CRYPTO_KEY.to_string(),
|
||||
))
|
||||
.expect("crypto"),
|
||||
)));
|
||||
let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(secrets);
|
||||
|
||||
let state = test_gateway_state(Some(ext_mgr));
|
||||
let app = test_oauth_router(state);
|
||||
@@ -2918,25 +2984,13 @@ mod tests {
|
||||
))
|
||||
.expect("crypto"),
|
||||
)));
|
||||
let tool_registry = Arc::new(ToolRegistry::new());
|
||||
let mcp_sm = Arc::new(crate::tools::mcp::session::McpSessionManager::new());
|
||||
let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(secrets.clone());
|
||||
let Some(created_at) = expired_flow_created_at() else {
|
||||
eprintln!("Skipping expired OAuth flow test: monotonic uptime below expiry window");
|
||||
return;
|
||||
};
|
||||
|
||||
let 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)
|
||||
// Insert an expired flow.
|
||||
let flow = crate::cli::oauth_defaults::PendingOAuthFlow {
|
||||
extension_name: "test_tool".to_string(),
|
||||
display_name: "Test Tool".to_string(),
|
||||
@@ -2954,9 +3008,9 @@ mod tests {
|
||||
secrets,
|
||||
sse_sender: None,
|
||||
gateway_token: None,
|
||||
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"),
|
||||
resource: None,
|
||||
client_id_secret_name: None,
|
||||
created_at,
|
||||
};
|
||||
|
||||
ext_mgr
|
||||
@@ -2986,6 +3040,80 @@ mod tests {
|
||||
assert!(html.contains("Authorization Failed"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_oauth_callback_expired_flow_broadcasts_auth_completed_failure() {
|
||||
use axum::body::Body;
|
||||
use tower::ServiceExt;
|
||||
|
||||
let secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync> =
|
||||
Arc::new(crate::secrets::InMemorySecretsStore::new(Arc::new(
|
||||
crate::secrets::SecretsCrypto::new(secrecy::SecretString::from(
|
||||
TEST_GATEWAY_CRYPTO_KEY.to_string(),
|
||||
))
|
||||
.expect("crypto"),
|
||||
)));
|
||||
let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(secrets.clone());
|
||||
|
||||
let (sender, mut receiver) = tokio::sync::broadcast::channel(4);
|
||||
let Some(created_at) = expired_flow_created_at() else {
|
||||
eprintln!("Skipping expired OAuth flow SSE test: monotonic uptime below expiry window");
|
||||
return;
|
||||
};
|
||||
let flow = crate::cli::oauth_defaults::PendingOAuthFlow {
|
||||
extension_name: "test_tool".to_string(),
|
||||
display_name: "Test Tool".to_string(),
|
||||
token_url: "https://example.com/token".to_string(),
|
||||
client_id: "client123".to_string(),
|
||||
client_secret: None,
|
||||
redirect_uri: "https://example.com/oauth/callback".to_string(),
|
||||
code_verifier: None,
|
||||
access_token_field: "access_token".to_string(),
|
||||
secret_name: "test_token".to_string(),
|
||||
provider: None,
|
||||
validation_endpoint: None,
|
||||
scopes: vec![],
|
||||
user_id: "test".to_string(),
|
||||
secrets,
|
||||
sse_sender: Some(sender),
|
||||
gateway_token: None,
|
||||
resource: None,
|
||||
client_id_secret_name: None,
|
||||
created_at,
|
||||
};
|
||||
|
||||
ext_mgr
|
||||
.pending_oauth_flows()
|
||||
.write()
|
||||
.await
|
||||
.insert("expired_state".to_string(), flow);
|
||||
|
||||
let state = test_gateway_state(Some(ext_mgr));
|
||||
let app = test_oauth_router(state);
|
||||
|
||||
let req = axum::http::Request::builder()
|
||||
.uri("/oauth/callback?code=test_code&state=expired_state")
|
||||
.body(Body::empty())
|
||||
.expect("request");
|
||||
|
||||
let resp = ServiceExt::<axum::http::Request<Body>>::oneshot(app, req)
|
||||
.await
|
||||
.expect("response");
|
||||
assert_eq!(resp.status(), StatusCode::OK);
|
||||
|
||||
match receiver.recv().await.expect("auth_completed event") {
|
||||
crate::channels::web::types::SseEvent::AuthCompleted {
|
||||
extension_name,
|
||||
success,
|
||||
message,
|
||||
} => {
|
||||
assert_eq!(extension_name, "test_tool");
|
||||
assert!(!success, "expired OAuth flow should broadcast failure");
|
||||
assert_eq!(message, "OAuth flow expired. Please try again.");
|
||||
}
|
||||
event => panic!("expected AuthCompleted event, got {event:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_oauth_callback_no_extension_manager() {
|
||||
use axum::body::Body;
|
||||
@@ -3024,28 +3152,16 @@ mod tests {
|
||||
))
|
||||
.expect("crypto"),
|
||||
)));
|
||||
let tool_registry = Arc::new(ToolRegistry::new());
|
||||
let mcp_sm = Arc::new(crate::tools::mcp::session::McpSessionManager::new());
|
||||
|
||||
let ext_mgr = Arc::new(ExtensionManager::new(
|
||||
mcp_sm,
|
||||
Arc::new(crate::tools::mcp::process::McpProcessManager::new()),
|
||||
secrets.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![],
|
||||
));
|
||||
let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(secrets.clone());
|
||||
|
||||
// Insert a flow keyed by raw nonce "test_nonce" (without instance prefix).
|
||||
// Use an expired flow so the handler exits before attempting a real HTTP
|
||||
// token exchange — we only need to verify that the instance prefix was
|
||||
// stripped and the flow was found by the raw nonce.
|
||||
let Some(created_at) = expired_flow_created_at() else {
|
||||
eprintln!("Skipping OAuth state-prefix test: monotonic uptime below expiry window");
|
||||
return;
|
||||
};
|
||||
let flow = crate::cli::oauth_defaults::PendingOAuthFlow {
|
||||
extension_name: "test_tool".to_string(),
|
||||
display_name: "Test Tool".to_string(),
|
||||
@@ -3063,10 +3179,10 @@ mod tests {
|
||||
secrets,
|
||||
sse_sender: None,
|
||||
gateway_token: None,
|
||||
resource: None,
|
||||
client_id_secret_name: None,
|
||||
// Expired — handler will reject after lookup (no network I/O)
|
||||
created_at: std::time::Instant::now()
|
||||
.checked_sub(std::time::Duration::from_secs(600))
|
||||
.expect("System uptime is too low to run expired flow test"),
|
||||
created_at,
|
||||
};
|
||||
|
||||
ext_mgr
|
||||
@@ -3137,24 +3253,27 @@ mod tests {
|
||||
|
||||
fn test_ext_mgr(
|
||||
secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync>,
|
||||
) -> Arc<ExtensionManager> {
|
||||
) -> (Arc<ExtensionManager>, tempfile::TempDir, tempfile::TempDir) {
|
||||
let tool_registry = Arc::new(ToolRegistry::new());
|
||||
let mcp_sm = Arc::new(crate::tools::mcp::session::McpSessionManager::new());
|
||||
let mcp_pm = Arc::new(crate::tools::mcp::process::McpProcessManager::new());
|
||||
Arc::new(ExtensionManager::new(
|
||||
let wasm_tools_dir = tempfile::tempdir().expect("temp wasm tools dir");
|
||||
let wasm_channels_dir = tempfile::tempdir().expect("temp wasm channels dir");
|
||||
let ext_mgr = Arc::new(ExtensionManager::new(
|
||||
mcp_sm,
|
||||
mcp_pm,
|
||||
secrets,
|
||||
tool_registry,
|
||||
None,
|
||||
None,
|
||||
std::path::PathBuf::from("/tmp/wasm_tools"),
|
||||
std::path::PathBuf::from("/tmp/wasm_channels"),
|
||||
wasm_tools_dir.path().to_path_buf(),
|
||||
wasm_channels_dir.path().to_path_buf(),
|
||||
None,
|
||||
"test".to_string(),
|
||||
None,
|
||||
vec![],
|
||||
))
|
||||
));
|
||||
(ext_mgr, wasm_tools_dir, wasm_channels_dir)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -3163,7 +3282,7 @@ mod tests {
|
||||
use tower::ServiceExt;
|
||||
|
||||
let secrets = test_secrets_store();
|
||||
let ext_mgr = test_ext_mgr(secrets);
|
||||
let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(secrets);
|
||||
let state = test_gateway_state(Some(ext_mgr));
|
||||
let app = test_relay_oauth_router(state);
|
||||
|
||||
@@ -3207,7 +3326,7 @@ mod tests {
|
||||
.await
|
||||
.expect("store nonce");
|
||||
|
||||
let ext_mgr = test_ext_mgr(secrets);
|
||||
let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(secrets);
|
||||
let state = test_gateway_state(Some(ext_mgr));
|
||||
let app = test_relay_oauth_router(state);
|
||||
|
||||
@@ -3252,7 +3371,7 @@ mod tests {
|
||||
.await
|
||||
.expect("store nonce");
|
||||
|
||||
let ext_mgr = test_ext_mgr(secrets.clone());
|
||||
let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(secrets.clone());
|
||||
let state = test_gateway_state(Some(ext_mgr));
|
||||
let app = test_relay_oauth_router(state);
|
||||
|
||||
|
||||
+323
-59
@@ -342,31 +342,27 @@ function connectSSE() {
|
||||
|
||||
eventSource.addEventListener('approval_needed', (e) => {
|
||||
const data = JSON.parse(e.data);
|
||||
if (!isCurrentThread(data.thread_id)) return;
|
||||
showApproval(data);
|
||||
const hasThread = !!data.thread_id;
|
||||
const forCurrentThread = !hasThread || isCurrentThread(data.thread_id);
|
||||
|
||||
if (forCurrentThread) {
|
||||
showApproval(data);
|
||||
} else {
|
||||
// Keep thread list fresh when approval is requested in a background thread.
|
||||
unreadThreads.set(data.thread_id, (unreadThreads.get(data.thread_id) || 0) + 1);
|
||||
debouncedLoadThreads();
|
||||
}
|
||||
|
||||
// Extension setup flows can surface approvals while user is on Extensions tab.
|
||||
if (currentTab === 'extensions') loadExtensions();
|
||||
});
|
||||
|
||||
eventSource.addEventListener('auth_required', (e) => {
|
||||
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);
|
||||
}
|
||||
handleAuthRequired(JSON.parse(e.data));
|
||||
});
|
||||
|
||||
eventSource.addEventListener('auth_completed', (e) => {
|
||||
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();
|
||||
handleAuthCompleted(JSON.parse(e.data));
|
||||
});
|
||||
|
||||
eventSource.addEventListener('extension_status', (e) => {
|
||||
@@ -438,7 +434,8 @@ function sendMessage() {
|
||||
const content = input.value.trim();
|
||||
if (!content && stagedImages.length === 0) return;
|
||||
|
||||
addMessage('user', content || '(images attached)');
|
||||
const imagesToDisplay = stagedImages.slice(); // snapshot before clearing
|
||||
addMessage('user', content, imagesToDisplay);
|
||||
input.value = '';
|
||||
autoResizeTextarea(input);
|
||||
input.focus();
|
||||
@@ -544,6 +541,24 @@ document.getElementById('chat-input').addEventListener('paste', (e) => {
|
||||
}
|
||||
});
|
||||
|
||||
document.getElementById('chat-input').addEventListener('dragover', (e) => {
|
||||
const hasImage = Array.from(e.dataTransfer.items || []).some(
|
||||
item => item.kind === 'file' && item.type.startsWith('image/')
|
||||
);
|
||||
if (hasImage) {
|
||||
e.preventDefault();
|
||||
e.dataTransfer.dropEffect = 'copy';
|
||||
}
|
||||
});
|
||||
|
||||
document.getElementById('chat-input').addEventListener('drop', (e) => {
|
||||
const files = Array.from(e.dataTransfer.files || []).filter(f => f.type.startsWith('image/'));
|
||||
if (files.length > 0) {
|
||||
e.preventDefault();
|
||||
handleImageFiles(files);
|
||||
}
|
||||
});
|
||||
|
||||
function addGeneratedImage(dataUrl, path) {
|
||||
const container = document.getElementById('chat-messages');
|
||||
const card = document.createElement('div');
|
||||
@@ -670,7 +685,7 @@ function renderMarkdown(text) {
|
||||
// Sanitize HTML output to prevent XSS from tool output or LLM responses.
|
||||
html = sanitizeRenderedHtml(html);
|
||||
// Inject copy buttons into <pre> blocks
|
||||
html = html.replace(/<pre>/g, '<pre class="code-block-wrapper"><button class="copy-btn" onclick="copyCodeBlock(this)">Copy</button>');
|
||||
html = html.replace(/<pre>/g, '<pre class="code-block-wrapper"><button class="copy-btn" data-action="copy-code">Copy</button>');
|
||||
return html;
|
||||
}
|
||||
return escapeHtml(text);
|
||||
@@ -702,12 +717,44 @@ function copyCodeBlock(btn) {
|
||||
});
|
||||
}
|
||||
|
||||
function addMessage(role, content) {
|
||||
function copyMessage(btn) {
|
||||
const message = btn.closest('.message');
|
||||
if (!message) return;
|
||||
const text = message.getAttribute('data-copy-text')
|
||||
|| message.getAttribute('data-raw')
|
||||
|| message.textContent
|
||||
|| '';
|
||||
navigator.clipboard.writeText(text).then(() => {
|
||||
btn.textContent = 'Copied';
|
||||
setTimeout(() => { btn.textContent = 'Copy'; }, 1200);
|
||||
}).catch(() => {
|
||||
btn.textContent = 'Failed';
|
||||
setTimeout(() => { btn.textContent = 'Copy'; }, 1200);
|
||||
});
|
||||
}
|
||||
|
||||
function addMessage(role, content, images) {
|
||||
const container = document.getElementById('chat-messages');
|
||||
const div = document.createElement('div');
|
||||
div.className = 'message ' + role;
|
||||
if (role === 'user') {
|
||||
div.textContent = content;
|
||||
if (images && images.length > 0) {
|
||||
const imgStrip = document.createElement('div');
|
||||
imgStrip.className = 'message-images';
|
||||
images.forEach(img => {
|
||||
const thumb = document.createElement('img');
|
||||
thumb.className = 'message-image-thumb';
|
||||
thumb.src = img.dataUrl;
|
||||
thumb.alt = 'Attached image';
|
||||
imgStrip.appendChild(thumb);
|
||||
});
|
||||
div.appendChild(imgStrip);
|
||||
}
|
||||
if (content) {
|
||||
const textDiv = document.createElement('div');
|
||||
textDiv.textContent = content;
|
||||
div.appendChild(textDiv);
|
||||
}
|
||||
} else {
|
||||
div.setAttribute('data-raw', content);
|
||||
div.innerHTML = renderMarkdown(content);
|
||||
@@ -723,7 +770,11 @@ function appendToLastAssistant(chunk) {
|
||||
const last = messages[messages.length - 1];
|
||||
const raw = (last.getAttribute('data-raw') || '') + chunk;
|
||||
last.setAttribute('data-raw', raw);
|
||||
last.innerHTML = renderMarkdown(raw);
|
||||
last.setAttribute('data-copy-text', raw);
|
||||
const content = last.querySelector('.message-content');
|
||||
if (content) {
|
||||
content.innerHTML = renderMarkdown(raw);
|
||||
}
|
||||
container.scrollTop = container.scrollHeight;
|
||||
} else {
|
||||
addMessage('assistant', chunk);
|
||||
@@ -977,7 +1028,26 @@ function finalizeActivityGroup() {
|
||||
_activeToolCards = {};
|
||||
}
|
||||
|
||||
function humanizeToolName(rawName) {
|
||||
if (!rawName) return '';
|
||||
return String(rawName)
|
||||
.replace(/[_-]+/g, ' ')
|
||||
.replace(/([a-z0-9])([A-Z])/g, '$1 $2')
|
||||
.replace(/^tool([a-zA-Z])/, 'tool $1')
|
||||
.replace(/\s+/g, ' ')
|
||||
.trim();
|
||||
}
|
||||
|
||||
function shouldShowChannelConnectedMessage(extensionName, success) {
|
||||
if (!success || !extensionName) return false;
|
||||
return String(extensionName).toLowerCase().includes('telegram');
|
||||
}
|
||||
|
||||
function showApproval(data) {
|
||||
// Avoid duplicate cards on reconnect/history refresh.
|
||||
const existing = document.querySelector('.approval-card[data-request-id="' + CSS.escape(data.request_id) + '"]');
|
||||
if (existing) return;
|
||||
|
||||
const container = document.getElementById('chat-messages');
|
||||
const card = document.createElement('div');
|
||||
card.className = 'approval-card';
|
||||
@@ -990,7 +1060,7 @@ function showApproval(data) {
|
||||
|
||||
const toolName = document.createElement('div');
|
||||
toolName.className = 'approval-tool-name';
|
||||
toolName.textContent = data.tool_name;
|
||||
toolName.textContent = humanizeToolName(data.tool_name);
|
||||
card.appendChild(toolName);
|
||||
|
||||
if (data.description) {
|
||||
@@ -1093,13 +1163,71 @@ function showJobCard(data) {
|
||||
|
||||
// --- Auth card ---
|
||||
|
||||
function handleAuthRequired(data) {
|
||||
if (data.auth_url) {
|
||||
// OAuth flow: show the global auth prompt with an OAuth button + optional token paste field.
|
||||
showAuthCard(data);
|
||||
} else {
|
||||
// Setup flow: fetch the extension's credential schema and show the multi-field
|
||||
// configure modal (the same UI used by the Extensions tab "Setup" button).
|
||||
showConfigureModal(data.extension_name);
|
||||
}
|
||||
}
|
||||
|
||||
function handleAuthCompleted(data) {
|
||||
// Dismiss only the matching extension's UI so unrelated setup work is not interrupted.
|
||||
removeAuthCard(data.extension_name);
|
||||
closeConfigureModal(data.extension_name);
|
||||
showToast(data.message, data.success ? 'success' : 'error');
|
||||
if (shouldShowChannelConnectedMessage(data.extension_name, data.success)) {
|
||||
addMessage('system', 'Telegram is now connected. You can message me there and I can send you notifications.');
|
||||
}
|
||||
if (currentTab === 'extensions') loadExtensions();
|
||||
enableChatInput();
|
||||
}
|
||||
|
||||
function queryByDataAttribute(selector, attributeName, attributeValue) {
|
||||
if (typeof attributeValue !== 'string') return document.querySelector(selector);
|
||||
|
||||
if (window.CSS && typeof window.CSS.escape === 'function') {
|
||||
return document.querySelector(
|
||||
selector + '[' + attributeName + '="' + window.CSS.escape(attributeValue) + '"]'
|
||||
);
|
||||
}
|
||||
|
||||
const candidates = document.querySelectorAll(selector);
|
||||
for (const candidate of candidates) {
|
||||
if (candidate.getAttribute(attributeName) === attributeValue) return candidate;
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
function getAuthOverlay(extensionName) {
|
||||
return queryByDataAttribute('.auth-overlay', 'data-extension-name', extensionName);
|
||||
}
|
||||
|
||||
function getAuthCard(extensionName) {
|
||||
return queryByDataAttribute('.auth-card', 'data-extension-name', extensionName);
|
||||
}
|
||||
|
||||
function getConfigureOverlay(extensionName) {
|
||||
return queryByDataAttribute('.configure-overlay', 'data-extension-name', extensionName);
|
||||
}
|
||||
|
||||
function showAuthCard(data) {
|
||||
// Remove any existing card for this extension first
|
||||
removeAuthCard(data.extension_name);
|
||||
// Keep a single global auth prompt so the experience is consistent across tabs.
|
||||
const existing = getAuthOverlay();
|
||||
if (existing) existing.remove();
|
||||
|
||||
const overlay = document.createElement('div');
|
||||
overlay.className = 'auth-overlay';
|
||||
overlay.setAttribute('data-extension-name', data.extension_name);
|
||||
overlay.addEventListener('click', (e) => {
|
||||
if (e.target === overlay) cancelAuth(data.extension_name);
|
||||
});
|
||||
|
||||
const container = document.getElementById('chat-messages');
|
||||
const card = document.createElement('div');
|
||||
card.className = 'auth-card';
|
||||
card.className = 'auth-card auth-modal';
|
||||
card.setAttribute('data-extension-name', data.extension_name);
|
||||
|
||||
const header = document.createElement('div');
|
||||
@@ -1178,21 +1306,30 @@ function showAuthCard(data) {
|
||||
actions.appendChild(cancelBtn);
|
||||
card.appendChild(actions);
|
||||
|
||||
container.appendChild(card);
|
||||
container.scrollTop = container.scrollHeight;
|
||||
overlay.appendChild(card);
|
||||
document.body.appendChild(overlay);
|
||||
tokenInput.focus();
|
||||
}
|
||||
|
||||
function removeAuthCard(extensionName) {
|
||||
const card = document.querySelector('.auth-card[data-extension-name="' + extensionName + '"]');
|
||||
if (card) card.remove();
|
||||
const overlay = getAuthOverlay(extensionName);
|
||||
if (overlay) {
|
||||
overlay.remove();
|
||||
return;
|
||||
}
|
||||
const card = getAuthCard(extensionName);
|
||||
if (card) {
|
||||
const parentOverlay = card.closest('.auth-overlay');
|
||||
if (parentOverlay) parentOverlay.remove();
|
||||
else card.remove();
|
||||
}
|
||||
}
|
||||
|
||||
function submitAuthToken(extensionName, tokenValue) {
|
||||
if (!tokenValue || !tokenValue.trim()) return;
|
||||
|
||||
// Disable submit button while in flight
|
||||
const card = document.querySelector('.auth-card[data-extension-name="' + extensionName + '"]');
|
||||
const card = getAuthCard(extensionName);
|
||||
if (card) {
|
||||
const btns = card.querySelectorAll('button');
|
||||
btns.forEach((b) => { b.disabled = true; });
|
||||
@@ -1203,8 +1340,10 @@ function submitAuthToken(extensionName, tokenValue) {
|
||||
body: { extension_name: extensionName, token: tokenValue.trim() },
|
||||
}).then((result) => {
|
||||
if (result.success) {
|
||||
// Close immediately for responsiveness; the authoritative success UX
|
||||
// (toast + extensions refresh) still comes from auth_completed SSE.
|
||||
removeAuthCard(extensionName);
|
||||
addMessage('system', result.message);
|
||||
enableChatInput();
|
||||
} else {
|
||||
showAuthCardError(extensionName, result.message);
|
||||
}
|
||||
@@ -1223,7 +1362,7 @@ function cancelAuth(extensionName) {
|
||||
}
|
||||
|
||||
function showAuthCardError(extensionName, message) {
|
||||
const card = document.querySelector('.auth-card[data-extension-name="' + extensionName + '"]');
|
||||
const card = getAuthCard(extensionName);
|
||||
if (!card) return;
|
||||
// Re-enable buttons
|
||||
const btns = card.querySelectorAll('button');
|
||||
@@ -1310,12 +1449,31 @@ function loadHistory(before) {
|
||||
function createMessageElement(role, content) {
|
||||
const div = document.createElement('div');
|
||||
div.className = 'message ' + role;
|
||||
if (role === 'user') {
|
||||
div.textContent = content;
|
||||
|
||||
if (role === 'assistant' || role === 'user') {
|
||||
div.classList.add('has-copy');
|
||||
div.setAttribute('data-copy-text', content);
|
||||
const copyBtn = document.createElement('button');
|
||||
copyBtn.className = 'message-copy-btn';
|
||||
copyBtn.type = 'button';
|
||||
copyBtn.setAttribute('aria-label', 'Copy message');
|
||||
copyBtn.textContent = 'Copy';
|
||||
copyBtn.addEventListener('click', (e) => {
|
||||
e.stopPropagation();
|
||||
copyMessage(copyBtn);
|
||||
});
|
||||
div.appendChild(copyBtn);
|
||||
}
|
||||
|
||||
const body = document.createElement('div');
|
||||
body.className = 'message-content';
|
||||
if (role === 'user' || role === 'system') {
|
||||
body.textContent = content;
|
||||
} else {
|
||||
div.setAttribute('data-raw', content);
|
||||
div.innerHTML = renderMarkdown(content);
|
||||
body.innerHTML = renderMarkdown(content);
|
||||
}
|
||||
div.appendChild(body);
|
||||
return div;
|
||||
}
|
||||
|
||||
@@ -1819,13 +1977,11 @@ function saveMemoryEdit() {
|
||||
|
||||
function buildBreadcrumb(path) {
|
||||
const parts = path.split('/');
|
||||
let html = '<a onclick="loadMemoryTree()">workspace</a>';
|
||||
let html = '<a data-action="breadcrumb-root" href="#">workspace</a>';
|
||||
let current = '';
|
||||
for (const part of parts) {
|
||||
current += (current ? '/' : '') + part;
|
||||
// Store the path in data-path (HTML-escaped) and read it back via this.dataset.path
|
||||
// to avoid single-quote injection in inline JS string literals.
|
||||
html += ' / <a onclick="readMemoryFile(this.dataset.path)" data-path="' + escapeHtml(current) + '">' + escapeHtml(part) + '</a>';
|
||||
html += ' / <a data-action="breadcrumb-file" data-path="' + escapeHtml(current) + '" href="#">' + escapeHtml(part) + '</a>';
|
||||
}
|
||||
return html;
|
||||
}
|
||||
@@ -2136,6 +2292,10 @@ function renderAvailableExtensionCard(entry) {
|
||||
showToast(I18n.t('extensions.installedSuccess', {name: entry.display_name}), 'success');
|
||||
// OAuth popup if auth started during install (builtin creds)
|
||||
if (res.auth_url) {
|
||||
showAuthCard({
|
||||
extension_name: entry.name,
|
||||
auth_url: res.auth_url,
|
||||
});
|
||||
showToast('Opening authentication for ' + entry.display_name, 'info');
|
||||
openOAuthUrl(res.auth_url);
|
||||
}
|
||||
@@ -2401,6 +2561,10 @@ function activateExtension(name) {
|
||||
if (res.success) {
|
||||
// Even on success, the tool may need OAuth (e.g., WASM loaded but no token yet)
|
||||
if (res.auth_url) {
|
||||
showAuthCard({
|
||||
extension_name: name,
|
||||
auth_url: res.auth_url,
|
||||
});
|
||||
showToast('Opening authentication for ' + name, 'info');
|
||||
openOAuthUrl(res.auth_url);
|
||||
}
|
||||
@@ -2409,6 +2573,10 @@ function activateExtension(name) {
|
||||
}
|
||||
|
||||
if (res.auth_url) {
|
||||
showAuthCard({
|
||||
extension_name: name,
|
||||
auth_url: res.auth_url,
|
||||
});
|
||||
showToast('Opening authentication for ' + name, 'info');
|
||||
openOAuthUrl(res.auth_url);
|
||||
} else if (res.awaiting_token) {
|
||||
@@ -2451,6 +2619,7 @@ function renderConfigureModal(name, secrets) {
|
||||
closeConfigureModal();
|
||||
const overlay = document.createElement('div');
|
||||
overlay.className = 'configure-overlay';
|
||||
overlay.setAttribute('data-extension-name', name);
|
||||
overlay.addEventListener('click', (e) => {
|
||||
if (e.target === overlay) closeConfigureModal();
|
||||
});
|
||||
@@ -2544,7 +2713,8 @@ function submitConfigureModal(name, fields) {
|
||||
}
|
||||
|
||||
// Disable buttons to prevent double-submit
|
||||
var btns = document.querySelectorAll('.configure-actions button');
|
||||
const overlay = getConfigureOverlay(name) || document.querySelector('.configure-overlay');
|
||||
var btns = overlay ? overlay.querySelectorAll('.configure-actions button') : [];
|
||||
btns.forEach(function(b) { b.disabled = true; });
|
||||
|
||||
apiFetch('/api/extensions/' + encodeURIComponent(name) + '/setup', {
|
||||
@@ -2555,8 +2725,10 @@ function submitConfigureModal(name, fields) {
|
||||
if (res.success) {
|
||||
closeConfigureModal();
|
||||
if (res.auth_url) {
|
||||
// OAuth flow started — open consent popup. The auth_completed SSE will
|
||||
// not arrive immediately (it fires after OAuth callback), so show a toast now.
|
||||
showAuthCard({
|
||||
extension_name: name,
|
||||
auth_url: res.auth_url,
|
||||
});
|
||||
showToast('Opening OAuth authorization for ' + name, 'info');
|
||||
openOAuthUrl(res.auth_url);
|
||||
loadExtensions();
|
||||
@@ -2575,8 +2747,9 @@ function submitConfigureModal(name, fields) {
|
||||
});
|
||||
}
|
||||
|
||||
function closeConfigureModal() {
|
||||
const existing = document.querySelector('.configure-overlay');
|
||||
function closeConfigureModal(extensionName) {
|
||||
if (typeof extensionName !== 'string') extensionName = null;
|
||||
const existing = getConfigureOverlay(extensionName);
|
||||
if (existing) existing.remove();
|
||||
}
|
||||
|
||||
@@ -2795,11 +2968,11 @@ function renderJobsList(jobs) {
|
||||
|
||||
let actionBtns = '';
|
||||
if (job.state === 'pending' || job.state === 'in_progress') {
|
||||
actionBtns = '<button class="btn-cancel" onclick="event.stopPropagation(); cancelJob(\'' + job.id + '\')">Cancel</button>';
|
||||
actionBtns = '<button class="btn-cancel" data-action="cancel-job" data-id="' + escapeHtml(job.id) + '">Cancel</button>';
|
||||
}
|
||||
// Retry is only shown in the detail view where can_restart is available.
|
||||
|
||||
return '<tr class="job-row" onclick="openJobDetail(\'' + job.id + '\')">'
|
||||
return '<tr class="job-row" data-action="open-job" data-id="' + escapeHtml(job.id) + '">'
|
||||
+ '<td title="' + escapeHtml(job.id) + '">' + shortId + '</td>'
|
||||
+ '<td>' + escapeHtml(job.title) + '</td>'
|
||||
+ '<td><span class="badge ' + stateClass + '">' + escapeHtml(job.state) + '</span></td>'
|
||||
@@ -2862,12 +3035,12 @@ function renderJobDetail(job) {
|
||||
const header = document.createElement('div');
|
||||
header.className = 'job-detail-header';
|
||||
|
||||
let headerHtml = '<button class="btn-back" onclick="closeJobDetail()">← Back</button>'
|
||||
let headerHtml = '<button class="btn-back" data-action="close-job-detail">← Back</button>'
|
||||
+ '<h2>' + escapeHtml(job.title) + '</h2>'
|
||||
+ '<span class="badge ' + stateClass + '">' + escapeHtml(job.state) + '</span>';
|
||||
|
||||
if ((job.state === 'failed' || job.state === 'interrupted') && job.can_restart === true) {
|
||||
headerHtml += '<button class="btn-restart" onclick="restartJob(\'' + job.id + '\')">Retry</button>';
|
||||
headerHtml += '<button class="btn-restart" data-action="restart-job" data-id="' + escapeHtml(job.id) + '">Retry</button>';
|
||||
}
|
||||
if (job.browse_url) {
|
||||
headerHtml += '<a class="btn-browse" href="' + escapeHtml(job.browse_url) + '" target="_blank">Browse Files</a>';
|
||||
@@ -3324,7 +3497,7 @@ function renderRoutinesList(routines) {
|
||||
const toggleLabel = r.enabled ? 'Disable' : 'Enable';
|
||||
const toggleClass = r.enabled ? 'btn-cancel' : 'btn-restart';
|
||||
|
||||
return '<tr class="routine-row" onclick="openRoutineDetail(\'' + r.id + '\')">'
|
||||
return '<tr class="routine-row" data-action="open-routine" data-id="' + escapeHtml(r.id) + '">'
|
||||
+ '<td>' + escapeHtml(r.name) + '</td>'
|
||||
+ '<td>' + escapeHtml(r.trigger_summary) + '</td>'
|
||||
+ '<td>' + escapeHtml(r.action_type) + '</td>'
|
||||
@@ -3333,9 +3506,9 @@ function renderRoutinesList(routines) {
|
||||
+ '<td>' + r.run_count + '</td>'
|
||||
+ '<td><span class="badge ' + statusClass + '">' + escapeHtml(r.status) + '</span></td>'
|
||||
+ '<td>'
|
||||
+ '<button class="' + toggleClass + '" onclick="event.stopPropagation(); toggleRoutine(\'' + r.id + '\')">' + toggleLabel + '</button> '
|
||||
+ '<button class="btn-restart" onclick="event.stopPropagation(); triggerRoutine(\'' + r.id + '\')">Run</button> '
|
||||
+ '<button class="btn-cancel" onclick="event.stopPropagation(); deleteRoutine(\'' + r.id + '\', \'' + escapeHtml(r.name) + '\')">Delete</button>'
|
||||
+ '<button class="' + toggleClass + '" data-action="toggle-routine" data-id="' + escapeHtml(r.id) + '">' + toggleLabel + '</button> '
|
||||
+ '<button class="btn-restart" data-action="trigger-routine" data-id="' + escapeHtml(r.id) + '">Run</button> '
|
||||
+ '<button class="btn-cancel" data-action="delete-routine" data-id="' + escapeHtml(r.id) + '" data-name="' + escapeHtml(r.name) + '">Delete</button>'
|
||||
+ '</td>'
|
||||
+ '</tr>';
|
||||
}).join('');
|
||||
@@ -3371,7 +3544,7 @@ function renderRoutineDetail(routine) {
|
||||
: 'active';
|
||||
|
||||
let html = '<div class="job-detail-header">'
|
||||
+ '<button class="btn-back" onclick="closeRoutineDetail()">← Back</button>'
|
||||
+ '<button class="btn-back" data-action="close-routine-detail">← Back</button>'
|
||||
+ '<h2>' + escapeHtml(routine.name) + '</h2>'
|
||||
+ '<span class="badge ' + statusClass + '">' + escapeHtml(statusLabel) + '</span>'
|
||||
+ '</div>';
|
||||
@@ -3418,7 +3591,7 @@ function renderRoutineDetail(routine) {
|
||||
+ '<td>' + formatDate(run.completed_at) + '</td>'
|
||||
+ '<td><span class="badge ' + runStatusClass + '">' + escapeHtml(run.status) + '</span></td>'
|
||||
+ '<td>' + escapeHtml(run.result_summary || '-')
|
||||
+ (run.job_id ? ' <a href="#" onclick="event.preventDefault(); switchTab(\'jobs\'); openJobDetail(\'' + run.job_id + '\')">[view job]</a>' : '')
|
||||
+ (run.job_id ? ' <a href="#" data-action="view-run-job" data-id="' + escapeHtml(run.job_id) + '">[view job]</a>' : '')
|
||||
+ '</td>'
|
||||
+ '<td>' + (run.tokens_used != null ? run.tokens_used : '-') + '</td>'
|
||||
+ '</tr>';
|
||||
@@ -3661,7 +3834,7 @@ function renderTeePopover(report) {
|
||||
+ '<div class="tee-field"><div class="tee-field-label">VM Config</div>'
|
||||
+ '<div class="tee-field-value">' + escapeHtml(vmConfig) + '</div></div>'
|
||||
+ '<div class="tee-popover-actions">'
|
||||
+ '<button class="tee-btn-copy" onclick="copyTeeReport()">Copy Full Report</button></div>';
|
||||
+ '<button class="tee-btn-copy" data-action="copy-tee-report">Copy Full Report</button></div>';
|
||||
}
|
||||
|
||||
function copyTeeReport() {
|
||||
@@ -4143,3 +4316,94 @@ function formatDate(isoString) {
|
||||
const d = new Date(isoString);
|
||||
return d.toLocaleString();
|
||||
}
|
||||
|
||||
// --- Event Listener Registration (CSP-safe, no inline handlers) ---
|
||||
|
||||
document.getElementById('auth-connect-btn').addEventListener('click', () => authenticate());
|
||||
document.getElementById('restart-overlay').addEventListener('click', () => cancelRestart());
|
||||
document.getElementById('restart-close-btn').addEventListener('click', () => cancelRestart());
|
||||
document.getElementById('restart-cancel-btn').addEventListener('click', () => cancelRestart());
|
||||
document.getElementById('restart-confirm-btn').addEventListener('click', () => confirmRestart());
|
||||
document.getElementById('restart-btn').addEventListener('click', () => triggerRestart());
|
||||
document.getElementById('thread-new-btn').addEventListener('click', () => createNewThread());
|
||||
document.getElementById('thread-toggle-btn').addEventListener('click', () => toggleThreadSidebar());
|
||||
document.getElementById('assistant-thread').addEventListener('click', () => switchToAssistant());
|
||||
document.getElementById('send-btn').addEventListener('click', () => sendMessage());
|
||||
document.getElementById('memory-edit-btn').addEventListener('click', () => startMemoryEdit());
|
||||
document.getElementById('memory-save-btn').addEventListener('click', () => saveMemoryEdit());
|
||||
document.getElementById('memory-cancel-btn').addEventListener('click', () => cancelMemoryEdit());
|
||||
document.getElementById('logs-server-level').addEventListener('change', (e) => setServerLogLevel(e.target.value));
|
||||
document.getElementById('logs-pause-btn').addEventListener('click', () => toggleLogsPause());
|
||||
document.getElementById('logs-clear-btn').addEventListener('click', () => clearLogs());
|
||||
document.getElementById('wasm-install-btn').addEventListener('click', () => installWasmExtension());
|
||||
document.getElementById('mcp-add-btn').addEventListener('click', () => addMcpServer());
|
||||
document.getElementById('skill-search-btn').addEventListener('click', () => searchClawHub());
|
||||
document.getElementById('skill-install-btn').addEventListener('click', () => installSkillFromForm());
|
||||
|
||||
// --- Delegated Event Handlers (for dynamically generated HTML) ---
|
||||
|
||||
document.addEventListener('click', function(e) {
|
||||
const el = e.target.closest('[data-action]');
|
||||
if (!el) return;
|
||||
const action = el.dataset.action;
|
||||
|
||||
switch (action) {
|
||||
case 'copy-code':
|
||||
copyCodeBlock(el);
|
||||
break;
|
||||
case 'breadcrumb-root':
|
||||
e.preventDefault();
|
||||
loadMemoryTree();
|
||||
break;
|
||||
case 'breadcrumb-file':
|
||||
e.preventDefault();
|
||||
readMemoryFile(el.dataset.path);
|
||||
break;
|
||||
case 'cancel-job':
|
||||
e.stopPropagation();
|
||||
cancelJob(el.dataset.id);
|
||||
break;
|
||||
case 'open-job':
|
||||
openJobDetail(el.dataset.id);
|
||||
break;
|
||||
case 'close-job-detail':
|
||||
closeJobDetail();
|
||||
break;
|
||||
case 'restart-job':
|
||||
restartJob(el.dataset.id);
|
||||
break;
|
||||
case 'open-routine':
|
||||
openRoutineDetail(el.dataset.id);
|
||||
break;
|
||||
case 'toggle-routine':
|
||||
e.stopPropagation();
|
||||
toggleRoutine(el.dataset.id);
|
||||
break;
|
||||
case 'trigger-routine':
|
||||
e.stopPropagation();
|
||||
triggerRoutine(el.dataset.id);
|
||||
break;
|
||||
case 'delete-routine':
|
||||
e.stopPropagation();
|
||||
deleteRoutine(el.dataset.id, el.dataset.name);
|
||||
break;
|
||||
case 'close-routine-detail':
|
||||
closeRoutineDetail();
|
||||
break;
|
||||
case 'view-run-job':
|
||||
e.preventDefault();
|
||||
switchTab('jobs');
|
||||
openJobDetail(el.dataset.id);
|
||||
break;
|
||||
case 'copy-tee-report':
|
||||
copyTeeReport();
|
||||
break;
|
||||
case 'switch-language':
|
||||
if (typeof switchLanguage === 'function') switchLanguage(el.dataset.lang);
|
||||
break;
|
||||
}
|
||||
});
|
||||
|
||||
document.getElementById('language-btn').addEventListener('click', function() {
|
||||
if (typeof toggleLanguageMenu === 'function') toggleLanguageMenu();
|
||||
});
|
||||
|
||||
@@ -9,12 +9,12 @@
|
||||
<link rel="preconnect" href="https://fonts.gstatic.com" crossorigin>
|
||||
<link href="https://fonts.googleapis.com/css2?family=DM+Sans:wght@400;500;600;700&family=IBM+Plex+Mono:wght@400;500;600&display=swap" rel="stylesheet">
|
||||
<link rel="stylesheet" href="/style.css">
|
||||
|
||||
|
||||
<!-- i18n Modules -->
|
||||
<script src="/i18n/index.js"></script>
|
||||
<script src="/i18n/en.js"></script>
|
||||
<script src="/i18n/zh-CN.js"></script>
|
||||
|
||||
|
||||
<script
|
||||
src="https://cdnjs.cloudflare.com/ajax/libs/dompurify/3.2.3/purify.min.js"
|
||||
integrity="sha384-osZDKVu4ipZP703HmPOhWdyBajcFyjX2Psjk//TG1Rc0AdwEtuToaylrmcK3LdAl"
|
||||
@@ -37,7 +37,7 @@
|
||||
<div class="auth-form">
|
||||
<label for="token-input" data-i18n="auth.tokenLabel">Gateway Token</label>
|
||||
<input type="password" id="token-input" data-i18n="auth.tokenPlaceholder" data-i18n-attr="placeholder" placeholder="Paste your auth token" autofocus>
|
||||
<button onclick="authenticate()" data-i18n="auth.connect">Connect</button>
|
||||
<button id="auth-connect-btn" data-i18n="auth.connect">Connect</button>
|
||||
</div>
|
||||
<div id="auth-error"></div>
|
||||
<p class="auth-hint" data-i18n="auth.hint">Enter the GATEWAY_AUTH_TOKEN from your .env configuration.</p>
|
||||
@@ -46,11 +46,11 @@
|
||||
|
||||
<!-- Restart Confirmation Modal -->
|
||||
<div id="restart-confirm-modal" class="restart-modal" style="display: none;">
|
||||
<div class="restart-modal-overlay" onclick="cancelRestart()"></div>
|
||||
<div class="restart-modal-overlay" id="restart-overlay"></div>
|
||||
<div class="restart-modal-content">
|
||||
<div class="restart-modal-header">
|
||||
<h2 data-i18n="restart.title">Restart IronClaw Instance</h2>
|
||||
<button class="restart-modal-close" onclick="cancelRestart()" data-i18n="restart.closeTooltip" data-i18n-attr="title"
|
||||
<button class="restart-modal-close" id="restart-close-btn" data-i18n="restart.closeTooltip" data-i18n-attr="title"
|
||||
title="Close">×</button>
|
||||
</div>
|
||||
<div class="restart-modal-body">
|
||||
@@ -63,8 +63,8 @@
|
||||
</div>
|
||||
</div>
|
||||
<div class="restart-modal-footer">
|
||||
<button class="restart-modal-btn cancel" onclick="cancelRestart()" data-i18n="restart.cancel">Cancel</button>
|
||||
<button class="restart-modal-btn confirm" onclick="confirmRestart()" data-i18n="restart.confirm">Confirm Restart</button>
|
||||
<button class="restart-modal-btn cancel" id="restart-cancel-btn" data-i18n="restart.cancel">Cancel</button>
|
||||
<button class="restart-modal-btn confirm" id="restart-confirm-btn" data-i18n="restart.confirm">Confirm Restart</button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
@@ -98,17 +98,17 @@
|
||||
<button data-tab="extensions" data-i18n="tab.extensions">Extensions</button>
|
||||
<button data-tab="skills" data-i18n="tab.skills">Skills</button>
|
||||
<div class="spacer"></div>
|
||||
|
||||
|
||||
<!-- Language Switcher -->
|
||||
<div class="language-switcher">
|
||||
<button class="language-btn" id="language-btn" type="button" onclick="toggleLanguageMenu()" title="Switch Language"
|
||||
<button class="language-btn" id="language-btn" type="button" title="Switch Language"
|
||||
aria-label="Switch language" aria-haspopup="true" aria-expanded="false" aria-controls="language-menu">🌐</button>
|
||||
<div class="language-menu" id="language-menu" style="display: none;">
|
||||
<button type="button" class="language-option" onclick="switchLanguage('en')" data-lang="en">English</button>
|
||||
<button type="button" class="language-option" onclick="switchLanguage('zh-CN')" data-lang="zh-CN">简体中文</button>
|
||||
<button type="button" class="language-option" data-action="switch-language" data-lang="en">English</button>
|
||||
<button type="button" class="language-option" data-action="switch-language" data-lang="zh-CN">简体中文</button>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
|
||||
<button class="status-logs-btn" data-tab="logs" data-i18n="tab.logs" title="Logs">Logs</button>
|
||||
<div class="tee-shield" id="tee-shield" style="display:none" title="Running in a Trusted Execution Environment">
|
||||
<svg width="16" height="16" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round">
|
||||
@@ -122,7 +122,7 @@
|
||||
<span id="sse-status" data-i18n="status.connected">Connected</span>
|
||||
<div class="gateway-popover" id="gateway-popover"></div>
|
||||
</div>
|
||||
<button class="restart-btn" id="restart-btn" onclick="triggerRestart()" data-i18n="status.restartTooltip"
|
||||
<button class="restart-btn" id="restart-btn" data-i18n="status.restartTooltip"
|
||||
data-i18n-attr="title" title="Gracefully restart the process" style="display: none;">
|
||||
<svg id="restart-icon" width="13" height="13" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round">
|
||||
<path d="M23 4v6h-6"></path>
|
||||
@@ -137,13 +137,13 @@
|
||||
<div class="tab-panel active" id="tab-chat">
|
||||
<div class="thread-sidebar" id="thread-sidebar">
|
||||
<div class="thread-sidebar-header">
|
||||
<button class="thread-new-btn" onclick="createNewThread()" data-i18n="chat.newThread" data-i18n-attr="title"
|
||||
<button class="thread-new-btn" id="thread-new-btn" data-i18n="chat.newThread" data-i18n-attr="title"
|
||||
title="New thread (Ctrl/Cmd+N)">+</button>
|
||||
<div class="spacer"></div>
|
||||
<button class="thread-toggle-btn" id="thread-toggle-btn" onclick="toggleThreadSidebar()" data-i18n="chat.toggleSidebar"
|
||||
<button class="thread-toggle-btn" id="thread-toggle-btn" data-i18n="chat.toggleSidebar"
|
||||
data-i18n-attr="title" title="Toggle sidebar">«</button>
|
||||
</div>
|
||||
<div class="assistant-item" id="assistant-thread" onclick="switchToAssistant()">
|
||||
<div class="assistant-item" id="assistant-thread">
|
||||
<span class="assistant-label" id="assistant-label" data-i18n="chat.assistant">Assistant</span>
|
||||
<span class="assistant-meta" id="assistant-meta"></span>
|
||||
</div>
|
||||
@@ -161,7 +161,7 @@
|
||||
<input type="file" id="image-file-input" accept="image/*" multiple style="display:none">
|
||||
<button id="attach-btn" class="attach-btn" data-i18n="chat.attachImages" data-i18n-attr="title" title="Attach images"
|
||||
aria-label="Attach images">📎</button>
|
||||
<button id="send-btn" onclick="sendMessage()" data-i18n="chat.send">Send</button>
|
||||
<button id="send-btn" data-i18n="chat.send">Send</button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
@@ -178,7 +178,7 @@
|
||||
<div class="memory-content">
|
||||
<div class="memory-breadcrumb" id="memory-breadcrumb">
|
||||
<span id="memory-breadcrumb-path">workspace /</span>
|
||||
<button class="memory-edit-btn" id="memory-edit-btn" style="display:none" onclick="startMemoryEdit()" data-i18n="memory.edit">Edit</button>
|
||||
<button class="memory-edit-btn" id="memory-edit-btn" style="display:none" data-i18n="memory.edit">Edit</button>
|
||||
</div>
|
||||
<div class="memory-viewer" id="memory-viewer">
|
||||
<div class="empty" data-i18n="memory.selectFile">Select a file to view its contents</div>
|
||||
@@ -186,8 +186,8 @@
|
||||
<div class="memory-editor" id="memory-editor" style="display:none">
|
||||
<textarea id="memory-edit-textarea"></textarea>
|
||||
<div class="memory-editor-actions">
|
||||
<button class="btn-save" onclick="saveMemoryEdit()" data-i18n="memory.save">Save</button>
|
||||
<button class="btn-cancel-edit" onclick="cancelMemoryEdit()" data-i18n="memory.cancel">Cancel</button>
|
||||
<button class="btn-save" id="memory-save-btn" data-i18n="memory.save">Save</button>
|
||||
<button class="btn-cancel-edit" id="memory-cancel-btn" data-i18n="memory.cancel">Cancel</button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
@@ -219,7 +219,7 @@
|
||||
<div class="tab-panel" id="tab-logs">
|
||||
<div class="logs-container">
|
||||
<div class="logs-toolbar">
|
||||
<select id="logs-server-level" onchange="setServerLogLevel(this.value)" title="Server-side log level (changes what the server emits)">
|
||||
<select id="logs-server-level" title="Server-side log level (changes what the server emits)">
|
||||
<option value="error">Server: ERROR</option>
|
||||
<option value="warn">Server: WARN</option>
|
||||
<option value="info" selected>Server: INFO</option>
|
||||
@@ -234,8 +234,8 @@
|
||||
</select>
|
||||
<input type="text" id="logs-target-filter" placeholder="Filter by target...">
|
||||
<label class="logs-checkbox"><input type="checkbox" id="logs-autoscroll" checked> <span data-i18n="logs.autoScroll">Auto-scroll</span></label>
|
||||
<button id="logs-pause-btn" onclick="toggleLogsPause()" data-i18n="logs.pause">Pause</button>
|
||||
<button onclick="clearLogs()" data-i18n="logs.clear">Clear</button>
|
||||
<button id="logs-pause-btn" data-i18n="logs.pause">Pause</button>
|
||||
<button id="logs-clear-btn" data-i18n="logs.clear">Clear</button>
|
||||
</div>
|
||||
<div class="logs-output" id="logs-output"></div>
|
||||
</div>
|
||||
@@ -287,7 +287,7 @@
|
||||
<div class="ext-install-form">
|
||||
<input type="text" id="wasm-install-name" data-i18n-placeholder="common.name" placeholder="Extension name">
|
||||
<input type="text" id="wasm-install-url" placeholder="URL to .tar.gz bundle">
|
||||
<button onclick="installWasmExtension()" data-i18n="extensions.install">Install</button>
|
||||
<button id="wasm-install-btn" data-i18n="extensions.install">Install</button>
|
||||
</div>
|
||||
</div>
|
||||
<div class="extensions-section">
|
||||
@@ -299,7 +299,7 @@
|
||||
<div class="ext-install-form">
|
||||
<input type="text" id="mcp-install-name" data-i18n-placeholder="common.name" placeholder="Server name">
|
||||
<input type="text" id="mcp-install-url" placeholder="MCP server URL (https://...)">
|
||||
<button onclick="addMcpServer()" data-i18n="mcp.add">Add</button>
|
||||
<button id="mcp-add-btn" data-i18n="mcp.add">Add</button>
|
||||
</div>
|
||||
</div>
|
||||
<div class="extensions-section">
|
||||
@@ -320,7 +320,7 @@
|
||||
<h3 data-i18n="skills.searchClawHub">Search ClawHub</h3>
|
||||
<div class="skill-search-box">
|
||||
<input type="text" id="skill-search-input" data-i18n-placeholder="skills.searchPlaceholder" placeholder="Search...">
|
||||
<button onclick="searchClawHub()" data-i18n="skills.search">Search</button>
|
||||
<button id="skill-search-btn" data-i18n="skills.search">Search</button>
|
||||
</div>
|
||||
<div class="extensions-list" id="skill-search-results"></div>
|
||||
</div>
|
||||
@@ -335,7 +335,7 @@
|
||||
<div class="ext-install-form">
|
||||
<input type="text" id="skill-install-name" data-i18n-placeholder="skills.namePlaceholder" placeholder="Skill name or slug">
|
||||
<input type="text" id="skill-install-url" data-i18n-placeholder="skills.urlPlaceholder" placeholder="HTTPS URL to SKILL.md (optional)">
|
||||
<button onclick="installSkillFromForm()" data-i18n="extensions.install">Install</button>
|
||||
<button id="skill-install-btn" data-i18n="extensions.install">Install</button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
@@ -666,6 +666,7 @@ body {
|
||||
font-size: 14px;
|
||||
line-height: 1.5;
|
||||
word-wrap: break-word;
|
||||
position: relative;
|
||||
}
|
||||
|
||||
.message.user {
|
||||
@@ -686,6 +687,58 @@ body {
|
||||
line-height: 1.6;
|
||||
}
|
||||
|
||||
.message.has-copy {
|
||||
padding-right: 52px;
|
||||
}
|
||||
|
||||
.message-content {
|
||||
min-width: 0;
|
||||
}
|
||||
|
||||
.message-copy-btn {
|
||||
position: absolute;
|
||||
top: 8px;
|
||||
right: 8px;
|
||||
z-index: 2;
|
||||
border: 1px solid var(--border);
|
||||
background: var(--bg-primary);
|
||||
color: var(--text-secondary);
|
||||
border-radius: 8px;
|
||||
font-size: 11px;
|
||||
padding: 2px 8px;
|
||||
opacity: 0;
|
||||
pointer-events: none;
|
||||
transition: opacity 0.15s ease;
|
||||
}
|
||||
|
||||
.message.user:hover .message-copy-btn,
|
||||
.message.assistant:hover .message-copy-btn,
|
||||
.message.user:focus-within .message-copy-btn,
|
||||
.message.assistant:focus-within .message-copy-btn {
|
||||
opacity: 1;
|
||||
pointer-events: auto;
|
||||
}
|
||||
|
||||
.message-copy-btn:focus-visible {
|
||||
opacity: 1;
|
||||
pointer-events: auto;
|
||||
outline: 2px solid var(--accent);
|
||||
outline-offset: 1px;
|
||||
}
|
||||
|
||||
.message-copy-btn:hover {
|
||||
background: var(--bg-secondary);
|
||||
color: var(--text-primary);
|
||||
}
|
||||
|
||||
@media (hover: none) {
|
||||
.message.user .message-copy-btn,
|
||||
.message.assistant .message-copy-btn {
|
||||
opacity: 1;
|
||||
pointer-events: auto;
|
||||
}
|
||||
}
|
||||
|
||||
.message.system {
|
||||
align-self: center;
|
||||
background: var(--bg-tertiary);
|
||||
@@ -1166,7 +1219,21 @@ body {
|
||||
color: var(--danger);
|
||||
}
|
||||
|
||||
/* Auth card (inline in chat) */
|
||||
/* Auth prompt */
|
||||
.auth-overlay {
|
||||
position: fixed;
|
||||
top: 0;
|
||||
left: 0;
|
||||
width: 100%;
|
||||
height: 100%;
|
||||
background: rgba(0, 0, 0, 0.6);
|
||||
z-index: 1001;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
padding: 16px;
|
||||
}
|
||||
|
||||
.auth-card {
|
||||
align-self: flex-start;
|
||||
max-width: 80%;
|
||||
@@ -1181,6 +1248,16 @@ body {
|
||||
transition: border-color 0.2s;
|
||||
}
|
||||
|
||||
.auth-overlay .auth-card {
|
||||
width: 460px;
|
||||
max-width: min(460px, 90vw);
|
||||
margin: 0;
|
||||
align-self: auto;
|
||||
background: var(--bg);
|
||||
border-color: rgba(52, 211, 153, 0.35);
|
||||
box-shadow: 0 24px 48px rgba(0, 0, 0, 0.35);
|
||||
}
|
||||
|
||||
.auth-card .auth-header {
|
||||
font-weight: 600;
|
||||
color: var(--accent);
|
||||
@@ -3826,8 +3903,9 @@ mark {
|
||||
display: flex;
|
||||
flex-direction: row;
|
||||
gap: 8px;
|
||||
padding: 4px;
|
||||
padding: 10px 4px 4px;
|
||||
overflow-x: auto;
|
||||
overflow-y: visible;
|
||||
min-height: 0;
|
||||
width: 100%;
|
||||
}
|
||||
@@ -3848,6 +3926,7 @@ mark {
|
||||
border-radius: 6px;
|
||||
object-fit: cover;
|
||||
display: block;
|
||||
border: 1px solid var(--border);
|
||||
}
|
||||
|
||||
.image-preview-remove {
|
||||
@@ -3856,21 +3935,41 @@ mark {
|
||||
right: -6px;
|
||||
width: 18px;
|
||||
height: 18px;
|
||||
padding: 4px !important;
|
||||
border-radius: 50%;
|
||||
background: var(--danger);
|
||||
color: #fff;
|
||||
border: none;
|
||||
font-size: 12px;
|
||||
line-height: 18px;
|
||||
border: 2px solid var(--bg-secondary);
|
||||
font-size: 10px;
|
||||
line-height: 14px;
|
||||
text-align: center;
|
||||
cursor: pointer;
|
||||
padding: 0;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
}
|
||||
|
||||
.image-preview-remove:hover {
|
||||
background: #c33;
|
||||
}
|
||||
|
||||
/* Images in user message bubbles */
|
||||
.message.user .message-images {
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
gap: 6px;
|
||||
margin-bottom: 4px;
|
||||
}
|
||||
|
||||
.message.user .message-image-thumb {
|
||||
max-width: 240px;
|
||||
max-height: 200px;
|
||||
border-radius: 8px;
|
||||
object-fit: cover;
|
||||
display: block;
|
||||
}
|
||||
|
||||
/* Generated Image */
|
||||
.generated-image-card {
|
||||
max-width: 512px;
|
||||
|
||||
+12
-4
@@ -176,8 +176,12 @@ async fn handle_client_message(
|
||||
incoming = incoming.with_attachments(attachments);
|
||||
}
|
||||
|
||||
let tx_guard = state.msg_tx.read().await;
|
||||
if let Some(ref tx) = *tx_guard {
|
||||
// Clone sender to avoid holding RwLock read guard across send().await
|
||||
let tx = {
|
||||
let tx_guard = state.msg_tx.read().await;
|
||||
tx_guard.as_ref().cloned()
|
||||
};
|
||||
if let Some(tx) = tx {
|
||||
if tx.send(incoming).await.is_err() {
|
||||
let _ = direct_tx
|
||||
.send(WsServerMessage::Error {
|
||||
@@ -245,8 +249,12 @@ async fn handle_client_message(
|
||||
if let Some(ref tid) = thread_id {
|
||||
msg = msg.with_thread(tid);
|
||||
}
|
||||
let tx_guard = state.msg_tx.read().await;
|
||||
if let Some(ref tx) = *tx_guard {
|
||||
// Clone sender to avoid holding RwLock read guard across send().await
|
||||
let tx = {
|
||||
let tx_guard = state.msg_tx.read().await;
|
||||
tx_guard.as_ref().cloned()
|
||||
};
|
||||
if let Some(tx) = tx {
|
||||
let _ = tx.send(msg).await;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,6 +7,7 @@
|
||||
//! - Managing WASM tools (`tool install`, `tool list`, `tool remove`)
|
||||
//! - Managing MCP servers (`mcp add`, `mcp auth`, `mcp list`, `mcp test`)
|
||||
//! - Querying workspace memory (`memory search`, `memory read`, `memory write`)
|
||||
//! - Managing routines (`routines list`, `routines create`, `routines edit`, ...)
|
||||
//! - Managing OS service (`service install`, `service start`, `service stop`)
|
||||
//! - Listing configured channels (`channels list`)
|
||||
//! - Active health diagnostics (`doctor`)
|
||||
@@ -23,6 +24,7 @@ pub mod memory;
|
||||
pub mod oauth_defaults;
|
||||
mod pairing;
|
||||
mod registry;
|
||||
mod routines;
|
||||
mod service;
|
||||
mod skills;
|
||||
pub mod status;
|
||||
@@ -39,6 +41,7 @@ pub use memory::MemoryCommand;
|
||||
pub use memory::run_memory_command_with_db;
|
||||
pub use pairing::{PairingCommand, run_pairing_command, run_pairing_command_with_store};
|
||||
pub use registry::{RegistryCommand, run_registry_command};
|
||||
pub use routines::{RoutinesCommand, run_routines_command};
|
||||
pub use service::{ServiceCommand, run_service_command};
|
||||
pub use skills::{SkillsCommand, run_skills_command};
|
||||
pub use status::run_status_command;
|
||||
@@ -147,6 +150,15 @@ pub enum Command {
|
||||
)]
|
||||
Channels(ChannelsCommand),
|
||||
|
||||
/// Manage routines (scheduled, event-driven, webhook, manual)
|
||||
#[command(
|
||||
subcommand,
|
||||
alias = "cron",
|
||||
about = "Manage routines",
|
||||
long_about = "List, create, edit, enable/disable, delete, and view history of routines.\nExamples:\n ironclaw routines list\n ironclaw routines create --name daily-digest --schedule '0 0 9 * * *' --prompt 'Summarize today'"
|
||||
)]
|
||||
Routines(RoutinesCommand),
|
||||
|
||||
/// Manage MCP servers (hosted tool providers)
|
||||
#[command(
|
||||
subcommand,
|
||||
@@ -281,6 +293,23 @@ pub async fn init_secrets_store()
|
||||
Ok(crate::db::create_secrets_store(&config.database, crypto).await?)
|
||||
}
|
||||
|
||||
/// Run the Routines CLI subcommand.
|
||||
pub async fn run_routines_cli(
|
||||
routines_cmd: &RoutinesCommand,
|
||||
config_path: Option<&std::path::Path>,
|
||||
) -> anyhow::Result<()> {
|
||||
let config = crate::config::Config::from_env_with_toml(config_path)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("{e:#}"))?;
|
||||
|
||||
let db: Arc<dyn crate::db::Database> = crate::db::connect_from_config(&config.database)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("{e:#}"))?;
|
||||
|
||||
let user_id = std::env::var("GATEWAY_USER_ID").unwrap_or_else(|_| "default".to_string());
|
||||
run_routines_command(routines_cmd.clone(), db, &user_id).await
|
||||
}
|
||||
|
||||
/// Run the Memory CLI subcommand.
|
||||
pub async fn run_memory_command(mem_cmd: &MemoryCommand) -> anyhow::Result<()> {
|
||||
let config = crate::config::Config::from_env()
|
||||
|
||||
@@ -172,6 +172,35 @@ pub async fn exchange_oauth_code(
|
||||
redirect_uri: &str,
|
||||
code_verifier: Option<&str>,
|
||||
access_token_field: &str,
|
||||
) -> Result<OAuthTokenResponse, OAuthCallbackError> {
|
||||
// Delegates to exchange_oauth_code_with_resource with resource=None.
|
||||
// Non-MCP OAuth flows don't need the RFC 8707 resource parameter.
|
||||
exchange_oauth_code_with_resource(
|
||||
token_url,
|
||||
client_id,
|
||||
client_secret,
|
||||
code,
|
||||
redirect_uri,
|
||||
code_verifier,
|
||||
access_token_field,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Exchange an OAuth authorization code for tokens, with optional RFC 8707 `resource` parameter.
|
||||
///
|
||||
/// The `resource` parameter scopes the issued token to a specific server (used by MCP OAuth).
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub async fn exchange_oauth_code_with_resource(
|
||||
token_url: &str,
|
||||
client_id: &str,
|
||||
client_secret: Option<&str>,
|
||||
code: &str,
|
||||
redirect_uri: &str,
|
||||
code_verifier: Option<&str>,
|
||||
access_token_field: &str,
|
||||
resource: Option<&str>,
|
||||
) -> Result<OAuthTokenResponse, OAuthCallbackError> {
|
||||
let client = reqwest::Client::new();
|
||||
let mut token_params = vec![
|
||||
@@ -184,6 +213,12 @@ pub async fn exchange_oauth_code(
|
||||
token_params.push(("code_verifier", verifier.to_string()));
|
||||
}
|
||||
|
||||
// RFC 8707: include the `resource` parameter so the authorization server
|
||||
// scopes the issued token to the specific MCP server (protected resource).
|
||||
if let Some(resource) = resource {
|
||||
token_params.push(("resource", resource.to_string()));
|
||||
}
|
||||
|
||||
let mut request = client.post(token_url);
|
||||
|
||||
if let Some(secret) = client_secret {
|
||||
@@ -388,6 +423,12 @@ pub struct PendingOAuthFlow {
|
||||
pub sse_sender: Option<tokio::sync::broadcast::Sender<crate::channels::web::types::SseEvent>>,
|
||||
/// Gateway auth token for authenticating with the platform token exchange proxy.
|
||||
pub gateway_token: Option<String>,
|
||||
/// RFC 8707 resource parameter (MCP OAuth only).
|
||||
/// Sent during token exchange to scope the token to a specific MCP server.
|
||||
pub resource: Option<String>,
|
||||
/// Secret name for persisting the client ID (MCP OAuth only).
|
||||
/// Needed so token refresh can find the client_id after the session ends.
|
||||
pub client_id_secret_name: Option<String>,
|
||||
/// When this flow was created (for expiry).
|
||||
pub created_at: std::time::Instant,
|
||||
}
|
||||
@@ -975,4 +1016,42 @@ mod tests {
|
||||
assert_eq!(strip_instance_prefix("abc123"), "abc123");
|
||||
assert_eq!(strip_instance_prefix(""), "");
|
||||
}
|
||||
|
||||
/// Verify that `build_oauth_url` includes the RFC 8707 `resource` parameter
|
||||
/// when passed through `extra_params`, which is how MCP OAuth gateway mode
|
||||
/// scopes tokens to a specific MCP server.
|
||||
#[test]
|
||||
fn test_build_oauth_url_includes_resource_via_extra_params() {
|
||||
use std::collections::HashMap;
|
||||
|
||||
use crate::cli::oauth_defaults::build_oauth_url;
|
||||
|
||||
let mut extra = HashMap::new();
|
||||
extra.insert(
|
||||
"resource".to_string(),
|
||||
"https://mcp.example.com".to_string(),
|
||||
);
|
||||
|
||||
let result = build_oauth_url(
|
||||
"https://auth.example.com/authorize",
|
||||
"client-123",
|
||||
"https://gateway.example.com/oauth/callback",
|
||||
&["read".to_string()],
|
||||
true,
|
||||
&extra,
|
||||
);
|
||||
|
||||
// The resource parameter should be URL-encoded in the auth URL
|
||||
assert!(
|
||||
result
|
||||
.url
|
||||
.contains("resource=https%3A%2F%2Fmcp.example.com"),
|
||||
"Expected resource param in URL: {}",
|
||||
result.url
|
||||
);
|
||||
// State and PKCE should be present
|
||||
assert!(result.url.contains("state="));
|
||||
assert!(result.url.contains("code_challenge="));
|
||||
assert!(result.code_verifier.is_some());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,732 @@
|
||||
//! `ironclaw routines` — manage scheduled routines from the CLI.
|
||||
//!
|
||||
//! Provides subcommands for listing, creating, editing, enabling/disabling,
|
||||
//! deleting, and viewing run history of routines without starting the full agent.
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use chrono::{DateTime, Utc};
|
||||
use clap::Subcommand;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::agent::routine::{
|
||||
NotifyConfig, Routine, RoutineAction, RoutineGuardrails, Trigger, next_cron_fire,
|
||||
};
|
||||
use crate::db::Database;
|
||||
|
||||
/// Routines subcommands.
|
||||
#[derive(Subcommand, Debug, Clone)]
|
||||
pub enum RoutinesCommand {
|
||||
/// List routines
|
||||
List {
|
||||
/// Filter by trigger type (e.g. "cron", "webhook", "event")
|
||||
#[arg(long)]
|
||||
trigger: Option<String>,
|
||||
|
||||
/// Include disabled routines
|
||||
#[arg(long)]
|
||||
disabled: bool,
|
||||
|
||||
/// Output as JSON (for scripting)
|
||||
#[arg(long)]
|
||||
json: bool,
|
||||
},
|
||||
|
||||
/// Create a new cron routine
|
||||
#[command(alias = "add")]
|
||||
Create {
|
||||
/// Routine name (must be unique per user)
|
||||
#[arg(long)]
|
||||
name: String,
|
||||
|
||||
/// Cron schedule (6-field: "sec min hour day month weekday")
|
||||
#[arg(long)]
|
||||
schedule: String,
|
||||
|
||||
/// Prompt for the LLM
|
||||
#[arg(long)]
|
||||
prompt: String,
|
||||
|
||||
/// Optional description
|
||||
#[arg(long, default_value = "")]
|
||||
description: String,
|
||||
|
||||
/// IANA timezone (e.g. "America/New_York")
|
||||
#[arg(long)]
|
||||
timezone: Option<String>,
|
||||
|
||||
/// Cooldown between fires in seconds
|
||||
#[arg(long, default_value = "300")]
|
||||
cooldown: u64,
|
||||
|
||||
/// Notification channel
|
||||
#[arg(long)]
|
||||
notify_channel: Option<String>,
|
||||
},
|
||||
|
||||
/// Edit an existing routine
|
||||
#[command(alias = "update")]
|
||||
Edit {
|
||||
/// Routine name
|
||||
#[arg(long)]
|
||||
name: String,
|
||||
|
||||
/// New schedule
|
||||
#[arg(long)]
|
||||
schedule: Option<String>,
|
||||
|
||||
/// New prompt
|
||||
#[arg(long)]
|
||||
prompt: Option<String>,
|
||||
|
||||
/// New description
|
||||
#[arg(long)]
|
||||
description: Option<String>,
|
||||
|
||||
/// New timezone
|
||||
#[arg(long)]
|
||||
timezone: Option<String>,
|
||||
|
||||
/// New cooldown in seconds
|
||||
#[arg(long)]
|
||||
cooldown: Option<u64>,
|
||||
},
|
||||
|
||||
/// Enable a routine
|
||||
Enable {
|
||||
/// Routine name
|
||||
name: String,
|
||||
},
|
||||
|
||||
/// Disable a routine
|
||||
Disable {
|
||||
/// Routine name
|
||||
name: String,
|
||||
},
|
||||
|
||||
/// Delete a routine
|
||||
#[command(alias = "rm")]
|
||||
Delete {
|
||||
/// Routine name
|
||||
name: String,
|
||||
|
||||
/// Skip confirmation prompt
|
||||
#[arg(short, long)]
|
||||
yes: bool,
|
||||
},
|
||||
|
||||
/// Show run history for a routine
|
||||
#[command(alias = "runs")]
|
||||
History {
|
||||
/// Routine name
|
||||
name: String,
|
||||
|
||||
/// Maximum number of runs to show
|
||||
#[arg(short, long, default_value = "10")]
|
||||
limit: i64,
|
||||
|
||||
/// Output as JSON (for scripting)
|
||||
#[arg(long)]
|
||||
json: bool,
|
||||
},
|
||||
}
|
||||
|
||||
/// Run a routines CLI command against the database.
|
||||
pub async fn run_routines_command(
|
||||
cmd: RoutinesCommand,
|
||||
db: Arc<dyn Database>,
|
||||
user_id: &str,
|
||||
) -> anyhow::Result<()> {
|
||||
match cmd {
|
||||
RoutinesCommand::List {
|
||||
trigger,
|
||||
disabled,
|
||||
json,
|
||||
} => list(&db, user_id, trigger.as_deref(), disabled, json).await,
|
||||
RoutinesCommand::Create {
|
||||
name,
|
||||
schedule,
|
||||
prompt,
|
||||
description,
|
||||
timezone,
|
||||
cooldown,
|
||||
notify_channel,
|
||||
} => {
|
||||
create(
|
||||
&db,
|
||||
user_id,
|
||||
&name,
|
||||
&schedule,
|
||||
&prompt,
|
||||
&description,
|
||||
timezone.as_deref(),
|
||||
cooldown,
|
||||
notify_channel,
|
||||
)
|
||||
.await
|
||||
}
|
||||
RoutinesCommand::Edit {
|
||||
name,
|
||||
schedule,
|
||||
prompt,
|
||||
description,
|
||||
timezone,
|
||||
cooldown,
|
||||
} => {
|
||||
edit(
|
||||
&db,
|
||||
user_id,
|
||||
&name,
|
||||
schedule.as_deref(),
|
||||
prompt.as_deref(),
|
||||
description.as_deref(),
|
||||
timezone.as_deref(),
|
||||
cooldown,
|
||||
)
|
||||
.await
|
||||
}
|
||||
RoutinesCommand::Enable { name } => set_enabled(&db, user_id, &name, true).await,
|
||||
RoutinesCommand::Disable { name } => set_enabled(&db, user_id, &name, false).await,
|
||||
RoutinesCommand::Delete { name, yes } => delete(&db, user_id, &name, yes).await,
|
||||
RoutinesCommand::History { name, limit, json } => {
|
||||
history(&db, user_id, &name, limit, json).await
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ── List ────────────────────────────────────────────────────
|
||||
|
||||
async fn list(
|
||||
db: &Arc<dyn Database>,
|
||||
user_id: &str,
|
||||
trigger_filter: Option<&str>,
|
||||
show_disabled: bool,
|
||||
json: bool,
|
||||
) -> anyhow::Result<()> {
|
||||
let routines = db.list_routines(user_id).await?;
|
||||
|
||||
let filtered: Vec<&Routine> = routines
|
||||
.iter()
|
||||
.filter(|r| {
|
||||
trigger_filter
|
||||
.map(|t| r.trigger.type_tag() == t)
|
||||
.unwrap_or(true)
|
||||
})
|
||||
.filter(|r| show_disabled || r.enabled)
|
||||
.collect();
|
||||
|
||||
if json {
|
||||
let items: Vec<serde_json::Value> = filtered
|
||||
.iter()
|
||||
.map(|r| {
|
||||
serde_json::json!({
|
||||
"id": r.id.to_string(),
|
||||
"name": r.name,
|
||||
"trigger": r.trigger.type_tag(),
|
||||
"enabled": r.enabled,
|
||||
"next_fire_at": r.next_fire_at,
|
||||
"last_run_at": r.last_run_at,
|
||||
"run_count": r.run_count,
|
||||
"consecutive_failures": r.consecutive_failures,
|
||||
})
|
||||
})
|
||||
.collect();
|
||||
println!("{}", serde_json::to_string_pretty(&items)?);
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
if filtered.is_empty() {
|
||||
if let Some(t) = trigger_filter {
|
||||
println!("No {t} routines found.");
|
||||
} else {
|
||||
println!("No routines found.");
|
||||
}
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
// Header
|
||||
println!(
|
||||
"{:<36} {:<20} {:<8} {:<8} {:<22} {:<22} {:>5}",
|
||||
"ID", "NAME", "TRIGGER", "STATUS", "NEXT FIRE", "LAST RUN", "RUNS"
|
||||
);
|
||||
println!("{}", "-".repeat(130));
|
||||
|
||||
for r in &filtered {
|
||||
let status = if r.enabled {
|
||||
if r.consecutive_failures > 0 {
|
||||
format!("err({})", r.consecutive_failures)
|
||||
} else {
|
||||
"active".to_string()
|
||||
}
|
||||
} else {
|
||||
"disabled".to_string()
|
||||
};
|
||||
|
||||
let next_fire = r
|
||||
.next_fire_at
|
||||
.map(format_relative)
|
||||
.unwrap_or_else(|| "-".to_string());
|
||||
|
||||
let last_run = r
|
||||
.last_run_at
|
||||
.map(format_relative)
|
||||
.unwrap_or_else(|| "-".to_string());
|
||||
|
||||
let name = truncate(&r.name, 20);
|
||||
|
||||
println!(
|
||||
"{:<36} {:<20} {:<8} {:<8} {:<22} {:<22} {:>5}",
|
||||
r.id,
|
||||
name,
|
||||
r.trigger.type_tag(),
|
||||
status,
|
||||
next_fire,
|
||||
last_run,
|
||||
r.run_count,
|
||||
);
|
||||
}
|
||||
|
||||
println!("\n{} routine(s)", filtered.len());
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ── Create ──────────────────────────────────────────────────
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
async fn create(
|
||||
db: &Arc<dyn Database>,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
schedule: &str,
|
||||
prompt: &str,
|
||||
description: &str,
|
||||
timezone: Option<&str>,
|
||||
cooldown_secs: u64,
|
||||
notify_channel: Option<String>,
|
||||
) -> anyhow::Result<()> {
|
||||
validate_timezone_arg(timezone)?;
|
||||
|
||||
// Validate the cron expression by computing next fire.
|
||||
let next_fire = next_cron_fire(schedule, timezone)
|
||||
.map_err(|e| anyhow::anyhow!("Invalid cron schedule: {e}"))?;
|
||||
|
||||
// Check for name conflict.
|
||||
if db.get_routine_by_name(user_id, name).await?.is_some() {
|
||||
anyhow::bail!("Routine '{}' already exists", name);
|
||||
}
|
||||
|
||||
let now = Utc::now();
|
||||
let routine = Routine {
|
||||
id: Uuid::new_v4(),
|
||||
name: name.to_string(),
|
||||
description: description.to_string(),
|
||||
user_id: user_id.to_string(),
|
||||
enabled: true,
|
||||
trigger: Trigger::Cron {
|
||||
schedule: schedule.to_string(),
|
||||
timezone: timezone.map(String::from),
|
||||
},
|
||||
action: RoutineAction::Lightweight {
|
||||
prompt: prompt.to_string(),
|
||||
context_paths: Vec::new(),
|
||||
max_tokens: 4096,
|
||||
use_tools: false,
|
||||
max_tool_rounds: 0,
|
||||
},
|
||||
guardrails: RoutineGuardrails {
|
||||
cooldown: std::time::Duration::from_secs(cooldown_secs),
|
||||
max_concurrent: 1,
|
||||
dedup_window: None,
|
||||
},
|
||||
notify: NotifyConfig {
|
||||
channel: notify_channel,
|
||||
user: user_id.to_string(),
|
||||
on_attention: true,
|
||||
on_failure: true,
|
||||
on_success: false,
|
||||
},
|
||||
last_run_at: None,
|
||||
next_fire_at: next_fire,
|
||||
run_count: 0,
|
||||
consecutive_failures: 0,
|
||||
state: serde_json::json!({}),
|
||||
created_at: now,
|
||||
updated_at: now,
|
||||
};
|
||||
|
||||
db.create_routine(&routine).await?;
|
||||
|
||||
println!("Created routine '{}'", name);
|
||||
println!(" ID: {}", routine.id);
|
||||
println!(" Schedule: {}", schedule);
|
||||
if let Some(tz) = timezone {
|
||||
println!(" Timezone: {}", tz);
|
||||
}
|
||||
if let Some(nf) = next_fire {
|
||||
println!(" Next fire: {}", format_relative(nf));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ── Edit ────────────────────────────────────────────────────
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
async fn edit(
|
||||
db: &Arc<dyn Database>,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
schedule: Option<&str>,
|
||||
prompt: Option<&str>,
|
||||
description: Option<&str>,
|
||||
timezone: Option<&str>,
|
||||
cooldown: Option<u64>,
|
||||
) -> anyhow::Result<()> {
|
||||
let mut routine = require_routine(db, user_id, name).await?;
|
||||
validate_timezone_arg(timezone)?;
|
||||
|
||||
let mut changed = false;
|
||||
|
||||
// Update schedule if provided (only valid for cron routines).
|
||||
if let Some(new_schedule) = schedule {
|
||||
let tz = timezone.or(match &routine.trigger {
|
||||
Trigger::Cron { timezone, .. } => timezone.as_deref(),
|
||||
_ => None,
|
||||
});
|
||||
let next_fire = next_cron_fire(new_schedule, tz)
|
||||
.map_err(|e| anyhow::anyhow!("Invalid cron schedule: {e}"))?;
|
||||
routine.trigger = Trigger::Cron {
|
||||
schedule: new_schedule.to_string(),
|
||||
timezone: tz.map(String::from),
|
||||
};
|
||||
routine.next_fire_at = next_fire;
|
||||
changed = true;
|
||||
} else if let Some(tz) = timezone {
|
||||
// Update only timezone, recompute next fire with existing schedule.
|
||||
if let Trigger::Cron { ref schedule, .. } = routine.trigger {
|
||||
let next_fire = next_cron_fire(schedule, Some(tz))
|
||||
.map_err(|e| anyhow::anyhow!("Invalid cron schedule: {e}"))?;
|
||||
routine.trigger = Trigger::Cron {
|
||||
schedule: schedule.clone(),
|
||||
timezone: Some(tz.to_string()),
|
||||
};
|
||||
routine.next_fire_at = next_fire;
|
||||
changed = true;
|
||||
} else {
|
||||
anyhow::bail!("Cannot set timezone on non-cron trigger");
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(new_prompt) = prompt {
|
||||
match &mut routine.action {
|
||||
RoutineAction::Lightweight { prompt: p, .. } => {
|
||||
*p = new_prompt.to_string();
|
||||
changed = true;
|
||||
}
|
||||
RoutineAction::FullJob { description: d, .. } => {
|
||||
*d = new_prompt.to_string();
|
||||
changed = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(new_desc) = description {
|
||||
routine.description = new_desc.to_string();
|
||||
changed = true;
|
||||
}
|
||||
|
||||
if let Some(cd) = cooldown {
|
||||
routine.guardrails.cooldown = std::time::Duration::from_secs(cd);
|
||||
changed = true;
|
||||
}
|
||||
|
||||
if !changed {
|
||||
println!("No changes specified.");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
routine.updated_at = Utc::now();
|
||||
db.update_routine(&routine).await?;
|
||||
println!("Updated routine '{}'", name);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ── Enable / Disable ────────────────────────────────────────
|
||||
|
||||
async fn set_enabled(
|
||||
db: &Arc<dyn Database>,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
enabled: bool,
|
||||
) -> anyhow::Result<()> {
|
||||
let mut routine = require_routine(db, user_id, name).await?;
|
||||
|
||||
if routine.enabled == enabled {
|
||||
println!(
|
||||
"Routine '{}' is already {}",
|
||||
name,
|
||||
if enabled { "enabled" } else { "disabled" }
|
||||
);
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
routine.enabled = enabled;
|
||||
|
||||
// Recompute next fire when enabling a cron routine.
|
||||
if enabled
|
||||
&& let Trigger::Cron {
|
||||
ref schedule,
|
||||
ref timezone,
|
||||
} = routine.trigger
|
||||
{
|
||||
routine.next_fire_at = next_cron_fire(schedule, timezone.as_deref())
|
||||
.map_err(|e| anyhow::anyhow!("Failed to compute next fire for stored schedule: {e}"))?;
|
||||
}
|
||||
|
||||
routine.updated_at = Utc::now();
|
||||
db.update_routine(&routine).await?;
|
||||
println!(
|
||||
"{} routine '{}'",
|
||||
if enabled { "Enabled" } else { "Disabled" },
|
||||
name
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ── Delete ──────────────────────────────────────────────────
|
||||
|
||||
async fn delete(
|
||||
db: &Arc<dyn Database>,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
skip_confirm: bool,
|
||||
) -> anyhow::Result<()> {
|
||||
let routine = require_routine(db, user_id, name).await?;
|
||||
|
||||
if !skip_confirm {
|
||||
println!("Routine: {}", routine.name);
|
||||
println!(" ID: {}", routine.id);
|
||||
println!(" Trigger: {}", routine.trigger.type_tag());
|
||||
if let Trigger::Cron { ref schedule, .. } = routine.trigger {
|
||||
println!("Schedule: {}", schedule);
|
||||
}
|
||||
println!(" Runs: {}", routine.run_count);
|
||||
print!("\nDelete this routine? [y/N] ");
|
||||
std::io::Write::flush(&mut std::io::stdout())?;
|
||||
|
||||
let mut input = String::new();
|
||||
std::io::stdin().read_line(&mut input)?;
|
||||
if !matches!(input.trim().to_lowercase().as_str(), "y" | "yes") {
|
||||
println!("Cancelled.");
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
|
||||
let deleted = db.delete_routine(routine.id).await?;
|
||||
if deleted {
|
||||
println!("Deleted routine '{}'", name);
|
||||
} else {
|
||||
anyhow::bail!("Failed to delete routine '{}'", name);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ── History ─────────────────────────────────────────────────
|
||||
|
||||
async fn history(
|
||||
db: &Arc<dyn Database>,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
limit: i64,
|
||||
json: bool,
|
||||
) -> anyhow::Result<()> {
|
||||
let routine = require_routine(db, user_id, name).await?;
|
||||
|
||||
let limit = limit.clamp(1, 50);
|
||||
let runs = db.list_routine_runs(routine.id, limit).await?;
|
||||
|
||||
if json {
|
||||
let items: Vec<serde_json::Value> = runs
|
||||
.iter()
|
||||
.map(|run| {
|
||||
serde_json::json!({
|
||||
"id": run.id.to_string(),
|
||||
"status": run.status.to_string(),
|
||||
"started_at": run.started_at,
|
||||
"completed_at": run.completed_at,
|
||||
"result_summary": run.result_summary,
|
||||
"tokens_used": run.tokens_used,
|
||||
})
|
||||
})
|
||||
.collect();
|
||||
println!("{}", serde_json::to_string_pretty(&items)?);
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
if runs.is_empty() {
|
||||
println!("No runs found for routine '{}'", name);
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
println!("Run history for '{}' (last {}):\n", name, runs.len());
|
||||
|
||||
println!(
|
||||
"{:<36} {:<8} {:<20} {:<12} SUMMARY",
|
||||
"RUN ID", "STATUS", "STARTED", "DURATION"
|
||||
);
|
||||
println!("{}", "-".repeat(100));
|
||||
|
||||
for run in &runs {
|
||||
let duration = run
|
||||
.completed_at
|
||||
.map(|end| {
|
||||
let secs = (end - run.started_at).num_seconds();
|
||||
if secs < 60 {
|
||||
format!("{}s", secs)
|
||||
} else {
|
||||
format!("{}m{}s", secs / 60, secs % 60)
|
||||
}
|
||||
})
|
||||
.unwrap_or_else(|| "running".to_string());
|
||||
|
||||
let summary = run
|
||||
.result_summary
|
||||
.as_deref()
|
||||
.map(|s| truncate(s, 40))
|
||||
.unwrap_or_else(|| "-".to_string());
|
||||
|
||||
println!(
|
||||
"{:<36} {:<8} {:<20} {:<12} {}",
|
||||
run.id,
|
||||
run.status,
|
||||
run.started_at.format("%Y-%m-%d %H:%M:%S"),
|
||||
duration,
|
||||
summary,
|
||||
);
|
||||
}
|
||||
|
||||
println!("\n{} run(s) shown", runs.len());
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ── Shared lookup ────────────────────────────────────────────
|
||||
|
||||
/// Look up a routine by name.
|
||||
async fn require_routine(
|
||||
db: &Arc<dyn Database>,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
) -> anyhow::Result<Routine> {
|
||||
db.get_routine_by_name(user_id, name)
|
||||
.await?
|
||||
.ok_or_else(|| anyhow::anyhow!("Routine '{}' not found", name))
|
||||
}
|
||||
|
||||
fn validate_timezone_arg(timezone: Option<&str>) -> anyhow::Result<()> {
|
||||
if let Some(tz) = timezone
|
||||
&& crate::timezone::parse_timezone(tz).is_none()
|
||||
{
|
||||
anyhow::bail!("Invalid timezone: '{tz}' is not a valid IANA timezone");
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ── Helpers ─────────────────────────────────────────────────
|
||||
|
||||
/// Format a datetime relative to now (e.g. "in 2h", "3m ago").
|
||||
fn format_relative(dt: DateTime<Utc>) -> String {
|
||||
let now = Utc::now();
|
||||
let diff = dt.signed_duration_since(now);
|
||||
let secs = diff.num_seconds();
|
||||
|
||||
if secs.abs() < 60 {
|
||||
if secs >= 0 {
|
||||
"in <1m".to_string()
|
||||
} else {
|
||||
"<1m ago".to_string()
|
||||
}
|
||||
} else if secs.abs() < 3600 {
|
||||
let mins = secs.abs() / 60;
|
||||
if secs >= 0 {
|
||||
format!("in {}m", mins)
|
||||
} else {
|
||||
format!("{}m ago", mins)
|
||||
}
|
||||
} else if secs.abs() < 86400 {
|
||||
let hours = secs.abs() / 3600;
|
||||
if secs >= 0 {
|
||||
format!("in {}h", hours)
|
||||
} else {
|
||||
format!("{}h ago", hours)
|
||||
}
|
||||
} else {
|
||||
let days = secs.abs() / 86400;
|
||||
if secs >= 0 {
|
||||
format!("in {}d", days)
|
||||
} else {
|
||||
format!("{}d ago", days)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Truncate a string to a maximum character length.
|
||||
fn truncate(s: &str, max_chars: usize) -> String {
|
||||
if s.chars().count() <= max_chars {
|
||||
s.to_string()
|
||||
} else {
|
||||
let truncated: String = s.chars().take(max_chars.saturating_sub(2)).collect();
|
||||
format!("{}..", truncated)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn format_relative_future() {
|
||||
let future = Utc::now() + chrono::Duration::hours(2);
|
||||
let result = format_relative(future);
|
||||
assert!(
|
||||
result.starts_with("in "),
|
||||
"expected 'in ...' for future time, got: {result}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn format_relative_past() {
|
||||
let past = Utc::now() - chrono::Duration::minutes(30);
|
||||
let result = format_relative(past);
|
||||
assert!(
|
||||
result.ends_with(" ago"),
|
||||
"expected '... ago' for past time, got: {result}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn format_relative_days() {
|
||||
let far_future = Utc::now() + chrono::Duration::days(3);
|
||||
let result = format_relative(far_future);
|
||||
assert!(result.contains('d'), "expected days in: {result}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn truncate_short_string() {
|
||||
assert_eq!(truncate("hello", 10), "hello");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn truncate_long_string() {
|
||||
let result = truncate("hello world", 7);
|
||||
assert_eq!(result, "hello..");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn truncate_multibyte_safe() {
|
||||
// Ensure no panic on multi-byte characters.
|
||||
let cjk = "你好世界测试";
|
||||
let result = truncate(cjk, 4);
|
||||
assert!(result.ends_with(".."), "got: {result}");
|
||||
// Must be valid UTF-8 (would have panicked otherwise).
|
||||
assert!(result.is_char_boundary(result.len()));
|
||||
}
|
||||
}
|
||||
@@ -1,33 +0,0 @@
|
||||
---
|
||||
source: src/cli/mod.rs
|
||||
assertion_line: 302
|
||||
expression: help
|
||||
---
|
||||
Secure personal AI assistant that protects your data and expands its capabilities
|
||||
|
||||
Usage: ironclaw [OPTIONS] [COMMAND]
|
||||
|
||||
Commands:
|
||||
run Run the AI agent
|
||||
onboard Run interactive setup wizard
|
||||
config Manage app configs
|
||||
tool Manage WASM tools
|
||||
registry Browse/install extensions
|
||||
mcp Manage MCP servers
|
||||
memory Manage workspace memory
|
||||
pairing Manage DM pairing
|
||||
service Manage OS service
|
||||
doctor Run diagnostics
|
||||
status Show system status
|
||||
completion Generate completions
|
||||
import Import from other AI systems
|
||||
help Print this message or the help of the given subcommand(s)
|
||||
|
||||
Options:
|
||||
--cli-only Run in interactive CLI mode only (disable other channels)
|
||||
--no-db Skip database connection (for testing)
|
||||
-m, --message <MESSAGE> Single message mode - send one message and exit
|
||||
-c, --config <CONFIG> Configuration file path (optional, uses env vars by default)
|
||||
--no-onboard Skip first-run onboarding check
|
||||
-h, --help Print help (see more with '--help')
|
||||
-V, --version Print version
|
||||
@@ -13,6 +13,7 @@ Commands:
|
||||
tool Manage WASM tools
|
||||
registry Browse/install extensions
|
||||
channels Manage channels
|
||||
routines Manage routines
|
||||
mcp Manage MCP servers
|
||||
memory Manage workspace memory
|
||||
pairing Manage DM pairing
|
||||
|
||||
@@ -1,49 +0,0 @@
|
||||
---
|
||||
source: src/cli/mod.rs
|
||||
assertion_line: 318
|
||||
expression: help
|
||||
---
|
||||
IronClaw is a secure AI assistant. Use 'ironclaw <subcommand> --help' for details.
|
||||
Examples:
|
||||
ironclaw run # Start the agent
|
||||
ironclaw config list # List configs
|
||||
|
||||
Usage: ironclaw [OPTIONS] [COMMAND]
|
||||
|
||||
Commands:
|
||||
run Run the AI agent
|
||||
onboard Run interactive setup wizard
|
||||
config Manage app configs
|
||||
tool Manage WASM tools
|
||||
registry Browse/install extensions
|
||||
mcp Manage MCP servers
|
||||
memory Manage workspace memory
|
||||
pairing Manage DM pairing
|
||||
service Manage OS service
|
||||
doctor Run diagnostics
|
||||
status Show system status
|
||||
completion Generate completions
|
||||
import Import from other AI systems
|
||||
help Print this message or the help of the given subcommand(s)
|
||||
|
||||
Options:
|
||||
--cli-only
|
||||
Run in interactive CLI mode only (disable other channels)
|
||||
|
||||
--no-db
|
||||
Skip database connection (for testing)
|
||||
|
||||
-m, --message <MESSAGE>
|
||||
Single message mode - send one message and exit
|
||||
|
||||
-c, --config <CONFIG>
|
||||
Configuration file path (optional, uses env vars by default)
|
||||
|
||||
--no-onboard
|
||||
Skip first-run onboarding check
|
||||
|
||||
-h, --help
|
||||
Print help (see a summary with '-h')
|
||||
|
||||
-V, --version
|
||||
Print version
|
||||
@@ -16,6 +16,7 @@ Commands:
|
||||
tool Manage WASM tools
|
||||
registry Browse/install extensions
|
||||
channels Manage channels
|
||||
routines Manage routines
|
||||
mcp Manage MCP servers
|
||||
memory Manage workspace memory
|
||||
pairing Manage DM pairing
|
||||
|
||||
@@ -23,6 +23,9 @@ pub struct EmbeddingsConfig {
|
||||
pub ollama_base_url: String,
|
||||
/// Embedding vector dimension. Inferred from the model name when not set explicitly.
|
||||
pub dimension: usize,
|
||||
/// Custom base URL for OpenAI-compatible embedding providers.
|
||||
/// When set, overrides the default `https://api.openai.com`.
|
||||
pub openai_base_url: Option<String>,
|
||||
}
|
||||
|
||||
impl Default for EmbeddingsConfig {
|
||||
@@ -36,6 +39,7 @@ impl Default for EmbeddingsConfig {
|
||||
model,
|
||||
ollama_base_url: "http://localhost:11434".to_string(),
|
||||
dimension,
|
||||
openai_base_url: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -74,6 +78,8 @@ impl EmbeddingsConfig {
|
||||
|
||||
let enabled = parse_bool_env("EMBEDDING_ENABLED", settings.embeddings.enabled)?;
|
||||
|
||||
let openai_base_url = optional_env("EMBEDDING_BASE_URL")?;
|
||||
|
||||
Ok(Self {
|
||||
enabled,
|
||||
provider,
|
||||
@@ -81,6 +87,7 @@ impl EmbeddingsConfig {
|
||||
model,
|
||||
ollama_base_url,
|
||||
dimension,
|
||||
openai_base_url,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -130,16 +137,27 @@ impl EmbeddingsConfig {
|
||||
}
|
||||
_ => {
|
||||
if let Some(api_key) = self.openai_api_key() {
|
||||
tracing::debug!(
|
||||
"Embeddings enabled via OpenAI (model: {}, dim: {})",
|
||||
self.model,
|
||||
self.dimension,
|
||||
);
|
||||
Some(Arc::new(crate::workspace::OpenAiEmbeddings::with_model(
|
||||
let mut provider = crate::workspace::OpenAiEmbeddings::with_model(
|
||||
api_key,
|
||||
&self.model,
|
||||
self.dimension,
|
||||
)))
|
||||
);
|
||||
if let Some(ref base_url) = self.openai_base_url {
|
||||
tracing::debug!(
|
||||
"Embeddings enabled via OpenAI (model: {}, base_url: {}, dim: {})",
|
||||
self.model,
|
||||
base_url,
|
||||
self.dimension,
|
||||
);
|
||||
provider = provider.with_base_url(base_url);
|
||||
} else {
|
||||
tracing::debug!(
|
||||
"Embeddings enabled via OpenAI (model: {}, dim: {})",
|
||||
self.model,
|
||||
self.dimension,
|
||||
);
|
||||
}
|
||||
Some(Arc::new(provider))
|
||||
} else {
|
||||
tracing::warn!("Embeddings configured but OPENAI_API_KEY not set");
|
||||
None
|
||||
@@ -164,6 +182,7 @@ mod tests {
|
||||
std::env::remove_var("EMBEDDING_PROVIDER");
|
||||
std::env::remove_var("EMBEDDING_MODEL");
|
||||
std::env::remove_var("OPENAI_API_KEY");
|
||||
std::env::remove_var("EMBEDDING_BASE_URL");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -247,4 +266,41 @@ mod tests {
|
||||
std::env::remove_var("EMBEDDING_ENABLED");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_base_url_parsed_from_env() {
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
clear_embedding_env();
|
||||
|
||||
// SAFETY: Under ENV_MUTEX, no concurrent env access.
|
||||
unsafe {
|
||||
std::env::set_var("EMBEDDING_BASE_URL", "https://custom.example.com");
|
||||
}
|
||||
|
||||
let settings = Settings::default();
|
||||
let config = EmbeddingsConfig::resolve(&settings).expect("resolve should succeed");
|
||||
assert_eq!(
|
||||
config.openai_base_url.as_deref(),
|
||||
Some("https://custom.example.com"),
|
||||
"EMBEDDING_BASE_URL env var should be parsed into openai_base_url"
|
||||
);
|
||||
|
||||
// SAFETY: Under ENV_MUTEX.
|
||||
unsafe {
|
||||
std::env::remove_var("EMBEDDING_BASE_URL");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embedding_base_url_defaults_to_none() {
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
clear_embedding_env();
|
||||
|
||||
let settings = Settings::default();
|
||||
let config = EmbeddingsConfig::resolve(&settings).expect("resolve should succeed");
|
||||
assert!(
|
||||
config.openai_base_url.is_none(),
|
||||
"openai_base_url should be None when EMBEDDING_BASE_URL is not set"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -18,6 +18,7 @@ pub mod relay;
|
||||
mod routines;
|
||||
mod safety;
|
||||
mod sandbox;
|
||||
mod search;
|
||||
mod secrets;
|
||||
mod skills;
|
||||
mod transcription;
|
||||
@@ -44,6 +45,7 @@ pub use self::routines::RoutineConfig;
|
||||
pub use self::safety::SafetyConfig;
|
||||
use self::safety::resolve_safety_config;
|
||||
pub use self::sandbox::{ClaudeCodeConfig, SandboxModeConfig};
|
||||
pub use self::search::WorkspaceSearchConfig;
|
||||
pub use self::secrets::SecretsConfig;
|
||||
pub use self::skills::SkillsConfig;
|
||||
pub use self::transcription::TranscriptionConfig;
|
||||
@@ -91,6 +93,7 @@ pub struct Config {
|
||||
pub claude_code: ClaudeCodeConfig,
|
||||
pub skills: SkillsConfig,
|
||||
pub transcription: TranscriptionConfig,
|
||||
pub search: WorkspaceSearchConfig,
|
||||
pub observability: crate::observability::ObservabilityConfig,
|
||||
/// Channel-relay integration (Slack via external relay service).
|
||||
/// Present only when both `CHANNEL_RELAY_URL` and `CHANNEL_RELAY_API_KEY` are set.
|
||||
@@ -166,6 +169,7 @@ impl Config {
|
||||
..SkillsConfig::default()
|
||||
},
|
||||
transcription: TranscriptionConfig::default(),
|
||||
search: WorkspaceSearchConfig::default(),
|
||||
observability: crate::observability::ObservabilityConfig::default(),
|
||||
relay: None,
|
||||
}
|
||||
@@ -318,6 +322,7 @@ impl Config {
|
||||
claude_code: ClaudeCodeConfig::resolve()?,
|
||||
skills: SkillsConfig::resolve()?,
|
||||
transcription: TranscriptionConfig::resolve(settings)?,
|
||||
search: WorkspaceSearchConfig::resolve()?,
|
||||
observability: crate::observability::ObservabilityConfig {
|
||||
backend: std::env::var("OBSERVABILITY_BACKEND").unwrap_or_else(|_| "none".into()),
|
||||
},
|
||||
|
||||
@@ -0,0 +1,211 @@
|
||||
use crate::config::helpers::{optional_env, parse_optional_env};
|
||||
use crate::error::ConfigError;
|
||||
use crate::workspace::FusionStrategy;
|
||||
|
||||
/// Workspace search configuration resolved from environment variables.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct WorkspaceSearchConfig {
|
||||
/// Fusion strategy: "rrf" or "weighted".
|
||||
pub fusion_strategy: FusionStrategy,
|
||||
/// RRF constant k (default 60).
|
||||
pub rrf_k: u32,
|
||||
/// FTS weight for fusion.
|
||||
///
|
||||
/// [`Default`] uses 0.5. When the configuration is resolved, per-strategy
|
||||
/// defaults are applied: 0.5 (RRF) or 0.3 (weighted).
|
||||
pub fts_weight: f32,
|
||||
/// Vector weight for fusion.
|
||||
///
|
||||
/// [`Default`] uses 0.5. When the configuration is resolved, per-strategy
|
||||
/// defaults are applied: 0.5 (RRF) or 0.7 (weighted).
|
||||
pub vector_weight: f32,
|
||||
}
|
||||
|
||||
impl Default for WorkspaceSearchConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
fusion_strategy: FusionStrategy::default(),
|
||||
rrf_k: 60,
|
||||
fts_weight: 0.5,
|
||||
vector_weight: 0.5,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl WorkspaceSearchConfig {
|
||||
pub(crate) fn resolve() -> Result<Self, ConfigError> {
|
||||
let fusion_strategy = match optional_env("SEARCH_FUSION_STRATEGY")? {
|
||||
Some(s) => match s.to_lowercase().as_str() {
|
||||
"rrf" => FusionStrategy::Rrf,
|
||||
"weighted" => FusionStrategy::WeightedScore,
|
||||
other => {
|
||||
return Err(ConfigError::InvalidValue {
|
||||
key: "SEARCH_FUSION_STRATEGY".to_string(),
|
||||
message: format!("must be 'rrf' or 'weighted', got '{other}'"),
|
||||
});
|
||||
}
|
||||
},
|
||||
None => FusionStrategy::default(),
|
||||
};
|
||||
|
||||
let rrf_k = parse_optional_env("SEARCH_RRF_K", 60u32)?;
|
||||
|
||||
// Per-strategy weight defaults: RRF uses 0.5/0.5, weighted uses 0.3/0.7 (vector-biased).
|
||||
let (default_fts, default_vec) = match fusion_strategy {
|
||||
FusionStrategy::Rrf => (0.5f32, 0.5f32),
|
||||
FusionStrategy::WeightedScore => (0.3f32, 0.7f32),
|
||||
};
|
||||
let fts_weight = parse_optional_env("SEARCH_FTS_WEIGHT", default_fts)?;
|
||||
let vector_weight = parse_optional_env("SEARCH_VECTOR_WEIGHT", default_vec)?;
|
||||
|
||||
if !fts_weight.is_finite() || fts_weight < 0.0 {
|
||||
return Err(ConfigError::InvalidValue {
|
||||
key: "SEARCH_FTS_WEIGHT".to_string(),
|
||||
message: "must be a finite, non-negative float".to_string(),
|
||||
});
|
||||
}
|
||||
if !vector_weight.is_finite() || vector_weight < 0.0 {
|
||||
return Err(ConfigError::InvalidValue {
|
||||
key: "SEARCH_VECTOR_WEIGHT".to_string(),
|
||||
message: "must be a finite, non-negative float".to_string(),
|
||||
});
|
||||
}
|
||||
if matches!(fusion_strategy, FusionStrategy::WeightedScore)
|
||||
&& fts_weight == 0.0
|
||||
&& vector_weight == 0.0
|
||||
{
|
||||
return Err(ConfigError::InvalidValue {
|
||||
key: "SEARCH_FTS_WEIGHT/SEARCH_VECTOR_WEIGHT".to_string(),
|
||||
message: "weighted fusion requires at least one non-zero weight".to_string(),
|
||||
});
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
fusion_strategy,
|
||||
rrf_k,
|
||||
fts_weight,
|
||||
vector_weight,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::config::helpers::ENV_MUTEX;
|
||||
|
||||
fn clear_search_env() {
|
||||
// SAFETY: Only called under ENV_MUTEX in tests.
|
||||
unsafe {
|
||||
std::env::remove_var("SEARCH_FUSION_STRATEGY");
|
||||
std::env::remove_var("SEARCH_RRF_K");
|
||||
std::env::remove_var("SEARCH_FTS_WEIGHT");
|
||||
std::env::remove_var("SEARCH_VECTOR_WEIGHT");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn defaults_when_no_env() {
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
clear_search_env();
|
||||
|
||||
let config = WorkspaceSearchConfig::resolve().expect("should resolve");
|
||||
assert_eq!(config.fusion_strategy, FusionStrategy::Rrf);
|
||||
assert_eq!(config.rrf_k, 60);
|
||||
assert!((config.fts_weight - 0.5).abs() < 0.001);
|
||||
assert!((config.vector_weight - 0.5).abs() < 0.001);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn env_overrides() {
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
clear_search_env();
|
||||
|
||||
// SAFETY: Under ENV_MUTEX.
|
||||
unsafe {
|
||||
std::env::set_var("SEARCH_FUSION_STRATEGY", "weighted");
|
||||
std::env::set_var("SEARCH_RRF_K", "30");
|
||||
std::env::set_var("SEARCH_FTS_WEIGHT", "0.9");
|
||||
std::env::set_var("SEARCH_VECTOR_WEIGHT", "0.1");
|
||||
}
|
||||
|
||||
let config = WorkspaceSearchConfig::resolve().expect("should resolve");
|
||||
assert_eq!(config.fusion_strategy, FusionStrategy::WeightedScore);
|
||||
assert_eq!(config.rrf_k, 30);
|
||||
assert!((config.fts_weight - 0.9).abs() < 0.001);
|
||||
assert!((config.vector_weight - 0.1).abs() < 0.001);
|
||||
|
||||
clear_search_env();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_strategy_rejected() {
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
clear_search_env();
|
||||
|
||||
// SAFETY: Under ENV_MUTEX.
|
||||
unsafe {
|
||||
std::env::set_var("SEARCH_FUSION_STRATEGY", "bm25");
|
||||
}
|
||||
|
||||
let result = WorkspaceSearchConfig::resolve();
|
||||
assert!(result.is_err());
|
||||
|
||||
clear_search_env();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn weighted_strategy_defaults() {
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
clear_search_env();
|
||||
|
||||
// SAFETY: Under ENV_MUTEX.
|
||||
unsafe {
|
||||
std::env::set_var("SEARCH_FUSION_STRATEGY", "weighted");
|
||||
}
|
||||
|
||||
let config = WorkspaceSearchConfig::resolve().expect("should resolve");
|
||||
assert_eq!(config.fusion_strategy, FusionStrategy::WeightedScore);
|
||||
// Weighted mode should default to 0.3 FTS / 0.7 vector
|
||||
assert!((config.fts_weight - 0.3).abs() < 0.001);
|
||||
assert!((config.vector_weight - 0.7).abs() < 0.001);
|
||||
|
||||
clear_search_env();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn weighted_both_zero_rejected() {
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
clear_search_env();
|
||||
|
||||
// SAFETY: Under ENV_MUTEX.
|
||||
unsafe {
|
||||
std::env::set_var("SEARCH_FUSION_STRATEGY", "weighted");
|
||||
std::env::set_var("SEARCH_FTS_WEIGHT", "0.0");
|
||||
std::env::set_var("SEARCH_VECTOR_WEIGHT", "0.0");
|
||||
}
|
||||
|
||||
let result = WorkspaceSearchConfig::resolve();
|
||||
assert!(result.is_err());
|
||||
|
||||
clear_search_env();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rrf_both_zero_allowed() {
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
clear_search_env();
|
||||
|
||||
// SAFETY: Under ENV_MUTEX.
|
||||
unsafe {
|
||||
std::env::set_var("SEARCH_FTS_WEIGHT", "0.0");
|
||||
std::env::set_var("SEARCH_VECTOR_WEIGHT", "0.0");
|
||||
}
|
||||
|
||||
// RRF ignores weights, so both=0 is fine
|
||||
let config = WorkspaceSearchConfig::resolve().expect("should resolve");
|
||||
assert_eq!(config.fusion_strategy, FusionStrategy::Rrf);
|
||||
|
||||
clear_search_env();
|
||||
}
|
||||
}
|
||||
+14
-8
@@ -16,6 +16,7 @@ mod workspace;
|
||||
|
||||
use std::path::Path;
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use chrono::{DateTime, NaiveDateTime, Utc};
|
||||
@@ -32,6 +33,8 @@ use crate::workspace::MemoryDocument;
|
||||
|
||||
use crate::db::libsql_migrations;
|
||||
|
||||
static NAIVE_TIMESTAMP_LOGGED: AtomicBool = AtomicBool::new(false);
|
||||
|
||||
/// Explicit column list for routines table (matches positional access in `row_to_routine_libsql`).
|
||||
pub(crate) const ROUTINE_COLUMNS: &str = "\
|
||||
id, name, description, user_id, enabled, \
|
||||
@@ -163,24 +166,27 @@ impl LibSqlBackend {
|
||||
///
|
||||
/// Returns an error if none of the formats match.
|
||||
pub(crate) fn parse_timestamp(s: &str) -> Result<DateTime<Utc>, String> {
|
||||
let log_naive_timestamp_once = || {
|
||||
if !NAIVE_TIMESTAMP_LOGGED.swap(true, Ordering::Relaxed) {
|
||||
tracing::debug!(
|
||||
timestamp = %s,
|
||||
"parsed naive timestamp without timezone; assuming UTC for backward compatibility"
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
// RFC 3339 (our canonical write format)
|
||||
if let Ok(dt) = DateTime::parse_from_rfc3339(s) {
|
||||
return Ok(dt.with_timezone(&Utc));
|
||||
}
|
||||
// Naive with fractional seconds (legacy or SQLite datetime() output)
|
||||
if let Ok(ndt) = NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S%.f") {
|
||||
tracing::warn!(
|
||||
timestamp = %s,
|
||||
"parsed naive timestamp without timezone; assuming UTC for backward compatibility"
|
||||
);
|
||||
log_naive_timestamp_once();
|
||||
return Ok(ndt.and_utc());
|
||||
}
|
||||
// Naive without fractional seconds (legacy format)
|
||||
if let Ok(ndt) = NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S") {
|
||||
tracing::warn!(
|
||||
timestamp = %s,
|
||||
"parsed naive timestamp without timezone; assuming UTC for backward compatibility"
|
||||
);
|
||||
log_naive_timestamp_once();
|
||||
return Ok(ndt.and_utc());
|
||||
}
|
||||
Err(format!("unparseable timestamp: {:?}", s))
|
||||
|
||||
@@ -14,7 +14,7 @@ use crate::db::WorkspaceStore;
|
||||
use crate::error::WorkspaceError;
|
||||
use crate::workspace::{
|
||||
MemoryChunk, MemoryDocument, RankedResult, SearchConfig, SearchResult, WorkspaceEntry,
|
||||
reciprocal_rank_fusion,
|
||||
fuse_results,
|
||||
};
|
||||
|
||||
use chrono::Utc;
|
||||
@@ -614,6 +614,6 @@ impl WorkspaceStore for LibSqlBackend {
|
||||
);
|
||||
}
|
||||
|
||||
Ok(reciprocal_rank_fusion(fts_results, vector_results, config))
|
||||
Ok(fuse_results(fts_results, vector_results, config))
|
||||
}
|
||||
}
|
||||
|
||||
+637
-66
@@ -27,7 +27,7 @@ use crate::secrets::{CreateSecretParams, SecretsStore};
|
||||
use crate::tools::ToolRegistry;
|
||||
use crate::tools::mcp::McpClient;
|
||||
use crate::tools::mcp::auth::{
|
||||
PkceChallenge, authorize_mcp_server, build_authorization_url, discover_full_oauth_metadata,
|
||||
authorize_mcp_server, canonical_resource_uri, discover_full_oauth_metadata,
|
||||
find_available_port, is_authenticated, register_client,
|
||||
};
|
||||
use crate::tools::mcp::config::McpServerConfig;
|
||||
@@ -108,6 +108,13 @@ pub struct ExtensionManager {
|
||||
/// Relay config captured at startup. Used by `auth_channel_relay` and
|
||||
/// `activate_channel_relay` instead of re-reading env vars.
|
||||
relay_config: Option<crate::config::RelayConfig>,
|
||||
/// When `true`, OAuth flows always return an auth URL to the caller
|
||||
/// instead of opening a browser on the server via `open::that()`.
|
||||
/// Set by the web gateway at startup via `enable_gateway_mode()`.
|
||||
gateway_mode: std::sync::atomic::AtomicBool,
|
||||
/// The gateway's own base URL for building OAuth redirect URIs.
|
||||
/// Set by the web gateway at startup via `enable_gateway_mode()`.
|
||||
gateway_base_url: RwLock<Option<String>>,
|
||||
}
|
||||
|
||||
/// Sanitize a URL for logging by removing query parameters and credentials.
|
||||
@@ -181,9 +188,75 @@ impl ExtensionManager {
|
||||
pending_oauth_flows: crate::cli::oauth_defaults::new_pending_oauth_registry(),
|
||||
gateway_token: std::env::var("GATEWAY_AUTH_TOKEN").ok(),
|
||||
relay_config: crate::config::RelayConfig::from_env(),
|
||||
gateway_mode: std::sync::atomic::AtomicBool::new(false),
|
||||
gateway_base_url: RwLock::new(None),
|
||||
}
|
||||
}
|
||||
|
||||
/// Enable gateway mode so OAuth flows return auth URLs to the frontend
|
||||
/// instead of calling `open::that()` on the server.
|
||||
///
|
||||
/// `base_url` is the gateway's own public URL (e.g. `https://my-gateway.example.com`),
|
||||
/// used to build OAuth redirect URIs when `IRONCLAW_OAUTH_CALLBACK_URL` is not set.
|
||||
pub async fn enable_gateway_mode(&self, base_url: String) {
|
||||
self.gateway_mode
|
||||
.store(true, std::sync::atomic::Ordering::Release);
|
||||
*self.gateway_base_url.write().await = Some(base_url);
|
||||
}
|
||||
|
||||
/// Returns `true` if OAuth should use gateway mode (return auth URL to
|
||||
/// frontend) rather than CLI mode (open browser on server via `open::that`).
|
||||
///
|
||||
/// Gateway mode is active when any of:
|
||||
/// - `enable_gateway_mode()` was called (web gateway is running), OR
|
||||
/// - `IRONCLAW_OAUTH_CALLBACK_URL` is set to a non-loopback URL, OR
|
||||
/// - `self.tunnel_url` is set to a non-loopback URL
|
||||
pub fn should_use_gateway_mode(&self) -> bool {
|
||||
if self.gateway_mode.load(std::sync::atomic::Ordering::Acquire) {
|
||||
return true;
|
||||
}
|
||||
if crate::cli::oauth_defaults::use_gateway_callback() {
|
||||
return true;
|
||||
}
|
||||
self.tunnel_url
|
||||
.as_ref()
|
||||
.filter(|u| !u.is_empty())
|
||||
.and_then(|raw| url::Url::parse(raw).ok())
|
||||
.and_then(|u| u.host_str().map(String::from))
|
||||
.map(|host| !crate::cli::oauth_defaults::is_loopback_host(&host))
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
/// Returns the OAuth redirect URI for gateway mode, or `None` for local mode.
|
||||
///
|
||||
/// Priority:
|
||||
/// 1. `IRONCLAW_OAUTH_CALLBACK_URL` env var (via `callback_url()`)
|
||||
/// 2. `gateway_base_url` (set by `enable_gateway_mode()`)
|
||||
/// 3. `tunnel_url` (from config)
|
||||
/// 4. `None` (local/CLI mode)
|
||||
async fn gateway_callback_redirect_uri(&self) -> Option<String> {
|
||||
use crate::cli::oauth_defaults;
|
||||
if oauth_defaults::use_gateway_callback() {
|
||||
return Some(format!("{}/oauth/callback", oauth_defaults::callback_url()));
|
||||
}
|
||||
// Use gateway_base_url from enable_gateway_mode()
|
||||
if let Some(ref base) = *self.gateway_base_url.read().await {
|
||||
let base = base.trim_end_matches('/');
|
||||
return Some(format!("{}/oauth/callback", base));
|
||||
}
|
||||
// Fall back to tunnel_url
|
||||
self.tunnel_url
|
||||
.as_ref()
|
||||
.filter(|u| !u.is_empty())
|
||||
.and_then(|raw| url::Url::parse(raw).ok())
|
||||
.and_then(|u| u.host_str().map(String::from))
|
||||
.filter(|host| !oauth_defaults::is_loopback_host(host))
|
||||
.map(|_| {
|
||||
let base = self.tunnel_url.as_ref().unwrap().trim_end_matches('/');
|
||||
format!("{}/oauth/callback", base)
|
||||
})
|
||||
}
|
||||
|
||||
/// Get the relay config stored at startup.
|
||||
fn relay_config(&self) -> Result<&crate::config::RelayConfig, ExtensionError> {
|
||||
self.relay_config.as_ref().ok_or_else(|| {
|
||||
@@ -193,6 +266,12 @@ impl ExtensionManager {
|
||||
})
|
||||
}
|
||||
|
||||
/// Inject a registry entry for testing. The entry is added to the discovery
|
||||
/// cache so it appears in search results alongside built-in entries.
|
||||
pub async fn inject_registry_entry(&self, entry: crate::extensions::RegistryEntry) {
|
||||
self.registry.cache_discovered(vec![entry]).await;
|
||||
}
|
||||
|
||||
/// Configure the channel runtime infrastructure for hot-activating WASM channels.
|
||||
///
|
||||
/// Call after construction (and after wrapping in `Arc`) once the channel
|
||||
@@ -707,6 +786,19 @@ impl ExtensionManager {
|
||||
Self::validate_extension_name(name)?;
|
||||
let kind = self.determine_installed_kind(name).await?;
|
||||
|
||||
// Clean up any in-progress OAuth flows for this extension.
|
||||
// TCP mode: abort the listener task so port 9876 is freed immediately.
|
||||
// Gateway mode: remove stale pending flow entries.
|
||||
if let Some(pending) = self.pending_auth.write().await.remove(name)
|
||||
&& let Some(handle) = pending.task_handle
|
||||
{
|
||||
handle.abort();
|
||||
}
|
||||
self.pending_oauth_flows
|
||||
.write()
|
||||
.await
|
||||
.retain(|_, flow| flow.extension_name != name);
|
||||
|
||||
match kind {
|
||||
ExtensionKind::McpServer => {
|
||||
// Unregister tools with this server's prefix
|
||||
@@ -740,6 +832,14 @@ impl ExtensionManager {
|
||||
// Unregister from tool registry
|
||||
self.tool_registry.unregister(name).await;
|
||||
|
||||
// Evict compiled module from runtime cache so reinstall uses fresh binary
|
||||
if let Some(ref rt) = self.wasm_tool_runtime {
|
||||
rt.remove(name).await;
|
||||
}
|
||||
|
||||
// Clear stale activation errors so reinstall starts clean
|
||||
self.activation_errors.write().await.remove(name);
|
||||
|
||||
// Revoke credential mappings from the shared registry
|
||||
let cap_path = self
|
||||
.wasm_tools_dir
|
||||
@@ -780,6 +880,9 @@ impl ExtensionManager {
|
||||
self.active_channel_names.write().await.remove(name);
|
||||
self.persist_active_channels().await;
|
||||
|
||||
// Clear stale activation errors so reinstall starts clean
|
||||
self.activation_errors.write().await.remove(name);
|
||||
|
||||
// Delete channel files
|
||||
let wasm_path = self.wasm_channels_dir.join(format!("{}.wasm", name));
|
||||
let cap_path = self
|
||||
@@ -1684,29 +1787,46 @@ impl ExtensionManager {
|
||||
return Ok(AuthResult::authenticated(name, ExtensionKind::McpServer));
|
||||
}
|
||||
|
||||
// Run the full OAuth flow (opens browser, waits for callback)
|
||||
// In gateway mode, build an auth URL and return it for the frontend to
|
||||
// open in the same browser. The gateway's /oauth/callback handler will
|
||||
// complete the token exchange.
|
||||
if self.should_use_gateway_mode() {
|
||||
return match self.auth_mcp_build_url(name, &server).await {
|
||||
Ok(result) => Ok(result),
|
||||
Err(ExtensionError::AuthNotSupported(_)) => Ok(AuthResult::awaiting_token(
|
||||
name,
|
||||
ExtensionKind::McpServer,
|
||||
format!(
|
||||
"Server '{}' does not support OAuth. \
|
||||
Please provide an API token/key for this server.",
|
||||
name
|
||||
),
|
||||
None,
|
||||
)),
|
||||
Err(e) => Err(e),
|
||||
};
|
||||
}
|
||||
|
||||
// CLI/local mode: run the full blocking OAuth flow (opens browser, waits for callback)
|
||||
match authorize_mcp_server(&server, &self.secrets, &self.user_id).await {
|
||||
Ok(_token) => {
|
||||
tracing::info!("MCP server '{}' authenticated via OAuth", name);
|
||||
Ok(AuthResult::authenticated(name, ExtensionKind::McpServer))
|
||||
}
|
||||
Err(crate::tools::mcp::auth::AuthError::NotSupported) => {
|
||||
// Server doesn't support OAuth, try building a URL first
|
||||
// Server doesn't support OAuth, try building a URL
|
||||
match self.auth_mcp_build_url(name, &server).await {
|
||||
Ok(result) => Ok(result),
|
||||
Err(_) => {
|
||||
// No OAuth, no DCR: fall back to manual token entry
|
||||
Ok(AuthResult::awaiting_token(
|
||||
name,
|
||||
ExtensionKind::McpServer,
|
||||
format!(
|
||||
"Server '{}' does not support OAuth. \
|
||||
Please provide an API token/key for this server.",
|
||||
name
|
||||
),
|
||||
None,
|
||||
))
|
||||
}
|
||||
Err(_) => Ok(AuthResult::awaiting_token(
|
||||
name,
|
||||
ExtensionKind::McpServer,
|
||||
format!(
|
||||
"Server '{}' does not support OAuth. \
|
||||
Please provide an API token/key for this server.",
|
||||
name
|
||||
),
|
||||
None,
|
||||
)),
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
@@ -1725,8 +1845,12 @@ impl ExtensionManager {
|
||||
}
|
||||
}
|
||||
|
||||
/// Build an auth URL for cases where non-interactive auth is needed
|
||||
/// (e.g., running via Telegram where we can't open a browser).
|
||||
/// Build an auth URL for MCP OAuth.
|
||||
///
|
||||
/// In gateway mode, stores a `PendingOAuthFlow` so the web gateway's
|
||||
/// `/oauth/callback` handler can complete the token exchange — the auth
|
||||
/// URL is sent to the frontend which opens it in the same browser.
|
||||
/// In local/CLI mode, builds the URL for the user to open manually.
|
||||
async fn auth_mcp_build_url(
|
||||
&self,
|
||||
name: &str,
|
||||
@@ -1735,60 +1859,153 @@ impl ExtensionManager {
|
||||
// Try to discover OAuth metadata and build a URL the user can open manually
|
||||
let metadata = discover_full_oauth_metadata(&server.url)
|
||||
.await
|
||||
.map_err(|e| ExtensionError::AuthFailed(e.to_string()))?;
|
||||
.map_err(|e| match e {
|
||||
crate::tools::mcp::auth::AuthError::NotSupported => {
|
||||
ExtensionError::AuthNotSupported(e.to_string())
|
||||
}
|
||||
_ => ExtensionError::AuthFailed(e.to_string()),
|
||||
})?;
|
||||
|
||||
use crate::cli::oauth_defaults;
|
||||
|
||||
let is_gateway = self.should_use_gateway_mode();
|
||||
|
||||
// Build redirect URI: gateway uses the public callback URL,
|
||||
// local mode binds a random port.
|
||||
let redirect_uri = if let Some(uri) = self.gateway_callback_redirect_uri().await {
|
||||
uri
|
||||
} else {
|
||||
let port = find_available_port()
|
||||
.await
|
||||
.map_err(|e| ExtensionError::AuthFailed(e.to_string()))?;
|
||||
format!("http://localhost:{}/callback", port.1)
|
||||
};
|
||||
|
||||
// Try DCR if no client_id configured
|
||||
let (client_id, redirect_uri) = if let Some(ref oauth) = server.oauth {
|
||||
let port = find_available_port()
|
||||
.await
|
||||
.map_err(|e| ExtensionError::AuthFailed(e.to_string()))?;
|
||||
let redirect = format!("http://localhost:{}/callback", port.1);
|
||||
(oauth.client_id.clone(), redirect)
|
||||
let (client_id, client_secret) = if let Some(ref oauth) = server.oauth {
|
||||
(oauth.client_id.clone(), None)
|
||||
} else if let Some(ref reg_endpoint) = metadata.registration_endpoint {
|
||||
let port = find_available_port()
|
||||
.await
|
||||
.map_err(|e| ExtensionError::AuthFailed(e.to_string()))?;
|
||||
let redirect = format!("http://localhost:{}/callback", port.1);
|
||||
|
||||
let registration = register_client(reg_endpoint, &redirect)
|
||||
let registration = register_client(reg_endpoint, &redirect_uri)
|
||||
.await
|
||||
.map_err(|e| ExtensionError::AuthFailed(e.to_string()))?;
|
||||
|
||||
(registration.client_id, redirect)
|
||||
(registration.client_id, None)
|
||||
} else {
|
||||
return Err(ExtensionError::AuthFailed(
|
||||
return Err(ExtensionError::AuthNotSupported(
|
||||
"Server doesn't support OAuth or Dynamic Client Registration".to_string(),
|
||||
));
|
||||
};
|
||||
|
||||
let pkce = PkceChallenge::generate();
|
||||
let auth_url = build_authorization_url(
|
||||
// RFC 8707: resource parameter to scope the token to this MCP server
|
||||
let resource = canonical_resource_uri(&server.url);
|
||||
|
||||
// Build authorization URL with CSRF state using the shared oauth_defaults
|
||||
// builder, which generates PKCE + state for us.
|
||||
let mut extra_params = server
|
||||
.oauth
|
||||
.as_ref()
|
||||
.map(|o| o.extra_params.clone())
|
||||
.unwrap_or_default();
|
||||
extra_params.insert("resource".to_string(), resource.clone());
|
||||
|
||||
let scopes = server
|
||||
.oauth
|
||||
.as_ref()
|
||||
.map(|o| o.scopes.clone())
|
||||
.unwrap_or_else(|| metadata.scopes_supported.clone());
|
||||
|
||||
let oauth_result = oauth_defaults::build_oauth_url(
|
||||
&metadata.authorization_endpoint,
|
||||
&client_id,
|
||||
&redirect_uri,
|
||||
&metadata.scopes_supported,
|
||||
Some(&pkce),
|
||||
&std::collections::HashMap::new(),
|
||||
None,
|
||||
&scopes,
|
||||
true, // Always use PKCE for MCP
|
||||
&extra_params,
|
||||
);
|
||||
let expected_state = oauth_result.state;
|
||||
let code_verifier = oauth_result.code_verifier;
|
||||
|
||||
// Store pending auth for later callback handling
|
||||
self.pending_auth.write().await.insert(
|
||||
name.to_string(),
|
||||
PendingAuth {
|
||||
_name: name.to_string(),
|
||||
_kind: ExtensionKind::McpServer,
|
||||
if is_gateway {
|
||||
// Gateway mode: store pending flow for the /oauth/callback handler.
|
||||
oauth_defaults::sweep_expired_flows(&self.pending_oauth_flows).await;
|
||||
|
||||
// Platform routing: prepend instance name to state
|
||||
let platform_state = oauth_defaults::build_platform_state(&expected_state);
|
||||
let auth_url = if platform_state != expected_state {
|
||||
oauth_result.url.replace(
|
||||
&format!("state={}", urlencoding::encode(&expected_state)),
|
||||
&format!("state={}", urlencoding::encode(&platform_state)),
|
||||
)
|
||||
} else {
|
||||
oauth_result.url
|
||||
};
|
||||
|
||||
let flow = oauth_defaults::PendingOAuthFlow {
|
||||
extension_name: name.to_string(),
|
||||
display_name: server.name.clone(),
|
||||
token_url: metadata.token_endpoint,
|
||||
client_id,
|
||||
client_secret,
|
||||
redirect_uri,
|
||||
code_verifier,
|
||||
access_token_field: "access_token".to_string(),
|
||||
secret_name: server.token_secret_name(),
|
||||
provider: Some(format!("mcp:{}", name)),
|
||||
validation_endpoint: None,
|
||||
scopes,
|
||||
user_id: self.user_id.clone(),
|
||||
secrets: Arc::clone(&self.secrets),
|
||||
sse_sender: self.sse_sender.read().await.clone(),
|
||||
gateway_token: self.gateway_token.clone(),
|
||||
resource: Some(resource),
|
||||
client_id_secret_name: if server.oauth.is_none() {
|
||||
Some(server.client_id_secret_name())
|
||||
} else {
|
||||
None
|
||||
},
|
||||
created_at: std::time::Instant::now(),
|
||||
task_handle: None,
|
||||
},
|
||||
);
|
||||
};
|
||||
|
||||
Ok(AuthResult::awaiting_authorization(
|
||||
name,
|
||||
ExtensionKind::McpServer,
|
||||
auth_url,
|
||||
"local".to_string(),
|
||||
))
|
||||
self.pending_oauth_flows
|
||||
.write()
|
||||
.await
|
||||
.insert(expected_state, flow);
|
||||
|
||||
self.pending_auth.write().await.insert(
|
||||
name.to_string(),
|
||||
PendingAuth {
|
||||
_name: name.to_string(),
|
||||
_kind: ExtensionKind::McpServer,
|
||||
created_at: std::time::Instant::now(),
|
||||
task_handle: None,
|
||||
},
|
||||
);
|
||||
|
||||
Ok(AuthResult::awaiting_authorization(
|
||||
name,
|
||||
ExtensionKind::McpServer,
|
||||
auth_url,
|
||||
"gateway".to_string(),
|
||||
))
|
||||
} else {
|
||||
// Local mode: return URL for manual opening
|
||||
self.pending_auth.write().await.insert(
|
||||
name.to_string(),
|
||||
PendingAuth {
|
||||
_name: name.to_string(),
|
||||
_kind: ExtensionKind::McpServer,
|
||||
created_at: std::time::Instant::now(),
|
||||
task_handle: None,
|
||||
},
|
||||
);
|
||||
|
||||
Ok(AuthResult::awaiting_authorization(
|
||||
name,
|
||||
ExtensionKind::McpServer,
|
||||
oauth_result.url,
|
||||
"local".to_string(),
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
async fn auth_wasm_tool(&self, name: &str) -> Result<AuthResult, ExtensionError> {
|
||||
@@ -2203,7 +2420,10 @@ impl ExtensionManager {
|
||||
flows.retain(|_, flow| flow.extension_name != name);
|
||||
}
|
||||
|
||||
let redirect_uri = format!("{}/callback", oauth_defaults::callback_url());
|
||||
let redirect_uri = self
|
||||
.gateway_callback_redirect_uri()
|
||||
.await
|
||||
.unwrap_or_else(|| format!("{}/callback", oauth_defaults::callback_url()));
|
||||
|
||||
// Merge scopes from all tools sharing this provider
|
||||
let merged_scopes = self
|
||||
@@ -2228,7 +2448,7 @@ impl ExtensionManager {
|
||||
.clone()
|
||||
.unwrap_or_else(|| name.to_string());
|
||||
|
||||
if oauth_defaults::use_gateway_callback() {
|
||||
if self.should_use_gateway_mode() {
|
||||
// Gateway mode: store pending flow state for the web gateway's
|
||||
// `/oauth/callback` handler to complete the exchange. No TCP listener
|
||||
// needed — the OAuth provider redirects to the gateway URL.
|
||||
@@ -2264,6 +2484,8 @@ impl ExtensionManager {
|
||||
secrets: Arc::clone(&self.secrets),
|
||||
sse_sender: self.sse_sender.read().await.clone(),
|
||||
gateway_token: self.gateway_token.clone(),
|
||||
resource: None,
|
||||
client_id_secret_name: None,
|
||||
created_at: std::time::Instant::now(),
|
||||
};
|
||||
|
||||
@@ -2605,11 +2827,17 @@ impl ExtensionManager {
|
||||
.await
|
||||
.map_err(|e| ExtensionError::ActivationFailed(e.to_string()))?;
|
||||
|
||||
// Try to list and create tools
|
||||
let mcp_tools = client
|
||||
.list_tools()
|
||||
.await
|
||||
.map_err(|e| ExtensionError::ActivationFailed(e.to_string()))?;
|
||||
// Try to list and create tools.
|
||||
// A 401/auth error means the server requires OAuth — surface as
|
||||
// AuthRequired so the activate handler triggers the OAuth flow.
|
||||
let mcp_tools = client.list_tools().await.map_err(|e| {
|
||||
let msg = e.to_string();
|
||||
if msg.contains("requires authentication") || msg.contains("401") {
|
||||
ExtensionError::AuthRequired
|
||||
} else {
|
||||
ExtensionError::ActivationFailed(msg)
|
||||
}
|
||||
})?;
|
||||
|
||||
let tool_impls = client
|
||||
.create_tools()
|
||||
@@ -2656,6 +2884,17 @@ impl ExtensionManager {
|
||||
});
|
||||
}
|
||||
|
||||
// Check auth status — block activation if required secrets are missing.
|
||||
// NeedsAuth (OAuth not yet completed) is allowed because configure() loads
|
||||
// the tool first, then starts the OAuth flow to obtain the token.
|
||||
let auth_state = self.check_tool_auth_status(name).await;
|
||||
if auth_state == ToolAuthState::NeedsSetup {
|
||||
return Err(ExtensionError::ActivationFailed(format!(
|
||||
"Tool '{}' requires configuration. Use the setup form to provide credentials.",
|
||||
name
|
||||
)));
|
||||
}
|
||||
|
||||
let runtime = self.wasm_tool_runtime.as_ref().ok_or_else(|| {
|
||||
ExtensionError::ActivationFailed("WASM runtime not available".to_string())
|
||||
})?;
|
||||
@@ -4291,14 +4530,18 @@ mod tests {
|
||||
// available" because the ExtensionManager had `wasm_tool_runtime: None`.
|
||||
|
||||
/// Build a minimal ExtensionManager suitable for unit tests.
|
||||
fn make_test_manager(
|
||||
fn make_test_manager_with_dirs(
|
||||
wasm_runtime: Option<Arc<crate::tools::wasm::WasmToolRuntime>>,
|
||||
tools_dir: std::path::PathBuf,
|
||||
channels_dir: std::path::PathBuf,
|
||||
) -> crate::extensions::manager::ExtensionManager {
|
||||
use crate::secrets::{InMemorySecretsStore, SecretsCrypto};
|
||||
use crate::tools::mcp::process::McpProcessManager;
|
||||
use crate::tools::mcp::session::McpSessionManager;
|
||||
|
||||
std::fs::create_dir_all(&tools_dir).ok();
|
||||
std::fs::create_dir_all(&channels_dir).ok();
|
||||
|
||||
let key = secrecy::SecretString::from(crate::secrets::keychain::generate_master_key_hex());
|
||||
let crypto = Arc::new(SecretsCrypto::new(key).expect("crypto"));
|
||||
let secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync> =
|
||||
@@ -4313,15 +4556,22 @@ mod tests {
|
||||
tools,
|
||||
None, // hooks
|
||||
wasm_runtime,
|
||||
tools_dir.clone(),
|
||||
tools_dir, // channels dir (unused here)
|
||||
None, // tunnel_url
|
||||
tools_dir,
|
||||
channels_dir,
|
||||
None, // tunnel_url
|
||||
"test".to_string(),
|
||||
None, // db
|
||||
vec![],
|
||||
)
|
||||
}
|
||||
|
||||
fn make_test_manager(
|
||||
wasm_runtime: Option<Arc<crate::tools::wasm::WasmToolRuntime>>,
|
||||
tools_dir: std::path::PathBuf,
|
||||
) -> crate::extensions::manager::ExtensionManager {
|
||||
make_test_manager_with_dirs(wasm_runtime, tools_dir.clone(), tools_dir)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_activate_wasm_tool_with_runtime_passes_runtime_check() {
|
||||
// When the ExtensionManager has a WASM runtime, activation should get
|
||||
@@ -4674,6 +4924,145 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_remove_wasm_tool_clears_pending_oauth_state_and_activation_error() {
|
||||
let dir = tempfile::tempdir().expect("temp dir");
|
||||
let mgr = make_test_manager(None, dir.path().to_path_buf());
|
||||
|
||||
std::fs::write(dir.path().join("gmail.wasm"), b"fake-tool").expect("write tool");
|
||||
|
||||
let listener = tokio::spawn(async {
|
||||
std::future::pending::<()>().await;
|
||||
});
|
||||
let abort_handle = listener.abort_handle();
|
||||
mgr.pending_auth.write().await.insert(
|
||||
"gmail".to_string(),
|
||||
super::PendingAuth {
|
||||
_name: "gmail".to_string(),
|
||||
_kind: ExtensionKind::WasmTool,
|
||||
created_at: std::time::Instant::now(),
|
||||
task_handle: Some(listener),
|
||||
},
|
||||
);
|
||||
|
||||
mgr.activation_errors
|
||||
.write()
|
||||
.await
|
||||
.insert("gmail".to_string(), "cached failure".to_string());
|
||||
|
||||
let secrets = Arc::clone(&mgr.secrets);
|
||||
mgr.pending_oauth_flows().write().await.insert(
|
||||
"gmail-state".to_string(),
|
||||
crate::cli::oauth_defaults::PendingOAuthFlow {
|
||||
extension_name: "gmail".to_string(),
|
||||
display_name: "Gmail".to_string(),
|
||||
token_url: "https://example.com/token".to_string(),
|
||||
client_id: "client123".to_string(),
|
||||
client_secret: None,
|
||||
redirect_uri: "https://example.com/oauth/callback".to_string(),
|
||||
code_verifier: None,
|
||||
access_token_field: "access_token".to_string(),
|
||||
secret_name: "google_oauth_token".to_string(),
|
||||
provider: None,
|
||||
validation_endpoint: None,
|
||||
scopes: vec![],
|
||||
user_id: "test".to_string(),
|
||||
secrets: Arc::clone(&secrets),
|
||||
sse_sender: None,
|
||||
gateway_token: None,
|
||||
resource: None,
|
||||
client_id_secret_name: None,
|
||||
created_at: std::time::Instant::now(),
|
||||
},
|
||||
);
|
||||
mgr.pending_oauth_flows().write().await.insert(
|
||||
"other-state".to_string(),
|
||||
crate::cli::oauth_defaults::PendingOAuthFlow {
|
||||
extension_name: "web-search".to_string(),
|
||||
display_name: "Web Search".to_string(),
|
||||
token_url: "https://example.com/token".to_string(),
|
||||
client_id: "client456".to_string(),
|
||||
client_secret: None,
|
||||
redirect_uri: "https://example.com/oauth/callback".to_string(),
|
||||
code_verifier: None,
|
||||
access_token_field: "access_token".to_string(),
|
||||
secret_name: "other_token".to_string(),
|
||||
provider: None,
|
||||
validation_endpoint: None,
|
||||
scopes: vec![],
|
||||
user_id: "test".to_string(),
|
||||
secrets,
|
||||
sse_sender: None,
|
||||
gateway_token: None,
|
||||
resource: None,
|
||||
client_id_secret_name: None,
|
||||
created_at: std::time::Instant::now(),
|
||||
},
|
||||
);
|
||||
|
||||
let result = mgr.remove("gmail").await;
|
||||
assert!(result.is_ok(), "remove should succeed: {:?}", result.err());
|
||||
|
||||
tokio::task::yield_now().await;
|
||||
|
||||
assert!(
|
||||
mgr.pending_auth.read().await.get("gmail").is_none(),
|
||||
"pending auth entry should be removed"
|
||||
);
|
||||
assert!(
|
||||
abort_handle.is_finished(),
|
||||
"pending auth listener should be aborted"
|
||||
);
|
||||
assert!(
|
||||
!mgr.activation_errors.read().await.contains_key("gmail"),
|
||||
"stale activation error should be cleared"
|
||||
);
|
||||
|
||||
let flows = mgr.pending_oauth_flows().read().await;
|
||||
assert!(
|
||||
!flows.contains_key("gmail-state"),
|
||||
"gateway OAuth flow for removed extension should be cleared"
|
||||
);
|
||||
assert!(
|
||||
flows.contains_key("other-state"),
|
||||
"unrelated pending OAuth flows should be retained"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_remove_wasm_channel_clears_activation_error_and_deletes_files() {
|
||||
let dir = tempfile::tempdir().expect("temp dir");
|
||||
let tools_dir = dir.path().join("tools");
|
||||
let channels_dir = dir.path().join("channels");
|
||||
let mgr = make_test_manager_with_dirs(None, tools_dir, channels_dir.clone());
|
||||
|
||||
let wasm_path = channels_dir.join("telegram.wasm");
|
||||
let cap_path = channels_dir.join("telegram.capabilities.json");
|
||||
std::fs::write(&wasm_path, b"fake-channel").expect("write channel");
|
||||
std::fs::write(&cap_path, b"{}").expect("write capabilities");
|
||||
|
||||
mgr.activation_errors
|
||||
.write()
|
||||
.await
|
||||
.insert("telegram".to_string(), "channel failed".to_string());
|
||||
|
||||
let result = mgr.remove("telegram").await;
|
||||
assert!(result.is_ok(), "remove should succeed: {:?}", result.err());
|
||||
|
||||
assert!(
|
||||
!mgr.activation_errors.read().await.contains_key("telegram"),
|
||||
"channel activation error should be cleared on remove"
|
||||
);
|
||||
assert!(
|
||||
!wasm_path.exists(),
|
||||
"channel wasm file should be deleted on remove"
|
||||
);
|
||||
assert!(
|
||||
!cap_path.exists(),
|
||||
"channel capabilities file should be deleted on remove"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_sanitize_url_with_query_params() {
|
||||
let url = "https://api.example.com/path?api_key=secret123&token=abc";
|
||||
@@ -4766,6 +5155,189 @@ mod tests {
|
||||
assert!(result.contains("/v1/users/123/profile"));
|
||||
}
|
||||
|
||||
// ---- gateway mode detection tests ----
|
||||
// Regression tests for a bug where MCP OAuth called `open::that()` on the
|
||||
// server machine instead of returning an auth URL to the gateway frontend.
|
||||
// The root cause was that `should_use_gateway_mode()` only checked the
|
||||
// `IRONCLAW_OAUTH_CALLBACK_URL` env var, ignoring `self.tunnel_url`.
|
||||
|
||||
/// Serializes env-mutating tests to prevent parallel races.
|
||||
static GATEWAY_ENV_MUTEX: std::sync::Mutex<()> = std::sync::Mutex::new(());
|
||||
|
||||
/// Build a minimal ExtensionManager with a custom tunnel_url.
|
||||
fn make_manager_with_tunnel(tunnel_url: Option<String>) -> ExtensionManager {
|
||||
use crate::secrets::{InMemorySecretsStore, SecretsCrypto};
|
||||
use crate::tools::mcp::process::McpProcessManager;
|
||||
use crate::tools::mcp::session::McpSessionManager;
|
||||
|
||||
let key = secrecy::SecretString::from(crate::secrets::keychain::generate_master_key_hex());
|
||||
let crypto = Arc::new(SecretsCrypto::new(key).expect("crypto"));
|
||||
let secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync> =
|
||||
Arc::new(InMemorySecretsStore::new(crypto));
|
||||
let tools = Arc::new(crate::tools::ToolRegistry::new());
|
||||
let mcp = Arc::new(McpSessionManager::new());
|
||||
let dir = std::env::temp_dir().join("ironclaw-test-gateway-mode");
|
||||
|
||||
ExtensionManager::new(
|
||||
mcp,
|
||||
Arc::new(McpProcessManager::new()),
|
||||
secrets,
|
||||
tools,
|
||||
None,
|
||||
None,
|
||||
dir.clone(),
|
||||
dir,
|
||||
tunnel_url,
|
||||
"test".to_string(),
|
||||
None,
|
||||
vec![],
|
||||
)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn should_use_gateway_mode_true_for_tunnel_url() {
|
||||
let _guard = GATEWAY_ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok();
|
||||
// SAFETY: Under GATEWAY_ENV_MUTEX, no concurrent env access.
|
||||
unsafe {
|
||||
std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL");
|
||||
}
|
||||
|
||||
let mgr = make_manager_with_tunnel(Some("https://my-gateway.example.com".into()));
|
||||
assert!(
|
||||
mgr.should_use_gateway_mode(),
|
||||
"should detect gateway mode from tunnel_url"
|
||||
);
|
||||
|
||||
unsafe {
|
||||
if let Some(val) = original {
|
||||
std::env::set_var("IRONCLAW_OAUTH_CALLBACK_URL", val);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn should_use_gateway_mode_false_without_tunnel() {
|
||||
let _guard = GATEWAY_ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok();
|
||||
unsafe {
|
||||
std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL");
|
||||
}
|
||||
|
||||
let mgr = make_manager_with_tunnel(None);
|
||||
assert!(
|
||||
!mgr.should_use_gateway_mode(),
|
||||
"should not detect gateway mode without tunnel_url or env var"
|
||||
);
|
||||
|
||||
unsafe {
|
||||
if let Some(val) = original {
|
||||
std::env::set_var("IRONCLAW_OAUTH_CALLBACK_URL", val);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn should_use_gateway_mode_false_for_loopback_tunnel() {
|
||||
let _guard = GATEWAY_ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok();
|
||||
unsafe {
|
||||
std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL");
|
||||
}
|
||||
|
||||
let mgr = make_manager_with_tunnel(Some("http://127.0.0.1:3001".into()));
|
||||
assert!(
|
||||
!mgr.should_use_gateway_mode(),
|
||||
"should not detect gateway mode for loopback tunnel_url"
|
||||
);
|
||||
|
||||
unsafe {
|
||||
if let Some(val) = original {
|
||||
std::env::set_var("IRONCLAW_OAUTH_CALLBACK_URL", val);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Helper to run an async test body while holding the env mutex.
|
||||
/// Clears `IRONCLAW_OAUTH_CALLBACK_URL` for the duration, restoring on drop.
|
||||
struct EnvGuard {
|
||||
original: Option<String>,
|
||||
_mutex: std::sync::MutexGuard<'static, ()>,
|
||||
}
|
||||
|
||||
impl EnvGuard {
|
||||
fn new() -> Self {
|
||||
let guard = GATEWAY_ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok();
|
||||
// SAFETY: Under GATEWAY_ENV_MUTEX, no concurrent env access.
|
||||
unsafe {
|
||||
std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL");
|
||||
}
|
||||
Self {
|
||||
original,
|
||||
_mutex: guard,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for EnvGuard {
|
||||
fn drop(&mut self) {
|
||||
// SAFETY: Under GATEWAY_ENV_MUTEX (still held by _mutex), no concurrent env access.
|
||||
unsafe {
|
||||
if let Some(ref val) = self.original {
|
||||
std::env::set_var("IRONCLAW_OAUTH_CALLBACK_URL", val);
|
||||
} else {
|
||||
std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL");
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_callback_redirect_uri_from_tunnel_url() {
|
||||
let _env = EnvGuard::new();
|
||||
|
||||
let mgr = make_manager_with_tunnel(Some("https://my-gateway.example.com".into()));
|
||||
assert_eq!(
|
||||
mgr.gateway_callback_redirect_uri().await,
|
||||
Some("https://my-gateway.example.com/oauth/callback".to_string()),
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_callback_redirect_uri_none_without_tunnel() {
|
||||
let _env = EnvGuard::new();
|
||||
|
||||
let mgr = make_manager_with_tunnel(None);
|
||||
assert_eq!(mgr.gateway_callback_redirect_uri().await, None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_callback_redirect_uri_trims_trailing_slash() {
|
||||
let _env = EnvGuard::new();
|
||||
|
||||
let mgr = make_manager_with_tunnel(Some("https://my-gateway.example.com/".into()));
|
||||
assert_eq!(
|
||||
mgr.gateway_callback_redirect_uri().await,
|
||||
Some("https://my-gateway.example.com/oauth/callback".to_string()),
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_mode_enabled_explicitly() {
|
||||
let _env = EnvGuard::new();
|
||||
|
||||
let mgr = make_manager_with_tunnel(None);
|
||||
assert!(!mgr.should_use_gateway_mode());
|
||||
|
||||
mgr.enable_gateway_mode("https://my-gateway.example.com".into())
|
||||
.await;
|
||||
assert!(mgr.should_use_gateway_mode());
|
||||
assert_eq!(
|
||||
mgr.gateway_callback_redirect_uri().await,
|
||||
Some("https://my-gateway.example.com/oauth/callback".to_string()),
|
||||
);
|
||||
}
|
||||
// ── Regression tests for PR #677 (unify-extension-lifecycle) ─────────
|
||||
|
||||
#[tokio::test]
|
||||
@@ -4915,7 +5487,6 @@ mod tests {
|
||||
"configure should have stored the relay stream token"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validation_failed_is_distinct_error_variant() {
|
||||
// Regression: ValidationFailed must be a distinct error variant so
|
||||
|
||||
@@ -517,6 +517,9 @@ pub enum ExtensionError {
|
||||
#[error("Authentication failed: {0}")]
|
||||
AuthFailed(String),
|
||||
|
||||
#[error("Server does not support OAuth: {0}")]
|
||||
AuthNotSupported(String),
|
||||
|
||||
#[error("Activation failed: {0}")]
|
||||
ActivationFailed(String),
|
||||
|
||||
|
||||
+17
-5
@@ -9,8 +9,8 @@ use ironclaw::{
|
||||
agent::{Agent, AgentDeps},
|
||||
app::{AppBuilder, AppBuilderFlags},
|
||||
channels::{
|
||||
ChannelManager, ChannelSecretUpdater, GatewayChannel, HttpChannel, ReplChannel,
|
||||
SignalChannel, WebhookServer, WebhookServerConfig,
|
||||
ChannelManager, GatewayChannel, HttpChannel, ReplChannel, SignalChannel, WebhookServer,
|
||||
WebhookServerConfig,
|
||||
wasm::{WasmChannelRouter, WasmChannelRuntime},
|
||||
web::log_layer::LogBroadcaster,
|
||||
},
|
||||
@@ -67,6 +67,10 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
)
|
||||
.await;
|
||||
}
|
||||
Some(Command::Routines(routines_cmd)) => {
|
||||
init_cli_tracing();
|
||||
return ironclaw::cli::run_routines_cli(routines_cmd, cli.config.as_deref()).await;
|
||||
}
|
||||
Some(Command::Mcp(mcp_cmd)) => {
|
||||
init_cli_tracing();
|
||||
return run_mcp_command(*mcp_cmd.clone()).await;
|
||||
@@ -429,9 +433,8 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
"Lifecycle hooks initialized"
|
||||
);
|
||||
|
||||
// Create session manager (shared between agent and web gateway)
|
||||
let session_manager =
|
||||
Arc::new(ironclaw::agent::SessionManager::new().with_hooks(components.hooks.clone()));
|
||||
// Reuse the shared agent session manager prepared by AppBuilder.
|
||||
let session_manager = Arc::clone(&components.agent_session_manager);
|
||||
|
||||
// Lazy scheduler slot — filled after Agent::new creates the Scheduler.
|
||||
// Allows CreateJobTool to dispatch local jobs via the Scheduler even though
|
||||
@@ -472,6 +475,14 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
gw = gw.with_log_level_handle(Arc::clone(&log_level_handle));
|
||||
gw = gw.with_tool_registry(Arc::clone(&components.tools));
|
||||
if let Some(ref ext_mgr) = components.extension_manager {
|
||||
// Enable gateway mode so MCP OAuth returns auth URLs to the frontend
|
||||
// instead of calling open::that() on the server.
|
||||
let gw_base = config
|
||||
.tunnel
|
||||
.public_url
|
||||
.clone()
|
||||
.unwrap_or_else(|| format!("http://{}:{}", gw_config.host, gw_config.port));
|
||||
ext_mgr.enable_gateway_mode(gw_base).await;
|
||||
gw = gw.with_extension_manager(Arc::clone(ext_mgr));
|
||||
}
|
||||
if !components.catalog_entries.is_empty() {
|
||||
@@ -725,6 +736,7 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
|
||||
#[cfg(unix)]
|
||||
{
|
||||
use ironclaw::channels::ChannelSecretUpdater;
|
||||
// Collect all channels that support secret updates
|
||||
let mut secret_updaters: Vec<Arc<dyn ChannelSecretUpdater>> = Vec::new();
|
||||
if let Some(ref state) = http_channel_state {
|
||||
|
||||
+30
-10
@@ -65,7 +65,20 @@ fn install_macos() -> Result<()> {
|
||||
let stdout = logs_dir.join("daemon.stdout.log");
|
||||
let stderr = logs_dir.join("daemon.stderr.log");
|
||||
|
||||
let plist = format!(
|
||||
let plist = macos_plist_content(
|
||||
&exe.display().to_string(),
|
||||
&stdout.display().to_string(),
|
||||
&stderr.display().to_string(),
|
||||
);
|
||||
|
||||
std::fs::write(&file, plist)?;
|
||||
println!("Installed launchd service: {}", file.display());
|
||||
println!(" Start with: ironclaw service start");
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn macos_plist_content(exe: &str, stdout: &str, stderr: &str) -> String {
|
||||
format!(
|
||||
r#"<?xml version="1.0" encoding="UTF-8"?>
|
||||
<!DOCTYPE plist PUBLIC "-//Apple//DTD PLIST 1.0//EN" "http://www.apple.com/DTDs/PropertyList-1.0.dtd">
|
||||
<plist version="1.0">
|
||||
@@ -81,6 +94,11 @@ fn install_macos() -> Result<()> {
|
||||
<true/>
|
||||
<key>KeepAlive</key>
|
||||
<true/>
|
||||
<key>EnvironmentVariables</key>
|
||||
<dict>
|
||||
<key>CLI_ENABLED</key>
|
||||
<string>false</string>
|
||||
</dict>
|
||||
<key>StandardOutPath</key>
|
||||
<string>{stdout}</string>
|
||||
<key>StandardErrorPath</key>
|
||||
@@ -89,15 +107,10 @@ fn install_macos() -> Result<()> {
|
||||
</plist>
|
||||
"#,
|
||||
label = SERVICE_LABEL,
|
||||
exe = xml_escape(&exe.display().to_string()),
|
||||
stdout = xml_escape(&stdout.display().to_string()),
|
||||
stderr = xml_escape(&stderr.display().to_string()),
|
||||
);
|
||||
|
||||
std::fs::write(&file, plist)?;
|
||||
println!("Installed launchd service: {}", file.display());
|
||||
println!(" Start with: ironclaw service start");
|
||||
Ok(())
|
||||
exe = xml_escape(exe),
|
||||
stdout = xml_escape(stdout),
|
||||
stderr = xml_escape(stderr),
|
||||
)
|
||||
}
|
||||
|
||||
fn install_linux() -> Result<()> {
|
||||
@@ -356,4 +369,11 @@ mod tests {
|
||||
let s = path.to_string_lossy();
|
||||
assert!(s.ends_with(".ironclaw/logs"), "unexpected path: {s}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn macos_plist_sets_cli_enabled_false() {
|
||||
let plist = macos_plist_content("/tmp/ironclaw", "/tmp/stdout.log", "/tmp/stderr.log");
|
||||
assert!(plist.contains("<key>EnvironmentVariables</key>"));
|
||||
assert!(plist.contains(" <key>CLI_ENABLED</key>\n <string>false</string>"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1067,6 +1067,8 @@ mod tests {
|
||||
prompt: "Check status".to_string(),
|
||||
context_paths: vec![],
|
||||
max_tokens: 500,
|
||||
use_tools: false,
|
||||
max_tool_rounds: 3,
|
||||
},
|
||||
guardrails: RoutineGuardrails {
|
||||
cooldown: std::time::Duration::from_secs(60),
|
||||
@@ -1198,6 +1200,8 @@ mod tests {
|
||||
prompt: "test".to_string(),
|
||||
context_paths: vec![],
|
||||
max_tokens: 100,
|
||||
use_tools: false,
|
||||
max_tool_rounds: 3,
|
||||
},
|
||||
guardrails: RoutineGuardrails {
|
||||
cooldown: std::time::Duration::from_secs(0),
|
||||
|
||||
@@ -256,7 +256,13 @@ impl Tool for ToolAuthTool {
|
||||
}
|
||||
|
||||
fn requires_approval(&self, _params: &serde_json::Value) -> ApprovalRequirement {
|
||||
ApprovalRequirement::UnlessAutoApproved
|
||||
// In gateway mode, tool_auth only returns an auth URL for the frontend
|
||||
// to open — no browser is launched server-side, so no approval needed.
|
||||
if self.manager.should_use_gateway_mode() {
|
||||
ApprovalRequirement::Never
|
||||
} else {
|
||||
ApprovalRequirement::UnlessAutoApproved
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -733,6 +739,22 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn tool_auth_no_approval_in_gateway_mode() {
|
||||
let manager = test_manager_stub();
|
||||
manager
|
||||
.enable_gateway_mode("http://localhost:3000".to_string())
|
||||
.await;
|
||||
let tool = ToolAuthTool {
|
||||
manager: manager.clone(),
|
||||
};
|
||||
assert_eq!(
|
||||
tool.requires_approval(&serde_json::json!({})),
|
||||
ApprovalRequirement::Never,
|
||||
"tool_auth should not require approval in gateway mode"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tool_upgrade_schema() {
|
||||
use crate::tools::tool::ApprovalRequirement;
|
||||
|
||||
@@ -397,10 +397,6 @@ impl Tool for ListDirTool {
|
||||
false // Directory listings are safe
|
||||
}
|
||||
|
||||
fn requires_approval(&self, _params: &serde_json::Value) -> ApprovalRequirement {
|
||||
ApprovalRequirement::UnlessAutoApproved
|
||||
}
|
||||
|
||||
fn domain(&self) -> ToolDomain {
|
||||
ToolDomain::Container
|
||||
}
|
||||
|
||||
+245
-68
@@ -31,6 +31,12 @@ const MAX_RESPONSE_SIZE: usize = 5 * 1024 * 1024;
|
||||
/// in memory for LLM context. Matches the WASM attachment size cap.
|
||||
const MAX_SAVE_TO_SIZE: usize = 50 * 1024 * 1024;
|
||||
|
||||
/// Default request timeout when the caller does not provide one.
|
||||
const DEFAULT_TIMEOUT_SECS: u64 = 30;
|
||||
|
||||
/// Maximum allowed request timeout to bound resource usage from LLM-controlled inputs.
|
||||
const MAX_TIMEOUT_SECS: u64 = 300;
|
||||
|
||||
/// Maximum number of redirects to follow for simple GET requests.
|
||||
const MAX_REDIRECTS: usize = 3;
|
||||
|
||||
@@ -244,43 +250,120 @@ fn is_html_response(headers: &HashMap<String, String>) -> bool {
|
||||
fn parse_headers_param(
|
||||
headers: Option<&serde_json::Value>,
|
||||
) -> Result<Vec<(String, String)>, ToolError> {
|
||||
fn parse_header_object(
|
||||
map: &serde_json::Map<String, serde_json::Value>,
|
||||
) -> Result<Vec<(String, String)>, ToolError> {
|
||||
let mut out = Vec::with_capacity(map.len());
|
||||
for (k, v) in map {
|
||||
let value = v.as_str().ok_or_else(|| {
|
||||
ToolError::InvalidParameters(format!("header '{}' must have a string value", k))
|
||||
})?;
|
||||
out.push((k.clone(), value.to_string()));
|
||||
}
|
||||
Ok(out)
|
||||
}
|
||||
|
||||
fn parse_header_array(items: &[serde_json::Value]) -> Result<Vec<(String, String)>, ToolError> {
|
||||
let mut out = Vec::with_capacity(items.len());
|
||||
for (idx, item) in items.iter().enumerate() {
|
||||
let obj = item.as_object().ok_or_else(|| {
|
||||
ToolError::InvalidParameters(format!(
|
||||
"headers[{}] must be an object with 'name' and 'value'",
|
||||
idx
|
||||
))
|
||||
})?;
|
||||
let name = obj.get("name").and_then(|v| v.as_str()).ok_or_else(|| {
|
||||
ToolError::InvalidParameters(format!("headers[{}].name must be a string", idx))
|
||||
})?;
|
||||
let value = obj.get("value").and_then(|v| v.as_str()).ok_or_else(|| {
|
||||
ToolError::InvalidParameters(format!("headers[{}].value must be a string", idx))
|
||||
})?;
|
||||
out.push((name.to_string(), value.to_string()));
|
||||
}
|
||||
Ok(out)
|
||||
}
|
||||
|
||||
match headers {
|
||||
None => Ok(Vec::new()),
|
||||
Some(serde_json::Value::Object(map)) => {
|
||||
let mut out = Vec::with_capacity(map.len());
|
||||
for (k, v) in map {
|
||||
let value = v.as_str().ok_or_else(|| {
|
||||
ToolError::InvalidParameters(format!("header '{}' must have a string value", k))
|
||||
})?;
|
||||
out.push((k.clone(), value.to_string()));
|
||||
Some(serde_json::Value::String(raw)) => {
|
||||
let trimmed = raw.trim();
|
||||
if trimmed.is_empty() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
Ok(out)
|
||||
}
|
||||
Some(serde_json::Value::Array(items)) => {
|
||||
let mut out = Vec::with_capacity(items.len());
|
||||
for (idx, item) in items.iter().enumerate() {
|
||||
let obj = item.as_object().ok_or_else(|| {
|
||||
ToolError::InvalidParameters(format!(
|
||||
"headers[{}] must be an object with 'name' and 'value'",
|
||||
idx
|
||||
))
|
||||
})?;
|
||||
let name = obj.get("name").and_then(|v| v.as_str()).ok_or_else(|| {
|
||||
ToolError::InvalidParameters(format!("headers[{}].name must be a string", idx))
|
||||
})?;
|
||||
let value = obj.get("value").and_then(|v| v.as_str()).ok_or_else(|| {
|
||||
ToolError::InvalidParameters(format!("headers[{}].value must be a string", idx))
|
||||
})?;
|
||||
out.push((name.to_string(), value.to_string()));
|
||||
let parsed = serde_json::from_str::<serde_json::Value>(trimmed).map_err(|e| {
|
||||
ToolError::InvalidParameters(format!(
|
||||
"headers string must contain valid JSON object/array: {}",
|
||||
e
|
||||
))
|
||||
})?;
|
||||
match parsed {
|
||||
serde_json::Value::Object(map) => parse_header_object(&map),
|
||||
serde_json::Value::Array(items) => parse_header_array(&items),
|
||||
_ => Err(ToolError::InvalidParameters(
|
||||
"headers string must decode to a JSON object or array".to_string(),
|
||||
)),
|
||||
}
|
||||
Ok(out)
|
||||
}
|
||||
Some(serde_json::Value::Object(map)) => parse_header_object(map),
|
||||
Some(serde_json::Value::Array(items)) => parse_header_array(items),
|
||||
Some(_) => Err(ToolError::InvalidParameters(
|
||||
"'headers' must be an object or an array of {name, value}".to_string(),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_timeout_secs_param(timeout: Option<&serde_json::Value>) -> Result<Option<u64>, ToolError> {
|
||||
let parsed = match timeout {
|
||||
None | Some(serde_json::Value::Null) => Ok(None),
|
||||
Some(serde_json::Value::Number(n)) => n.as_u64().map(Some).ok_or_else(|| {
|
||||
ToolError::InvalidParameters("timeout_secs must be a non-negative integer".to_string())
|
||||
}),
|
||||
Some(serde_json::Value::String(raw)) => {
|
||||
let trimmed = raw.trim();
|
||||
if trimmed.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
let secs = trimmed.parse::<u64>().map_err(|_| {
|
||||
ToolError::InvalidParameters(
|
||||
"timeout_secs string must contain a non-negative integer".to_string(),
|
||||
)
|
||||
})?;
|
||||
Ok(Some(secs))
|
||||
}
|
||||
Some(_) => Err(ToolError::InvalidParameters(
|
||||
"timeout_secs must be an integer".to_string(),
|
||||
)),
|
||||
}?;
|
||||
|
||||
if let Some(secs) = parsed
|
||||
&& secs > MAX_TIMEOUT_SECS
|
||||
{
|
||||
return Err(ToolError::InvalidParameters(format!(
|
||||
"timeout_secs must be <= {}",
|
||||
MAX_TIMEOUT_SECS
|
||||
)));
|
||||
}
|
||||
|
||||
Ok(parsed)
|
||||
}
|
||||
|
||||
fn parse_save_to_param(save_to: Option<&serde_json::Value>) -> Result<Option<String>, ToolError> {
|
||||
match save_to {
|
||||
None | Some(serde_json::Value::Null) => Ok(None),
|
||||
Some(serde_json::Value::String(path)) => {
|
||||
let trimmed = path.trim();
|
||||
if trimmed.is_empty() {
|
||||
Ok(None)
|
||||
} else {
|
||||
Ok(Some(trimmed.to_string()))
|
||||
}
|
||||
}
|
||||
Some(_) => Err(ToolError::InvalidParameters(
|
||||
"save_to must be a string".to_string(),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
/// Extract host from URL in params (for approval checks).
|
||||
fn extract_host_from_params(params: &serde_json::Value) -> Option<String> {
|
||||
params
|
||||
@@ -315,7 +398,7 @@ impl Tool for HttpTool {
|
||||
"method": {
|
||||
"type": "string",
|
||||
"enum": ["GET", "POST", "PUT", "DELETE", "PATCH"],
|
||||
"description": "HTTP method"
|
||||
"description": "HTTP method (default: GET)"
|
||||
},
|
||||
"url": {
|
||||
"type": "string",
|
||||
@@ -346,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/."
|
||||
}
|
||||
},
|
||||
"required": ["method", "url"]
|
||||
"required": ["url"]
|
||||
})
|
||||
}
|
||||
|
||||
@@ -357,7 +440,8 @@ impl Tool for HttpTool {
|
||||
) -> Result<ToolOutput, ToolError> {
|
||||
let start = std::time::Instant::now();
|
||||
|
||||
let method = require_str(¶ms, "method")?;
|
||||
let method = params["method"].as_str().unwrap_or("GET");
|
||||
let method_upper = method.to_uppercase();
|
||||
|
||||
let url = require_str(¶ms, "url")?;
|
||||
let mut parsed_url = validate_url(url)?;
|
||||
@@ -379,6 +463,9 @@ impl Tool for HttpTool {
|
||||
|
||||
// Parse headers
|
||||
let mut headers_vec = parse_headers_param(params.get("headers"))?;
|
||||
let timeout_secs = parse_timeout_secs_param(params.get("timeout_secs"))?;
|
||||
let save_to = parse_save_to_param(params.get("save_to"))?;
|
||||
let effective_timeout = Duration::from_secs(timeout_secs.unwrap_or(DEFAULT_TIMEOUT_SECS));
|
||||
|
||||
// Build request
|
||||
let mut request = match method.to_uppercase().as_str() {
|
||||
@@ -395,6 +482,8 @@ impl Tool for HttpTool {
|
||||
}
|
||||
};
|
||||
|
||||
request = request.timeout(effective_timeout);
|
||||
|
||||
// Add headers
|
||||
for (key, value) in &headers_vec {
|
||||
request = request.header(key.as_str(), value.as_str());
|
||||
@@ -403,7 +492,9 @@ impl Tool for HttpTool {
|
||||
// Add body if present
|
||||
let body_bytes = if let Some(body) = params.get("body") {
|
||||
if let Some(body_str) = body.as_str() {
|
||||
if let Ok(json_body) = serde_json::from_str::<serde_json::Value>(body_str) {
|
||||
if body_str.is_empty() {
|
||||
None
|
||||
} else if let Ok(json_body) = serde_json::from_str::<serde_json::Value>(body_str) {
|
||||
let bytes = serde_json::to_vec(&json_body).map_err(|e| {
|
||||
ToolError::InvalidParameters(format!("invalid body JSON: {}", e))
|
||||
})?;
|
||||
@@ -468,7 +559,7 @@ impl Tool for HttpTool {
|
||||
|
||||
// Build the interceptor request descriptor for recording/replay
|
||||
let intercept_req = crate::llm::recording::HttpExchangeRequest {
|
||||
method: method.to_uppercase(),
|
||||
method: method_upper,
|
||||
url: parsed_url.to_string(),
|
||||
headers: headers_vec.clone(),
|
||||
body: body_bytes
|
||||
@@ -510,7 +601,7 @@ impl Tool for HttpTool {
|
||||
let hop_client = build_pinned_client(
|
||||
&hop_host,
|
||||
&hop_addrs,
|
||||
Duration::from_secs(30),
|
||||
effective_timeout,
|
||||
reqwest::redirect::Policy::none(),
|
||||
)?;
|
||||
|
||||
@@ -524,7 +615,7 @@ impl Tool for HttpTool {
|
||||
.await
|
||||
.map_err(|e| {
|
||||
if e.is_timeout() {
|
||||
ToolError::Timeout(Duration::from_secs(30))
|
||||
ToolError::Timeout(effective_timeout)
|
||||
} else {
|
||||
ToolError::ExternalService(e.to_string())
|
||||
}
|
||||
@@ -588,7 +679,7 @@ impl Tool for HttpTool {
|
||||
} else {
|
||||
let resp = request.send().await.map_err(|e| {
|
||||
if e.is_timeout() {
|
||||
ToolError::Timeout(Duration::from_secs(30))
|
||||
ToolError::Timeout(effective_timeout)
|
||||
} else {
|
||||
ToolError::ExternalService(e.to_string())
|
||||
}
|
||||
@@ -616,7 +707,7 @@ impl Tool for HttpTool {
|
||||
.collect();
|
||||
|
||||
// Use a larger size limit when saving to disk (file downloads)
|
||||
let saving_to_disk = params.get("save_to").is_some();
|
||||
let saving_to_disk = save_to.is_some();
|
||||
let max_size = if saving_to_disk {
|
||||
MAX_SAVE_TO_SIZE
|
||||
} else {
|
||||
@@ -661,11 +752,11 @@ impl Tool for HttpTool {
|
||||
let body_bytes = bytes::Bytes::from(body);
|
||||
|
||||
// If save_to is specified, write raw bytes to file and return metadata.
|
||||
if let Some(save_to) = params.get("save_to").and_then(|v| v.as_str()) {
|
||||
let save_to_owned = save_to.to_string();
|
||||
if let Some(save_to) = save_to {
|
||||
let saved_to = save_to.clone();
|
||||
let bytes_clone = body_bytes.clone();
|
||||
tokio::task::spawn_blocking(move || {
|
||||
let canonical = validate_save_to_path(&save_to_owned)?;
|
||||
let canonical = validate_save_to_path(&save_to)?;
|
||||
std::fs::write(&canonical, &bytes_clone).map_err(|e| {
|
||||
ToolError::ExecutionFailed(format!("failed to write file: {}", e))
|
||||
})?;
|
||||
@@ -676,7 +767,7 @@ impl Tool for HttpTool {
|
||||
.map_err(|e: ToolError| e)?;
|
||||
let result = serde_json::json!({
|
||||
"status": status,
|
||||
"saved_to": save_to,
|
||||
"saved_to": saved_to,
|
||||
"size_bytes": body_bytes.len(),
|
||||
"headers": headers,
|
||||
});
|
||||
@@ -738,18 +829,22 @@ impl Tool for HttpTool {
|
||||
}
|
||||
|
||||
fn requires_approval(&self, params: &serde_json::Value) -> ApprovalRequirement {
|
||||
// 1. Manual auth headers/query params in LLM params
|
||||
if crate::safety::params_contain_manual_credentials(params) {
|
||||
let has_credentials = crate::safety::params_contain_manual_credentials(params)
|
||||
|| (self.credential_registry.as_ref().is_some_and(|registry| {
|
||||
extract_host_from_params(params)
|
||||
.is_some_and(|host| registry.has_credentials_for_host(&host))
|
||||
}));
|
||||
|
||||
if has_credentials {
|
||||
return ApprovalRequirement::Always;
|
||||
}
|
||||
// 2. Target host has credential mappings (will be auto-injected)
|
||||
if let Some(ref registry) = self.credential_registry
|
||||
&& let Some(host) = extract_host_from_params(params)
|
||||
&& registry.has_credentials_for_host(&host)
|
||||
{
|
||||
return ApprovalRequirement::Always;
|
||||
|
||||
// GET requests (or missing method, since GET is the default) are low-risk
|
||||
let method = params["method"].as_str().unwrap_or("GET");
|
||||
if method.eq_ignore_ascii_case("GET") {
|
||||
return ApprovalRequirement::Never;
|
||||
}
|
||||
// Default: outbound HTTP still needs approval unless auto-approved
|
||||
|
||||
ApprovalRequirement::UnlessAutoApproved
|
||||
}
|
||||
|
||||
@@ -887,6 +982,71 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_headers_param_accepts_stringified_array() {
|
||||
let headers =
|
||||
serde_json::json!("[{\"name\":\"Authorization\",\"value\":\"Bearer token\"}]");
|
||||
let parsed = parse_headers_param(Some(&headers)).unwrap();
|
||||
assert_eq!(
|
||||
parsed,
|
||||
vec![("Authorization".to_string(), "Bearer token".to_string())]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_headers_param_rejects_double_string_encoding() {
|
||||
let headers = serde_json::json!("\"hello\"");
|
||||
let err = parse_headers_param(Some(&headers)).unwrap_err();
|
||||
assert!(
|
||||
err.to_string()
|
||||
.contains("headers string must decode to a JSON object or array"),
|
||||
"unexpected error: {}",
|
||||
err
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_timeout_secs_param_accepts_string_integer() {
|
||||
let timeout = serde_json::json!("30");
|
||||
assert_eq!(parse_timeout_secs_param(Some(&timeout)).unwrap(), Some(30));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_timeout_secs_param_treats_empty_string_as_none() {
|
||||
let timeout = serde_json::json!("");
|
||||
assert_eq!(parse_timeout_secs_param(Some(&timeout)).unwrap(), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_timeout_secs_param_rejects_value_above_cap() {
|
||||
let timeout = serde_json::json!(MAX_TIMEOUT_SECS + 1);
|
||||
let err = parse_timeout_secs_param(Some(&timeout)).unwrap_err();
|
||||
assert!(
|
||||
err.to_string()
|
||||
.contains(&format!("timeout_secs must be <= {}", MAX_TIMEOUT_SECS)),
|
||||
"unexpected error: {}",
|
||||
err
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_timeout_secs_param_rejects_string_value_above_cap() {
|
||||
let timeout = serde_json::json!((MAX_TIMEOUT_SECS + 1).to_string());
|
||||
let err = parse_timeout_secs_param(Some(&timeout)).unwrap_err();
|
||||
assert!(
|
||||
err.to_string()
|
||||
.contains(&format!("timeout_secs must be <= {}", MAX_TIMEOUT_SECS)),
|
||||
"unexpected error: {}",
|
||||
err
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_save_to_param_treats_empty_string_as_none() {
|
||||
let save_to = serde_json::json!("");
|
||||
assert_eq!(parse_save_to_param(Some(&save_to)).unwrap(), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_http_tool_schema_body_is_freeform() {
|
||||
let schema = HttpTool::new().parameters_schema();
|
||||
@@ -907,12 +1067,22 @@ mod tests {
|
||||
// ── Approval requirement tests ──────────────────────────────────────
|
||||
|
||||
#[test]
|
||||
fn test_no_auth_headers_returns_unless_auto_approved() {
|
||||
fn test_get_no_auth_headers_returns_never() {
|
||||
let tool = HttpTool::new();
|
||||
let params = serde_json::json!({
|
||||
"method": "GET",
|
||||
"url": "https://api.example.com/data"
|
||||
});
|
||||
assert_eq!(tool.requires_approval(¶ms), ApprovalRequirement::Never);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_post_no_auth_headers_returns_unless_auto_approved() {
|
||||
let tool = HttpTool::new();
|
||||
let params = serde_json::json!({
|
||||
"method": "POST",
|
||||
"url": "https://api.example.com/data"
|
||||
});
|
||||
assert_eq!(
|
||||
tool.requires_approval(¶ms),
|
||||
ApprovalRequirement::UnlessAutoApproved
|
||||
@@ -996,21 +1166,18 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_non_auth_headers_return_unless_auto_approved() {
|
||||
fn test_get_non_auth_headers_return_never() {
|
||||
let tool = HttpTool::new();
|
||||
let params = serde_json::json!({
|
||||
"method": "GET",
|
||||
"url": "https://example.com",
|
||||
"headers": {"Content-Type": "application/json", "Accept": "text/html"}
|
||||
});
|
||||
assert_eq!(
|
||||
tool.requires_approval(¶ms),
|
||||
ApprovalRequirement::UnlessAutoApproved
|
||||
);
|
||||
assert_eq!(tool.requires_approval(¶ms), ApprovalRequirement::Never);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_empty_headers_return_unless_auto_approved() {
|
||||
fn test_get_empty_headers_return_never() {
|
||||
let tool = HttpTool::new();
|
||||
|
||||
// Empty object
|
||||
@@ -1019,10 +1186,7 @@ mod tests {
|
||||
"url": "https://example.com",
|
||||
"headers": {}
|
||||
});
|
||||
assert_eq!(
|
||||
tool.requires_approval(¶ms),
|
||||
ApprovalRequirement::UnlessAutoApproved
|
||||
);
|
||||
assert_eq!(tool.requires_approval(¶ms), ApprovalRequirement::Never);
|
||||
|
||||
// Empty array
|
||||
let params = serde_json::json!({
|
||||
@@ -1030,10 +1194,7 @@ mod tests {
|
||||
"url": "https://example.com",
|
||||
"headers": []
|
||||
});
|
||||
assert_eq!(
|
||||
tool.requires_approval(¶ms),
|
||||
ApprovalRequirement::UnlessAutoApproved
|
||||
);
|
||||
assert_eq!(tool.requires_approval(¶ms), ApprovalRequirement::Never);
|
||||
}
|
||||
|
||||
// ── Credential registry approval tests ─────────────────────────────
|
||||
@@ -1063,7 +1224,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_host_without_credential_mapping_returns_unless_auto_approved() {
|
||||
fn test_get_host_without_credential_mapping_returns_never() {
|
||||
use crate::tools::wasm::SharedCredentialRegistry;
|
||||
|
||||
let registry = Arc::new(SharedCredentialRegistry::new());
|
||||
@@ -1075,10 +1236,7 @@ mod tests {
|
||||
"method": "GET",
|
||||
"url": "https://api.example.com/data"
|
||||
});
|
||||
assert_eq!(
|
||||
tool.requires_approval(¶ms),
|
||||
ApprovalRequirement::UnlessAutoApproved
|
||||
);
|
||||
assert_eq!(tool.requires_approval(¶ms), ApprovalRequirement::Never);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -1119,6 +1277,25 @@ mod tests {
|
||||
assert_eq!(extract_host_from_params(¶ms), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_requires_approval_with_stringified_http_params() {
|
||||
use crate::tools::wasm::SharedCredentialRegistry;
|
||||
|
||||
let tool = HttpTool::new().with_credentials(
|
||||
Arc::new(SharedCredentialRegistry::new()),
|
||||
Arc::new(test_secrets_store()),
|
||||
);
|
||||
let req = serde_json::json!({
|
||||
"body": "",
|
||||
"headers": "[]",
|
||||
"method": "GET",
|
||||
"save_to": "",
|
||||
"timeout_secs": "30",
|
||||
"url": "https://r.jina.ai/http://news.baidu.com/"
|
||||
});
|
||||
let _ = tool.requires_approval(&req);
|
||||
}
|
||||
|
||||
// ── DNS pinning tests ─────────────────────────────────────────────
|
||||
|
||||
#[tokio::test]
|
||||
|
||||
@@ -8,7 +8,7 @@ use secrecy::{ExposeSecret, SecretString};
|
||||
|
||||
use crate::context::JobContext;
|
||||
use crate::tools::builtin::path_utils::validate_path;
|
||||
use crate::tools::tool::{ApprovalRequirement, Tool, ToolError, ToolOutput};
|
||||
use crate::tools::tool::{Tool, ToolError, ToolOutput};
|
||||
|
||||
/// Tool for analyzing images using a vision-capable model.
|
||||
pub struct ImageAnalyzeTool {
|
||||
@@ -86,10 +86,6 @@ impl Tool for ImageAnalyzeTool {
|
||||
})
|
||||
}
|
||||
|
||||
fn requires_approval(&self, _params: &serde_json::Value) -> ApprovalRequirement {
|
||||
ApprovalRequirement::UnlessAutoApproved
|
||||
}
|
||||
|
||||
fn requires_sanitization(&self) -> bool {
|
||||
true
|
||||
}
|
||||
@@ -185,6 +181,7 @@ impl Tool for ImageAnalyzeTool {
|
||||
mod tests {
|
||||
use super::super::media_type_from_path;
|
||||
use super::*;
|
||||
use crate::tools::tool::ApprovalRequirement;
|
||||
use tempfile::TempDir;
|
||||
|
||||
#[test]
|
||||
@@ -199,7 +196,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_requires_approval_returns_unless_auto_approved() {
|
||||
fn test_requires_approval_returns_never() {
|
||||
let tool = ImageAnalyzeTool::new(
|
||||
"https://api.example.com".to_string(),
|
||||
"test-key".to_string(),
|
||||
@@ -208,7 +205,7 @@ mod tests {
|
||||
);
|
||||
assert_eq!(
|
||||
tool.requires_approval(&serde_json::json!({})),
|
||||
ApprovalRequirement::UnlessAutoApproved
|
||||
ApprovalRequirement::Never
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@ use secrecy::{ExposeSecret, SecretString};
|
||||
|
||||
use crate::context::JobContext;
|
||||
use crate::tools::builtin::path_utils::validate_path;
|
||||
use crate::tools::tool::{ApprovalRequirement, Tool, ToolError, ToolOutput};
|
||||
use crate::tools::tool::{Tool, ToolError, ToolOutput};
|
||||
|
||||
/// Tool for editing images using an AI image editing API.
|
||||
pub struct ImageEditTool {
|
||||
@@ -85,10 +85,6 @@ impl Tool for ImageEditTool {
|
||||
})
|
||||
}
|
||||
|
||||
fn requires_approval(&self, _params: &serde_json::Value) -> ApprovalRequirement {
|
||||
ApprovalRequirement::UnlessAutoApproved
|
||||
}
|
||||
|
||||
fn requires_sanitization(&self) -> bool {
|
||||
false
|
||||
}
|
||||
@@ -266,6 +262,7 @@ impl ImageEditTool {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::tools::tool::ApprovalRequirement;
|
||||
use tempfile::TempDir;
|
||||
|
||||
#[test]
|
||||
@@ -280,7 +277,7 @@ mod tests {
|
||||
assert!(!tool.requires_sanitization());
|
||||
assert_eq!(
|
||||
tool.requires_approval(&serde_json::json!({})),
|
||||
ApprovalRequirement::UnlessAutoApproved
|
||||
ApprovalRequirement::Never
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
@@ -5,7 +5,6 @@ use secrecy::{ExposeSecret, SecretString};
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::context::JobContext;
|
||||
use crate::tools::tool::ApprovalRequirement;
|
||||
use crate::tools::{Tool, ToolError, ToolOutput};
|
||||
|
||||
/// Tool for generating images using FLUX or compatible image generation APIs.
|
||||
@@ -87,10 +86,6 @@ impl Tool for ImageGenerateTool {
|
||||
})
|
||||
}
|
||||
|
||||
fn requires_approval(&self, _params: &serde_json::Value) -> ApprovalRequirement {
|
||||
ApprovalRequirement::UnlessAutoApproved
|
||||
}
|
||||
|
||||
fn requires_sanitization(&self) -> bool {
|
||||
false
|
||||
}
|
||||
@@ -186,6 +181,7 @@ impl Tool for ImageGenerateTool {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::tools::tool::ApprovalRequirement;
|
||||
|
||||
#[test]
|
||||
fn test_tool_metadata() {
|
||||
@@ -197,7 +193,7 @@ mod tests {
|
||||
assert_eq!(tool.name(), "image_generate");
|
||||
assert_eq!(
|
||||
tool.requires_approval(&serde_json::json!({})),
|
||||
ApprovalRequirement::UnlessAutoApproved
|
||||
ApprovalRequirement::Never
|
||||
);
|
||||
|
||||
let schema = tool.parameters_schema();
|
||||
|
||||
+122
-61
@@ -12,6 +12,7 @@
|
||||
//! Use `memory_write` to persist important facts that should be remembered
|
||||
//! across sessions.
|
||||
|
||||
use std::path::Path;
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
@@ -26,6 +27,28 @@ use crate::workspace::{Workspace, paths};
|
||||
const PROTECTED_IDENTITY_FILES: &[&str] =
|
||||
&[paths::IDENTITY, paths::SOUL, paths::AGENTS, paths::USER];
|
||||
|
||||
/// Detect paths that are clearly local filesystem references, not workspace-memory docs.
|
||||
///
|
||||
/// Examples:
|
||||
/// - `/Users/.../file.md` (Unix absolute)
|
||||
/// - `C:\Users\...` or `D:/work/...` (Windows absolute)
|
||||
/// - `~/notes.md` (home expansion shorthand)
|
||||
fn looks_like_filesystem_path(path: &str) -> bool {
|
||||
if path.is_empty() {
|
||||
return false;
|
||||
}
|
||||
|
||||
if Path::new(path).is_absolute() || path.starts_with("~/") {
|
||||
return true;
|
||||
}
|
||||
|
||||
let bytes = path.as_bytes();
|
||||
bytes.len() >= 3
|
||||
&& bytes[0].is_ascii_alphabetic()
|
||||
&& bytes[1] == b':'
|
||||
&& (bytes[2] == b'\\' || bytes[2] == b'/')
|
||||
}
|
||||
|
||||
/// Tool for searching workspace memory.
|
||||
///
|
||||
/// Performs hybrid search (FTS + semantic) across all memory documents.
|
||||
@@ -143,7 +166,8 @@ impl Tool for MemoryWriteTool {
|
||||
be remembered across sessions. Targets: 'memory' for curated long-term facts, \
|
||||
'daily_log' for timestamped session notes, 'heartbeat' for the periodic \
|
||||
checklist (HEARTBEAT.md), 'bootstrap' to clear the first-run ritual file, \
|
||||
or provide a custom path for arbitrary file creation."
|
||||
or provide a custom workspace path for arbitrary file creation. \
|
||||
Never pass absolute filesystem paths like '/Users/...' or 'C:\\...'."
|
||||
}
|
||||
|
||||
fn parameters_schema(&self) -> serde_json::Value {
|
||||
@@ -183,6 +207,14 @@ impl Tool for MemoryWriteTool {
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("daily_log");
|
||||
|
||||
if looks_like_filesystem_path(target) {
|
||||
return Err(ToolError::InvalidParameters(format!(
|
||||
"'{}' looks like a local filesystem path. memory_write only works with workspace-memory paths. \
|
||||
Use write_file for filesystem writes. For opening files in an editor, use shell with: open \"<absolute_path>\".",
|
||||
target
|
||||
)));
|
||||
}
|
||||
|
||||
// Bootstrap target: clear BOOTSTRAP.md to mark first-run ritual complete.
|
||||
// Handled early because it accepts empty content (unlike other targets).
|
||||
if target == "bootstrap" {
|
||||
@@ -332,7 +364,8 @@ impl Tool for MemoryReadTool {
|
||||
fn description(&self) -> &str {
|
||||
"Read a file from the workspace memory (database-backed storage). \
|
||||
Use this to read files shown by memory_tree. NOT for local filesystem files \
|
||||
(use read_file for those). Works with identity files, heartbeat checklist, \
|
||||
(use read_file for those). Do not pass absolute paths like '/Users/...' or 'C:\\...'. \
|
||||
Works with identity files, heartbeat checklist, \
|
||||
memory, daily logs, or any custom workspace path."
|
||||
}
|
||||
|
||||
@@ -358,6 +391,14 @@ impl Tool for MemoryReadTool {
|
||||
|
||||
let path = require_str(¶ms, "path")?;
|
||||
|
||||
if looks_like_filesystem_path(path) {
|
||||
return Err(ToolError::InvalidParameters(format!(
|
||||
"'{}' looks like a local filesystem path. memory_read only works with workspace-memory paths. \
|
||||
Use read_file for filesystem reads. For opening files in an editor, use shell with: open \"<absolute_path>\".",
|
||||
path
|
||||
)));
|
||||
}
|
||||
|
||||
let doc = self
|
||||
.workspace
|
||||
.read(path)
|
||||
@@ -498,80 +539,100 @@ impl Tool for MemoryTreeTool {
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(all(test, feature = "postgres"))]
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn make_test_workspace() -> Arc<Workspace> {
|
||||
Arc::new(Workspace::new(
|
||||
"test_user",
|
||||
deadpool_postgres::Pool::builder(deadpool_postgres::Manager::new(
|
||||
tokio_postgres::Config::new(),
|
||||
tokio_postgres::NoTls,
|
||||
#[test]
|
||||
fn detects_filesystem_paths() {
|
||||
assert!(looks_like_filesystem_path("/Users/nige/file.md"));
|
||||
assert!(looks_like_filesystem_path("C:\\Users\\nige\\file.md"));
|
||||
assert!(looks_like_filesystem_path("D:/work/file.md"));
|
||||
assert!(looks_like_filesystem_path("~/notes.md"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn allows_workspace_memory_paths() {
|
||||
assert!(!looks_like_filesystem_path("MEMORY.md"));
|
||||
assert!(!looks_like_filesystem_path("daily/2026-03-11.md"));
|
||||
assert!(!looks_like_filesystem_path("projects/alpha/notes.md"));
|
||||
}
|
||||
|
||||
#[cfg(feature = "postgres")]
|
||||
mod postgres_schema_tests {
|
||||
use super::*;
|
||||
|
||||
fn make_test_workspace() -> Arc<Workspace> {
|
||||
Arc::new(Workspace::new(
|
||||
"test_user",
|
||||
deadpool_postgres::Pool::builder(deadpool_postgres::Manager::new(
|
||||
tokio_postgres::Config::new(),
|
||||
tokio_postgres::NoTls,
|
||||
))
|
||||
.build()
|
||||
.unwrap(),
|
||||
))
|
||||
.build()
|
||||
.unwrap(),
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_memory_search_schema() {
|
||||
let workspace = make_test_workspace();
|
||||
let tool = MemorySearchTool::new(workspace);
|
||||
#[test]
|
||||
fn test_memory_search_schema() {
|
||||
let workspace = make_test_workspace();
|
||||
let tool = MemorySearchTool::new(workspace);
|
||||
|
||||
assert_eq!(tool.name(), "memory_search");
|
||||
assert!(!tool.requires_sanitization());
|
||||
assert_eq!(tool.name(), "memory_search");
|
||||
assert!(!tool.requires_sanitization());
|
||||
|
||||
let schema = tool.parameters_schema();
|
||||
assert!(schema["properties"]["query"].is_object());
|
||||
assert!(
|
||||
schema["required"]
|
||||
.as_array()
|
||||
.unwrap()
|
||||
.contains(&"query".into())
|
||||
);
|
||||
}
|
||||
let schema = tool.parameters_schema();
|
||||
assert!(schema["properties"]["query"].is_object());
|
||||
assert!(
|
||||
schema["required"]
|
||||
.as_array()
|
||||
.unwrap()
|
||||
.contains(&"query".into())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_memory_write_schema() {
|
||||
let workspace = make_test_workspace();
|
||||
let tool = MemoryWriteTool::new(workspace);
|
||||
#[test]
|
||||
fn test_memory_write_schema() {
|
||||
let workspace = make_test_workspace();
|
||||
let tool = MemoryWriteTool::new(workspace);
|
||||
|
||||
assert_eq!(tool.name(), "memory_write");
|
||||
assert_eq!(tool.name(), "memory_write");
|
||||
|
||||
let schema = tool.parameters_schema();
|
||||
assert!(schema["properties"]["content"].is_object());
|
||||
assert!(schema["properties"]["target"].is_object());
|
||||
assert!(schema["properties"]["append"].is_object());
|
||||
}
|
||||
let schema = tool.parameters_schema();
|
||||
assert!(schema["properties"]["content"].is_object());
|
||||
assert!(schema["properties"]["target"].is_object());
|
||||
assert!(schema["properties"]["append"].is_object());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_memory_read_schema() {
|
||||
let workspace = make_test_workspace();
|
||||
let tool = MemoryReadTool::new(workspace);
|
||||
#[test]
|
||||
fn test_memory_read_schema() {
|
||||
let workspace = make_test_workspace();
|
||||
let tool = MemoryReadTool::new(workspace);
|
||||
|
||||
assert_eq!(tool.name(), "memory_read");
|
||||
assert_eq!(tool.name(), "memory_read");
|
||||
|
||||
let schema = tool.parameters_schema();
|
||||
assert!(schema["properties"]["path"].is_object());
|
||||
assert!(
|
||||
schema["required"]
|
||||
.as_array()
|
||||
.unwrap()
|
||||
.contains(&"path".into())
|
||||
);
|
||||
}
|
||||
let schema = tool.parameters_schema();
|
||||
assert!(schema["properties"]["path"].is_object());
|
||||
assert!(
|
||||
schema["required"]
|
||||
.as_array()
|
||||
.unwrap()
|
||||
.contains(&"path".into())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_memory_tree_schema() {
|
||||
let workspace = make_test_workspace();
|
||||
let tool = MemoryTreeTool::new(workspace);
|
||||
#[test]
|
||||
fn test_memory_tree_schema() {
|
||||
let workspace = make_test_workspace();
|
||||
let tool = MemoryTreeTool::new(workspace);
|
||||
|
||||
assert_eq!(tool.name(), "memory_tree");
|
||||
assert_eq!(tool.name(), "memory_tree");
|
||||
|
||||
let schema = tool.parameters_schema();
|
||||
assert!(schema["properties"]["path"].is_object());
|
||||
assert!(schema["properties"]["depth"].is_object());
|
||||
assert_eq!(schema["properties"]["depth"]["default"], 1);
|
||||
let schema = tool.parameters_schema();
|
||||
assert!(schema["properties"]["path"].is_object());
|
||||
assert!(schema["properties"]["depth"].is_object());
|
||||
assert_eq!(schema["properties"]["depth"]["default"], 1);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -15,6 +15,7 @@ pub mod secrets_tools;
|
||||
pub(crate) mod shell;
|
||||
pub mod skill_tools;
|
||||
mod time;
|
||||
mod tool_info;
|
||||
|
||||
pub use echo::EchoTool;
|
||||
pub use extension_tools::{
|
||||
@@ -39,6 +40,7 @@ pub use secrets_tools::{SecretDeleteTool, SecretListTool};
|
||||
pub use shell::ShellTool;
|
||||
pub use skill_tools::{SkillInstallTool, SkillListTool, SkillRemoveTool, SkillSearchTool};
|
||||
pub use time::TimeTool;
|
||||
pub use tool_info::ToolInfoTool;
|
||||
mod html_converter;
|
||||
pub mod image_analyze;
|
||||
pub mod image_edit;
|
||||
|
||||
@@ -104,6 +104,14 @@ impl Tool for RoutineCreateTool {
|
||||
"enum": ["lightweight", "full_job"],
|
||||
"description": "Execution mode: 'lightweight' (single LLM call, default) or 'full_job' (multi-turn with tools)"
|
||||
},
|
||||
"use_tools": {
|
||||
"type": "boolean",
|
||||
"description": "Enable tool access in lightweight mode (default: false). Only safe tools (no approval required) are available. Ignored for full_job mode."
|
||||
},
|
||||
"max_tool_rounds": {
|
||||
"type": "integer",
|
||||
"description": "Max tool call rounds in lightweight mode (default: 3). Only used when use_tools is true."
|
||||
},
|
||||
"cooldown_secs": {
|
||||
"type": "integer",
|
||||
"description": "Minimum seconds between fires (default: 300)"
|
||||
@@ -262,11 +270,24 @@ impl Tool for RoutineCreateTool {
|
||||
})
|
||||
.unwrap_or_default();
|
||||
|
||||
let use_tools = params
|
||||
.get("use_tools")
|
||||
.and_then(|v| v.as_bool())
|
||||
.unwrap_or(false);
|
||||
|
||||
let max_tool_rounds = params
|
||||
.get("max_tool_rounds")
|
||||
.and_then(|v| v.as_u64())
|
||||
.map(|v| v.clamp(1, crate::agent::routine::MAX_TOOL_ROUNDS_LIMIT as u64) as u32)
|
||||
.unwrap_or(3);
|
||||
|
||||
let action = match action_type {
|
||||
"lightweight" => RoutineAction::Lightweight {
|
||||
prompt: prompt.to_string(),
|
||||
context_paths,
|
||||
max_tokens: 4096,
|
||||
use_tools,
|
||||
max_tool_rounds,
|
||||
},
|
||||
"full_job" => {
|
||||
let tool_permissions = crate::agent::routine::parse_tool_permissions(¶ms);
|
||||
|
||||
@@ -0,0 +1,183 @@
|
||||
//! 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(¶ms, "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(®istry));
|
||||
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(®istry));
|
||||
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(®istry));
|
||||
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(®istry));
|
||||
drop(registry);
|
||||
|
||||
let ctx = JobContext::default();
|
||||
let result = tool
|
||||
.execute(serde_json::json!({"name": "echo"}), &ctx)
|
||||
.await;
|
||||
assert!(matches!(result, Err(ToolError::ExecutionFailed(_))));
|
||||
}
|
||||
}
|
||||
+77
-2
@@ -669,7 +669,7 @@ pub async fn authorize_mcp_server(
|
||||
}
|
||||
|
||||
// Determine client_id and endpoints
|
||||
let (client_id, authorization_url, token_url, use_pkce, scopes, extra_params) =
|
||||
let (client_id, authorization_url, token_url, use_pkce, scopes, mut extra_params) =
|
||||
if let Some(oauth) = &server_config.oauth {
|
||||
// Pre-configured OAuth
|
||||
let (auth_url, tok_url) = discover_oauth_endpoints(server_config).await?;
|
||||
@@ -711,6 +711,13 @@ pub async fn authorize_mcp_server(
|
||||
None
|
||||
};
|
||||
|
||||
// Generate OAuth state parameter. While optional in OAuth 2.1 with PKCE,
|
||||
// some MCP servers (e.g. Attio) require it.
|
||||
let mut state_bytes = [0u8; 16];
|
||||
rand::rngs::OsRng.fill_bytes(&mut state_bytes);
|
||||
let state = URL_SAFE_NO_PAD.encode(state_bytes);
|
||||
extra_params.insert("state".to_string(), state);
|
||||
|
||||
// Compute canonical resource URI for RFC 8707
|
||||
let resource = canonical_resource_uri(&server_config.url);
|
||||
|
||||
@@ -741,7 +748,10 @@ pub async fn authorize_mcp_server(
|
||||
|
||||
println!(" Waiting for authorization...");
|
||||
|
||||
// Wait for callback
|
||||
// Wait for callback. State is sent in the URL for servers that require it
|
||||
// (e.g. Attio), but we don't enforce validation on the callback because MCP
|
||||
// servers use PKCE which already binds the request to the token exchange,
|
||||
// and some servers may not echo state back.
|
||||
let code = wait_for_authorization_callback(listener, &server_config.name).await?;
|
||||
|
||||
println!(" Exchanging code for token...");
|
||||
@@ -1711,4 +1721,69 @@ mod tests {
|
||||
|
||||
assert!(!url.contains("resource="));
|
||||
}
|
||||
|
||||
/// Regression test: MCP OAuth authorization URLs must include a `state`
|
||||
/// parameter. While OAuth 2.1 makes `state` optional when PKCE is used,
|
||||
/// some MCP servers (e.g. Attio) require it and reject requests without it:
|
||||
/// {"error":"invalid_request","error_description":"Invalid value provided
|
||||
/// for: state"}
|
||||
///
|
||||
/// Including `state` is harmless for servers that don't require it, since
|
||||
/// it is a standard OAuth parameter that compliant servers will echo back
|
||||
/// or ignore.
|
||||
///
|
||||
/// The state is generated in `authorize_mcp_server` and injected into
|
||||
/// `extra_params` before `build_authorization_url` is called. This test
|
||||
/// verifies that `build_authorization_url` correctly propagates state from
|
||||
/// extra_params into the URL, and that each generated state is unique.
|
||||
#[test]
|
||||
fn test_authorization_url_includes_state_parameter() {
|
||||
// Simulate what authorize_mcp_server does: generate state and
|
||||
// insert it into extra_params.
|
||||
let mut extra_params = HashMap::new();
|
||||
let mut state_bytes = [0u8; 16];
|
||||
rand::rngs::OsRng.fill_bytes(&mut state_bytes);
|
||||
let state = URL_SAFE_NO_PAD.encode(state_bytes);
|
||||
extra_params.insert("state".to_string(), state.clone());
|
||||
|
||||
let pkce = PkceChallenge::generate();
|
||||
let url = build_authorization_url(
|
||||
"https://app.attio.com/oidc/authorize",
|
||||
"test-client",
|
||||
"http://127.0.0.1:9876/callback",
|
||||
&[
|
||||
"mcp".to_string(),
|
||||
"offline_access".to_string(),
|
||||
"openid".to_string(),
|
||||
],
|
||||
Some(&pkce),
|
||||
&extra_params,
|
||||
Some("https://mcp.attio.com/mcp"),
|
||||
);
|
||||
|
||||
// State must be present in the URL
|
||||
assert!(
|
||||
url.contains(&format!("state={}", state)),
|
||||
"Authorization URL must include the state parameter, got: {}",
|
||||
url,
|
||||
);
|
||||
|
||||
// State must be base64url-encoded (no padding, no +/)
|
||||
assert!(!state.contains('+'), "State must be base64url-safe");
|
||||
assert!(!state.contains('/'), "State must be base64url-safe");
|
||||
assert!(!state.contains('='), "State must not have padding");
|
||||
|
||||
// State must have sufficient entropy (16 bytes -> 22 base64url chars)
|
||||
assert!(
|
||||
state.len() >= 22,
|
||||
"State must have at least 128 bits of entropy, got {} chars",
|
||||
state.len(),
|
||||
);
|
||||
|
||||
// Two generated states must differ
|
||||
let mut state_bytes_2 = [0u8; 16];
|
||||
rand::rngs::OsRng.fill_bytes(&mut state_bytes_2);
|
||||
let state_2 = URL_SAFE_NO_PAD.encode(state_bytes_2);
|
||||
assert_ne!(state, state_2, "State must be unique per request");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -12,6 +12,7 @@ pub mod builtin;
|
||||
pub mod execute;
|
||||
pub mod mcp;
|
||||
pub mod rate_limiter;
|
||||
pub mod redaction;
|
||||
pub mod schema_validator;
|
||||
pub mod wasm;
|
||||
|
||||
|
||||
@@ -0,0 +1,251 @@
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
const REDACTED: &str = "[REDACTED]";
|
||||
const SENSITIVE_EXACT: &[&str] = &[
|
||||
"authorization",
|
||||
"proxy-authorization",
|
||||
"cookie",
|
||||
"set-cookie",
|
||||
"x-api-key",
|
||||
"api-key",
|
||||
"api_key",
|
||||
"access_token",
|
||||
"refresh_token",
|
||||
"session_token",
|
||||
"id_token",
|
||||
"token",
|
||||
"password",
|
||||
"passwd",
|
||||
"secret",
|
||||
"client_secret",
|
||||
"private_key",
|
||||
"apikey",
|
||||
"apisecret",
|
||||
];
|
||||
|
||||
const SENSITIVE_PARTS: &[&str] = &[
|
||||
"password",
|
||||
"passwd",
|
||||
"secret",
|
||||
"credential",
|
||||
"authorization",
|
||||
"cookie",
|
||||
"apikey",
|
||||
"apisecret",
|
||||
];
|
||||
const TOKEN_PARTS: &[&str] = &["token", "jwt"];
|
||||
const KEY_PARTS: &[&str] = &["key"];
|
||||
const CONTEXT_PARTS: &[&str] = &[
|
||||
"auth",
|
||||
"oauth",
|
||||
"authorization",
|
||||
"api",
|
||||
"access",
|
||||
"refresh",
|
||||
"session",
|
||||
"bearer",
|
||||
"private",
|
||||
"client",
|
||||
"id",
|
||||
"app",
|
||||
"user",
|
||||
"application",
|
||||
"account",
|
||||
];
|
||||
|
||||
fn split_camel_case_key_parts(key: &str) -> Vec<String> {
|
||||
if key.is_empty() {
|
||||
return Vec::new();
|
||||
}
|
||||
|
||||
let chars: Vec<char> = key.chars().collect();
|
||||
let mut parts = Vec::new();
|
||||
let mut start = 0;
|
||||
|
||||
for i in 1..chars.len() {
|
||||
let prev = chars[i - 1];
|
||||
let cur = chars[i];
|
||||
let next = chars.get(i + 1).copied();
|
||||
|
||||
let boundary = (prev.is_ascii_lowercase() && cur.is_ascii_uppercase())
|
||||
|| (prev.is_ascii_alphabetic() && cur.is_ascii_digit())
|
||||
|| (prev.is_ascii_digit() && cur.is_ascii_alphabetic())
|
||||
|| (prev.is_ascii_uppercase()
|
||||
&& cur.is_ascii_uppercase()
|
||||
&& next.map(|n| n.is_ascii_lowercase()).unwrap_or(false));
|
||||
|
||||
if boundary {
|
||||
parts.push(chars[start..i].iter().collect::<String>());
|
||||
start = i;
|
||||
}
|
||||
}
|
||||
|
||||
parts.push(chars[start..].iter().collect::<String>());
|
||||
parts
|
||||
}
|
||||
|
||||
fn tokenize_key_parts(key: &str) -> Vec<String> {
|
||||
let mut parts = Vec::new();
|
||||
|
||||
for segment in key.split(|c: char| !c.is_ascii_alphanumeric()) {
|
||||
if segment.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
parts.extend(split_camel_case_key_parts(segment));
|
||||
}
|
||||
|
||||
parts.into_iter().map(|p| p.to_ascii_lowercase()).collect()
|
||||
}
|
||||
|
||||
fn has_exact(parts: &[String], candidates: &[&str]) -> bool {
|
||||
parts
|
||||
.iter()
|
||||
.any(|part| candidates.iter().any(|candidate| part == candidate))
|
||||
}
|
||||
|
||||
fn has_candidate_or_numbered_variant(parts: &[String], candidates: &[&str]) -> bool {
|
||||
parts.iter().any(|part| {
|
||||
candidates.iter().any(|candidate| {
|
||||
if part == candidate {
|
||||
return true;
|
||||
}
|
||||
let Some(suffix) = part.strip_prefix(candidate) else {
|
||||
return false;
|
||||
};
|
||||
!suffix.is_empty() && suffix.chars().all(|c| c.is_ascii_digit())
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn has_contextual_suffix(parts: &[String], candidates: &[&str]) -> bool {
|
||||
parts.iter().any(|part| {
|
||||
candidates.iter().any(|candidate| {
|
||||
let Some(prefix) = part.strip_suffix(candidate) else {
|
||||
return false;
|
||||
};
|
||||
!prefix.is_empty() && CONTEXT_PARTS.contains(&prefix)
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn is_sensitive_key(key: &str) -> bool {
|
||||
let lower = key.to_ascii_lowercase();
|
||||
if SENSITIVE_EXACT.contains(&lower.as_str()) {
|
||||
return true;
|
||||
}
|
||||
|
||||
let parts = tokenize_key_parts(key);
|
||||
if parts.is_empty() {
|
||||
return false;
|
||||
}
|
||||
|
||||
if has_candidate_or_numbered_variant(&parts, SENSITIVE_PARTS) {
|
||||
return true;
|
||||
}
|
||||
|
||||
let has_token = has_candidate_or_numbered_variant(&parts, TOKEN_PARTS);
|
||||
let has_key = has_candidate_or_numbered_variant(&parts, KEY_PARTS);
|
||||
|
||||
if has_token && has_key {
|
||||
return true;
|
||||
}
|
||||
|
||||
if has_contextual_suffix(&parts, TOKEN_PARTS) || has_contextual_suffix(&parts, KEY_PARTS) {
|
||||
return true;
|
||||
}
|
||||
|
||||
let has_context = has_exact(&parts, CONTEXT_PARTS);
|
||||
has_context && (has_token || has_key)
|
||||
}
|
||||
|
||||
fn redact_in_place(value: &mut Value) {
|
||||
match value {
|
||||
Value::Object(map) => redact_object(map),
|
||||
Value::Array(items) => {
|
||||
for item in items {
|
||||
redact_in_place(item);
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
fn redact_object(map: &mut Map<String, Value>) {
|
||||
for (key, val) in map {
|
||||
if is_sensitive_key(key) {
|
||||
*val = Value::String(REDACTED.to_string());
|
||||
} else {
|
||||
redact_in_place(val);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn redact_sensitive_json(value: &Value) -> Value {
|
||||
let mut cloned = value.clone();
|
||||
redact_in_place(&mut cloned);
|
||||
cloned
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{is_sensitive_key, redact_sensitive_json};
|
||||
|
||||
#[test]
|
||||
fn redacts_exact_sensitive_keys() {
|
||||
let input = serde_json::json!({
|
||||
"headers": {
|
||||
"Authorization": "Bearer abc",
|
||||
"x-api-key": "k-123",
|
||||
"content-type": "application/json"
|
||||
},
|
||||
"password": "p@ss"
|
||||
});
|
||||
let out = redact_sensitive_json(&input);
|
||||
assert_eq!(out["headers"]["Authorization"], "[REDACTED]");
|
||||
assert_eq!(out["headers"]["x-api-key"], "[REDACTED]");
|
||||
assert_eq!(out["headers"]["content-type"], "application/json");
|
||||
assert_eq!(out["password"], "[REDACTED]");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn redacts_nested_sensitive_keys() {
|
||||
let input = serde_json::json!({
|
||||
"body": {
|
||||
"clientSecret": "xyz",
|
||||
"nested": [{"authToken": "123"}, {"query": "ok"}]
|
||||
}
|
||||
});
|
||||
let out = redact_sensitive_json(&input);
|
||||
assert_eq!(out["body"]["clientSecret"], "[REDACTED]");
|
||||
assert_eq!(out["body"]["nested"][0]["authToken"], "[REDACTED]");
|
||||
assert_eq!(out["body"]["nested"][1]["query"], "ok");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn does_not_over_redact_common_non_sensitive_keys() {
|
||||
assert!(!is_sensitive_key("author"));
|
||||
assert!(!is_sensitive_key("authorize_user"));
|
||||
assert!(!is_sensitive_key("token_count"));
|
||||
assert!(!is_sensitive_key("tokenize"));
|
||||
assert!(!is_sensitive_key("oauth_redirect_uri"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn still_redacts_expected_token_keys() {
|
||||
assert!(is_sensitive_key("auth_token"));
|
||||
assert!(is_sensitive_key("oauth_token"));
|
||||
assert!(is_sensitive_key("accessToken"));
|
||||
assert!(is_sensitive_key("apiKey"));
|
||||
assert!(is_sensitive_key("token_key"));
|
||||
assert!(is_sensitive_key("appTokenKey"));
|
||||
assert!(is_sensitive_key("userJwt"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn redacts_lowercase_digit_suffix_segments() {
|
||||
assert!(is_sensitive_key("password123"));
|
||||
assert!(is_sensitive_key("secret99"));
|
||||
assert!(is_sensitive_key("accounttoken2"));
|
||||
}
|
||||
}
|
||||
+45
-1
@@ -23,7 +23,7 @@ use crate::tools::builtin::{
|
||||
ToolUpgradeTool, WriteFileTool,
|
||||
};
|
||||
use crate::tools::rate_limiter::RateLimiter;
|
||||
use crate::tools::tool::{Tool, ToolDomain};
|
||||
use crate::tools::tool::{ApprovalRequirement, Tool, ToolDomain};
|
||||
use crate::tools::wasm::{
|
||||
Capabilities, OAuthRefreshConfig, ResourceLimits, SharedCredentialRegistry, WasmError,
|
||||
WasmStorageError, WasmToolRuntime, WasmToolStore, WasmToolWrapper,
|
||||
@@ -75,6 +75,7 @@ const PROTECTED_TOOL_NAMES: &[&str] = &[
|
||||
"image_generate",
|
||||
"image_edit",
|
||||
"image_analyze",
|
||||
"tool_info",
|
||||
];
|
||||
|
||||
/// Registry of available tools.
|
||||
@@ -245,6 +246,17 @@ impl ToolRegistry {
|
||||
tracing::debug!("Registered {} built-in tools", self.count());
|
||||
}
|
||||
|
||||
/// Register the `tool_info` discovery tool.
|
||||
///
|
||||
/// Requires `Arc<Self>` so the tool can query the registry for other tools'
|
||||
/// schemas at runtime. Call after `register_builtin_tools()`.
|
||||
pub fn register_tool_info(self: &Arc<Self>) {
|
||||
use crate::tools::builtin::ToolInfoTool;
|
||||
let tool = ToolInfoTool::new(Arc::downgrade(self));
|
||||
self.register_sync(Arc::new(tool));
|
||||
tracing::debug!("Registered tool_info discovery tool");
|
||||
}
|
||||
|
||||
/// Register only orchestrator-domain tools (safe for the main process).
|
||||
///
|
||||
/// This registers tools that don't touch the filesystem or run shell commands:
|
||||
@@ -278,6 +290,38 @@ impl ToolRegistry {
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Get tool definitions excluding specific tools by name.
|
||||
///
|
||||
/// Used by lightweight routines to filter out denylisted and approval-gated tools
|
||||
/// so the LLM only sees tools it is actually allowed to call.
|
||||
pub async fn tool_definitions_excluding(&self, deny: &[&str]) -> Vec<ToolDefinition> {
|
||||
let empty_params = serde_json::Value::Object(serde_json::Map::new());
|
||||
let mut defs: Vec<ToolDefinition> = self
|
||||
.tools
|
||||
.read()
|
||||
.await
|
||||
.values()
|
||||
.filter(|tool| {
|
||||
// Exclude denylisted tools
|
||||
if deny.contains(&tool.name()) {
|
||||
return false;
|
||||
}
|
||||
// Exclude tools that require approval
|
||||
matches!(
|
||||
tool.requires_approval(&empty_params),
|
||||
ApprovalRequirement::Never
|
||||
)
|
||||
})
|
||||
.map(|tool| ToolDefinition {
|
||||
name: tool.name().to_string(),
|
||||
description: tool.description().to_string(),
|
||||
parameters: tool.parameters_schema(),
|
||||
})
|
||||
.collect();
|
||||
defs.sort_unstable_by(|a, b| a.name.cmp(&b.name));
|
||||
defs
|
||||
}
|
||||
|
||||
/// Register development tools for building software.
|
||||
///
|
||||
/// These tools provide shell access, file operations, and code editing
|
||||
|
||||
@@ -336,6 +336,17 @@ pub trait Tool: Send + Sync {
|
||||
None
|
||||
}
|
||||
|
||||
/// Full parameter schema for discovery and coercion purposes.
|
||||
///
|
||||
/// Unlike `parameters_schema()` (which may be permissive to keep the tools
|
||||
/// array compact), this returns the complete typed schema. Used by the
|
||||
/// `tool_info` built-in and by WASM parameter coercion.
|
||||
///
|
||||
/// Default: delegates to `parameters_schema()`.
|
||||
fn discovery_schema(&self) -> serde_json::Value {
|
||||
self.parameters_schema()
|
||||
}
|
||||
|
||||
/// Get the tool schema for LLM function calling.
|
||||
fn schema(&self) -> ToolSchema {
|
||||
ToolSchema {
|
||||
|
||||
+6
-84
@@ -1,7 +1,5 @@
|
||||
//! WASM sandbox error types.
|
||||
|
||||
use std::fmt;
|
||||
|
||||
use thiserror::Error;
|
||||
|
||||
/// Errors that can occur during WASM tool execution.
|
||||
@@ -68,13 +66,13 @@ pub enum WasmError {
|
||||
Timeout(std::time::Duration),
|
||||
|
||||
/// Component returned an error response.
|
||||
/// When `hint` is non-empty it carries the tool's description and parameter
|
||||
/// schema so the LLM can retry with correct arguments.
|
||||
/// When `hint` is non-empty it points the LLM to `tool_info` so it can
|
||||
/// fetch the tool's full parameter schema on demand.
|
||||
#[error("Tool error: {message}{}", if hint.is_empty() { String::new() } else { format!("\n\nTool usage hint:\n{hint}") })]
|
||||
ToolReturnedError {
|
||||
/// The error message from the WASM tool.
|
||||
message: String,
|
||||
/// Optional description + schema hint (empty when unavailable).
|
||||
/// Optional retry hint (empty when unavailable).
|
||||
hint: String,
|
||||
},
|
||||
|
||||
@@ -99,73 +97,9 @@ impl From<WasmError> for crate::tools::ToolError {
|
||||
}
|
||||
}
|
||||
|
||||
/// Details about a trap that occurred during execution.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct TrapInfo {
|
||||
/// Human-readable trap message.
|
||||
pub message: String,
|
||||
/// Trap code if available.
|
||||
pub code: Option<TrapCode>,
|
||||
}
|
||||
|
||||
impl fmt::Display for TrapInfo {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match &self.code {
|
||||
Some(code) => write!(f, "{}: {}", code, self.message),
|
||||
None => write!(f, "{}", self.message),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Known trap codes from Wasmtime.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum TrapCode {
|
||||
/// Out of bounds memory access.
|
||||
MemoryOutOfBounds,
|
||||
/// Out of bounds table access.
|
||||
TableOutOfBounds,
|
||||
/// Indirect call type mismatch.
|
||||
IndirectCallToNull,
|
||||
/// Signature mismatch on indirect call.
|
||||
BadSignature,
|
||||
/// Integer overflow.
|
||||
IntegerOverflow,
|
||||
/// Integer division by zero.
|
||||
IntegerDivisionByZero,
|
||||
/// Invalid conversion to integer.
|
||||
BadConversionToInteger,
|
||||
/// Unreachable instruction executed.
|
||||
UnreachableCodeReached,
|
||||
/// Call stack exhausted.
|
||||
StackOverflow,
|
||||
/// Out of fuel.
|
||||
OutOfFuel,
|
||||
/// Unknown trap code.
|
||||
Unknown,
|
||||
}
|
||||
|
||||
impl fmt::Display for TrapCode {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
let s = match self {
|
||||
TrapCode::MemoryOutOfBounds => "memory out of bounds",
|
||||
TrapCode::TableOutOfBounds => "table out of bounds",
|
||||
TrapCode::IndirectCallToNull => "indirect call to null",
|
||||
TrapCode::BadSignature => "bad signature",
|
||||
TrapCode::IntegerOverflow => "integer overflow",
|
||||
TrapCode::IntegerDivisionByZero => "integer division by zero",
|
||||
TrapCode::BadConversionToInteger => "bad conversion to integer",
|
||||
TrapCode::UnreachableCodeReached => "unreachable code reached",
|
||||
TrapCode::StackOverflow => "stack overflow",
|
||||
TrapCode::OutOfFuel => "out of fuel",
|
||||
TrapCode::Unknown => "unknown trap",
|
||||
};
|
||||
write!(f, "{}", s)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::tools::wasm::error::{TrapCode, TrapInfo, WasmError};
|
||||
use crate::tools::wasm::error::WasmError;
|
||||
|
||||
#[test]
|
||||
fn test_error_display() {
|
||||
@@ -180,17 +114,6 @@ mod tests {
|
||||
assert!(err.to_string().contains("10000000"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_trap_info_display() {
|
||||
let info = TrapInfo {
|
||||
message: "access at offset 0x1000".to_string(),
|
||||
code: Some(TrapCode::MemoryOutOfBounds),
|
||||
};
|
||||
let s = info.to_string();
|
||||
assert!(s.contains("memory out of bounds"));
|
||||
assert!(s.contains("access at offset"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_conversion_to_tool_error() {
|
||||
let wasm_err = WasmError::Trapped("test trap".to_string());
|
||||
@@ -218,12 +141,11 @@ mod tests {
|
||||
fn test_tool_returned_error_with_hint() {
|
||||
let err = WasmError::ToolReturnedError {
|
||||
message: "unknown action: foobar".to_string(),
|
||||
hint: "Description: Gmail tool\nParameters schema: {\"type\":\"object\"}".to_string(),
|
||||
hint: "Tip: call tool_info(name: \"gmail\", include_schema: true) for the full parameter schema.".to_string(),
|
||||
};
|
||||
let display = err.to_string();
|
||||
assert!(display.contains("unknown action: foobar"));
|
||||
assert!(display.contains("Tool usage hint"));
|
||||
assert!(display.contains("Gmail tool"));
|
||||
assert!(display.contains("Parameters schema"));
|
||||
assert!(display.contains("tool_info"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -67,14 +67,8 @@ pub struct WasmResourceLimiter {
|
||||
memory_used: u64,
|
||||
/// Maximum tables allowed.
|
||||
max_tables: u32,
|
||||
/// Current table count.
|
||||
#[allow(dead_code)] // Reserved for table limit enforcement
|
||||
tables_created: u32,
|
||||
/// Maximum instances allowed.
|
||||
max_instances: u32,
|
||||
/// Current instance count.
|
||||
#[allow(dead_code)] // Reserved for instance limit enforcement
|
||||
instances_created: u32,
|
||||
}
|
||||
|
||||
impl WasmResourceLimiter {
|
||||
@@ -87,9 +81,7 @@ impl WasmResourceLimiter {
|
||||
memory_limit,
|
||||
memory_used: 0,
|
||||
max_tables: 10,
|
||||
tables_created: 0,
|
||||
max_instances: 10, // Component model needs multiple instances for WASI
|
||||
instances_created: 0,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -96,7 +96,7 @@ pub(crate) mod storage;
|
||||
mod wrapper;
|
||||
|
||||
// Core types
|
||||
pub use error::{TrapCode, TrapInfo, WasmError};
|
||||
pub use error::WasmError;
|
||||
pub use host::{HostState, LogEntry, LogLevel};
|
||||
pub use limits::{
|
||||
DEFAULT_FUEL_LIMIT, DEFAULT_MEMORY_LIMIT, DEFAULT_TIMEOUT, FuelConfig, ResourceLimits,
|
||||
|
||||
+26
-41
@@ -123,7 +123,9 @@ pub struct PreparedModule {
|
||||
pub name: String,
|
||||
/// Tool description (cached from component).
|
||||
pub description: String,
|
||||
/// Parameter schema JSON (cached from component).
|
||||
/// Full parameter schema JSON extracted from the component.
|
||||
/// Used for discovery and coercion, not necessarily for the compact
|
||||
/// schema advertised in the main tools array.
|
||||
pub schema: serde_json::Value,
|
||||
/// Pre-compiled component (cheaply cloneable via internal Arc).
|
||||
component: wasmtime::component::Component,
|
||||
@@ -265,11 +267,29 @@ impl WasmToolRuntime {
|
||||
let component = wasmtime::component::Component::new(&engine, &wasm_bytes)
|
||||
.map_err(|e| WasmError::CompilationFailed(e.to_string()))?;
|
||||
|
||||
// We need to instantiate briefly to extract metadata.
|
||||
// In a full implementation, we'd use WIT bindgen to get typed access.
|
||||
// For now, we extract what we can from the component.
|
||||
let description = extract_tool_description(&engine, &component)?;
|
||||
let schema = extract_tool_schema(&engine, &component)?;
|
||||
// Briefly instantiate to extract metadata (description + schema)
|
||||
// from the tool's exports, analogous to MCP's list_tools().
|
||||
let effective_limits = limits.clone().unwrap_or(default_limits.clone());
|
||||
let (description, schema) = crate::tools::wasm::wrapper::extract_wasm_metadata(
|
||||
&engine,
|
||||
&component,
|
||||
&effective_limits,
|
||||
)
|
||||
.unwrap_or_else(|e| {
|
||||
tracing::warn!(
|
||||
name = %name,
|
||||
error = %e,
|
||||
"WASM metadata extraction failed, using fallbacks"
|
||||
);
|
||||
(
|
||||
"WASM sandboxed tool".to_string(),
|
||||
serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {},
|
||||
"additionalProperties": true
|
||||
}),
|
||||
)
|
||||
});
|
||||
|
||||
Ok::<_, WasmError>(PreparedModule {
|
||||
name: name.clone(),
|
||||
@@ -321,41 +341,6 @@ impl WasmToolRuntime {
|
||||
}
|
||||
}
|
||||
|
||||
/// Extract tool description from a compiled component.
|
||||
///
|
||||
/// Returns a generic fallback. Callers should prefer loading the description
|
||||
/// from the sidecar `*.capabilities.json` file and overriding via
|
||||
/// `WasmToolWrapper::with_description()` or the `WasmToolRegistration::description` field.
|
||||
fn extract_tool_description(
|
||||
_engine: &Engine,
|
||||
_component: &wasmtime::component::Component,
|
||||
) -> Result<String, WasmError> {
|
||||
// WIT bindgen extraction is not yet implemented (see TODO #4 in CLAUDE.md).
|
||||
// Real descriptions come from the capabilities.json sidecar file, which is
|
||||
// loaded by the WasmToolLoader and passed as an override at registration time.
|
||||
Ok("WASM sandboxed tool".to_string())
|
||||
}
|
||||
|
||||
/// Extract tool parameter schema from a compiled component.
|
||||
///
|
||||
/// Returns a permissive fallback that accepts any JSON object. Callers should
|
||||
/// prefer loading the schema from the sidecar `*.capabilities.json` file and
|
||||
/// overriding via `WasmToolWrapper::with_schema()` or the
|
||||
/// `WasmToolRegistration::schema` field.
|
||||
fn extract_tool_schema(
|
||||
_engine: &Engine,
|
||||
_component: &wasmtime::component::Component,
|
||||
) -> Result<serde_json::Value, WasmError> {
|
||||
// WIT bindgen extraction is not yet implemented (see TODO #4 in CLAUDE.md).
|
||||
// Real schemas come from the capabilities.json sidecar file, which is
|
||||
// loaded by the WasmToolLoader and passed as an override at registration time.
|
||||
Ok(serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {},
|
||||
"additionalProperties": true
|
||||
}))
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for WasmToolRuntime {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("WasmToolRuntime")
|
||||
|
||||
+271
-53
@@ -464,9 +464,10 @@ pub struct WasmToolWrapper {
|
||||
/// Capabilities to grant to this tool.
|
||||
capabilities: Capabilities,
|
||||
/// Cached description (from PreparedModule or override).
|
||||
/// Stored without any tool_info hints — hints are composed at display time.
|
||||
description: String,
|
||||
/// Cached schema (from PreparedModule or override).
|
||||
schema: serde_json::Value,
|
||||
/// Compact and discovery schemas for this tool.
|
||||
schemas: WasmToolSchemas,
|
||||
/// Injected credentials for HTTP requests (e.g., OAuth tokens).
|
||||
/// Keys are placeholder names like "GOOGLE_ACCESS_TOKEN".
|
||||
credentials: HashMap<String, String>,
|
||||
@@ -477,6 +478,84 @@ pub struct WasmToolWrapper {
|
||||
oauth_refresh: Option<OAuthRefreshConfig>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
struct WasmToolSchemas {
|
||||
/// Compact schema advertised in the main tools array.
|
||||
///
|
||||
/// This stays permissive by default to avoid serializing full exported
|
||||
/// WASM schemas on every LLM call. Sidecars can override it explicitly.
|
||||
advertised: serde_json::Value,
|
||||
/// Full schema available for discovery and coercion.
|
||||
///
|
||||
/// Seeded from the WASM `schema()` export at registration time, unless a
|
||||
/// sidecar explicitly overrides it.
|
||||
discovery: serde_json::Value,
|
||||
}
|
||||
|
||||
impl WasmToolSchemas {
|
||||
fn permissive_schema() -> serde_json::Value {
|
||||
serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {},
|
||||
"additionalProperties": true
|
||||
})
|
||||
}
|
||||
|
||||
fn is_permissive_schema(schema: &serde_json::Value) -> bool {
|
||||
schema
|
||||
.get("properties")
|
||||
.and_then(|p| p.as_object())
|
||||
.is_none_or(|p| p.is_empty())
|
||||
}
|
||||
|
||||
fn new(discovery: serde_json::Value) -> Self {
|
||||
Self {
|
||||
advertised: Self::permissive_schema(),
|
||||
discovery,
|
||||
}
|
||||
}
|
||||
|
||||
fn with_override(&self, schema: serde_json::Value) -> Self {
|
||||
Self {
|
||||
advertised: schema.clone(),
|
||||
discovery: schema,
|
||||
}
|
||||
}
|
||||
|
||||
fn is_advertised_permissive(&self) -> bool {
|
||||
Self::is_permissive_schema(&self.advertised)
|
||||
}
|
||||
|
||||
fn advertised(&self) -> serde_json::Value {
|
||||
self.advertised.clone()
|
||||
}
|
||||
|
||||
fn discovery(&self) -> serde_json::Value {
|
||||
self.discovery.clone()
|
||||
}
|
||||
|
||||
/// Return the best schema available for type coercion.
|
||||
///
|
||||
/// Prefers the discovery schema when it has typed properties. Falls back
|
||||
/// to the `PreparedModule` schema extracted at load time rather than
|
||||
/// re-calling the WASM `schema()` export mid-execution, which could
|
||||
/// interact with mutable linear memory state.
|
||||
fn effective_for_coercion(&self, prepared_schema: &serde_json::Value) -> serde_json::Value {
|
||||
if !Self::is_permissive_schema(&self.discovery) {
|
||||
return self.discovery.clone();
|
||||
}
|
||||
|
||||
// Fall back to the load-time extracted schema from PreparedModule.
|
||||
// This avoids calling schema() on the already-running WASM instance
|
||||
// where mutable state could produce inconsistent results.
|
||||
if !Self::is_permissive_schema(prepared_schema) {
|
||||
return prepared_schema.clone();
|
||||
}
|
||||
|
||||
self.discovery.clone()
|
||||
}
|
||||
}
|
||||
|
||||
impl WasmToolWrapper {
|
||||
/// Create a new WASM tool wrapper.
|
||||
pub fn new(
|
||||
@@ -486,7 +565,7 @@ impl WasmToolWrapper {
|
||||
) -> Self {
|
||||
Self {
|
||||
description: prepared.description.clone(),
|
||||
schema: prepared.schema.clone(),
|
||||
schemas: WasmToolSchemas::new(prepared.schema.clone()),
|
||||
runtime,
|
||||
prepared,
|
||||
capabilities,
|
||||
@@ -504,7 +583,7 @@ impl WasmToolWrapper {
|
||||
|
||||
/// Override the parameter schema.
|
||||
pub fn with_schema(mut self, schema: serde_json::Value) -> Self {
|
||||
self.schema = schema;
|
||||
self.schemas = self.schemas.with_override(schema);
|
||||
self
|
||||
}
|
||||
|
||||
@@ -615,9 +694,18 @@ impl WasmToolWrapper {
|
||||
}
|
||||
})?;
|
||||
|
||||
// Get typed interface — used for execute.
|
||||
let tool_iface = instance.near_agent_tool();
|
||||
|
||||
// Determine effective schema for type coercion.
|
||||
// Prefer the discovery schema when typed; fall back to the load-time
|
||||
// extracted schema from PreparedModule rather than re-calling the WASM
|
||||
// export on the already-running instance.
|
||||
let effective_schema = self.schemas.effective_for_coercion(&self.prepared.schema);
|
||||
|
||||
// Coerce string-encoded values to their schema-declared types.
|
||||
// LLMs frequently pass numeric values as strings (e.g. "5" instead of 5).
|
||||
let params = coerce_params_to_schema(params, &self.schema);
|
||||
let params = coerce_params_to_schema(params, &effective_schema);
|
||||
|
||||
// Prepare the request
|
||||
let params_json = serde_json::to_string(¶ms)
|
||||
@@ -629,7 +717,6 @@ impl WasmToolWrapper {
|
||||
};
|
||||
|
||||
// Call execute using the generated typed interface
|
||||
let tool_iface = instance.near_agent_tool();
|
||||
let response = tool_iface.call_execute(&mut store, &request).map_err(|e| {
|
||||
let error_str = e.to_string();
|
||||
if error_str.contains("out of fuel") {
|
||||
@@ -644,12 +731,13 @@ impl WasmToolWrapper {
|
||||
// Get logs from host state
|
||||
let logs = store.data_mut().host_state.take_logs();
|
||||
|
||||
// Check for tool-level error — on failure, call the WASM module's
|
||||
// description() and schema() exports so the LLM can retry with the
|
||||
// correct parameters without us having to include the (large) schema
|
||||
// in every request's tools array.
|
||||
// Check for tool-level error — point the LLM to tool_info for the
|
||||
// full schema instead of dumping ~3.5KB inline.
|
||||
if let Some(err) = response.error {
|
||||
let hint = build_tool_hint(tool_iface, &mut store);
|
||||
let hint = format!(
|
||||
"Tip: call tool_info(name: \"{}\", include_schema: true) for the full parameter schema.",
|
||||
self.prepared.name
|
||||
);
|
||||
return Err(WasmError::ToolReturnedError { message: err, hint });
|
||||
}
|
||||
|
||||
@@ -658,47 +746,55 @@ impl WasmToolWrapper {
|
||||
}
|
||||
}
|
||||
|
||||
/// Maximum characters for the description portion of a tool hint.
|
||||
const HINT_DESC_MAX: usize = 500;
|
||||
/// Maximum characters for the schema portion of a tool hint.
|
||||
const HINT_SCHEMA_MAX: usize = 3000;
|
||||
/// Extract metadata (description + schema) from a WASM tool by briefly
|
||||
/// instantiating it and calling its `description()` and `schema()` exports.
|
||||
/// Analogous to MCP's `list_tools()` — discovers tool capabilities at load time.
|
||||
///
|
||||
/// Falls back to generic description and permissive schema on failure.
|
||||
pub(super) fn extract_wasm_metadata(
|
||||
engine: &wasmtime::Engine,
|
||||
component: &wasmtime::component::Component,
|
||||
limits: &ResourceLimits,
|
||||
) -> Result<(String, serde_json::Value), WasmError> {
|
||||
let store_data = StoreData::new(
|
||||
limits.memory_bytes,
|
||||
Capabilities::default(),
|
||||
HashMap::new(),
|
||||
vec![],
|
||||
);
|
||||
let mut store = Store::new(engine, store_data);
|
||||
|
||||
/// Call the WASM module's `description()` and `schema()` exports to build a
|
||||
/// hint string. Returns an empty string if both calls fail or return empty.
|
||||
/// Description is capped at [`HINT_DESC_MAX`] chars, schema at
|
||||
/// [`HINT_SCHEMA_MAX`] chars.
|
||||
fn build_tool_hint(tool_iface: &wit_tool::Guest, store: &mut Store<StoreData>) -> String {
|
||||
let desc = tool_iface
|
||||
.call_description(&mut *store)
|
||||
// Configure fuel + epoch deadline so extraction can't hang
|
||||
if let Err(e) = store.set_fuel(limits.fuel) {
|
||||
tracing::debug!("Fuel not enabled for metadata extraction: {e}");
|
||||
}
|
||||
store.epoch_deadline_trap();
|
||||
let ticks = (limits.timeout.as_millis() / EPOCH_TICK_INTERVAL.as_millis()).max(1) as u64;
|
||||
store.set_epoch_deadline(ticks);
|
||||
store.limiter(|data| &mut data.limiter);
|
||||
|
||||
// Instantiate with minimal linker
|
||||
let mut linker = Linker::new(engine);
|
||||
WasmToolWrapper::add_host_functions(&mut linker)?;
|
||||
let instance = SandboxedTool::instantiate(&mut store, component, &linker)
|
||||
.map_err(|e| WasmError::InstantiationFailed(e.to_string()))?;
|
||||
let tool_iface = instance.near_agent_tool();
|
||||
|
||||
// Extract description (fall back to generic)
|
||||
let description = tool_iface
|
||||
.call_description(&mut store)
|
||||
.unwrap_or_else(|_| "WASM sandboxed tool".to_string());
|
||||
|
||||
// Extract and parse schema (fall back to permissive)
|
||||
let schema = tool_iface
|
||||
.call_schema(&mut store)
|
||||
.ok()
|
||||
.unwrap_or_default();
|
||||
let schema = tool_iface.call_schema(&mut *store).ok().unwrap_or_default();
|
||||
if desc.is_empty() && schema.is_empty() {
|
||||
return String::new();
|
||||
}
|
||||
let mut hint = String::new();
|
||||
if !desc.is_empty() {
|
||||
hint.push_str("Description: ");
|
||||
if desc.len() > HINT_DESC_MAX {
|
||||
let end = crate::util::floor_char_boundary(&desc, HINT_DESC_MAX);
|
||||
hint.push_str(&desc[..end]);
|
||||
hint.push('…');
|
||||
} else {
|
||||
hint.push_str(&desc);
|
||||
}
|
||||
hint.push('\n');
|
||||
}
|
||||
if !schema.is_empty() {
|
||||
hint.push_str("Parameters schema: ");
|
||||
if schema.len() > HINT_SCHEMA_MAX {
|
||||
let end = crate::util::floor_char_boundary(&schema, HINT_SCHEMA_MAX);
|
||||
hint.push_str(&schema[..end]);
|
||||
hint.push('…');
|
||||
} else {
|
||||
hint.push_str(&schema);
|
||||
}
|
||||
}
|
||||
hint
|
||||
.and_then(|s| serde_json::from_str::<serde_json::Value>(&s).ok())
|
||||
.unwrap_or_else(|| {
|
||||
serde_json::json!({"type": "object", "properties": {}, "additionalProperties": true})
|
||||
});
|
||||
|
||||
Ok((description, schema))
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
@@ -712,7 +808,33 @@ impl Tool for WasmToolWrapper {
|
||||
}
|
||||
|
||||
fn parameters_schema(&self) -> serde_json::Value {
|
||||
self.schema.clone()
|
||||
self.schemas.advertised()
|
||||
}
|
||||
|
||||
fn discovery_schema(&self) -> serde_json::Value {
|
||||
self.schemas.discovery()
|
||||
}
|
||||
|
||||
/// Compose the tool schema for LLM function calling.
|
||||
///
|
||||
/// When the advertised schema is permissive (no typed properties), appends
|
||||
/// a hint to the description directing the LLM to call `tool_info` for the
|
||||
/// full parameter schema. This keeps the raw description clean while still
|
||||
/// guiding the LLM.
|
||||
fn schema(&self) -> crate::tools::tool::ToolSchema {
|
||||
let description = if self.schemas.is_advertised_permissive() {
|
||||
format!(
|
||||
"{} (call tool_info(name: \"{}\", include_schema: true) for parameter schema)",
|
||||
self.description, self.prepared.name
|
||||
)
|
||||
} else {
|
||||
self.description.clone()
|
||||
};
|
||||
crate::tools::tool::ToolSchema {
|
||||
name: self.prepared.name.clone(),
|
||||
description,
|
||||
parameters: self.schemas.advertised(),
|
||||
}
|
||||
}
|
||||
|
||||
async fn execute(
|
||||
@@ -749,7 +871,7 @@ impl Tool for WasmToolWrapper {
|
||||
let prepared = Arc::clone(&self.prepared);
|
||||
let capabilities = self.capabilities.clone();
|
||||
let description = self.description.clone();
|
||||
let schema = self.schema.clone();
|
||||
let schemas = self.schemas.clone();
|
||||
let credentials = self.credentials.clone();
|
||||
|
||||
// Execute in blocking task with timeout
|
||||
@@ -759,7 +881,7 @@ impl Tool for WasmToolWrapper {
|
||||
prepared,
|
||||
capabilities,
|
||||
description,
|
||||
schema,
|
||||
schemas,
|
||||
credentials,
|
||||
secrets_store: None, // Not needed in blocking task
|
||||
oauth_refresh: None, // Already used above for pre-refresh
|
||||
@@ -1232,6 +1354,7 @@ mod tests {
|
||||
TEST_GOOGLE_OAUTH_TOKEN, TEST_OAUTH_CLIENT_ID, TEST_OAUTH_CLIENT_SECRET,
|
||||
test_secrets_store,
|
||||
};
|
||||
use crate::tools::tool::Tool;
|
||||
use crate::tools::wasm::capabilities::Capabilities;
|
||||
use crate::tools::wasm::runtime::{WasmRuntimeConfig, WasmToolRuntime};
|
||||
|
||||
@@ -1246,6 +1369,84 @@ mod tests {
|
||||
assert!(runtime.config().fuel_config.enabled);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_advertised_schema_stays_permissive_until_sidecar_override() {
|
||||
let discovery_schema = serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"query": { "type": "string" },
|
||||
"limit": { "type": "integer" }
|
||||
},
|
||||
"required": ["query"]
|
||||
});
|
||||
|
||||
let runtime = Arc::new(WasmToolRuntime::new(WasmRuntimeConfig::for_testing()).unwrap());
|
||||
let prepared = runtime
|
||||
.prepare("search", b"\0asm\x0d\0\x01\0", None)
|
||||
.await
|
||||
.unwrap();
|
||||
let mut wrapper =
|
||||
super::WasmToolWrapper::new(Arc::clone(&runtime), prepared, Capabilities::default());
|
||||
wrapper.schemas = super::WasmToolSchemas::new(discovery_schema.clone());
|
||||
wrapper.description = "Search documents".to_string();
|
||||
|
||||
// Advertised schema stays permissive; discovery holds the typed schema
|
||||
assert_eq!(
|
||||
wrapper.parameters_schema(),
|
||||
serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {},
|
||||
"additionalProperties": true
|
||||
})
|
||||
);
|
||||
assert_eq!(wrapper.discovery_schema(), discovery_schema);
|
||||
|
||||
// Raw description is clean — no tool_info hint baked in
|
||||
assert!(!wrapper.description().contains("tool_info"));
|
||||
|
||||
// But schema() composes the hint at display time when advertised is permissive
|
||||
let schema = wrapper.schema();
|
||||
assert!(
|
||||
schema.description.contains("tool_info"),
|
||||
"schema().description should contain tool_info hint: {}",
|
||||
schema.description
|
||||
);
|
||||
assert!(
|
||||
schema.description.contains("include_schema: true"),
|
||||
"hint should mention include_schema: true: {}",
|
||||
schema.description
|
||||
);
|
||||
|
||||
// After sidecar override, both schemas match and hint disappears
|
||||
let wrapper = wrapper.with_schema(serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"query": { "type": "string" }
|
||||
},
|
||||
"required": ["query"]
|
||||
}));
|
||||
|
||||
assert_eq!(
|
||||
wrapper.parameters_schema(),
|
||||
serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"query": { "type": "string" }
|
||||
},
|
||||
"required": ["query"]
|
||||
})
|
||||
);
|
||||
assert_eq!(wrapper.discovery_schema(), wrapper.parameters_schema());
|
||||
|
||||
// With typed schema, schema() should NOT include tool_info hint
|
||||
let schema = wrapper.schema();
|
||||
assert!(
|
||||
!schema.description.contains("tool_info"),
|
||||
"schema().description should not contain tool_info hint when typed: {}",
|
||||
schema.description
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_capabilities_default() {
|
||||
let caps = Capabilities::default();
|
||||
@@ -1788,6 +1989,23 @@ mod tests {
|
||||
assert_eq!(result["count"], serde_json::json!("not-a-number"));
|
||||
}
|
||||
|
||||
/// Regression: permissive fallback schema (empty properties) must NOT coerce.
|
||||
/// This documents the bug where WASM tools with no sidecar `parameters` field
|
||||
/// got the permissive fallback, causing coercion to be a no-op and LLM-provided
|
||||
/// string integers to reach the WASM tool un-coerced.
|
||||
#[test]
|
||||
fn test_coerce_noop_with_permissive_schema() {
|
||||
let permissive = serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {},
|
||||
"additionalProperties": true
|
||||
});
|
||||
let params = serde_json::json!({"query": "test", "count": "10"});
|
||||
let result = super::coerce_params_to_schema(params, &permissive);
|
||||
// With empty properties, no coercion happens — string stays string
|
||||
assert_eq!(result["count"], serde_json::json!("10"));
|
||||
}
|
||||
|
||||
/// Regression test: leak scan must run on raw headers (before credential
|
||||
/// injection), not after. If it ran post-injection, the host-injected
|
||||
/// Slack bot token (`xoxb-...`) would trigger a Block and reject the
|
||||
|
||||
@@ -60,12 +60,18 @@ pub trait EmbeddingProvider: Send + Sync {
|
||||
}
|
||||
}
|
||||
|
||||
/// Default base URL for the OpenAI API.
|
||||
const OPENAI_API_BASE_URL: &str = "https://api.openai.com";
|
||||
|
||||
/// OpenAI embedding provider using text-embedding-ada-002 or text-embedding-3-small.
|
||||
///
|
||||
/// Supports any OpenAI-compatible embedding endpoint via [`with_base_url`](Self::with_base_url).
|
||||
pub struct OpenAiEmbeddings {
|
||||
client: reqwest::Client,
|
||||
api_key: String,
|
||||
model: String,
|
||||
dimension: usize,
|
||||
base_url: String,
|
||||
}
|
||||
|
||||
impl OpenAiEmbeddings {
|
||||
@@ -78,6 +84,7 @@ impl OpenAiEmbeddings {
|
||||
api_key: api_key.into(),
|
||||
model: "text-embedding-3-small".to_string(),
|
||||
dimension: 1536,
|
||||
base_url: OPENAI_API_BASE_URL.to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -88,6 +95,7 @@ impl OpenAiEmbeddings {
|
||||
api_key: api_key.into(),
|
||||
model: "text-embedding-ada-002".to_string(),
|
||||
dimension: 1536,
|
||||
base_url: OPENAI_API_BASE_URL.to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -98,6 +106,7 @@ impl OpenAiEmbeddings {
|
||||
api_key: api_key.into(),
|
||||
model: "text-embedding-3-large".to_string(),
|
||||
dimension: 3072,
|
||||
base_url: OPENAI_API_BASE_URL.to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -112,8 +121,35 @@ impl OpenAiEmbeddings {
|
||||
api_key: api_key.into(),
|
||||
model: model.into(),
|
||||
dimension,
|
||||
base_url: OPENAI_API_BASE_URL.to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Set a custom base URL for OpenAI-compatible embedding providers.
|
||||
///
|
||||
/// The URL must use `http://` or `https://` scheme. If no scheme is present,
|
||||
/// `https://` is prepended automatically. Trailing slashes are stripped.
|
||||
pub fn with_base_url(mut self, base_url: &str) -> Self {
|
||||
let url = base_url.trim();
|
||||
|
||||
// Auto-prepend https:// if no scheme is present.
|
||||
let mut url = if !url.starts_with("http://") && !url.starts_with("https://") {
|
||||
tracing::debug!(
|
||||
"No scheme in embedding base URL '{}', prepending https://",
|
||||
url
|
||||
);
|
||||
format!("https://{url}")
|
||||
} else {
|
||||
url.to_string()
|
||||
};
|
||||
|
||||
while url.ends_with('/') {
|
||||
url.pop();
|
||||
}
|
||||
|
||||
self.base_url = url;
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
@@ -173,9 +209,11 @@ impl EmbeddingProvider for OpenAiEmbeddings {
|
||||
input: texts,
|
||||
};
|
||||
|
||||
let url = format!("{}/v1/embeddings", self.base_url);
|
||||
|
||||
let response = self
|
||||
.client
|
||||
.post("https://api.openai.com/v1/embeddings")
|
||||
.post(&url)
|
||||
.header("Authorization", format!("Bearer {}", self.api_key))
|
||||
.json(&request)
|
||||
.send()
|
||||
@@ -575,9 +613,37 @@ mod tests {
|
||||
let provider = OpenAiEmbeddings::new("test-key");
|
||||
assert_eq!(provider.dimension(), 1536);
|
||||
assert_eq!(provider.model_name(), "text-embedding-3-small");
|
||||
assert_eq!(provider.base_url, OPENAI_API_BASE_URL);
|
||||
|
||||
let provider = OpenAiEmbeddings::large("test-key");
|
||||
assert_eq!(provider.dimension(), 3072);
|
||||
assert_eq!(provider.model_name(), "text-embedding-3-large");
|
||||
assert_eq!(provider.base_url, OPENAI_API_BASE_URL);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_openai_with_base_url_valid() {
|
||||
let provider =
|
||||
OpenAiEmbeddings::new("test-key").with_base_url("https://custom.example.com");
|
||||
assert_eq!(provider.base_url, "https://custom.example.com");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_openai_with_base_url_strips_trailing_slashes() {
|
||||
let provider =
|
||||
OpenAiEmbeddings::new("test-key").with_base_url("https://custom.example.com///");
|
||||
assert_eq!(provider.base_url, "https://custom.example.com");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_openai_with_base_url_http_scheme() {
|
||||
let provider = OpenAiEmbeddings::new("test-key").with_base_url("http://localhost:8080");
|
||||
assert_eq!(provider.base_url, "http://localhost:8080");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_openai_with_base_url_schemeless_prepends_https() {
|
||||
let provider = OpenAiEmbeddings::new("test-key").with_base_url("custom.example.com/v1");
|
||||
assert_eq!(provider.base_url, "https://custom.example.com/v1");
|
||||
}
|
||||
}
|
||||
|
||||
+19
-3
@@ -55,7 +55,9 @@ pub use embeddings::{
|
||||
};
|
||||
#[cfg(feature = "postgres")]
|
||||
pub use repository::Repository;
|
||||
pub use search::{RankedResult, SearchConfig, SearchResult, reciprocal_rank_fusion};
|
||||
pub use search::{
|
||||
FusionStrategy, RankedResult, SearchConfig, SearchResult, fuse_results, reciprocal_rank_fusion,
|
||||
};
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
@@ -332,6 +334,8 @@ pub struct Workspace {
|
||||
storage: WorkspaceStorage,
|
||||
/// Embedding provider for semantic search.
|
||||
embeddings: Option<Arc<dyn EmbeddingProvider>>,
|
||||
/// Default search configuration applied to all queries.
|
||||
search_defaults: SearchConfig,
|
||||
}
|
||||
|
||||
impl Workspace {
|
||||
@@ -343,6 +347,7 @@ impl Workspace {
|
||||
agent_id: None,
|
||||
storage: WorkspaceStorage::Repo(Repository::new(pool)),
|
||||
embeddings: None,
|
||||
search_defaults: SearchConfig::default(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -355,6 +360,7 @@ impl Workspace {
|
||||
agent_id: None,
|
||||
storage: WorkspaceStorage::Db(db),
|
||||
embeddings: None,
|
||||
search_defaults: SearchConfig::default(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -370,6 +376,16 @@ impl Workspace {
|
||||
self
|
||||
}
|
||||
|
||||
/// Set the default search configuration from workspace search config.
|
||||
pub fn with_search_config(mut self, config: &crate::config::WorkspaceSearchConfig) -> Self {
|
||||
self.search_defaults = SearchConfig::default()
|
||||
.with_fusion_strategy(config.fusion_strategy)
|
||||
.with_rrf_k(config.rrf_k)
|
||||
.with_fts_weight(config.fts_weight)
|
||||
.with_vector_weight(config.vector_weight);
|
||||
self
|
||||
}
|
||||
|
||||
/// Get the user ID.
|
||||
pub fn user_id(&self) -> &str {
|
||||
&self.user_id
|
||||
@@ -709,13 +725,13 @@ impl Workspace {
|
||||
/// Hybrid search across all memory documents.
|
||||
///
|
||||
/// Combines full-text search (BM25) with semantic search (vector similarity)
|
||||
/// using Reciprocal Rank Fusion (RRF).
|
||||
/// using the configured fusion strategy.
|
||||
pub async fn search(
|
||||
&self,
|
||||
query: &str,
|
||||
limit: usize,
|
||||
) -> Result<Vec<SearchResult>, WorkspaceError> {
|
||||
self.search_with_config(query, SearchConfig::default().with_limit(limit))
|
||||
self.search_with_config(query, self.search_defaults.clone().with_limit(limit))
|
||||
.await
|
||||
}
|
||||
|
||||
|
||||
@@ -12,7 +12,7 @@ use uuid::Uuid;
|
||||
use crate::error::WorkspaceError;
|
||||
|
||||
use crate::workspace::document::{MemoryChunk, MemoryDocument, WorkspaceEntry};
|
||||
use crate::workspace::search::{RankedResult, SearchConfig, SearchResult, reciprocal_rank_fusion};
|
||||
use crate::workspace::search::{RankedResult, SearchConfig, SearchResult, fuse_results};
|
||||
|
||||
/// Database repository for workspace operations.
|
||||
pub struct Repository {
|
||||
@@ -415,7 +415,7 @@ impl Repository {
|
||||
Vec::new()
|
||||
};
|
||||
|
||||
Ok(reciprocal_rank_fusion(fts_results, vector_results, config))
|
||||
Ok(fuse_results(fts_results, vector_results, config))
|
||||
}
|
||||
|
||||
/// Full-text search using PostgreSQL ts_rank_cd.
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user