mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-25 14:53:34 +00:00
Compare commits
57
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
413ecf8697 | ||
|
|
e1d364c636 | ||
|
|
7806273aa6 | ||
|
|
26d274ac79 | ||
|
|
b425213c53 | ||
|
|
37bba72397 | ||
|
|
2df9602d56 | ||
|
|
06c84a5c77 | ||
|
|
04c5c3fe9f | ||
|
|
a516e92156 | ||
|
|
de7f503df9 | ||
|
|
fe4c3c5fe6 | ||
|
|
14de4c1b57 | ||
|
|
2d332f12f0 | ||
|
|
46218ec794 | ||
|
|
6a2a6cd050 | ||
|
|
df49b17d0f | ||
|
|
c87525d81f | ||
|
|
9ae04f14e3 | ||
|
|
470de5bd2d | ||
|
|
69cddb10fd | ||
|
|
b4b19738a8 | ||
|
|
a1f0208956 | ||
|
|
3615967f92 | ||
|
|
704d63f16a | ||
|
|
902492bcdb | ||
|
|
13697976db | ||
|
|
cbcd5adcc0 | ||
|
|
e24c33ff90 | ||
|
|
f99991d27b | ||
|
|
89600e2b5c | ||
|
|
e4e78d8a87 | ||
|
|
b9446712e9 | ||
|
|
31a4330f24 | ||
|
|
9b47dbbaed | ||
|
|
ac3c928853 | ||
|
|
bf2a08be94 | ||
|
|
308758c27c | ||
|
|
f60c91e9a7 | ||
|
|
a22d44f2b2 | ||
|
|
a181c8b384 | ||
|
|
35a79caf87 | ||
|
|
85999b25a8 | ||
|
|
b60e5e907a | ||
|
|
d562dc8d90 | ||
|
|
c239a4fc2a | ||
|
|
944968bf76 | ||
|
|
18b59ae9a7 | ||
|
|
f4855962fc | ||
|
|
f18fb5173b | ||
|
|
78878ad7ef | ||
|
|
5f841554d5 | ||
|
|
6adf95b6d1 | ||
|
|
8530f44630 | ||
|
|
5257fecca1 | ||
|
|
20073ccf57 | ||
|
|
1a26b1e57f |
@@ -115,5 +115,12 @@ HEARTBEAT_NOTIFY_USER=default
|
||||
SAFETY_MAX_OUTPUT_LENGTH=100000
|
||||
SAFETY_INJECTION_CHECK_ENABLED=true
|
||||
|
||||
# Restart Feature (Docker containers only)
|
||||
# Set IRONCLAW_IN_DOCKER=true in the container entrypoint to enable the restart feature.
|
||||
# Without this, the restart tool and /restart command will be disabled.
|
||||
# IRONCLAW_IN_DOCKER=false
|
||||
# IRONCLAW_RESTART_DELAY=5 # default wait before exit (seconds, range: 1-30)
|
||||
# IRONCLAW_MAX_FAILURES=10 # max consecutive failures before container exits
|
||||
|
||||
# Logging
|
||||
RUST_LOG=ironclaw=debug,tower_http=debug
|
||||
|
||||
@@ -62,6 +62,9 @@ create "scope: ci" "546E7A" "CI/CD workflows"
|
||||
create "scope: docs" "78909C" "Documentation"
|
||||
create "scope: dependencies" "90A4AE" "Dependency updates"
|
||||
|
||||
echo "==> Creating workflow labels..."
|
||||
create "skip-regression-check" "9E9E9E" "Acknowledged: fix without regression test"
|
||||
|
||||
echo "==> Creating contributor labels..."
|
||||
create "contributor: new" "FFF9C4" "First-time contributor"
|
||||
create "contributor: regular" "FFE082" "2-5 merged PRs"
|
||||
|
||||
@@ -0,0 +1,198 @@
|
||||
name: Code Coverage
|
||||
on:
|
||||
push:
|
||||
branches: [main]
|
||||
|
||||
permissions:
|
||||
id-token: write
|
||||
contents: read
|
||||
|
||||
jobs:
|
||||
coverage:
|
||||
name: Coverage (${{ matrix.name }})
|
||||
runs-on: ubuntu-latest
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
include:
|
||||
- name: all-features
|
||||
flags: "--all-features"
|
||||
has_postgres: true
|
||||
- name: default
|
||||
flags: ""
|
||||
has_postgres: true
|
||||
- name: libsql-only
|
||||
flags: "--no-default-features --features libsql"
|
||||
has_postgres: false
|
||||
services:
|
||||
postgres:
|
||||
image: pgvector/pgvector:pg16
|
||||
env:
|
||||
POSTGRES_USER: postgres
|
||||
POSTGRES_PASSWORD: postgres
|
||||
POSTGRES_DB: ironclaw_test
|
||||
ports:
|
||||
- 5432:5432
|
||||
options: >-
|
||||
--health-cmd "pg_isready -U postgres"
|
||||
--health-interval 10s
|
||||
--health-timeout 5s
|
||||
--health-retries 5
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
|
||||
- uses: dtolnay/rust-toolchain@stable
|
||||
with:
|
||||
components: llvm-tools-preview
|
||||
targets: wasm32-wasip2
|
||||
|
||||
- uses: Swatinem/rust-cache@v2
|
||||
with:
|
||||
key: coverage-${{ matrix.name }}
|
||||
|
||||
- name: Install cargo-llvm-cov
|
||||
uses: taiki-e/install-action@cargo-llvm-cov
|
||||
|
||||
- name: Install cargo-component
|
||||
run: |
|
||||
if ! command -v cargo-component >/dev/null 2>&1; then
|
||||
cargo install cargo-component --locked
|
||||
fi
|
||||
|
||||
- name: Build WASM channels (for integration tests)
|
||||
run: ./scripts/build-wasm-extensions.sh --channels
|
||||
|
||||
- name: Run database migrations
|
||||
if: matrix.has_postgres
|
||||
run: |
|
||||
set -euo pipefail
|
||||
readarray -t migration_files < <(printf '%s\n' migrations/V*.sql | sort -V)
|
||||
for f in "${migration_files[@]}"; do
|
||||
echo "Applying $f..."
|
||||
psql -v ON_ERROR_STOP=1 -f "$f"
|
||||
done
|
||||
env:
|
||||
PGHOST: localhost
|
||||
PGUSER: postgres
|
||||
PGPASSWORD: postgres
|
||||
PGDATABASE: ironclaw_test
|
||||
|
||||
- name: Set DATABASE_URL for postgres configs
|
||||
if: matrix.has_postgres
|
||||
run: echo "DATABASE_URL=postgres://postgres:postgres@localhost/ironclaw_test" >> "$GITHUB_ENV"
|
||||
|
||||
- name: Generate coverage
|
||||
run: cargo llvm-cov ${{ matrix.flags }} --workspace --lcov --output-path lcov.info
|
||||
|
||||
- name: Upload to Codecov
|
||||
uses: codecov/codecov-action@v5
|
||||
with:
|
||||
files: lcov.info
|
||||
flags: ${{ matrix.name }}
|
||||
disable_search: true
|
||||
use_oidc: true
|
||||
fail_ci_if_error: true
|
||||
|
||||
e2e-coverage:
|
||||
name: E2E Coverage
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
|
||||
- uses: dtolnay/rust-toolchain@stable
|
||||
with:
|
||||
components: llvm-tools-preview
|
||||
targets: wasm32-wasip2
|
||||
|
||||
- uses: Swatinem/rust-cache@v2
|
||||
with:
|
||||
key: e2e-coverage
|
||||
|
||||
- name: Install cargo-llvm-cov
|
||||
uses: taiki-e/install-action@cargo-llvm-cov
|
||||
|
||||
- name: Install cargo-component
|
||||
run: |
|
||||
if ! command -v cargo-component >/dev/null 2>&1; then
|
||||
cargo install cargo-component --locked
|
||||
fi
|
||||
|
||||
- name: Build WASM channels
|
||||
run: ./scripts/build-wasm-extensions.sh --channels
|
||||
|
||||
- name: Set up coverage instrumentation
|
||||
run: |
|
||||
# show-env outputs shell-quoted values (KEY='value') but GITHUB_ENV
|
||||
# expects unquoted KEY=value. Strip only the wrapping single quotes
|
||||
# from KEY='value' lines without altering any internal characters.
|
||||
cargo llvm-cov show-env | sed -E "s/^([A-Za-z_][A-Za-z0-9_]*)='(.*)'$/\1=\2/" >> "$GITHUB_ENV"
|
||||
|
||||
- name: Clean coverage workspace
|
||||
run: cargo llvm-cov clean --workspace
|
||||
|
||||
- name: Build instrumented binary
|
||||
run: cargo build --no-default-features --features libsql
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Install E2E dependencies
|
||||
run: |
|
||||
cd tests/e2e
|
||||
pip install -e .
|
||||
playwright install --with-deps chromium
|
||||
|
||||
- name: Run E2E tests
|
||||
run: |
|
||||
pytest tests/e2e/ -v -x --timeout=120
|
||||
env:
|
||||
RUST_LOG: ironclaw=info
|
||||
RUST_BACKTRACE: "1"
|
||||
|
||||
- name: Verify profraw files exist
|
||||
if: always()
|
||||
run: |
|
||||
echo "LLVM_PROFILE_FILE=${LLVM_PROFILE_FILE}"
|
||||
echo "CARGO_LLVM_COV_TARGET_DIR=${CARGO_LLVM_COV_TARGET_DIR}"
|
||||
profraw_count=$(find target/ -name '*.profraw' 2>/dev/null | wc -l)
|
||||
echo "Found ${profraw_count} .profraw files under target/"
|
||||
find target/ -name '*.profraw' 2>/dev/null || true
|
||||
if [ "$profraw_count" -eq 0 ]; then
|
||||
echo "::warning::No .profraw files found — coverage report will fail"
|
||||
fi
|
||||
|
||||
- name: Generate coverage report
|
||||
if: always()
|
||||
run: cargo llvm-cov report --lcov --output-path e2e-coverage.info
|
||||
|
||||
- name: Upload to Codecov
|
||||
if: always()
|
||||
uses: codecov/codecov-action@v5
|
||||
with:
|
||||
files: e2e-coverage.info
|
||||
flags: e2e
|
||||
disable_search: true
|
||||
use_oidc: true
|
||||
fail_ci_if_error: true
|
||||
|
||||
- name: Upload screenshots on failure
|
||||
if: failure()
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: e2e-screenshots
|
||||
path: tests/e2e/screenshots/
|
||||
if-no-files-found: ignore
|
||||
|
||||
coverage-gate:
|
||||
name: Coverage
|
||||
runs-on: ubuntu-latest
|
||||
if: always()
|
||||
needs: [coverage, e2e-coverage]
|
||||
steps:
|
||||
- run: |
|
||||
if [[ "${{ needs.coverage.result }}" != "success" || "${{ needs.e2e-coverage.result }}" != "success" ]]; then
|
||||
echo "One or more coverage jobs failed"
|
||||
exit 1
|
||||
fi
|
||||
@@ -9,8 +9,9 @@ on:
|
||||
- "tests/e2e/**"
|
||||
|
||||
jobs:
|
||||
e2e:
|
||||
name: Browser E2E
|
||||
# ── Step 1: compile once ──────────────────────────────────────────────────
|
||||
build:
|
||||
name: Build ironclaw (libsql)
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
steps:
|
||||
@@ -25,9 +26,44 @@ jobs:
|
||||
~/.cargo/registry
|
||||
key: e2e-${{ runner.os }}-${{ hashFiles('Cargo.lock') }}
|
||||
|
||||
- name: Build ironclaw (libsql)
|
||||
- name: Build
|
||||
run: cargo build --no-default-features --features libsql
|
||||
|
||||
- name: Upload binary
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: ironclaw-e2e-binary
|
||||
path: target/debug/ironclaw
|
||||
retention-days: 1
|
||||
|
||||
# ── Step 2: run test slices in parallel ───────────────────────────────────
|
||||
test:
|
||||
name: E2E (${{ matrix.group }})
|
||||
needs: build
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
strategy:
|
||||
fail-fast: false
|
||||
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"
|
||||
- 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"
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
|
||||
- name: Download binary
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: ironclaw-e2e-binary
|
||||
path: target/debug/
|
||||
|
||||
- name: Make binary executable
|
||||
run: chmod +x target/debug/ironclaw
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
@@ -38,13 +74,26 @@ jobs:
|
||||
pip install -e .
|
||||
playwright install --with-deps chromium
|
||||
|
||||
- name: Run E2E tests
|
||||
run: pytest tests/e2e/ -v -x --timeout=120
|
||||
- name: Run E2E tests (${{ matrix.group }})
|
||||
run: pytest ${{ matrix.files }} -v --timeout=120
|
||||
|
||||
- name: Upload screenshots on failure
|
||||
if: failure()
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: e2e-screenshots
|
||||
name: e2e-screenshots-${{ matrix.group }}
|
||||
path: tests/e2e/screenshots/
|
||||
if-no-files-found: ignore
|
||||
|
||||
# ── Roll-up for branch protection ────────────────────────────────────────
|
||||
e2e:
|
||||
name: E2E Tests
|
||||
runs-on: ubuntu-latest
|
||||
if: always()
|
||||
needs: [test]
|
||||
steps:
|
||||
- run: |
|
||||
if [[ "${{ needs.test.result }}" != "success" ]]; then
|
||||
echo "One or more E2E jobs failed"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
@@ -0,0 +1,107 @@
|
||||
name: Regression Test Check
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
|
||||
jobs:
|
||||
regression-test:
|
||||
name: Regression test enforcement
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 0
|
||||
|
||||
- name: Check for regression tests
|
||||
env:
|
||||
PR_TITLE: ${{ github.event.pull_request.title }}
|
||||
PR_LABELS: ${{ join(github.event.pull_request.labels.*.name, ',') }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
|
||||
BASE_REF="origin/${{ github.event.pull_request.base.ref }}"
|
||||
|
||||
# --- 1. Is this a fix PR? Check title first, then commit messages ---
|
||||
IS_FIX=false
|
||||
|
||||
if grep -qiE '^(fix(\(.*\))?|hotfix|bugfix):' <<< "$PR_TITLE"; then
|
||||
IS_FIX=true
|
||||
fi
|
||||
|
||||
if [ "$IS_FIX" = false ]; then
|
||||
COMMITS=$(git log --format='%s' "${BASE_REF}..HEAD")
|
||||
if grep -qiE '^(fix(\(.*\))?|hotfix|bugfix):' <<< "$COMMITS"; then
|
||||
IS_FIX=true
|
||||
fi
|
||||
fi
|
||||
|
||||
if [ "$IS_FIX" = false ]; then
|
||||
echo "Not a fix PR — skipping regression test check."
|
||||
exit 0
|
||||
fi
|
||||
|
||||
echo "Fix PR detected."
|
||||
|
||||
# --- 2. Skip label or commit message marker ---
|
||||
if grep -qF ',skip-regression-check,' <<< ",$PR_LABELS,"; then
|
||||
echo "skip-regression-check label present — skipping."
|
||||
exit 0
|
||||
fi
|
||||
|
||||
COMMIT_BODIES=$(git log --format='%B' "${BASE_REF}..HEAD")
|
||||
if grep -qF '[skip-regression-check]' <<< "$COMMIT_BODIES"; then
|
||||
echo "[skip-regression-check] found in commit message — skipping."
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# --- 3. Exempt static-only / docs-only changes ---
|
||||
CHANGED_FILES=$(git diff --name-only "${BASE_REF}...HEAD")
|
||||
|
||||
if [ -z "$CHANGED_FILES" ]; then
|
||||
echo "No changed files — skipping."
|
||||
exit 0
|
||||
fi
|
||||
|
||||
ALL_EXEMPT=true
|
||||
while IFS= read -r file; do
|
||||
case "$file" in
|
||||
src/channels/web/static/*) ;;
|
||||
*.md) ;;
|
||||
*) ALL_EXEMPT=false; break ;;
|
||||
esac
|
||||
done <<< "$CHANGED_FILES"
|
||||
|
||||
if [ "$ALL_EXEMPT" = true ]; then
|
||||
echo "All changes are static assets or docs — skipping."
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# --- 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
|
||||
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 (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 }
|
||||
/^\+[^+]/ { has_add=1 }
|
||||
END { if (has_test && has_add) found=1; exit !found }
|
||||
'; then
|
||||
echo "Test changes found in existing test functions."
|
||||
exit 0
|
||||
fi
|
||||
|
||||
if grep -qE '^tests/' <<< "$CHANGED_FILES"; then
|
||||
echo "Test file changes found under tests/."
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# --- 5. No tests found ---
|
||||
echo "::warning::This PR looks like a bug fix but contains no test changes. Every fix should include a regression test. Add a #[test] or #[tokio::test], or apply the 'skip-regression-check' label if not feasible."
|
||||
exit 1
|
||||
@@ -413,6 +413,9 @@ jobs:
|
||||
- build-wasm-extensions
|
||||
if: ${{ always() && needs.host.result == 'success' && needs.build-wasm-extensions.result == 'success' }}
|
||||
runs-on: "ubuntu-22.04"
|
||||
permissions:
|
||||
contents: write
|
||||
pull-requests: write
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
steps:
|
||||
@@ -445,7 +448,7 @@ jobs:
|
||||
fi
|
||||
done
|
||||
done < "$CHECKSUMS"
|
||||
- name: Commit updated manifests
|
||||
- name: Create PR with updated manifests
|
||||
run: |
|
||||
git config user.name "github-actions[bot]"
|
||||
git config user.email "github-actions[bot]@users.noreply.github.com"
|
||||
@@ -453,8 +456,15 @@ jobs:
|
||||
if git diff --cached --quiet; then
|
||||
echo "No manifest changes to commit"
|
||||
else
|
||||
BRANCH="chore/update-checksums-$(date +%s)"
|
||||
git checkout -b "$BRANCH"
|
||||
git commit -m "chore: update WASM artifact SHA256 checksums [skip ci]"
|
||||
git push
|
||||
git push origin "$BRANCH"
|
||||
gh pr create \
|
||||
--title "chore: update WASM artifact SHA256 checksums" \
|
||||
--body "Auto-generated by release CI. Updates SHA256 checksums in registry manifests to match the released WASM artifacts." \
|
||||
--base main \
|
||||
--head "$BRANCH"
|
||||
fi
|
||||
|
||||
announce:
|
||||
|
||||
@@ -26,9 +26,14 @@ jobs:
|
||||
uses: dtolnay/rust-toolchain@stable
|
||||
with:
|
||||
profile: minimal
|
||||
targets: wasm32-wasip2
|
||||
- uses: Swatinem/rust-cache@v2
|
||||
with:
|
||||
key: ${{ matrix.name }}
|
||||
- name: Install cargo-component
|
||||
run: cargo install cargo-component --locked || true
|
||||
- name: Build WASM channels (for integration tests)
|
||||
run: ./scripts/build-wasm-extensions.sh --channels
|
||||
- name: Run Tests
|
||||
run: cargo test ${{ matrix.flags }} -- --nocapture
|
||||
|
||||
@@ -46,6 +51,27 @@ jobs:
|
||||
- name: Run Telegram Channel Tests
|
||||
run: cargo test --manifest-path channels-src/telegram/Cargo.toml -- --nocapture
|
||||
|
||||
wasm-wit-compat:
|
||||
name: WASM WIT Compatibility
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@v6
|
||||
- name: Install Rust
|
||||
uses: dtolnay/rust-toolchain@stable
|
||||
with:
|
||||
profile: minimal
|
||||
targets: wasm32-wasip2
|
||||
- uses: Swatinem/rust-cache@v2
|
||||
with:
|
||||
key: wasm-extensions
|
||||
- name: Install cargo-component
|
||||
run: cargo install cargo-component --locked || true
|
||||
- name: Build all WASM extensions against current WIT
|
||||
run: ./scripts/build-wasm-extensions.sh
|
||||
- name: Instantiation test (host linker compatibility)
|
||||
run: cargo test --all-features wit_compat -- --nocapture
|
||||
|
||||
docker-build:
|
||||
name: Docker Build
|
||||
runs-on: ubuntu-latest
|
||||
@@ -55,15 +81,34 @@ jobs:
|
||||
- name: Build Docker image
|
||||
run: docker build -t ironclaw-test:ci .
|
||||
|
||||
version-check:
|
||||
name: Version Bump Check
|
||||
runs-on: ubuntu-latest
|
||||
if: github.event_name == 'pull_request'
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@v6
|
||||
with:
|
||||
fetch-depth: 0
|
||||
- name: Check version bumps for changed extensions
|
||||
env:
|
||||
PR_LABELS: ${{ join(github.event.pull_request.labels.*.name, ',') }}
|
||||
run: ./scripts/check-version-bumps.sh
|
||||
|
||||
# Roll-up job for branch protection
|
||||
run-tests:
|
||||
name: Run Tests
|
||||
runs-on: ubuntu-latest
|
||||
if: always()
|
||||
needs: [tests, telegram-tests, docker-build]
|
||||
needs: [tests, telegram-tests, wasm-wit-compat, docker-build, version-check]
|
||||
steps:
|
||||
- run: |
|
||||
if [[ "${{ needs.tests.result }}" != "success" || "${{ needs.telegram-tests.result }}" != "success" || "${{ needs.docker-build.result }}" != "success" ]]; then
|
||||
if [[ "${{ needs.tests.result }}" != "success" || "${{ needs.telegram-tests.result }}" != "success" || "${{ needs.wasm-wit-compat.result }}" != "success" || "${{ needs.docker-build.result }}" != "success" ]]; then
|
||||
echo "One or more jobs failed"
|
||||
exit 1
|
||||
fi
|
||||
# version-check only runs on PRs, so skip/success are both acceptable
|
||||
if [[ "${{ needs.version-check.result }}" == "failure" ]]; then
|
||||
echo "Version bump check failed"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
@@ -16,6 +16,9 @@ target/
|
||||
# Benchmark results (local runs, not committed)
|
||||
bench-results/
|
||||
|
||||
# Coverage reports (local runs, not committed)
|
||||
/coverage/
|
||||
|
||||
# WASM build artifacts (loaded from disk, not bundled)
|
||||
*.wasm
|
||||
|
||||
|
||||
@@ -7,6 +7,97 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
||||
|
||||
## [Unreleased]
|
||||
|
||||
## [0.16.0](https://github.com/nearai/ironclaw/compare/v0.15.0...v0.16.0) - 2026-03-06
|
||||
|
||||
### Added
|
||||
|
||||
- *(e2e)* extensions tab tests, CI parallelization, and 3 production bug fixes ([#584](https://github.com/nearai/ironclaw/pull/584))
|
||||
- WASM extension versioning with WIT compat checks ([#592](https://github.com/nearai/ironclaw/pull/592))
|
||||
- Add HMAC-SHA256 webhook signature validation for Slack ([#588](https://github.com/nearai/ironclaw/pull/588))
|
||||
- restart ([#531](https://github.com/nearai/ironclaw/pull/531))
|
||||
- merge http/web_fetch tools, add tool output stash for large responses ([#578](https://github.com/nearai/ironclaw/pull/578))
|
||||
- integrate 13-dimension complexity scorer into smart routing ([#529](https://github.com/nearai/ironclaw/pull/529))
|
||||
|
||||
### Fixed
|
||||
|
||||
- *(llm)* fix reasoning model response parsing bugs ([#564](https://github.com/nearai/ironclaw/pull/564)) ([#580](https://github.com/nearai/ironclaw/pull/580))
|
||||
- *(ci)* fix three coverage workflow failures ([#597](https://github.com/nearai/ironclaw/pull/597))
|
||||
- Telegram channel accepts group messages from all users if owner_… ([#590](https://github.com/nearai/ironclaw/pull/590))
|
||||
- *(ci)* anchor coverage/ gitignore rule to repo root ([#591](https://github.com/nearai/ironclaw/pull/591))
|
||||
- *(security)* use OsRng for all security-critical key and token generation ([#519](https://github.com/nearai/ironclaw/pull/519))
|
||||
- prevent concurrent memory hygiene passes and Windows file lock errors ([#535](https://github.com/nearai/ironclaw/pull/535))
|
||||
- sort tool_definitions() for deterministic LLM tool ordering ([#582](https://github.com/nearai/ironclaw/pull/582))
|
||||
- *(ci)* persist all cargo-llvm-cov env vars for E2E coverage ([#559](https://github.com/nearai/ironclaw/pull/559))
|
||||
|
||||
### Other
|
||||
|
||||
- *(llm)* complete response cache — set_model invalidation, stats logging, sync mutex ([#290](https://github.com/nearai/ironclaw/pull/290))
|
||||
- add 29 E2E trace tests for issues #571-575 ([#593](https://github.com/nearai/ironclaw/pull/593))
|
||||
- add 26 tests for multi-thread safety, db CRUD, concurrency, errors ([#442](https://github.com/nearai/ironclaw/pull/442))
|
||||
- update WASM artifact SHA256 checksums [skip ci] ([#560](https://github.com/nearai/ironclaw/pull/560))
|
||||
- add WIT compatibility tests for WASM extensions ([#586](https://github.com/nearai/ironclaw/pull/586))
|
||||
- Trajectory benchmarks and e2e trace test rig ([#553](https://github.com/nearai/ironclaw/pull/553))
|
||||
|
||||
## [0.15.0](https://github.com/nearai/ironclaw/compare/v0.14.0...v0.15.0) - 2026-03-04
|
||||
|
||||
### Added
|
||||
|
||||
- *(oauth)* route callbacks through web gateway for hosted instances ([#555](https://github.com/nearai/ironclaw/pull/555))
|
||||
- *(web)* show error details for failed tool calls ([#490](https://github.com/nearai/ironclaw/pull/490))
|
||||
- *(extensions)* improve auth UX and add load-time validation ([#536](https://github.com/nearai/ironclaw/pull/536))
|
||||
- add local-test skill and Dockerfile.test for web gateway testing ([#524](https://github.com/nearai/ironclaw/pull/524))
|
||||
|
||||
### Fixed
|
||||
|
||||
- *(security)* restrict query-token auth to SSE endpoints only ([#528](https://github.com/nearai/ironclaw/pull/528))
|
||||
- *(ci)* flush profraw coverage data in E2E teardown ([#550](https://github.com/nearai/ironclaw/pull/550))
|
||||
- *(wasm)* coerce string parameters to schema-declared types ([#498](https://github.com/nearai/ironclaw/pull/498))
|
||||
- *(agent)* strip leaked [Called tool ...] text from responses ([#497](https://github.com/nearai/ironclaw/pull/497))
|
||||
- *(web)* reset job list UI on restart failure ([#499](https://github.com/nearai/ironclaw/pull/499))
|
||||
- *(security)* replace .unwrap() panics in pairing store with proper error handling ([#515](https://github.com/nearai/ironclaw/pull/515))
|
||||
|
||||
### Other
|
||||
|
||||
- Fix UTF-8 unsafe truncation in sandbox log capture ([#359](https://github.com/nearai/ironclaw/pull/359))
|
||||
- enhance coverage with feature matrix, postgres, and E2E ([#523](https://github.com/nearai/ironclaw/pull/523))
|
||||
|
||||
## [0.14.0](https://github.com/nearai/ironclaw/compare/v0.13.1...v0.14.0) - 2026-03-04
|
||||
|
||||
### Added
|
||||
|
||||
- remove the okta tool ([#506](https://github.com/nearai/ironclaw/pull/506))
|
||||
- add OAuth support for WASM tools in web gateway ([#489](https://github.com/nearai/ironclaw/pull/489))
|
||||
- *(web)* fix jobs UI parity for non-sandbox mode ([#491](https://github.com/nearai/ironclaw/pull/491))
|
||||
- *(workspace)* add TOOLS.md, BOOTSTRAP.md, and disk-to-DB import ([#477](https://github.com/nearai/ironclaw/pull/477))
|
||||
|
||||
### Fixed
|
||||
|
||||
- *(web)* mobile browser bar obscures chat input ([#508](https://github.com/nearai/ironclaw/pull/508))
|
||||
- *(web)* assign unique thread_id to manual routine triggers ([#500](https://github.com/nearai/ironclaw/pull/500))
|
||||
- *(web)* refresh routine UI after Run Now trigger ([#501](https://github.com/nearai/ironclaw/pull/501))
|
||||
- *(skills)* use slug for skill download URL from ClawHub ([#502](https://github.com/nearai/ironclaw/pull/502))
|
||||
- *(workspace)* thread document path through search results ([#503](https://github.com/nearai/ironclaw/pull/503))
|
||||
- *(workspace)* import custom templates before seeding defaults ([#505](https://github.com/nearai/ironclaw/pull/505))
|
||||
- use std::sync::RwLock in MessageTool to avoid runtime panic ([#411](https://github.com/nearai/ironclaw/pull/411))
|
||||
- wire secrets store into all WASM runtime activation paths ([#479](https://github.com/nearai/ironclaw/pull/479))
|
||||
|
||||
### Other
|
||||
|
||||
- enforce regression tests for fix commits ([#517](https://github.com/nearai/ironclaw/pull/517))
|
||||
- add code coverage with cargo-llvm-cov and Codecov ([#511](https://github.com/nearai/ironclaw/pull/511))
|
||||
- Remove restart infrastructure, generalize WASM channel setup ([#493](https://github.com/nearai/ironclaw/pull/493))
|
||||
|
||||
## [0.13.1](https://github.com/nearai/ironclaw/compare/v0.13.0...v0.13.1) - 2026-03-02
|
||||
|
||||
### Added
|
||||
|
||||
- add Brave Web Search WASM tool ([#474](https://github.com/nearai/ironclaw/pull/474))
|
||||
|
||||
### Fixed
|
||||
|
||||
- *(web)* auto-scroll and Enter key completion for slash command autocomplete ([#475](https://github.com/nearai/ironclaw/pull/475))
|
||||
- correct download URLs for telegram-mtproto and slack-tool extensions ([#470](https://github.com/nearai/ironclaw/pull/470))
|
||||
|
||||
## [0.13.0](https://github.com/nearai/ironclaw/compare/v0.12.0...v0.13.0) - 2026-03-02
|
||||
|
||||
### Added
|
||||
|
||||
@@ -321,6 +321,8 @@ cargo check --all-features # all features
|
||||
```
|
||||
Dead code behind the wrong `#[cfg]` gate will only show up when building with a single feature.
|
||||
|
||||
**Regression test with every fix:** Every bug fix must include a test that would have caught the bug. Add a `#[test]` or `#[tokio::test]` that reproduces the original failure. Exempt: changes limited to `src/channels/web/static/` or `.md` files. Use `[skip-regression-check]` in commit message or PR label if genuinely not feasible. The `commit-msg` hook and CI workflow enforce this automatically.
|
||||
|
||||
**Zero clippy warnings policy:** Fix ALL clippy warnings before committing, including pre-existing ones in files you didn't change. Never leave warnings behind — treat `cargo clippy` output as a zero-tolerance gate.
|
||||
|
||||
**Mechanical verification before committing:** Run these checks on changed files before committing:
|
||||
@@ -328,6 +330,7 @@ Dead code behind the wrong `#[cfg]` gate will only show up when building with a
|
||||
- `grep -rnE '\.unwrap\(|\.expect\(' <files>` -- no panics in production
|
||||
- `grep -rn 'super::' <files>` -- use `crate::` imports
|
||||
- If you fixed a pattern bug, `grep` for other instances of that pattern across `src/`
|
||||
- Fix commits must include regression tests (enforced by `commit-msg` hook; bypass with `[skip-regression-check]`)
|
||||
|
||||
## Configuration
|
||||
|
||||
|
||||
Generated
+25
-1
@@ -2828,7 +2828,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "ironclaw"
|
||||
version = "0.13.0"
|
||||
version = "0.16.0"
|
||||
dependencies = [
|
||||
"aes-gcm",
|
||||
"aho-corasick",
|
||||
@@ -2853,6 +2853,7 @@ dependencies = [
|
||||
"futures",
|
||||
"hex",
|
||||
"hkdf",
|
||||
"hmac",
|
||||
"html-to-markdown-rs",
|
||||
"http-body-util",
|
||||
"hyper 1.8.1",
|
||||
@@ -2879,6 +2880,7 @@ dependencies = [
|
||||
"secrecy",
|
||||
"secret-service",
|
||||
"security-framework",
|
||||
"semver",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"serde_yml",
|
||||
@@ -2900,6 +2902,7 @@ dependencies = [
|
||||
"tower-http 0.6.8",
|
||||
"tracing",
|
||||
"tracing-subscriber",
|
||||
"tracing-test",
|
||||
"url",
|
||||
"urlencoding",
|
||||
"uuid",
|
||||
@@ -6227,6 +6230,27 @@ dependencies = [
|
||||
"tracing-serde",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tracing-test"
|
||||
version = "0.2.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "19a4c448db514d4f24c5ddb9f73f2ee71bfb24c526cf0c570ba142d1119e0051"
|
||||
dependencies = [
|
||||
"tracing-core",
|
||||
"tracing-subscriber",
|
||||
"tracing-test-macro",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "tracing-test-macro"
|
||||
version = "0.2.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "ad06847b7afb65c7866a36664b75c40b895e318cea4f71299f013fb22965329d"
|
||||
dependencies = [
|
||||
"quote",
|
||||
"syn 2.0.117",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "try-lock"
|
||||
version = "0.2.5"
|
||||
|
||||
+6
-2
@@ -12,14 +12,13 @@ exclude = [
|
||||
"tools-src/google-drive",
|
||||
"tools-src/google-sheets",
|
||||
"tools-src/google-slides",
|
||||
"tools-src/okta",
|
||||
"tools-src/slack",
|
||||
"tools-src/telegram",
|
||||
]
|
||||
|
||||
[package]
|
||||
name = "ironclaw"
|
||||
version = "0.13.0"
|
||||
version = "0.16.0"
|
||||
edition = "2024"
|
||||
rust-version = "1.92"
|
||||
description = "Secure personal AI assistant that protects your data and expands its capabilities on the fly"
|
||||
@@ -107,6 +106,9 @@ serde_yml = "0.0.12"
|
||||
dirs = "6"
|
||||
fs4 = "0.6"
|
||||
|
||||
# Semantic versioning
|
||||
semver = "1"
|
||||
|
||||
# Secrecy for sensitive values
|
||||
secrecy = { version = "0.10", features = ["serde"] }
|
||||
|
||||
@@ -129,6 +131,7 @@ wasmparser = "0.220" # WASM binary parsing for validation
|
||||
# Cryptography for secrets management
|
||||
aes-gcm = "0.10"
|
||||
hkdf = "0.12"
|
||||
hmac = "0.12"
|
||||
sha2 = "0.10"
|
||||
blake3 = "1"
|
||||
rand = "0.8"
|
||||
@@ -171,6 +174,7 @@ zbus = "4"
|
||||
|
||||
[dev-dependencies]
|
||||
tokio-test = "0.4"
|
||||
tracing-test = "0.2"
|
||||
tokio-tungstenite = "0.26"
|
||||
testcontainers-modules = { version = "0.11", features = ["postgres"] }
|
||||
pretty_assertions = "1"
|
||||
|
||||
@@ -0,0 +1,57 @@
|
||||
# Lightweight test Dockerfile for IronClaw web gateway testing.
|
||||
#
|
||||
# Build:
|
||||
# docker build --platform linux/amd64 -f Dockerfile.test -t ironclaw-test .
|
||||
#
|
||||
# Run (each on a different port):
|
||||
# docker run --rm -p 3003:3003 ironclaw-test
|
||||
# docker run --rm -p 3004:3003 ironclaw-test
|
||||
# docker run --rm -p 3005:3003 ironclaw-test
|
||||
|
||||
# Stage 1: Build (libsql only — no PostgreSQL dependency)
|
||||
FROM rust:1.92-slim-bookworm AS builder
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
pkg-config libssl-dev cmake gcc g++ \
|
||||
&& rm -rf /var/lib/apt/lists/* \
|
||||
&& rustup target add wasm32-wasip2 \
|
||||
&& cargo install wasm-tools
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
COPY Cargo.toml Cargo.lock ./
|
||||
COPY build.rs build.rs
|
||||
COPY src/ src/
|
||||
COPY tests/ tests/
|
||||
COPY migrations/ migrations/
|
||||
COPY registry/ registry/
|
||||
COPY channels-src/ channels-src/
|
||||
COPY wit/ wit/
|
||||
|
||||
RUN cargo build --release --no-default-features --features libsql --bin ironclaw
|
||||
|
||||
# Stage 2: Runtime
|
||||
FROM debian:bookworm-slim
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
ca-certificates libssl3 \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
COPY --from=builder /app/target/release/ironclaw /usr/local/bin/ironclaw
|
||||
|
||||
RUN useradd -m -u 1000 -s /bin/bash ironclaw
|
||||
USER ironclaw
|
||||
WORKDIR /home/ironclaw
|
||||
|
||||
EXPOSE 3003
|
||||
|
||||
ENV RUST_LOG=ironclaw=info \
|
||||
GATEWAY_ENABLED=true \
|
||||
GATEWAY_HOST=0.0.0.0 \
|
||||
GATEWAY_PORT=3003 \
|
||||
GATEWAY_AUTH_TOKEN=test \
|
||||
DATABASE_BACKEND=libsql \
|
||||
LIBSQL_PATH=/home/ironclaw/test.db \
|
||||
SANDBOX_ENABLED=false
|
||||
|
||||
ENTRYPOINT ["ironclaw", "--no-onboard"]
|
||||
@@ -1,4 +1,6 @@
|
||||
{
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"type": "channel",
|
||||
"name": "discord",
|
||||
"description": "Discord Gateway/Webhook channel for handling slash commands, buttons, and messages",
|
||||
@@ -6,15 +8,16 @@
|
||||
"required_secrets": [
|
||||
{
|
||||
"name": "discord_bot_token",
|
||||
"prompt": "Enter your Discord Bot Token (from Developer Portal)",
|
||||
"prompt": "Enter your Discord Bot Token. Find it under Bot > Token in your Discord Application settings.",
|
||||
"optional": false
|
||||
},
|
||||
{
|
||||
"name": "discord_public_key",
|
||||
"prompt": "Enter your Discord Application Public Key (from Developer Portal > General Information)",
|
||||
"prompt": "Enter your Discord Application Public Key (found under General Information in your Discord Application settings).",
|
||||
"optional": false
|
||||
}
|
||||
]
|
||||
],
|
||||
"setup_url": "https://discord.com/developers/applications"
|
||||
},
|
||||
"capabilities": {
|
||||
"http": {
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
{
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"type": "channel",
|
||||
"name": "slack",
|
||||
"description": "Slack Events API channel for receiving and responding to Slack messages",
|
||||
@@ -6,15 +8,16 @@
|
||||
"required_secrets": [
|
||||
{
|
||||
"name": "slack_bot_token",
|
||||
"prompt": "Enter your Slack Bot OAuth Token (xoxb-...)",
|
||||
"prompt": "Enter your Slack Bot User OAuth Token (starts with xoxb-). Find it under OAuth & Permissions in your Slack App settings.",
|
||||
"optional": false
|
||||
},
|
||||
{
|
||||
"name": "slack_signing_secret",
|
||||
"prompt": "Enter your Slack Signing Secret (from App Credentials)",
|
||||
"prompt": "Enter your Slack App Signing Secret (found under Basic Information > App Credentials in your Slack App settings).",
|
||||
"optional": false
|
||||
}
|
||||
]
|
||||
],
|
||||
"setup_url": "https://api.slack.com/apps"
|
||||
},
|
||||
"capabilities": {
|
||||
"http": {
|
||||
@@ -43,6 +46,9 @@
|
||||
"emit_rate_limit": {
|
||||
"messages_per_minute": 100,
|
||||
"messages_per_hour": 5000
|
||||
},
|
||||
"webhook": {
|
||||
"hmac_secret_name": "slack_signing_secret"
|
||||
}
|
||||
}
|
||||
},
|
||||
|
||||
@@ -1032,11 +1032,14 @@ fn handle_message(message: TelegramMessage) {
|
||||
return;
|
||||
}
|
||||
}
|
||||
} else if is_private {
|
||||
// No owner_id: apply dm_policy for private chats
|
||||
} else {
|
||||
// No owner_id: apply authorization based on dm_policy and allow_from
|
||||
// This applies to both private and group chats when owner_id is null
|
||||
let dm_policy =
|
||||
channel_host::workspace_read(DM_POLICY_PATH).unwrap_or_else(|| "pairing".to_string());
|
||||
|
||||
// For private chats with non-open policy, check allowlist
|
||||
// For group chats with non-open policy, also check allowlist
|
||||
if dm_policy != "open" {
|
||||
// Build effective allow list: config allow_from + pairing store
|
||||
let mut allowed: Vec<String> = channel_host::workspace_read(ALLOW_FROM_PATH)
|
||||
@@ -1054,8 +1057,8 @@ fn handle_message(message: TelegramMessage) {
|
||||
|| username_opt.map_or(false, |u| allowed.contains(&u.to_string()));
|
||||
|
||||
if !is_allowed {
|
||||
if dm_policy == "pairing" {
|
||||
// Upsert pairing request and send reply
|
||||
if is_private && dm_policy == "pairing" {
|
||||
// Upsert pairing request and send reply (only for private chats)
|
||||
let meta = serde_json::json!({
|
||||
"chat_id": message.chat.id,
|
||||
"user_id": from.id,
|
||||
@@ -1083,6 +1086,15 @@ fn handle_message(message: TelegramMessage) {
|
||||
);
|
||||
}
|
||||
}
|
||||
} else if !is_private {
|
||||
// For group chats with non-open dm_policy, just log and drop
|
||||
channel_host::log(
|
||||
channel_host::LogLevel::Debug,
|
||||
&format!(
|
||||
"Dropping message from unauthorized user {} in group chat",
|
||||
from.id
|
||||
),
|
||||
);
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
{
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"type": "channel",
|
||||
"name": "telegram",
|
||||
"description": "Telegram Bot API channel for receiving and responding to Telegram messages",
|
||||
@@ -9,7 +11,8 @@
|
||||
"prompt": "Enter your Telegram Bot API token (from @BotFather)",
|
||||
"optional": false
|
||||
}
|
||||
]
|
||||
],
|
||||
"setup_url": "https://t.me/BotFather"
|
||||
},
|
||||
"capabilities": {
|
||||
"http": {
|
||||
@@ -39,6 +42,10 @@
|
||||
"emit_rate_limit": {
|
||||
"messages_per_minute": 100,
|
||||
"messages_per_hour": 5000
|
||||
},
|
||||
"webhook": {
|
||||
"secret_header": "X-Telegram-Bot-Api-Secret-Token",
|
||||
"secret_name": "telegram_webhook_secret"
|
||||
}
|
||||
}
|
||||
},
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
{
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"type": "channel",
|
||||
"name": "whatsapp",
|
||||
"description": "WhatsApp Cloud API channel for receiving and responding to WhatsApp messages",
|
||||
@@ -6,7 +8,7 @@
|
||||
"required_secrets": [
|
||||
{
|
||||
"name": "whatsapp_access_token",
|
||||
"prompt": "Enter your WhatsApp Cloud API access token (from Meta Developer Portal)",
|
||||
"prompt": "Enter your WhatsApp Cloud API permanent access token (from the Meta Developer Portal under your app's WhatsApp > API Setup).",
|
||||
"validation": "^[A-Za-z0-9_-]+$"
|
||||
},
|
||||
{
|
||||
@@ -16,7 +18,8 @@
|
||||
"auto_generate": { "length": 32 }
|
||||
}
|
||||
],
|
||||
"validation_endpoint": "https://graph.facebook.com/v18.0/me?access_token={whatsapp_access_token}"
|
||||
"validation_endpoint": "https://graph.facebook.com/v18.0/me?access_token={whatsapp_access_token}",
|
||||
"setup_url": "https://developers.facebook.com/apps"
|
||||
},
|
||||
"capabilities": {
|
||||
"http": {
|
||||
|
||||
+10
@@ -0,0 +1,10 @@
|
||||
coverage:
|
||||
status:
|
||||
project:
|
||||
default:
|
||||
target: auto
|
||||
threshold: 1%
|
||||
patch:
|
||||
default:
|
||||
target: 80%
|
||||
threshold: 5%
|
||||
@@ -24,6 +24,15 @@ GATEWAY_HOST=0.0.0.0
|
||||
GATEWAY_PORT=3000
|
||||
GATEWAY_AUTH_TOKEN=CHANGE_ME
|
||||
|
||||
# Restart Feature (Docker containers only)
|
||||
# IMPORTANT: Set this in the container entrypoint or docker-compose to enable restart.
|
||||
# The Docker entrypoint loop monitors exit codes:
|
||||
# - Exit code 0 = clean restart: reset failure counter, wait IRONCLAW_RESTART_DELAY, restart
|
||||
# - Exit code ≠ 0 = failure: increment counter, exit after IRONCLAW_MAX_FAILURES
|
||||
IRONCLAW_IN_DOCKER=false
|
||||
IRONCLAW_RESTART_DELAY=5 # seconds to wait before restarting (range: 1-30)
|
||||
IRONCLAW_MAX_FAILURES=10 # max consecutive failures before container exits
|
||||
|
||||
# Disabled for initial deploy
|
||||
SANDBOX_ENABLED=false
|
||||
HEARTBEAT_ENABLED=false
|
||||
|
||||
@@ -0,0 +1,195 @@
|
||||
# Smart Model Routing for IronClaw
|
||||
|
||||
**Status:** Implemented
|
||||
**Author:** Microwave
|
||||
**Date:** 2026-02-19
|
||||
|
||||
## What
|
||||
|
||||
Automatic model selection based on request complexity. The router analyzes each user message and selects an appropriate model tier (flash/standard/pro/frontier), then maps that tier to a configured model.
|
||||
|
||||
## Why
|
||||
|
||||
1. **Cost optimization** — Simple requests ("hi", "what time is it") don't need expensive models
|
||||
2. **User experience** — Simple requests return faster with lightweight models
|
||||
3. **NEAR AI native** — Default backend uses NEAR AI inference where costs vary by model
|
||||
4. **Zero-config value** — Users benefit immediately without configuration
|
||||
5. **Not just power users** — Everyone gets smart defaults, power users can override
|
||||
|
||||
## How
|
||||
|
||||
### Architecture
|
||||
|
||||
```
|
||||
User Message
|
||||
│
|
||||
▼
|
||||
┌──────────────────┐
|
||||
│ Pattern Overrides │ ← Fast-path for obvious cases (greetings, security audits)
|
||||
└────────┬─────────┘
|
||||
│ no match
|
||||
▼
|
||||
┌──────────────────┐
|
||||
│ Complexity Scorer │ ← 13-dimension analysis
|
||||
└────────┬─────────┘
|
||||
│ score 0-100
|
||||
▼
|
||||
┌──────────────────┐
|
||||
│ Tier Mapping │ ← 0-15: flash, 16-40: standard, 41-65: pro, 66+: frontier
|
||||
└────────┬─────────┘
|
||||
│ tier
|
||||
▼
|
||||
┌──────────────────┐
|
||||
│ Model Selection │ ← Currently: cheap provider (Flash/Standard/Pro) vs primary (Frontier)
|
||||
└────────┬─────────┘ Target: per-tier model mapping via config
|
||||
│
|
||||
▼
|
||||
LLM Provider
|
||||
```
|
||||
|
||||
### Complexity Scorer (13 Dimensions)
|
||||
|
||||
Each dimension produces a 0-100 score. Weighted sum determines total.
|
||||
|
||||
| Dimension | Weight | Signals |
|
||||
|-----------|--------|---------|
|
||||
| Reasoning Words | 14% | "why", "explain", "compare", "trade-offs" |
|
||||
| Token Estimate | 12% | Prompt length |
|
||||
| Code Indicators | 10% | Backticks, syntax, "implement", "PR" |
|
||||
| Multi-Step | 10% | "first", "then", "after", "steps" |
|
||||
| Domain Specific | 10% | Technical terms (configurable) |
|
||||
| Creativity | 7% | "write", "summarize", "tweet", "blog" |
|
||||
| Question Complexity | 7% | Multiple questions, open-ended starters |
|
||||
| Precision | 6% | Numbers, "exactly", "calculate" |
|
||||
| Ambiguity | 5% | Vague references |
|
||||
| Context Dependency | 5% | "previous", "you said" |
|
||||
| Sentence Complexity | 5% | Commas, conjunctions, clause depth |
|
||||
| Tool Likelihood | 5% | "read", "deploy", "install" |
|
||||
| Safety Sensitivity | 4% | "password", "auth", "vulnerability" |
|
||||
|
||||
**Multi-dimensional boost:** +30% when 3+ dimensions score above threshold.
|
||||
|
||||
### Tier Boundaries
|
||||
|
||||
| Score | Tier | Typical Use Case |
|
||||
|-------|------|------------------|
|
||||
| 0-15 | flash | Greetings, acknowledgments, quick lookups |
|
||||
| 16-40 | standard | Writing, comparisons, defined tasks |
|
||||
| 41-65 | pro | Multi-step analysis, code review |
|
||||
| 66+ | frontier | Critical decisions, security audits |
|
||||
|
||||
### Pattern Overrides
|
||||
|
||||
Fast-path rules that bypass scoring for obvious cases:
|
||||
|
||||
```yaml
|
||||
# Force flash tier
|
||||
- "^(hi|hello|hey|thanks|ok|sure|yes|no)$"
|
||||
- "^what.*(time|date|day)"
|
||||
|
||||
# Force frontier tier
|
||||
- "security.*(audit|review|scan)"
|
||||
- "vulnerabilit(y|ies).*(review|scan|check|audit)"
|
||||
|
||||
# Force pro tier
|
||||
- "deploy.*(mainnet|production)"
|
||||
```
|
||||
|
||||
### Configuration
|
||||
|
||||
> **Note:** The current implementation supports smart routing via
|
||||
> `NEARAI_CHEAP_MODEL` and `SMART_ROUTING_CASCADE` env vars, plus
|
||||
> `domain_keywords` on `SmartRoutingConfig`. The full `llm.routing` YAML
|
||||
> schema below is the target design — not all knobs are wired yet.
|
||||
|
||||
**Default (zero-config):**
|
||||
```yaml
|
||||
llm:
|
||||
routing:
|
||||
enabled: true # default
|
||||
```
|
||||
|
||||
**Power user overrides (target schema):**
|
||||
```yaml
|
||||
llm:
|
||||
routing:
|
||||
enabled: true
|
||||
tiers:
|
||||
flash: "claude-3-5-haiku-latest"
|
||||
standard: "claude-sonnet-4-5-latest"
|
||||
pro: "claude-sonnet-4-5-latest"
|
||||
frontier: "claude-opus-4-5-latest"
|
||||
thinking:
|
||||
pro: "low"
|
||||
frontier: "medium"
|
||||
overrides:
|
||||
- pattern: "my-custom-pattern"
|
||||
tier: "pro"
|
||||
domain_keywords: # Custom keywords for your domain
|
||||
- "mycompany"
|
||||
- "myproduct"
|
||||
- "internal-tool"
|
||||
```
|
||||
|
||||
If `domain_keywords` is not set, uses `DEFAULT_DOMAIN_KEYWORDS` which covers common web3/infra terms.
|
||||
|
||||
**Disable routing (pin model):**
|
||||
```yaml
|
||||
llm:
|
||||
routing:
|
||||
enabled: false
|
||||
model: "claude-opus-4-5"
|
||||
```
|
||||
|
||||
**Bring your own keys:**
|
||||
```yaml
|
||||
llm:
|
||||
backend: anthropic
|
||||
api_key: "sk-..."
|
||||
routing:
|
||||
enabled: true # still works with external providers
|
||||
```
|
||||
|
||||
### Integration Points
|
||||
|
||||
1. **RoutingProvider** — New wrapper implementing `LlmProvider` trait (like `FailoverProvider`)
|
||||
2. **Scorer** — Pure function, no I/O, fast (~1ms)
|
||||
3. **Config schema** — Extend `LlmConfig` with `routing` section
|
||||
4. **Telemetry** — Log routing decisions for observability
|
||||
|
||||
### Model Agnosticism
|
||||
|
||||
**Critical:** No hardcoded model names in the router logic itself.
|
||||
|
||||
- Tier→model mappings come from config
|
||||
- Default mappings use `-latest` patterns where supported
|
||||
- NEAR AI backend handles actual model resolution
|
||||
- Router only knows about tiers
|
||||
|
||||
### Layers of Control
|
||||
|
||||
| Layer | User Type | Config |
|
||||
|-------|-----------|--------|
|
||||
| 1. Zero-config | Everyone | `routing.enabled: true` (default) |
|
||||
| 2. Tier tuning | Power users | Custom `routing.tiers` mapping |
|
||||
| 3. Pattern overrides | Power users | Custom `routing.overrides` |
|
||||
| 4. Model pinning | Power users | `routing.enabled: false` + `model: X` |
|
||||
| 5. Own API keys | Power users | `backend: anthropic` + `api_key` |
|
||||
|
||||
## Implementation Plan
|
||||
|
||||
1. [x] Port scorer to Rust (`src/llm/smart_routing.rs`)
|
||||
2. [x] Implement router wrapper (`src/llm/smart_routing.rs`)
|
||||
3. [x] Extend config schema (`src/config.rs`)
|
||||
4. [x] Wire into provider creation (`src/llm/mod.rs`)
|
||||
5. [x] Add telemetry/logging
|
||||
6. [x] Tests with real conversation samples
|
||||
7. [x] Codex + Gemini security review
|
||||
8. [x] Documentation updated (this spec)
|
||||
|
||||
## Expected Outcomes
|
||||
|
||||
- **50-70% cost reduction** for typical usage patterns
|
||||
- **Faster responses** for simple requests
|
||||
- **Zero config required** for default benefits
|
||||
- **Full control** for power users who want it
|
||||
@@ -0,0 +1,19 @@
|
||||
-- Add wit_version column to wasm_tools for WIT interface version tracking
|
||||
ALTER TABLE wasm_tools ADD COLUMN IF NOT EXISTS wit_version TEXT NOT NULL DEFAULT '0.1.0';
|
||||
|
||||
-- Create wasm_channels table for DB-stored channel extensions
|
||||
CREATE TABLE IF NOT EXISTS wasm_channels (
|
||||
id UUID PRIMARY KEY,
|
||||
user_id TEXT NOT NULL,
|
||||
name TEXT NOT NULL,
|
||||
version TEXT NOT NULL DEFAULT '0.1.0',
|
||||
wit_version TEXT NOT NULL DEFAULT '0.1.0',
|
||||
description TEXT NOT NULL DEFAULT '',
|
||||
wasm_binary BYTEA NOT NULL,
|
||||
binary_hash BYTEA NOT NULL,
|
||||
capabilities_json TEXT NOT NULL DEFAULT '{}',
|
||||
status TEXT NOT NULL DEFAULT 'active',
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
CONSTRAINT unique_wasm_channel UNIQUE (user_id, name)
|
||||
);
|
||||
@@ -3,29 +3,35 @@
|
||||
"display_name": "Discord Channel",
|
||||
"kind": "channel",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"description": "Talk to your agent in Discord",
|
||||
"keywords": ["messaging", "chat", "discord", "bot"],
|
||||
|
||||
"keywords": [
|
||||
"messaging",
|
||||
"chat",
|
||||
"discord",
|
||||
"bot"
|
||||
],
|
||||
"source": {
|
||||
"dir": "channels-src/discord",
|
||||
"capabilities": "discord.capabilities.json",
|
||||
"crate_name": "discord-channel"
|
||||
},
|
||||
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/discord-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "56b8d92e3e0f32d9cfaf0dc1ac8aa8c42b98cda4d92cd42a74e0af01eeee60a1"
|
||||
}
|
||||
},
|
||||
|
||||
"auth_summary": {
|
||||
"method": "manual",
|
||||
"provider": "Discord",
|
||||
"secrets": ["discord_bot_token"],
|
||||
"secrets": [
|
||||
"discord_bot_token"
|
||||
],
|
||||
"shared_auth": null,
|
||||
"setup_url": "https://discord.com/developers/applications"
|
||||
},
|
||||
|
||||
"tags": ["messaging"]
|
||||
"tags": [
|
||||
"messaging"
|
||||
]
|
||||
}
|
||||
|
||||
@@ -3,29 +3,37 @@
|
||||
"display_name": "Slack Channel",
|
||||
"kind": "channel",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"description": "Talk to your agent in Slack",
|
||||
"keywords": ["messaging", "chat", "workspace", "slack"],
|
||||
|
||||
"keywords": [
|
||||
"messaging",
|
||||
"chat",
|
||||
"workspace",
|
||||
"slack"
|
||||
],
|
||||
"source": {
|
||||
"dir": "channels-src/slack",
|
||||
"capabilities": "slack.capabilities.json",
|
||||
"crate_name": "slack-channel"
|
||||
},
|
||||
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/slack-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "536b52179094d228e18d51b203286c80f800e0132b42fa08ce292c76a13a70e8"
|
||||
}
|
||||
},
|
||||
|
||||
"auth_summary": {
|
||||
"method": "manual",
|
||||
"provider": "Slack",
|
||||
"secrets": ["slack_bot_token", "slack_signing_secret"],
|
||||
"secrets": [
|
||||
"slack_bot_token",
|
||||
"slack_signing_secret"
|
||||
],
|
||||
"shared_auth": null,
|
||||
"setup_url": "https://api.slack.com/apps"
|
||||
},
|
||||
|
||||
"tags": ["default", "messaging"]
|
||||
"tags": [
|
||||
"default",
|
||||
"messaging"
|
||||
]
|
||||
}
|
||||
|
||||
@@ -3,29 +3,36 @@
|
||||
"display_name": "Telegram Channel",
|
||||
"kind": "channel",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"description": "Talk to your agent through a Telegram bot",
|
||||
"keywords": ["messaging", "bot", "chat", "telegram"],
|
||||
|
||||
"keywords": [
|
||||
"messaging",
|
||||
"bot",
|
||||
"chat",
|
||||
"telegram"
|
||||
],
|
||||
"source": {
|
||||
"dir": "channels-src/telegram",
|
||||
"capabilities": "telegram.capabilities.json",
|
||||
"crate_name": "telegram-channel"
|
||||
},
|
||||
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/telegram-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "9bcc39d717c2b7e4e3327d31fc0032fbe4bf666852622c69b88dbaee236f77ee"
|
||||
}
|
||||
},
|
||||
|
||||
"auth_summary": {
|
||||
"method": "manual",
|
||||
"provider": "Telegram",
|
||||
"secrets": ["telegram_bot_token"],
|
||||
"secrets": [
|
||||
"telegram_bot_token"
|
||||
],
|
||||
"shared_auth": null,
|
||||
"setup_url": "https://t.me/BotFather"
|
||||
},
|
||||
|
||||
"tags": ["default", "messaging"]
|
||||
"tags": [
|
||||
"default",
|
||||
"messaging"
|
||||
]
|
||||
}
|
||||
|
||||
@@ -3,29 +3,36 @@
|
||||
"display_name": "WhatsApp Channel",
|
||||
"kind": "channel",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"description": "Talk to your agent through WhatsApp",
|
||||
"keywords": ["messaging", "chat", "whatsapp", "meta"],
|
||||
|
||||
"keywords": [
|
||||
"messaging",
|
||||
"chat",
|
||||
"whatsapp",
|
||||
"meta"
|
||||
],
|
||||
"source": {
|
||||
"dir": "channels-src/whatsapp",
|
||||
"capabilities": "whatsapp.capabilities.json",
|
||||
"crate_name": "whatsapp-channel"
|
||||
},
|
||||
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/whatsapp-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "54dc4a6c2b4e07bce6de54cab415920be518c94a096df45f1104f0fa5c8f6382"
|
||||
}
|
||||
},
|
||||
|
||||
"auth_summary": {
|
||||
"method": "manual",
|
||||
"provider": "Meta",
|
||||
"secrets": ["whatsapp_access_token", "whatsapp_verify_token"],
|
||||
"secrets": [
|
||||
"whatsapp_access_token",
|
||||
"whatsapp_verify_token"
|
||||
],
|
||||
"shared_auth": null,
|
||||
"setup_url": "https://developers.facebook.com/apps/"
|
||||
},
|
||||
|
||||
"tags": ["messaging"]
|
||||
"tags": [
|
||||
"messaging"
|
||||
]
|
||||
}
|
||||
|
||||
@@ -3,29 +3,37 @@
|
||||
"display_name": "GitHub",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"description": "GitHub integration for issues, PRs, repos, and code search",
|
||||
"keywords": ["git", "code", "issues", "pull-requests", "repositories"],
|
||||
|
||||
"keywords": [
|
||||
"git",
|
||||
"code",
|
||||
"issues",
|
||||
"pull-requests",
|
||||
"repositories"
|
||||
],
|
||||
"source": {
|
||||
"dir": "tools-src/github",
|
||||
"capabilities": "github-tool.capabilities.json",
|
||||
"crate_name": "github-tool"
|
||||
},
|
||||
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/github-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "a068208d454be34585e8809816991c1853c255cd6f04c6ca9277e21219b745bc"
|
||||
}
|
||||
},
|
||||
|
||||
"auth_summary": {
|
||||
"method": "manual",
|
||||
"provider": "GitHub",
|
||||
"secrets": ["github_token"],
|
||||
"secrets": [
|
||||
"github_token"
|
||||
],
|
||||
"shared_auth": null,
|
||||
"setup_url": "https://github.com/settings/tokens"
|
||||
},
|
||||
|
||||
"tags": ["default", "development"]
|
||||
"tags": [
|
||||
"default",
|
||||
"development"
|
||||
]
|
||||
}
|
||||
|
||||
@@ -3,29 +3,37 @@
|
||||
"display_name": "Gmail",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"description": "Read, send, and manage Gmail messages and threads",
|
||||
"keywords": ["email", "google", "mail", "messaging"],
|
||||
|
||||
"keywords": [
|
||||
"email",
|
||||
"google",
|
||||
"mail",
|
||||
"messaging"
|
||||
],
|
||||
"source": {
|
||||
"dir": "tools-src/gmail",
|
||||
"capabilities": "gmail-tool.capabilities.json",
|
||||
"crate_name": "gmail-tool"
|
||||
},
|
||||
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/gmail-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "557e3d485de337340e059948b3099dcacae1c2e937e2bf66bad52b105c5be9bc"
|
||||
}
|
||||
},
|
||||
|
||||
"auth_summary": {
|
||||
"method": "oauth",
|
||||
"provider": "Google",
|
||||
"secrets": ["google_oauth_token"],
|
||||
"secrets": [
|
||||
"google_oauth_token"
|
||||
],
|
||||
"shared_auth": "google_oauth_token",
|
||||
"setup_url": "https://console.cloud.google.com/apis/credentials"
|
||||
},
|
||||
|
||||
"tags": ["default", "google", "messaging"]
|
||||
"tags": [
|
||||
"default",
|
||||
"google",
|
||||
"messaging"
|
||||
]
|
||||
}
|
||||
|
||||
@@ -3,29 +3,37 @@
|
||||
"display_name": "Google Calendar",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"description": "Create, read, update, and delete Google Calendar events",
|
||||
"keywords": ["calendar", "google", "scheduling", "events"],
|
||||
|
||||
"keywords": [
|
||||
"calendar",
|
||||
"google",
|
||||
"scheduling",
|
||||
"events"
|
||||
],
|
||||
"source": {
|
||||
"dir": "tools-src/google-calendar",
|
||||
"capabilities": "google-calendar-tool.capabilities.json",
|
||||
"crate_name": "google-calendar-tool"
|
||||
},
|
||||
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-calendar-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "787d57cb55cf492af013b69bfed8472a1cdd349fbb137f6084c0adcb51a467a6"
|
||||
}
|
||||
},
|
||||
|
||||
"auth_summary": {
|
||||
"method": "oauth",
|
||||
"provider": "Google",
|
||||
"secrets": ["google_oauth_token"],
|
||||
"secrets": [
|
||||
"google_oauth_token"
|
||||
],
|
||||
"shared_auth": "google_oauth_token",
|
||||
"setup_url": "https://console.cloud.google.com/apis/credentials"
|
||||
},
|
||||
|
||||
"tags": ["default", "google", "productivity"]
|
||||
"tags": [
|
||||
"default",
|
||||
"google",
|
||||
"productivity"
|
||||
]
|
||||
}
|
||||
|
||||
@@ -3,29 +3,36 @@
|
||||
"display_name": "Google Docs",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"description": "Create and edit Google Docs documents",
|
||||
"keywords": ["documents", "google", "writing", "docs"],
|
||||
|
||||
"keywords": [
|
||||
"documents",
|
||||
"google",
|
||||
"writing",
|
||||
"docs"
|
||||
],
|
||||
"source": {
|
||||
"dir": "tools-src/google-docs",
|
||||
"capabilities": "google-docs-tool.capabilities.json",
|
||||
"crate_name": "google-docs-tool"
|
||||
},
|
||||
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-docs-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "64407e650b6b8fcf9892255ef237a7683469fd465e3ba84e709fd71cd77fcbab"
|
||||
}
|
||||
},
|
||||
|
||||
"auth_summary": {
|
||||
"method": "oauth",
|
||||
"provider": "Google",
|
||||
"secrets": ["google_oauth_token"],
|
||||
"secrets": [
|
||||
"google_oauth_token"
|
||||
],
|
||||
"shared_auth": "google_oauth_token",
|
||||
"setup_url": "https://console.cloud.google.com/apis/credentials"
|
||||
},
|
||||
|
||||
"tags": ["google", "productivity"]
|
||||
"tags": [
|
||||
"google",
|
||||
"productivity"
|
||||
]
|
||||
}
|
||||
|
||||
@@ -3,29 +3,37 @@
|
||||
"display_name": "Google Drive",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"description": "Upload, download, search, and manage Google Drive files and folders",
|
||||
"keywords": ["storage", "google", "files", "drive"],
|
||||
|
||||
"keywords": [
|
||||
"storage",
|
||||
"google",
|
||||
"files",
|
||||
"drive"
|
||||
],
|
||||
"source": {
|
||||
"dir": "tools-src/google-drive",
|
||||
"capabilities": "google-drive-tool.capabilities.json",
|
||||
"crate_name": "google-drive-tool"
|
||||
},
|
||||
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-drive-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "3618689278d9e8546e37490c7ba05ee4dffb0acfaa89f0865161616326c41ee8"
|
||||
}
|
||||
},
|
||||
|
||||
"auth_summary": {
|
||||
"method": "oauth",
|
||||
"provider": "Google",
|
||||
"secrets": ["google_oauth_token"],
|
||||
"secrets": [
|
||||
"google_oauth_token"
|
||||
],
|
||||
"shared_auth": "google_oauth_token",
|
||||
"setup_url": "https://console.cloud.google.com/apis/credentials"
|
||||
},
|
||||
|
||||
"tags": ["default", "google", "storage"]
|
||||
"tags": [
|
||||
"default",
|
||||
"google",
|
||||
"storage"
|
||||
]
|
||||
}
|
||||
|
||||
@@ -3,29 +3,36 @@
|
||||
"display_name": "Google Sheets",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"description": "Read and write Google Sheets spreadsheet data",
|
||||
"keywords": ["spreadsheets", "google", "data", "sheets"],
|
||||
|
||||
"keywords": [
|
||||
"spreadsheets",
|
||||
"google",
|
||||
"data",
|
||||
"sheets"
|
||||
],
|
||||
"source": {
|
||||
"dir": "tools-src/google-sheets",
|
||||
"capabilities": "google-sheets-tool.capabilities.json",
|
||||
"crate_name": "google-sheets-tool"
|
||||
},
|
||||
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-sheets-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "e486d228891a5b431b9993cebe0820c5331130916555ac3ba33b5b23bfd4ad75"
|
||||
}
|
||||
},
|
||||
|
||||
"auth_summary": {
|
||||
"method": "oauth",
|
||||
"provider": "Google",
|
||||
"secrets": ["google_oauth_token"],
|
||||
"secrets": [
|
||||
"google_oauth_token"
|
||||
],
|
||||
"shared_auth": "google_oauth_token",
|
||||
"setup_url": "https://console.cloud.google.com/apis/credentials"
|
||||
},
|
||||
|
||||
"tags": ["google", "productivity"]
|
||||
"tags": [
|
||||
"google",
|
||||
"productivity"
|
||||
]
|
||||
}
|
||||
|
||||
@@ -3,29 +3,35 @@
|
||||
"display_name": "Google Slides",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"description": "Create and edit Google Slides presentations",
|
||||
"keywords": ["presentations", "google", "slides"],
|
||||
|
||||
"keywords": [
|
||||
"presentations",
|
||||
"google",
|
||||
"slides"
|
||||
],
|
||||
"source": {
|
||||
"dir": "tools-src/google-slides",
|
||||
"capabilities": "google-slides-tool.capabilities.json",
|
||||
"crate_name": "google-slides-tool"
|
||||
},
|
||||
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-slides-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "9f2ba32c6d43cf87f50e2ae6720a1e9d89c00666aab2047080aadb503968c66a"
|
||||
}
|
||||
},
|
||||
|
||||
"auth_summary": {
|
||||
"method": "oauth",
|
||||
"provider": "Google",
|
||||
"secrets": ["google_oauth_token"],
|
||||
"secrets": [
|
||||
"google_oauth_token"
|
||||
],
|
||||
"shared_auth": "google_oauth_token",
|
||||
"setup_url": "https://console.cloud.google.com/apis/credentials"
|
||||
},
|
||||
|
||||
"tags": ["google", "productivity"]
|
||||
"tags": [
|
||||
"google",
|
||||
"productivity"
|
||||
]
|
||||
}
|
||||
|
||||
@@ -1,31 +0,0 @@
|
||||
{
|
||||
"name": "okta",
|
||||
"display_name": "Okta",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"description": "Okta SSO for user profile, app catalog, and SSO launch links",
|
||||
"keywords": ["sso", "identity", "authentication", "okta"],
|
||||
|
||||
"source": {
|
||||
"dir": "tools-src/okta",
|
||||
"capabilities": "okta-tool.capabilities.json",
|
||||
"crate_name": "okta-tool"
|
||||
},
|
||||
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/okta-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
}
|
||||
},
|
||||
|
||||
"auth_summary": {
|
||||
"method": "oauth",
|
||||
"provider": "Okta",
|
||||
"secrets": ["okta_oauth_token"],
|
||||
"shared_auth": null,
|
||||
"setup_url": "https://developer.okta.com/docs/guides/implement-oauth-for-okta/main/"
|
||||
},
|
||||
|
||||
"tags": ["identity"]
|
||||
}
|
||||
@@ -3,29 +3,35 @@
|
||||
"display_name": "Slack Tool",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"description": "Your agent uses Slack to post and read messages in your workspace",
|
||||
"keywords": ["messaging", "chat", "workspace"],
|
||||
|
||||
"keywords": [
|
||||
"messaging",
|
||||
"chat",
|
||||
"workspace"
|
||||
],
|
||||
"source": {
|
||||
"dir": "tools-src/slack",
|
||||
"capabilities": "slack-tool.capabilities.json",
|
||||
"crate_name": "slack-tool"
|
||||
},
|
||||
|
||||
"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/latest/download/slack-tool-wasm32-wasip2.tar.gz",
|
||||
"sha256": "536b52179094d228e18d51b203286c80f800e0132b42fa08ce292c76a13a70e8"
|
||||
}
|
||||
},
|
||||
|
||||
"auth_summary": {
|
||||
"method": "oauth",
|
||||
"provider": "Slack",
|
||||
"secrets": ["slack_bot_token"],
|
||||
"secrets": [
|
||||
"slack_bot_token"
|
||||
],
|
||||
"shared_auth": null,
|
||||
"setup_url": "https://api.slack.com/apps"
|
||||
},
|
||||
|
||||
"tags": ["default", "messaging"]
|
||||
"tags": [
|
||||
"default",
|
||||
"messaging"
|
||||
]
|
||||
}
|
||||
|
||||
@@ -3,29 +3,36 @@
|
||||
"display_name": "Telegram Tool",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"description": "Your agent uses your Telegram account to read and send messages",
|
||||
"keywords": ["messaging", "chat", "telegram", "mtproto"],
|
||||
|
||||
"keywords": [
|
||||
"messaging",
|
||||
"chat",
|
||||
"telegram",
|
||||
"mtproto"
|
||||
],
|
||||
"source": {
|
||||
"dir": "tools-src/telegram",
|
||||
"capabilities": "telegram-tool.capabilities.json",
|
||||
"crate_name": "telegram-tool"
|
||||
},
|
||||
|
||||
"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/latest/download/telegram-mtproto-wasm32-wasip2.tar.gz",
|
||||
"sha256": "9bcc39d717c2b7e4e3327d31fc0032fbe4bf666852622c69b88dbaee236f77ee"
|
||||
}
|
||||
},
|
||||
|
||||
"auth_summary": {
|
||||
"method": "manual",
|
||||
"provider": "Telegram",
|
||||
"secrets": ["telegram_api_id", "telegram_api_hash"],
|
||||
"secrets": [
|
||||
"telegram_api_id",
|
||||
"telegram_api_hash"
|
||||
],
|
||||
"shared_auth": null,
|
||||
"setup_url": "https://my.telegram.org/apps"
|
||||
},
|
||||
|
||||
"tags": ["messaging"]
|
||||
"tags": [
|
||||
"messaging"
|
||||
]
|
||||
}
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
{
|
||||
"name": "web-search",
|
||||
"display_name": "Web Search",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"description": "Search the web using Brave Search API",
|
||||
"keywords": [
|
||||
"search",
|
||||
"web",
|
||||
"brave",
|
||||
"internet"
|
||||
],
|
||||
"source": {
|
||||
"dir": "tools-src/web-search",
|
||||
"capabilities": "web-search-tool.capabilities.json",
|
||||
"crate_name": "web-search-tool"
|
||||
},
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/web-search-wasm32-wasip2.tar.gz",
|
||||
"sha256": "003b0667878035ba0e6029ec322bf5e6e7bab06c3022d5fe897ce288d2d156e0"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
"method": "manual",
|
||||
"provider": "Brave",
|
||||
"secrets": [
|
||||
"brave_api_key"
|
||||
],
|
||||
"shared_auth": null,
|
||||
"setup_url": "https://brave.com/search/api/"
|
||||
},
|
||||
"tags": [
|
||||
"default",
|
||||
"search"
|
||||
]
|
||||
}
|
||||
Executable
+74
@@ -0,0 +1,74 @@
|
||||
#!/usr/bin/env bash
|
||||
# Build all WASM tools and channels from source.
|
||||
#
|
||||
# Verifies that every tool/channel in the registry compiles against the
|
||||
# current WIT definitions. Used by CI and can be run locally.
|
||||
#
|
||||
# Prerequisites:
|
||||
# rustup target add wasm32-wasip2
|
||||
# cargo install cargo-component --locked
|
||||
#
|
||||
# Usage:
|
||||
# ./scripts/build-wasm-extensions.sh # build all
|
||||
# ./scripts/build-wasm-extensions.sh --tools # tools only
|
||||
# ./scripts/build-wasm-extensions.sh --channels # channels only
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
cd "$(dirname "$0")/.."
|
||||
|
||||
BUILD_TOOLS=true
|
||||
BUILD_CHANNELS=true
|
||||
FAILED=()
|
||||
|
||||
if [[ "${1:-}" == "--tools" ]]; then
|
||||
BUILD_CHANNELS=false
|
||||
elif [[ "${1:-}" == "--channels" ]]; then
|
||||
BUILD_TOOLS=false
|
||||
fi
|
||||
|
||||
build_extension() {
|
||||
local manifest_path="$1"
|
||||
local source_dir
|
||||
local crate_name
|
||||
|
||||
source_dir=$(jq -r '.source.dir' "$manifest_path")
|
||||
crate_name=$(jq -r '.source.crate_name' "$manifest_path")
|
||||
local name
|
||||
name=$(basename "$manifest_path" .json)
|
||||
|
||||
if [ ! -d "$source_dir" ]; then
|
||||
echo " SKIP $name (source dir $source_dir not found)"
|
||||
return 0
|
||||
fi
|
||||
|
||||
echo " BUILD $name ($crate_name) from $source_dir"
|
||||
if ! cargo component build --release --manifest-path "$source_dir/Cargo.toml" 2>&1; then
|
||||
echo " FAIL $name"
|
||||
FAILED+=("$name")
|
||||
return 1
|
||||
fi
|
||||
echo " OK $name"
|
||||
}
|
||||
|
||||
if $BUILD_TOOLS; then
|
||||
echo "Building WASM tools..."
|
||||
for manifest in registry/tools/*.json; do
|
||||
build_extension "$manifest" || true
|
||||
done
|
||||
fi
|
||||
|
||||
if $BUILD_CHANNELS; then
|
||||
echo "Building WASM channels..."
|
||||
for manifest in registry/channels/*.json; do
|
||||
build_extension "$manifest" || true
|
||||
done
|
||||
fi
|
||||
|
||||
echo ""
|
||||
if [ ${#FAILED[@]} -gt 0 ]; then
|
||||
echo "FAILED: ${FAILED[*]}"
|
||||
exit 1
|
||||
else
|
||||
echo "All WASM extensions built successfully."
|
||||
fi
|
||||
Executable
+251
@@ -0,0 +1,251 @@
|
||||
#!/usr/bin/env bash
|
||||
set -euo pipefail
|
||||
|
||||
# CI script: check that version bumps accompany WIT or extension source changes.
|
||||
# Exit 0 if all checks pass, exit 1 if any version wasn't bumped.
|
||||
|
||||
ERRORS=0
|
||||
|
||||
# --- Skip mechanism -----------------------------------------------------------
|
||||
|
||||
if [[ "${PR_LABELS:-}" == *"skip-version-check"* ]]; then
|
||||
echo "skip-version-check label detected — skipping all version checks."
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# Check commit messages for [skip-version-check]
|
||||
if git log "origin/${GITHUB_BASE_REF:-main}...HEAD" --pretty=format:"%s %b" 2>/dev/null \
|
||||
| grep -qF '[skip-version-check]'; then
|
||||
echo "[skip-version-check] found in commit message — skipping all version checks."
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# --- Determine base branch and changed files ----------------------------------
|
||||
|
||||
BASE_BRANCH="${GITHUB_BASE_REF:-main}"
|
||||
echo "Base branch: $BASE_BRANCH"
|
||||
|
||||
# Ensure the base branch ref is available
|
||||
if ! git rev-parse "origin/${BASE_BRANCH}" >/dev/null 2>&1; then
|
||||
echo "Fetching origin/${BASE_BRANCH}..."
|
||||
git fetch origin "$BASE_BRANCH" --depth=1
|
||||
fi
|
||||
|
||||
CHANGED_FILES=$(git diff --name-only "origin/${BASE_BRANCH}...HEAD")
|
||||
|
||||
if [[ -z "$CHANGED_FILES" ]]; then
|
||||
echo "No changed files detected. Nothing to check."
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# --- Helper functions ---------------------------------------------------------
|
||||
|
||||
# Extract the version from a WIT package line like: package near:[email protected];
|
||||
extract_wit_version() {
|
||||
local file="$1"
|
||||
if [[ ! -f "$file" ]]; then
|
||||
echo ""
|
||||
return
|
||||
fi
|
||||
sed -n 's/^[[:space:]]*package[[:space:]]\+[^@]*@\([0-9][0-9.]*[0-9]\)[[:space:]]*;.*/\1/p' "$file" \
|
||||
| head -n1
|
||||
}
|
||||
|
||||
# Extract version from the base branch copy of a file
|
||||
extract_wit_version_base() {
|
||||
local file="$1"
|
||||
git show "origin/${BASE_BRANCH}:${file}" 2>/dev/null \
|
||||
| sed -n 's/^[[:space:]]*package[[:space:]]\+[^@]*@\([0-9][0-9.]*[0-9]\)[[:space:]]*;.*/\1/p' \
|
||||
| head -n1 || true
|
||||
}
|
||||
|
||||
# Extract a Rust string constant value: pub const NAME: &str = "value";
|
||||
extract_rust_const() {
|
||||
local file="$1"
|
||||
local const_name="$2"
|
||||
if [[ ! -f "$file" ]]; then
|
||||
echo ""
|
||||
return
|
||||
fi
|
||||
sed -n "s/^.*${const_name}[[:space:]]*:[[:space:]]*&str[[:space:]]*=[[:space:]]*\"\([^\"]*\)\".*/\1/p" "$file" \
|
||||
| head -n1
|
||||
}
|
||||
|
||||
# Extract JSON "version" field using jq
|
||||
extract_json_version() {
|
||||
local file="$1"
|
||||
if [[ ! -f "$file" ]]; then
|
||||
echo ""
|
||||
return
|
||||
fi
|
||||
jq -r '.version // empty' "$file" 2>/dev/null || true
|
||||
}
|
||||
|
||||
# Extract JSON "version" from the base branch copy of a file
|
||||
extract_json_version_base() {
|
||||
local file="$1"
|
||||
git show "origin/${BASE_BRANCH}:${file}" 2>/dev/null | jq -r '.version // empty' 2>/dev/null || true
|
||||
}
|
||||
|
||||
# Return 0 if $1 (new) is strictly greater than $2 (old) via sort -V, or old is empty.
|
||||
version_was_bumped() {
|
||||
local new="$1"
|
||||
local old="$2"
|
||||
if [[ -z "$old" ]]; then
|
||||
# No prior version — treat as new, no bump required
|
||||
return 0
|
||||
fi
|
||||
if [[ -z "$new" ]]; then
|
||||
# Version was removed — that's a problem
|
||||
return 1
|
||||
fi
|
||||
if [[ "$new" == "$old" ]]; then
|
||||
return 1
|
||||
fi
|
||||
# Check new > old via sort -V
|
||||
local highest
|
||||
highest=$(printf '%s\n%s\n' "$new" "$old" | sort -V | tail -n1)
|
||||
[[ "$highest" == "$new" ]]
|
||||
}
|
||||
|
||||
# --- 1. WIT changes ----------------------------------------------------------
|
||||
|
||||
WIT_TOOL_CHANGED=false
|
||||
WIT_CHANNEL_CHANGED=false
|
||||
|
||||
if echo "$CHANGED_FILES" | grep -qx 'wit/tool\.wit'; then
|
||||
WIT_TOOL_CHANGED=true
|
||||
fi
|
||||
if echo "$CHANGED_FILES" | grep -qx 'wit/channel\.wit'; then
|
||||
WIT_CHANNEL_CHANGED=true
|
||||
fi
|
||||
|
||||
if $WIT_TOOL_CHANGED; then
|
||||
echo ""
|
||||
echo "=== wit/tool.wit changed ==="
|
||||
|
||||
NEW_VER=$(extract_wit_version "wit/tool.wit")
|
||||
OLD_VER=$(extract_wit_version_base "wit/tool.wit")
|
||||
echo " WIT package version: ${OLD_VER:-<none>} -> ${NEW_VER:-<missing>}"
|
||||
|
||||
if ! version_was_bumped "${NEW_VER}" "${OLD_VER}"; then
|
||||
echo " ERROR: wit/tool.wit package version was not bumped (${OLD_VER} -> ${NEW_VER:-<missing>})."
|
||||
ERRORS=$((ERRORS + 1))
|
||||
else
|
||||
echo " OK: WIT package version bumped."
|
||||
fi
|
||||
|
||||
# Check WIT_TOOL_VERSION constant matches
|
||||
CONST_VER=$(extract_rust_const "src/tools/wasm/mod.rs" "WIT_TOOL_VERSION")
|
||||
if [[ -n "$NEW_VER" && "$CONST_VER" != "$NEW_VER" ]]; then
|
||||
echo " ERROR: WIT_TOOL_VERSION in src/tools/wasm/mod.rs is '${CONST_VER}' but wit/tool.wit has '${NEW_VER}'. They must match."
|
||||
ERRORS=$((ERRORS + 1))
|
||||
elif [[ -n "$NEW_VER" ]]; then
|
||||
echo " OK: WIT_TOOL_VERSION matches wit/tool.wit."
|
||||
fi
|
||||
fi
|
||||
|
||||
if $WIT_CHANNEL_CHANGED; then
|
||||
echo ""
|
||||
echo "=== wit/channel.wit changed ==="
|
||||
|
||||
NEW_VER=$(extract_wit_version "wit/channel.wit")
|
||||
OLD_VER=$(extract_wit_version_base "wit/channel.wit")
|
||||
echo " WIT package version: ${OLD_VER:-<none>} -> ${NEW_VER:-<missing>}"
|
||||
|
||||
if ! version_was_bumped "${NEW_VER}" "${OLD_VER}"; then
|
||||
echo " ERROR: wit/channel.wit package version was not bumped (${OLD_VER} -> ${NEW_VER:-<missing>})."
|
||||
ERRORS=$((ERRORS + 1))
|
||||
else
|
||||
echo " OK: WIT package version bumped."
|
||||
fi
|
||||
|
||||
# Check WIT_CHANNEL_VERSION constant matches
|
||||
CONST_VER=$(extract_rust_const "src/tools/wasm/mod.rs" "WIT_CHANNEL_VERSION")
|
||||
if [[ -n "$NEW_VER" && "$CONST_VER" != "$NEW_VER" ]]; then
|
||||
echo " ERROR: WIT_CHANNEL_VERSION in src/tools/wasm/mod.rs is '${CONST_VER}' but wit/channel.wit has '${NEW_VER}'. They must match."
|
||||
ERRORS=$((ERRORS + 1))
|
||||
elif [[ -n "$NEW_VER" ]]; then
|
||||
echo " OK: WIT_CHANNEL_VERSION matches wit/channel.wit."
|
||||
fi
|
||||
fi
|
||||
|
||||
if $WIT_TOOL_CHANGED || $WIT_CHANNEL_CHANGED; then
|
||||
echo ""
|
||||
echo " WARNING: WIT interface changed. All published registry extensions should bump their versions for compatibility."
|
||||
fi
|
||||
|
||||
# --- 2. Tool source changes ---------------------------------------------------
|
||||
|
||||
TOOL_NAMES=$(echo "$CHANGED_FILES" | sed -n 's|^tools-src/\([^/]*\)/.*|\1|p' | sort -u)
|
||||
|
||||
if [[ -n "$TOOL_NAMES" ]]; then
|
||||
echo ""
|
||||
echo "=== Tool source changes ==="
|
||||
fi
|
||||
|
||||
for tool in $TOOL_NAMES; do
|
||||
REGISTRY_FILE="registry/tools/${tool}.json"
|
||||
echo ""
|
||||
echo " --- tools-src/${tool}/ changed ---"
|
||||
|
||||
if [[ ! -f "$REGISTRY_FILE" ]]; then
|
||||
echo " SKIP: ${REGISTRY_FILE} does not exist yet (new extension?)."
|
||||
continue
|
||||
fi
|
||||
|
||||
NEW_VER=$(extract_json_version "$REGISTRY_FILE")
|
||||
OLD_VER=$(extract_json_version_base "$REGISTRY_FILE")
|
||||
|
||||
echo " Registry version: ${OLD_VER:-<none>} -> ${NEW_VER:-<missing>}"
|
||||
|
||||
if ! version_was_bumped "${NEW_VER}" "${OLD_VER}"; then
|
||||
echo " ERROR: ${REGISTRY_FILE} version was not bumped (${OLD_VER} -> ${NEW_VER:-<missing>}). Bump the version when changing tools-src/${tool}/."
|
||||
ERRORS=$((ERRORS + 1))
|
||||
else
|
||||
echo " OK: version bumped."
|
||||
fi
|
||||
done
|
||||
|
||||
# --- 3. Channel source changes ------------------------------------------------
|
||||
|
||||
CHANNEL_NAMES=$(echo "$CHANGED_FILES" | sed -n 's|^channels-src/\([^/]*\)/.*|\1|p' | sort -u)
|
||||
|
||||
if [[ -n "$CHANNEL_NAMES" ]]; then
|
||||
echo ""
|
||||
echo "=== Channel source changes ==="
|
||||
fi
|
||||
|
||||
for channel in $CHANNEL_NAMES; do
|
||||
REGISTRY_FILE="registry/channels/${channel}.json"
|
||||
echo ""
|
||||
echo " --- channels-src/${channel}/ changed ---"
|
||||
|
||||
if [[ ! -f "$REGISTRY_FILE" ]]; then
|
||||
echo " SKIP: ${REGISTRY_FILE} does not exist yet (new extension?)."
|
||||
continue
|
||||
fi
|
||||
|
||||
NEW_VER=$(extract_json_version "$REGISTRY_FILE")
|
||||
OLD_VER=$(extract_json_version_base "$REGISTRY_FILE")
|
||||
|
||||
echo " Registry version: ${OLD_VER:-<none>} -> ${NEW_VER:-<missing>}"
|
||||
|
||||
if ! version_was_bumped "${NEW_VER}" "${OLD_VER}"; then
|
||||
echo " ERROR: ${REGISTRY_FILE} version was not bumped (${OLD_VER} -> ${NEW_VER:-<missing>}). Bump the version when changing channels-src/${channel}/."
|
||||
ERRORS=$((ERRORS + 1))
|
||||
else
|
||||
echo " OK: version bumped."
|
||||
fi
|
||||
done
|
||||
|
||||
# --- Summary ------------------------------------------------------------------
|
||||
|
||||
echo ""
|
||||
if [[ $ERRORS -gt 0 ]]; then
|
||||
echo "FAILED: ${ERRORS} version check(s) did not pass. See errors above."
|
||||
exit 1
|
||||
else
|
||||
echo "All version checks passed."
|
||||
exit 0
|
||||
fi
|
||||
Executable
+81
@@ -0,0 +1,81 @@
|
||||
#!/usr/bin/env bash
|
||||
# commit-msg hook: require regression tests for fix commits.
|
||||
#
|
||||
# Installed by scripts/dev-setup.sh as .git/hooks/commit-msg.
|
||||
# Bypass with [skip-regression-check] in the commit message.
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
MSG_FILE="$1"
|
||||
FIRST_LINE=$(head -1 "$MSG_FILE")
|
||||
|
||||
# --- 1. Is this a fix commit? ---
|
||||
if ! grep -qiE '^(fix(\(.*\))?|hotfix|bugfix):' <<< "$FIRST_LINE"; then
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# --- 2. Skip marker ---
|
||||
if grep -qF '[skip-regression-check]' "$MSG_FILE"; then
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# --- 3. Exempt static-only / docs-only changes ---
|
||||
# Get staged files (commit-msg runs after staging is finalized).
|
||||
STAGED_FILES=$(git diff --cached --name-only --diff-filter=ACMR)
|
||||
|
||||
if [ -z "$STAGED_FILES" ]; then
|
||||
exit 0
|
||||
fi
|
||||
|
||||
ALL_EXEMPT=true
|
||||
while IFS= read -r file; do
|
||||
case "$file" in
|
||||
src/channels/web/static/*) ;;
|
||||
*.md) ;;
|
||||
*) ALL_EXEMPT=false; break ;;
|
||||
esac
|
||||
done <<< "$STAGED_FILES"
|
||||
|
||||
if [ "$ALL_EXEMPT" = true ]; then
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# --- 4. Look for test changes in staged .rs files ---
|
||||
|
||||
# Fast path: new test attributes or test modules in added lines.
|
||||
if git diff --cached -U0 -- '*.rs' | grep -qE '^\+.*(#\[test\]|#\[tokio::test\]|#\[cfg\(test\)\]|mod tests)'; then
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# Whole-function context: detect edits inside existing test functions.
|
||||
# -W shows the full enclosing function, so #[test] appears in context
|
||||
# lines when changes are inside a test function.
|
||||
if git diff --cached -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 }
|
||||
/^\+[^+]/ { has_add=1 }
|
||||
END { if (has_test && has_add) found=1; exit !found }
|
||||
'; then
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# Also check for new/modified files under tests/
|
||||
if grep -qE '^tests/' <<< "$STAGED_FILES"; then
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# --- 5. No test found — block the commit ---
|
||||
echo ""
|
||||
echo "╔══════════════════════════════════════════════════════════════╗"
|
||||
echo "║ REGRESSION TEST REQUIRED ║"
|
||||
echo "║ ║"
|
||||
echo "║ This commit looks like a bug fix but has no test changes. ║"
|
||||
echo "║ Every fix should include a test that reproduces the bug. ║"
|
||||
echo "║ ║"
|
||||
echo "║ Options: ║"
|
||||
echo "║ • Add a #[test] or #[tokio::test] that catches the bug ║"
|
||||
echo "║ • Add [skip-regression-check] to your commit message ║"
|
||||
echo "╚══════════════════════════════════════════════════════════════╝"
|
||||
echo ""
|
||||
exit 1
|
||||
Executable
+101
@@ -0,0 +1,101 @@
|
||||
#!/usr/bin/env bash
|
||||
# Generate an HTML coverage report for a given set of tests.
|
||||
#
|
||||
# Usage:
|
||||
# ./scripts/coverage.sh # all tests (lib only)
|
||||
# ./scripts/coverage.sh safety # tests matching "safety"
|
||||
# ./scripts/coverage.sh safety::sanitizer # specific module tests
|
||||
# ./scripts/coverage.sh test_a test_b test_c # multiple test filters
|
||||
#
|
||||
# Options (env vars):
|
||||
# COV_OPEN=1 Auto-open the report in a browser (default: 1)
|
||||
# COV_FORMAT=html Output format: html, text, json, lcov (default: html)
|
||||
# COV_OUT=coverage Output directory (default: coverage/)
|
||||
# COV_FEATURES="" Extra --features to pass (default: none)
|
||||
# COV_ALL_TARGETS=0 Set to 1 to include integration tests (default: lib only)
|
||||
#
|
||||
# Requires: cargo-llvm-cov (install: cargo install cargo-llvm-cov)
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
COV_OPEN="${COV_OPEN:-1}"
|
||||
COV_FORMAT="${COV_FORMAT:-html}"
|
||||
COV_OUT="${COV_OUT:-coverage}"
|
||||
COV_FEATURES="${COV_FEATURES:-}"
|
||||
COV_ALL_TARGETS="${COV_ALL_TARGETS:-0}"
|
||||
|
||||
cd "$(git rev-parse --show-toplevel)"
|
||||
|
||||
if ! command -v cargo-llvm-cov &>/dev/null; then
|
||||
echo "ERROR: cargo-llvm-cov not found. Install with: cargo install cargo-llvm-cov"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# Clean stale profiling data to avoid "mismatched data" warnings.
|
||||
cargo llvm-cov clean --workspace 2>/dev/null || true
|
||||
|
||||
# Build the cargo llvm-cov command
|
||||
cmd=(cargo llvm-cov)
|
||||
|
||||
# Features
|
||||
if [[ -n "$COV_FEATURES" ]]; then
|
||||
cmd+=(--features "$COV_FEATURES")
|
||||
else
|
||||
cmd+=(--all-features)
|
||||
fi
|
||||
|
||||
# By default, only run the lib unit tests (fast, no integration test compilation).
|
||||
# Set COV_ALL_TARGETS=1 to include integration tests.
|
||||
if [[ "$COV_ALL_TARGETS" != "1" ]]; then
|
||||
cmd+=(--lib)
|
||||
fi
|
||||
|
||||
# Output format
|
||||
case "$COV_FORMAT" in
|
||||
html)
|
||||
cmd+=(--html --output-dir "$COV_OUT")
|
||||
;;
|
||||
text)
|
||||
cmd+=(--text)
|
||||
;;
|
||||
json)
|
||||
cmd+=(--json --output-path "$COV_OUT/coverage.json")
|
||||
;;
|
||||
lcov)
|
||||
cmd+=(--lcov --output-path "$COV_OUT/lcov.info")
|
||||
;;
|
||||
*)
|
||||
echo "ERROR: Unknown format '$COV_FORMAT'. Use: html, text, json, lcov"
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
|
||||
# Test name filters (passed after -- to cargo test)
|
||||
if [[ $# -gt 0 ]]; then
|
||||
if [[ $# -eq 1 ]]; then
|
||||
cmd+=(-- "$1")
|
||||
else
|
||||
# Join filters with | for regex matching
|
||||
filter=$(IFS='|'; echo "$*")
|
||||
cmd+=(-- "$filter")
|
||||
fi
|
||||
fi
|
||||
|
||||
echo "Running: ${cmd[*]}"
|
||||
echo ""
|
||||
|
||||
"${cmd[@]}"
|
||||
|
||||
# Open report
|
||||
if [[ "$COV_FORMAT" == "html" && "$COV_OPEN" == "1" ]]; then
|
||||
index="$COV_OUT/html/index.html"
|
||||
if [[ -f "$index" ]]; then
|
||||
echo ""
|
||||
echo "Report: $index"
|
||||
if command -v open &>/dev/null; then
|
||||
open "$index"
|
||||
elif command -v xdg-open &>/dev/null; then
|
||||
xdg-open "$index"
|
||||
fi
|
||||
fi
|
||||
fi
|
||||
+17
-5
@@ -24,14 +24,14 @@ if ! command -v rustup &>/dev/null; then
|
||||
echo "ERROR: rustup not found. Install from https://rustup.rs"
|
||||
exit 1
|
||||
fi
|
||||
echo "[1/5] rustup found: $(rustup --version 2>/dev/null | head -1)"
|
||||
echo "[1/6] rustup found: $(rustup --version 2>/dev/null | head -1)"
|
||||
|
||||
# 2. Add WASM target (required by build.rs for channel compilation)
|
||||
echo "[2/5] Adding wasm32-wasip2 target..."
|
||||
echo "[2/6] Adding wasm32-wasip2 target..."
|
||||
rustup target add wasm32-wasip2
|
||||
|
||||
# 3. Install wasm-tools (required by build.rs for WASM component model)
|
||||
echo "[3/5] Installing wasm-tools..."
|
||||
echo "[3/6] Installing wasm-tools..."
|
||||
if command -v wasm-tools &>/dev/null; then
|
||||
echo " wasm-tools already installed: $(wasm-tools --version)"
|
||||
else
|
||||
@@ -39,13 +39,25 @@ else
|
||||
fi
|
||||
|
||||
# 4. Verify the project compiles
|
||||
echo "[4/5] Running cargo check..."
|
||||
echo "[4/6] Running cargo check..."
|
||||
cargo check
|
||||
|
||||
# 5. Run tests using libsql temp DB (no Docker/external DB needed)
|
||||
echo "[5/5] Running tests (no external DB required)..."
|
||||
echo "[5/6] Running tests (no external DB required)..."
|
||||
cargo test
|
||||
|
||||
# 6. Install git hooks
|
||||
echo "[6/6] Installing git hooks..."
|
||||
HOOKS_DIR=$(git rev-parse --git-path hooks 2>/dev/null) || true
|
||||
if [ -n "$HOOKS_DIR" ]; then
|
||||
mkdir -p "$HOOKS_DIR"
|
||||
SCRIPT_ABS="$(cd "$(dirname "$0")" && pwd)/commit-msg-regression.sh"
|
||||
ln -sf "$SCRIPT_ABS" "$HOOKS_DIR/commit-msg"
|
||||
echo " commit-msg hook installed (regression test enforcement)"
|
||||
else
|
||||
echo " Skipped: not a git repository"
|
||||
fi
|
||||
|
||||
echo ""
|
||||
echo "=== Setup complete ==="
|
||||
echo ""
|
||||
|
||||
@@ -0,0 +1,225 @@
|
||||
---
|
||||
name: local-test
|
||||
version: 0.1.0
|
||||
description: Build, run, and test IronClaw locally using Docker containers and Chrome MCP browser automation.
|
||||
activation:
|
||||
keywords:
|
||||
- test locally
|
||||
- local test
|
||||
- docker test
|
||||
- test my changes
|
||||
- test in docker
|
||||
- test web gateway
|
||||
- spin up test
|
||||
- test container
|
||||
patterns:
|
||||
- "test.*local"
|
||||
- "docker.*test"
|
||||
- "spin.*up.*test"
|
||||
- "test.*changes.*docker"
|
||||
max_context_tokens: 3000
|
||||
---
|
||||
|
||||
# Local Testing with Docker + Chrome MCP
|
||||
|
||||
Use this skill to build, run, and test IronClaw web gateway changes locally using `Dockerfile.test` and Chrome MCP browser automation tools.
|
||||
|
||||
## Quick Start
|
||||
|
||||
```bash
|
||||
# Build the test image (libsql-only, no PostgreSQL needed)
|
||||
docker build --platform linux/amd64 -f Dockerfile.test -t ironclaw-test .
|
||||
|
||||
# Run on port 3003 (default)
|
||||
docker run --rm -p 3003:3003 \
|
||||
-e ONBOARD_COMPLETED=true \
|
||||
-e CLI_ENABLED=false \
|
||||
-e NEARAI_API_KEY=<key> \
|
||||
ironclaw-test
|
||||
|
||||
# Open in browser
|
||||
# http://localhost:3003/?token=test
|
||||
```
|
||||
|
||||
## Building the Image
|
||||
|
||||
The test Dockerfile uses a two-stage build: Rust compilation with `--features libsql` (no PostgreSQL dependency), then a minimal Debian runtime image.
|
||||
|
||||
```bash
|
||||
docker build --platform linux/amd64 -f Dockerfile.test -t ironclaw-test .
|
||||
```
|
||||
|
||||
Build takes ~5-10 minutes on first run (cached subsequent builds are faster). The `--platform linux/amd64` flag avoids QEMU warnings on Apple Silicon but can be omitted if targeting native architecture.
|
||||
|
||||
## Running Containers
|
||||
|
||||
### Required Environment Variables
|
||||
|
||||
| Variable | Purpose | Default in Dockerfile |
|
||||
|----------|---------|----------------------|
|
||||
| `ONBOARD_COMPLETED=true` | Skip onboarding wizard (exits immediately otherwise) | not set |
|
||||
| `CLI_ENABLED=false` | Disable TUI/REPL (causes EOF shutdown otherwise) | not set |
|
||||
|
||||
### LLM Backend Configuration
|
||||
|
||||
Pick ONE of these configurations:
|
||||
|
||||
**NEAR AI (API key mode):**
|
||||
```bash
|
||||
docker run --rm -p 3003:3003 \
|
||||
-e ONBOARD_COMPLETED=true \
|
||||
-e CLI_ENABLED=false \
|
||||
-e NEARAI_API_KEY=<your-key> \
|
||||
ironclaw-test
|
||||
```
|
||||
|
||||
**NEAR AI (session token mode):**
|
||||
```bash
|
||||
docker run --rm -p 3003:3003 \
|
||||
-e ONBOARD_COMPLETED=true \
|
||||
-e CLI_ENABLED=false \
|
||||
-e NEARAI_SESSION_TOKEN=<sess_xxx> \
|
||||
-e NEARAI_BASE_URL=https://private.near.ai \
|
||||
ironclaw-test
|
||||
```
|
||||
|
||||
**OpenAI:**
|
||||
```bash
|
||||
docker run --rm -p 3003:3003 \
|
||||
-e ONBOARD_COMPLETED=true \
|
||||
-e CLI_ENABLED=false \
|
||||
-e LLM_BACKEND=openai \
|
||||
-e OPENAI_API_KEY=<your-key> \
|
||||
ironclaw-test
|
||||
```
|
||||
|
||||
**Anthropic:**
|
||||
```bash
|
||||
docker run --rm -p 3003:3003 \
|
||||
-e ONBOARD_COMPLETED=true \
|
||||
-e CLI_ENABLED=false \
|
||||
-e LLM_BACKEND=anthropic \
|
||||
-e ANTHROPIC_API_KEY=<your-key> \
|
||||
ironclaw-test
|
||||
```
|
||||
|
||||
**Dummy run (no LLM, just test the UI loads):**
|
||||
```bash
|
||||
docker run --rm -p 3003:3003 \
|
||||
-e ONBOARD_COMPLETED=true \
|
||||
-e CLI_ENABLED=false \
|
||||
-e NEARAI_API_KEY=dummy \
|
||||
ironclaw-test
|
||||
```
|
||||
|
||||
### Common Overrides
|
||||
|
||||
| Variable | Purpose | Example |
|
||||
|----------|---------|---------|
|
||||
| `GATEWAY_PORT` | Change the listen port | `3003` (default) |
|
||||
| `GATEWAY_AUTH_TOKEN` | Auth token for API | `test` (default) |
|
||||
| `NEARAI_MODEL` | Override LLM model | `claude-3-5-sonnet-20241022` |
|
||||
| `RUST_LOG` | Logging verbosity | `ironclaw=debug` |
|
||||
| `ROUTINES_ENABLED` | Enable routines | `true`/`false` |
|
||||
| `SKILLS_ENABLED` | Enable skills system | `true` (default) |
|
||||
|
||||
### Multi-Instance Testing
|
||||
|
||||
Run multiple containers on different host ports:
|
||||
|
||||
```bash
|
||||
docker run --rm -d --name ic-test-a -p 3003:3003 -e ONBOARD_COMPLETED=true -e CLI_ENABLED=false -e NEARAI_API_KEY=dummy ironclaw-test
|
||||
docker run --rm -d --name ic-test-b -p 3004:3003 -e ONBOARD_COMPLETED=true -e CLI_ENABLED=false -e NEARAI_API_KEY=dummy ironclaw-test
|
||||
```
|
||||
|
||||
## Chrome MCP Testing Workflow
|
||||
|
||||
Use the Claude for Chrome browser automation tools to test the web UI.
|
||||
|
||||
### Step 1: Get Browser Context
|
||||
|
||||
```
|
||||
mcp__claude-in-chrome__tabs_context_mcp
|
||||
```
|
||||
|
||||
Always start here to see current tabs and get fresh tab IDs.
|
||||
|
||||
### Step 2: Open the Gateway
|
||||
|
||||
```
|
||||
mcp__claude-in-chrome__tabs_create_mcp url=http://localhost:3003/?token=test
|
||||
```
|
||||
|
||||
### Step 3: Verify the Page
|
||||
|
||||
```
|
||||
mcp__claude-in-chrome__read_page
|
||||
```
|
||||
|
||||
Check for:
|
||||
- "Connected" indicator in top-right
|
||||
- All tabs visible: Chat, Memory, Jobs, Routines, Extensions, Skills
|
||||
|
||||
### Step 4: Take Screenshots
|
||||
|
||||
```
|
||||
mcp__claude-in-chrome__computer action=screenshot
|
||||
```
|
||||
|
||||
### Step 5: Test Mobile Viewport
|
||||
|
||||
```
|
||||
mcp__claude-in-chrome__resize_window width=375 height=812
|
||||
mcp__claude-in-chrome__computer action=screenshot
|
||||
```
|
||||
|
||||
Reset to desktop:
|
||||
```
|
||||
mcp__claude-in-chrome__resize_window width=1280 height=800
|
||||
```
|
||||
|
||||
### Step 6: Run JavaScript Checks
|
||||
|
||||
```
|
||||
mcp__claude-in-chrome__javascript_tool script="document.querySelector('.connection-status')?.textContent"
|
||||
```
|
||||
|
||||
### Step 7: Test Interactions
|
||||
|
||||
Click tabs, send messages, search skills — use `computer` tool with `action=click` and coordinate-based clicks, or use `find` + `form_input` for text entry.
|
||||
|
||||
## Cleanup
|
||||
|
||||
```bash
|
||||
# Stop a specific container
|
||||
docker stop ic-test-a
|
||||
|
||||
# Stop all test containers
|
||||
docker ps --filter ancestor=ironclaw-test -q | xargs -r docker stop
|
||||
|
||||
# Remove the test image
|
||||
docker rmi ironclaw-test
|
||||
```
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
### Container exits immediately
|
||||
- **Missing `ONBOARD_COMPLETED=true`**: The onboarding wizard tries to read stdin, gets EOF, and exits.
|
||||
- **Missing `CLI_ENABLED=false`**: The REPL channel reads stdin, gets EOF, and shuts down the agent.
|
||||
|
||||
### "Model not found" or LLM errors
|
||||
- Check that your API key/token is valid and the model name is correct.
|
||||
- For NEAR AI session token mode, you also need `NEARAI_BASE_URL=https://private.near.ai`.
|
||||
|
||||
### Platform mismatch warnings on Apple Silicon
|
||||
- The `--platform linux/amd64` flag causes QEMU emulation warnings — these are harmless.
|
||||
- Alternatively, omit the flag and build natively if your dependencies support ARM64.
|
||||
|
||||
### Port already in use
|
||||
- The dev server defaults to port 3001; the test Dockerfile defaults to 3003 to avoid conflicts.
|
||||
- Use a different host port: `-p 3005:3003`.
|
||||
|
||||
### Cannot connect from browser
|
||||
- Verify `GATEWAY_HOST=0.0.0.0` (set by default in Dockerfile).
|
||||
- Check the container logs: `docker logs <container-id>`.
|
||||
- Make sure you include the token query param: `?token=test`.
|
||||
+22
-3
@@ -73,6 +73,10 @@ pub struct AgentDeps {
|
||||
pub hooks: Arc<HookRegistry>,
|
||||
/// Cost enforcement guardrails (daily budget, hourly rate limits).
|
||||
pub cost_guard: Arc<crate::agent::cost_guard::CostGuard>,
|
||||
/// SSE broadcast sender for live job event streaming to the web gateway.
|
||||
pub sse_tx: Option<tokio::sync::broadcast::Sender<crate::channels::web::types::SseEvent>>,
|
||||
/// HTTP interceptor for trace recording/replay.
|
||||
pub http_interceptor: Option<Arc<dyn crate::llm::recording::HttpInterceptor>>,
|
||||
}
|
||||
|
||||
/// The main agent that coordinates all components.
|
||||
@@ -111,7 +115,7 @@ impl Agent {
|
||||
|
||||
let session_manager = session_manager.unwrap_or_else(|| Arc::new(SessionManager::new()));
|
||||
|
||||
let scheduler = Arc::new(Scheduler::new(
|
||||
let mut scheduler = Scheduler::new(
|
||||
config.clone(),
|
||||
context_manager.clone(),
|
||||
deps.llm.clone(),
|
||||
@@ -119,7 +123,11 @@ impl Agent {
|
||||
deps.tools.clone(),
|
||||
deps.store.clone(),
|
||||
deps.hooks.clone(),
|
||||
));
|
||||
);
|
||||
if let Some(ref tx) = deps.sse_tx {
|
||||
scheduler.set_sse_sender(tx.clone());
|
||||
}
|
||||
let scheduler = Arc::new(scheduler);
|
||||
|
||||
Self {
|
||||
config,
|
||||
@@ -627,6 +635,10 @@ impl Agent {
|
||||
|
||||
// Parse submission type first
|
||||
let mut submission = SubmissionParser::parse(&message.content);
|
||||
tracing::debug!(
|
||||
"[agent_loop] Parsed submission: {:?}",
|
||||
std::any::type_name_of_val(&submission)
|
||||
);
|
||||
|
||||
// Hook: BeforeInbound — allow hooks to modify or reject user input
|
||||
if let Submission::UserInput { ref content } = submission {
|
||||
@@ -711,7 +723,14 @@ impl Agent {
|
||||
.await
|
||||
}
|
||||
Submission::SystemCommand { command, args } => {
|
||||
self.handle_system_command(&command, &args).await
|
||||
tracing::debug!(
|
||||
"[agent_loop] SystemCommand: command={}, channel={}",
|
||||
command,
|
||||
message.channel
|
||||
);
|
||||
// Authorization checks (including restart channel check) are enforced in handle_system_command
|
||||
self.handle_system_command(&command, &args, &message.channel)
|
||||
.await
|
||||
}
|
||||
Submission::Undo => self.process_undo(session, thread_id).await,
|
||||
Submission::Redo => self.process_redo(session, thread_id).await,
|
||||
|
||||
+70
-2
@@ -68,7 +68,10 @@ impl Agent {
|
||||
self.handle_help_job(&message.user_id, &job_id).await?
|
||||
}
|
||||
MessageIntent::Command { command, args } => {
|
||||
match self.handle_command(&command, &args).await? {
|
||||
match self
|
||||
.handle_command(&command, &args, &message.channel)
|
||||
.await?
|
||||
{
|
||||
Some(s) => s,
|
||||
None => return Ok(SubmissionResult::Ok { message: None }), // Shutdown signal
|
||||
}
|
||||
@@ -466,6 +469,7 @@ impl Agent {
|
||||
&self,
|
||||
command: &str,
|
||||
args: &[String],
|
||||
channel: &str,
|
||||
) -> Result<SubmissionResult, Error> {
|
||||
match command {
|
||||
"help" => Ok(SubmissionResult::response(concat!(
|
||||
@@ -501,12 +505,75 @@ impl Agent {
|
||||
" /heartbeat Run heartbeat check\n",
|
||||
" /summarize Summarize current thread\n",
|
||||
" /suggest Suggest next steps\n",
|
||||
" /restart Gracefully restart the process\n",
|
||||
"\n",
|
||||
" /quit Exit",
|
||||
))),
|
||||
|
||||
"ping" => Ok(SubmissionResult::response("pong!")),
|
||||
|
||||
"restart" => {
|
||||
tracing::info!("[commands::restart] Restart command received");
|
||||
// Channel authorization check: restart is only available via web interface
|
||||
if channel != "gateway" {
|
||||
tracing::warn!(
|
||||
"[commands::restart] Restart rejected: not from gateway channel (from: {})",
|
||||
channel
|
||||
);
|
||||
return Ok(SubmissionResult::error(
|
||||
"Restart is only available through the web interface with explicit user confirmation. \
|
||||
Use the Restart button in the UI."
|
||||
.to_string(),
|
||||
));
|
||||
}
|
||||
// Environment check: restart is only available in Docker containers
|
||||
let in_docker = std::env::var("IRONCLAW_IN_DOCKER")
|
||||
.map(|v| v.to_lowercase() == "true")
|
||||
.unwrap_or(false);
|
||||
|
||||
tracing::debug!("[commands::restart] IRONCLAW_IN_DOCKER={}", in_docker);
|
||||
|
||||
if !in_docker {
|
||||
tracing::warn!(
|
||||
"[commands::restart] Restart rejected: not in Docker environment"
|
||||
);
|
||||
return Ok(SubmissionResult::error(
|
||||
"Restart is not available in this environment. \
|
||||
The IRONCLAW_IN_DOCKER environment variable must be set to 'true' for Docker deployments."
|
||||
.to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
// Execute restart tool directly (don't dispatch as a job for LLM planning)
|
||||
// This ensures the tool runs immediately without LLM involvement
|
||||
use crate::tools::Tool;
|
||||
let tool = crate::tools::builtin::RestartTool;
|
||||
let params = serde_json::json!({});
|
||||
|
||||
// Create a minimal JobContext for the tool
|
||||
let dummy_ctx =
|
||||
crate::context::JobContext::with_user("system", "Restart", "Graceful restart");
|
||||
|
||||
match tool.execute(params, &dummy_ctx).await {
|
||||
Ok(output) => {
|
||||
tracing::info!("[commands::restart] RestartTool executed successfully");
|
||||
// Extract text from the ToolOutput result
|
||||
let response = match output.result {
|
||||
serde_json::Value::String(s) => s,
|
||||
_ => output.result.to_string(),
|
||||
};
|
||||
Ok(SubmissionResult::response(response))
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!(
|
||||
"[commands::restart] RestartTool execution failed: {:?}",
|
||||
e
|
||||
);
|
||||
Ok(SubmissionResult::error(format!("Restart failed: {}", e)))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
"version" => Ok(SubmissionResult::response(format!(
|
||||
"{} v{}",
|
||||
env!("CARGO_PKG_NAME"),
|
||||
@@ -744,10 +811,11 @@ impl Agent {
|
||||
&self,
|
||||
command: &str,
|
||||
args: &[String],
|
||||
channel: &str,
|
||||
) -> Result<Option<String>, Error> {
|
||||
// System commands are now handled directly via Submission::SystemCommand,
|
||||
// but the router may still send us unknown /commands.
|
||||
match self.handle_system_command(command, args).await? {
|
||||
match self.handle_system_command(command, args, channel).await? {
|
||||
SubmissionResult::Response { content } => Ok(Some(content)),
|
||||
SubmissionResult::Ok { message } => Ok(message),
|
||||
SubmissionResult::Error { message } => Ok(Some(format!("Error: {}", message))),
|
||||
|
||||
+148
-19
@@ -15,6 +15,7 @@ use crate::channels::{IncomingMessage, StatusUpdate};
|
||||
use crate::context::JobContext;
|
||||
use crate::error::Error;
|
||||
use crate::llm::{ChatMessage, Reasoning, ReasoningContext, RespondResult};
|
||||
use crate::tools::redact_params;
|
||||
|
||||
/// Result of the agentic loop execution.
|
||||
pub(super) enum AgenticLoopResult {
|
||||
@@ -126,7 +127,9 @@ impl Agent {
|
||||
let mut context_messages = initial_messages;
|
||||
|
||||
// Create a JobContext for tool execution (chat doesn't have a real job)
|
||||
let job_ctx = JobContext::with_user(&message.user_id, "chat", "Interactive chat session");
|
||||
let mut job_ctx =
|
||||
JobContext::with_user(&message.user_id, "chat", "Interactive chat session");
|
||||
job_ctx.http_interceptor = self.deps.http_interceptor.clone();
|
||||
|
||||
let max_tool_iterations = self.config.max_tool_iterations;
|
||||
// Force a text-only response on the last iteration to guarantee termination
|
||||
@@ -291,7 +294,11 @@ impl Agent {
|
||||
|
||||
match output.result {
|
||||
RespondResult::Text(text) => {
|
||||
return Ok(AgenticLoopResult::Response(text));
|
||||
// Strip internal "[Called tool ...]" text that can leak when
|
||||
// provider flattening (e.g. NEAR AI) converts tool_calls to
|
||||
// plain text and the LLM echoes it back.
|
||||
let sanitized = strip_internal_tool_call_text(&text);
|
||||
return Ok(AgenticLoopResult::Response(sanitized));
|
||||
}
|
||||
RespondResult::ToolCalls {
|
||||
tool_calls,
|
||||
@@ -317,14 +324,25 @@ impl Agent {
|
||||
)
|
||||
.await;
|
||||
|
||||
// Record tool calls in the thread
|
||||
// Record tool calls in the thread with sensitive params redacted.
|
||||
// Look up each tool's sensitive_params before acquiring the session lock.
|
||||
{
|
||||
let mut redacted_args: Vec<serde_json::Value> =
|
||||
Vec::with_capacity(tool_calls.len());
|
||||
for tc in &tool_calls {
|
||||
let safe = if let Some(tool) = self.tools().get(&tc.name).await {
|
||||
redact_params(&tc.arguments, tool.sensitive_params())
|
||||
} else {
|
||||
tc.arguments.clone()
|
||||
};
|
||||
redacted_args.push(safe);
|
||||
}
|
||||
let mut sess = session.lock().await;
|
||||
if let Some(thread) = sess.threads.get_mut(&thread_id)
|
||||
&& let Some(turn) = thread.last_turn_mut()
|
||||
{
|
||||
for tc in &tool_calls {
|
||||
turn.record_tool_call(&tc.name, tc.arguments.clone());
|
||||
for (tc, safe_args) in tool_calls.iter().zip(redacted_args) {
|
||||
turn.record_tool_call(&tc.name, safe_args);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -353,11 +371,22 @@ impl Agent {
|
||||
for (idx, original_tc) in tool_calls.iter().enumerate() {
|
||||
let mut tc = original_tc.clone();
|
||||
|
||||
// Fetch the tool upfront so we can redact sensitive params
|
||||
// before they touch hooks or approval display.
|
||||
let tool_opt = self.tools().get(&tc.name).await;
|
||||
let sensitive = tool_opt
|
||||
.as_ref()
|
||||
.map(|t| t.sensitive_params())
|
||||
.unwrap_or(&[]);
|
||||
|
||||
// Hook: BeforeToolCall (runs before approval so hooks can
|
||||
// modify parameters — approval is checked on final params)
|
||||
// modify parameters — approval is checked on final params).
|
||||
// Hooks receive redacted params so sensitive values are not
|
||||
// exposed to hook handlers or their logs.
|
||||
let hook_params = redact_params(&tc.arguments, sensitive);
|
||||
let event = crate::hooks::HookEvent::ToolCall {
|
||||
tool_name: tc.name.clone(),
|
||||
parameters: tc.arguments.clone(),
|
||||
parameters: hook_params,
|
||||
user_id: message.user_id.clone(),
|
||||
context: "chat".to_string(),
|
||||
};
|
||||
@@ -384,8 +413,20 @@ impl Agent {
|
||||
}
|
||||
Ok(crate::hooks::HookOutcome::Continue {
|
||||
modified: Some(new_params),
|
||||
}) => match serde_json::from_str(&new_params) {
|
||||
Ok(parsed) => tc.arguments = parsed,
|
||||
}) => match serde_json::from_str::<serde_json::Value>(&new_params) {
|
||||
Ok(mut parsed) => {
|
||||
// Restore original sensitive param values so a hook
|
||||
// cannot overwrite them (they were sent as [REDACTED]).
|
||||
if let Some(obj) = parsed.as_object_mut() {
|
||||
for key in sensitive {
|
||||
if let Some(orig_val) = original_tc.arguments.get(*key)
|
||||
{
|
||||
obj.insert((*key).to_string(), orig_val.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
tc.arguments = parsed;
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
tool = %tc.name,
|
||||
@@ -400,7 +441,7 @@ impl Agent {
|
||||
// Check if tool requires approval on the final (post-hook)
|
||||
// parameters. Skipped when auto_approve_tools is set.
|
||||
if !self.config.auto_approve_tools
|
||||
&& let Some(tool) = self.tools().get(&tc.name).await
|
||||
&& let Some(tool) = tool_opt
|
||||
{
|
||||
use crate::tools::ApprovalRequirement;
|
||||
let needs_approval = match tool.requires_approval(&tc.arguments) {
|
||||
@@ -447,14 +488,17 @@ impl Agent {
|
||||
.execute_chat_tool(&tc.name, &tc.arguments, &job_ctx)
|
||||
.await;
|
||||
|
||||
let disp_tool = self.tools().get(&tc.name).await;
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::ToolCompleted {
|
||||
name: tc.name.clone(),
|
||||
success: result.is_ok(),
|
||||
},
|
||||
StatusUpdate::tool_completed(
|
||||
tc.name.clone(),
|
||||
&result,
|
||||
&tc.arguments,
|
||||
disp_tool.as_deref(),
|
||||
),
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
@@ -495,13 +539,16 @@ impl Agent {
|
||||
)
|
||||
.await;
|
||||
|
||||
let par_tool = tools.get(&tc.name).await;
|
||||
let _ = channels
|
||||
.send_status(
|
||||
&channel,
|
||||
StatusUpdate::ToolCompleted {
|
||||
name: tc.name.clone(),
|
||||
success: result.is_ok(),
|
||||
},
|
||||
StatusUpdate::tool_completed(
|
||||
tc.name.clone(),
|
||||
&result,
|
||||
&tc.arguments,
|
||||
par_tool.as_deref(),
|
||||
),
|
||||
&metadata,
|
||||
)
|
||||
.await;
|
||||
@@ -641,6 +688,15 @@ impl Agent {
|
||||
deferred_auth = Some(instructions);
|
||||
}
|
||||
|
||||
// Stash full output so subsequent tools can reference it
|
||||
if let Ok(ref output) = tool_result {
|
||||
job_ctx
|
||||
.tool_output_stash
|
||||
.write()
|
||||
.await
|
||||
.insert(tc.id.clone(), output.clone());
|
||||
}
|
||||
|
||||
// Sanitize and add tool result to context
|
||||
let result_content = match tool_result {
|
||||
Ok(output) => {
|
||||
@@ -671,10 +727,15 @@ impl Agent {
|
||||
|
||||
// Handle approval if a tool needed it
|
||||
if let Some((approval_idx, tc, tool)) = approval_needed {
|
||||
// Show redacted params in the approval UI — the user already knows
|
||||
// the sensitive value (they provided it); showing it again is
|
||||
// unnecessary and creates a leakage path through channel logs.
|
||||
let display_params = redact_params(&tc.arguments, tool.sensitive_params());
|
||||
let pending = PendingApproval {
|
||||
request_id: Uuid::new_v4(),
|
||||
tool_name: tc.name.clone(),
|
||||
parameters: tc.arguments.clone(),
|
||||
display_parameters: display_params,
|
||||
description: tool.description().to_string(),
|
||||
tool_call_id: tc.id.clone(),
|
||||
context_messages: context_messages.clone(),
|
||||
@@ -734,9 +795,10 @@ pub(super) async fn execute_chat_tool_standalone(
|
||||
.into());
|
||||
}
|
||||
|
||||
let safe_params = redact_params(params, tool.sensitive_params());
|
||||
tracing::debug!(
|
||||
tool = %tool_name,
|
||||
params = %params,
|
||||
params = %safe_params,
|
||||
"Tool call started"
|
||||
);
|
||||
|
||||
@@ -900,6 +962,38 @@ fn compact_messages_for_retry(messages: &[ChatMessage]) -> Vec<ChatMessage> {
|
||||
compacted
|
||||
}
|
||||
|
||||
/// Strip internal `[Called tool ...]` and `[Tool ... returned: ...]` markers
|
||||
/// from a response string. These markers are inserted by provider-level message
|
||||
/// flattening (e.g. NEAR AI) and can leak into the user-visible response when
|
||||
/// the LLM echoes them back.
|
||||
fn strip_internal_tool_call_text(text: &str) -> String {
|
||||
// Remove lines that are purely internal tool-call markers.
|
||||
// Pattern: lines matching `[Called tool <name>(...)]` or `[Tool <name> returned: ...]`
|
||||
let result = text
|
||||
.lines()
|
||||
.filter(|line| {
|
||||
let trimmed = line.trim();
|
||||
!((trimmed.starts_with("[Called tool ") && trimmed.ends_with(']'))
|
||||
|| (trimmed.starts_with("[Tool ")
|
||||
&& trimmed.contains(" returned:")
|
||||
&& trimmed.ends_with(']')))
|
||||
})
|
||||
.fold(String::new(), |mut acc, s| {
|
||||
if !acc.is_empty() {
|
||||
acc.push('\n');
|
||||
}
|
||||
acc.push_str(s);
|
||||
acc
|
||||
});
|
||||
|
||||
let result = result.trim();
|
||||
if result.is_empty() {
|
||||
"I wasn't able to complete that request. Could you try rephrasing or providing more details?".to_string()
|
||||
} else {
|
||||
result.to_string()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
@@ -982,6 +1076,8 @@ mod tests {
|
||||
skills_config: SkillsConfig::default(),
|
||||
hooks: Arc::new(HookRegistry::new()),
|
||||
cost_guard: Arc::new(CostGuard::new(CostGuardConfig::default())),
|
||||
sse_tx: None,
|
||||
http_interceptor: None,
|
||||
};
|
||||
|
||||
Agent::new(
|
||||
@@ -1085,6 +1181,7 @@ mod tests {
|
||||
request_id: uuid::Uuid::new_v4(),
|
||||
tool_name: "shell".to_string(),
|
||||
parameters: serde_json::json!({"command": "echo hi"}),
|
||||
display_parameters: serde_json::json!({"command": "echo hi"}),
|
||||
description: "Run shell command".to_string(),
|
||||
tool_call_id: "call_1".to_string(),
|
||||
context_messages: vec![],
|
||||
@@ -1719,6 +1816,8 @@ mod tests {
|
||||
skills_config: SkillsConfig::default(),
|
||||
hooks: Arc::new(HookRegistry::new()),
|
||||
cost_guard: Arc::new(CostGuard::new(CostGuardConfig::default())),
|
||||
sse_tx: None,
|
||||
http_interceptor: None,
|
||||
};
|
||||
|
||||
Agent::new(
|
||||
@@ -1830,6 +1929,8 @@ mod tests {
|
||||
skills_config: SkillsConfig::default(),
|
||||
hooks: Arc::new(HookRegistry::new()),
|
||||
cost_guard: Arc::new(CostGuard::new(CostGuardConfig::default())),
|
||||
sse_tx: None,
|
||||
http_interceptor: None,
|
||||
};
|
||||
|
||||
Agent::new(
|
||||
@@ -1899,4 +2000,32 @@ mod tests {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_strip_internal_tool_call_text_removes_markers() {
|
||||
let input = "[Called tool search({\"query\": \"test\"})]\nHere is the answer.";
|
||||
let result = super::strip_internal_tool_call_text(input);
|
||||
assert_eq!(result, "Here is the answer.");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_strip_internal_tool_call_text_removes_returned_markers() {
|
||||
let input = "[Tool search returned: some result]\nSummary of findings.";
|
||||
let result = super::strip_internal_tool_call_text(input);
|
||||
assert_eq!(result, "Summary of findings.");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_strip_internal_tool_call_text_all_markers_yields_fallback() {
|
||||
let input = "[Called tool search({\"query\": \"test\"})]\n[Tool search returned: error]";
|
||||
let result = super::strip_internal_tool_call_text(input);
|
||||
assert!(result.contains("wasn't able to complete"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_strip_internal_tool_call_text_preserves_normal_text() {
|
||||
let input = "This is a normal response with [brackets] inside.";
|
||||
let result = super::strip_internal_tool_call_text(input);
|
||||
assert_eq!(result, input);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -10,6 +10,7 @@ use uuid::Uuid;
|
||||
|
||||
use crate::agent::task::{Task, TaskContext, TaskOutput};
|
||||
use crate::agent::worker::{Worker, WorkerDeps};
|
||||
use crate::channels::web::types::SseEvent;
|
||||
use crate::config::AgentConfig;
|
||||
use crate::context::{ContextManager, JobContext, JobState};
|
||||
use crate::db::Database;
|
||||
@@ -28,6 +29,8 @@ pub enum WorkerMessage {
|
||||
Stop,
|
||||
/// Check health.
|
||||
Ping,
|
||||
/// Inject a follow-up user message into the worker's reasoning context.
|
||||
UserMessage(String),
|
||||
}
|
||||
|
||||
/// Status of a scheduled job.
|
||||
@@ -51,6 +54,8 @@ pub struct Scheduler {
|
||||
tools: Arc<ToolRegistry>,
|
||||
store: Option<Arc<dyn Database>>,
|
||||
hooks: Arc<HookRegistry>,
|
||||
/// SSE broadcast sender for live job event streaming.
|
||||
sse_tx: Option<tokio::sync::broadcast::Sender<SseEvent>>,
|
||||
/// Running jobs (main LLM-driven jobs).
|
||||
jobs: Arc<RwLock<HashMap<Uuid, ScheduledJob>>>,
|
||||
/// Running sub-tasks (tool executions, background tasks).
|
||||
@@ -76,11 +81,17 @@ impl Scheduler {
|
||||
tools,
|
||||
store,
|
||||
hooks,
|
||||
sse_tx: None,
|
||||
jobs: Arc::new(RwLock::new(HashMap::new())),
|
||||
subtasks: Arc::new(RwLock::new(HashMap::new())),
|
||||
}
|
||||
}
|
||||
|
||||
/// Set the SSE broadcast sender for live job event streaming.
|
||||
pub fn set_sse_sender(&mut self, tx: tokio::sync::broadcast::Sender<SseEvent>) {
|
||||
self.sse_tx = Some(tx);
|
||||
}
|
||||
|
||||
/// Create, persist, and schedule a job in one shot.
|
||||
///
|
||||
/// This is the preferred entry point for dispatching new jobs. It:
|
||||
@@ -169,6 +180,7 @@ impl Scheduler {
|
||||
hooks: self.hooks.clone(),
|
||||
timeout: self.config.job_timeout,
|
||||
use_planning: self.config.use_planning,
|
||||
sse_tx: self.sse_tx.clone(),
|
||||
};
|
||||
let worker = Worker::new(job_id, deps);
|
||||
|
||||
@@ -500,6 +512,26 @@ impl Scheduler {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Send a follow-up user message to a running job.
|
||||
///
|
||||
/// Returns `Ok(())` if the message was queued, `Err` if the job is not running.
|
||||
pub async fn send_message(&self, job_id: Uuid, content: String) -> Result<(), JobError> {
|
||||
// Clone the sender while holding the lock, then release before the
|
||||
// async send to avoid blocking scheduler writes during backpressure.
|
||||
let tx = {
|
||||
let jobs = self.jobs.read().await;
|
||||
let scheduled = jobs.get(&job_id).ok_or(JobError::NotFound { id: job_id })?;
|
||||
scheduled.tx.clone()
|
||||
};
|
||||
tx.send(WorkerMessage::UserMessage(content))
|
||||
.await
|
||||
.map_err(|_| JobError::Failed {
|
||||
id: job_id,
|
||||
reason: "Worker channel closed".to_string(),
|
||||
})?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Check if a job is running.
|
||||
pub async fn is_running(&self, job_id: Uuid) -> bool {
|
||||
self.jobs.read().await.contains_key(&job_id)
|
||||
|
||||
@@ -148,8 +148,12 @@ pub struct PendingApproval {
|
||||
pub request_id: Uuid,
|
||||
/// Tool name requiring approval.
|
||||
pub tool_name: String,
|
||||
/// Tool parameters.
|
||||
/// Tool parameters (original values, used for execution).
|
||||
pub parameters: serde_json::Value,
|
||||
/// Redacted tool parameters (sensitive values replaced with `[REDACTED]`).
|
||||
/// Used for display in approval UI, logs, and SSE broadcasts.
|
||||
#[serde(default)]
|
||||
pub display_parameters: serde_json::Value,
|
||||
/// Description of what the tool will do.
|
||||
pub description: String,
|
||||
/// Tool call ID from LLM (for proper context continuation).
|
||||
@@ -950,6 +954,7 @@ mod tests {
|
||||
request_id: Uuid::new_v4(),
|
||||
tool_name: "shell".to_string(),
|
||||
parameters: serde_json::json!({"command": "rm -rf /"}),
|
||||
display_parameters: serde_json::json!({"command": "rm -rf /"}),
|
||||
description: "dangerous command".to_string(),
|
||||
tool_call_id: "call_123".to_string(),
|
||||
context_messages: vec![ChatMessage::user("do it")],
|
||||
@@ -974,6 +979,7 @@ mod tests {
|
||||
request_id: Uuid::new_v4(),
|
||||
tool_name: "http".to_string(),
|
||||
parameters: serde_json::json!({}),
|
||||
display_parameters: serde_json::json!({}),
|
||||
description: "test".to_string(),
|
||||
tool_call_id: "call_456".to_string(),
|
||||
context_messages: vec![],
|
||||
|
||||
@@ -14,6 +14,7 @@ impl SubmissionParser {
|
||||
pub fn parse(content: &str) -> Submission {
|
||||
let trimmed = content.trim();
|
||||
let lower = trimmed.to_lowercase();
|
||||
tracing::debug!("[SubmissionParser::parse] Parsing input: {:?}", trimmed);
|
||||
|
||||
// Control commands (exact match or prefix)
|
||||
if lower == "/undo" {
|
||||
@@ -91,6 +92,13 @@ impl SubmissionParser {
|
||||
args: vec![],
|
||||
};
|
||||
}
|
||||
if lower == "/restart" {
|
||||
tracing::debug!("[SubmissionParser::parse] Recognized /restart command");
|
||||
return Submission::SystemCommand {
|
||||
command: "restart".to_string(),
|
||||
args: vec![],
|
||||
};
|
||||
}
|
||||
if lower.starts_with("/model") {
|
||||
let args: Vec<String> = trimmed
|
||||
.split_whitespace()
|
||||
|
||||
+33
-21
@@ -21,6 +21,7 @@ use crate::channels::{IncomingMessage, StatusUpdate};
|
||||
use crate::context::JobContext;
|
||||
use crate::error::Error;
|
||||
use crate::llm::ChatMessage;
|
||||
use crate::tools::redact_params;
|
||||
|
||||
impl Agent {
|
||||
/// Hydrate a historical thread from DB into memory if not already present.
|
||||
@@ -357,7 +358,7 @@ impl Agent {
|
||||
let request_id = pending.request_id;
|
||||
let tool_name = pending.tool_name.clone();
|
||||
let description = pending.description.clone();
|
||||
let parameters = pending.parameters.clone();
|
||||
let parameters = pending.display_parameters.clone();
|
||||
thread.await_approval(pending);
|
||||
let _ = self
|
||||
.channels
|
||||
@@ -733,8 +734,9 @@ impl Agent {
|
||||
}
|
||||
|
||||
// Execute the approved tool and continue the loop
|
||||
let job_ctx =
|
||||
let mut job_ctx =
|
||||
JobContext::with_user(&message.user_id, "chat", "Interactive chat session");
|
||||
job_ctx.http_interceptor = self.deps.http_interceptor.clone();
|
||||
|
||||
let _ = self
|
||||
.channels
|
||||
@@ -751,14 +753,17 @@ impl Agent {
|
||||
.execute_chat_tool(&pending.tool_name, &pending.parameters, &job_ctx)
|
||||
.await;
|
||||
|
||||
let tool_ref = self.tools().get(&pending.tool_name).await;
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::ToolCompleted {
|
||||
name: pending.tool_name.clone(),
|
||||
success: tool_result.is_ok(),
|
||||
},
|
||||
StatusUpdate::tool_completed(
|
||||
pending.tool_name.clone(),
|
||||
&tool_result,
|
||||
&pending.display_parameters,
|
||||
tool_ref.as_deref(),
|
||||
),
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
@@ -908,14 +913,17 @@ impl Agent {
|
||||
.execute_chat_tool(&tc.name, &tc.arguments, &job_ctx)
|
||||
.await;
|
||||
|
||||
let deferred_tool = self.tools().get(&tc.name).await;
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::ToolCompleted {
|
||||
name: tc.name.clone(),
|
||||
success: result.is_ok(),
|
||||
},
|
||||
StatusUpdate::tool_completed(
|
||||
tc.name.clone(),
|
||||
&result,
|
||||
&tc.arguments,
|
||||
deferred_tool.as_deref(),
|
||||
),
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
@@ -957,13 +965,16 @@ impl Agent {
|
||||
)
|
||||
.await;
|
||||
|
||||
let par_tool = tools.get(&tc.name).await;
|
||||
let _ = channels
|
||||
.send_status(
|
||||
&channel,
|
||||
StatusUpdate::ToolCompleted {
|
||||
name: tc.name.clone(),
|
||||
success: result.is_ok(),
|
||||
},
|
||||
StatusUpdate::tool_completed(
|
||||
tc.name.clone(),
|
||||
&result,
|
||||
&tc.arguments,
|
||||
par_tool.as_deref(),
|
||||
),
|
||||
&metadata,
|
||||
)
|
||||
.await;
|
||||
@@ -1086,6 +1097,7 @@ impl Agent {
|
||||
request_id: Uuid::new_v4(),
|
||||
tool_name: tc.name.clone(),
|
||||
parameters: tc.arguments.clone(),
|
||||
display_parameters: redact_params(&tc.arguments, tool.sensitive_params()),
|
||||
description: tool.description().to_string(),
|
||||
tool_call_id: tc.id.clone(),
|
||||
context_messages: context_messages.clone(),
|
||||
@@ -1095,7 +1107,7 @@ impl Agent {
|
||||
let request_id = new_pending.request_id;
|
||||
let tool_name = new_pending.tool_name.clone();
|
||||
let description = new_pending.description.clone();
|
||||
let parameters = new_pending.parameters.clone();
|
||||
let parameters = new_pending.display_parameters.clone();
|
||||
|
||||
{
|
||||
let mut sess = session.lock().await;
|
||||
@@ -1162,7 +1174,7 @@ impl Agent {
|
||||
let request_id = new_pending.request_id;
|
||||
let tool_name = new_pending.tool_name.clone();
|
||||
let description = new_pending.description.clone();
|
||||
let parameters = new_pending.parameters.clone();
|
||||
let parameters = new_pending.display_parameters.clone();
|
||||
thread.await_approval(new_pending);
|
||||
let _ = self
|
||||
.channels
|
||||
@@ -1284,7 +1296,7 @@ impl Agent {
|
||||
};
|
||||
|
||||
match ext_mgr.auth(&pending.extension_name, Some(token)).await {
|
||||
Ok(result) if result.status == "authenticated" => {
|
||||
Ok(result) if result.is_authenticated() => {
|
||||
tracing::info!(
|
||||
"Extension '{}' authenticated via auth mode",
|
||||
pending.extension_name
|
||||
@@ -1353,8 +1365,8 @@ impl Agent {
|
||||
}
|
||||
}
|
||||
let msg = result
|
||||
.instructions
|
||||
.clone()
|
||||
.instructions()
|
||||
.map(String::from)
|
||||
.unwrap_or_else(|| "Invalid token. Please try again.".to_string());
|
||||
// Re-emit AuthRequired so web UI re-shows the card
|
||||
let _ = self
|
||||
@@ -1364,8 +1376,8 @@ impl Agent {
|
||||
StatusUpdate::AuthRequired {
|
||||
extension_name: pending.extension_name.clone(),
|
||||
instructions: Some(msg.clone()),
|
||||
auth_url: result.auth_url,
|
||||
setup_url: result.setup_url,
|
||||
auth_url: result.auth_url().map(String::from),
|
||||
setup_url: result.setup_url().map(String::from),
|
||||
},
|
||||
&message.metadata,
|
||||
)
|
||||
|
||||
+220
-25
@@ -9,6 +9,7 @@ use uuid::Uuid;
|
||||
|
||||
use crate::agent::scheduler::WorkerMessage;
|
||||
use crate::agent::task::TaskOutput;
|
||||
use crate::channels::web::types::SseEvent;
|
||||
use crate::context::{ContextManager, JobState};
|
||||
use crate::db::Database;
|
||||
use crate::error::Error;
|
||||
@@ -17,8 +18,8 @@ use crate::llm::{
|
||||
ActionPlan, ChatMessage, LlmProvider, Reasoning, ReasoningContext, RespondResult, ToolSelection,
|
||||
};
|
||||
use crate::safety::SafetyLayer;
|
||||
use crate::tools::ToolRegistry;
|
||||
use crate::tools::rate_limiter::RateLimitResult;
|
||||
use crate::tools::{ToolRegistry, redact_params};
|
||||
|
||||
/// Shared dependencies for worker execution.
|
||||
///
|
||||
@@ -34,6 +35,8 @@ pub struct WorkerDeps {
|
||||
pub hooks: Arc<HookRegistry>,
|
||||
pub timeout: Duration,
|
||||
pub use_planning: bool,
|
||||
/// SSE broadcast sender for live job event streaming to the web gateway.
|
||||
pub sse_tx: Option<tokio::sync::broadcast::Sender<SseEvent>>,
|
||||
}
|
||||
|
||||
/// Worker that executes a single job.
|
||||
@@ -98,18 +101,90 @@ impl Worker {
|
||||
}
|
||||
}
|
||||
|
||||
/// Fire-and-forget persistence of a job event.
|
||||
/// Fire-and-forget persistence of a job event and SSE broadcast.
|
||||
fn log_event(&self, event_type: &str, data: serde_json::Value) {
|
||||
let job_id = self.job_id;
|
||||
|
||||
// Persist to DB
|
||||
if let Some(store) = self.store() {
|
||||
let store = store.clone();
|
||||
let job_id = self.job_id;
|
||||
let event_type = event_type.to_string();
|
||||
let et = event_type.to_string();
|
||||
let d = data.clone();
|
||||
tokio::spawn(async move {
|
||||
if let Err(e) = store.save_job_event(job_id, &event_type, &data).await {
|
||||
if let Err(e) = store.save_job_event(job_id, &et, &d).await {
|
||||
tracing::warn!("Failed to persist event for job {}: {}", job_id, e);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
// Broadcast SSE for live web UI updates
|
||||
if let Some(ref tx) = self.deps.sse_tx {
|
||||
let job_id_str = job_id.to_string();
|
||||
let event = match event_type {
|
||||
"message" => Some(SseEvent::JobMessage {
|
||||
job_id: job_id_str,
|
||||
role: data
|
||||
.get("role")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("assistant")
|
||||
.to_string(),
|
||||
content: data
|
||||
.get("content")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("")
|
||||
.to_string(),
|
||||
}),
|
||||
"tool_use" => Some(SseEvent::JobToolUse {
|
||||
job_id: job_id_str,
|
||||
tool_name: data
|
||||
.get("tool_name")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("unknown")
|
||||
.to_string(),
|
||||
input: data
|
||||
.get("input")
|
||||
.cloned()
|
||||
.unwrap_or(serde_json::Value::Null),
|
||||
}),
|
||||
"tool_result" => Some(SseEvent::JobToolResult {
|
||||
job_id: job_id_str,
|
||||
tool_name: data
|
||||
.get("tool_name")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("unknown")
|
||||
.to_string(),
|
||||
output: data
|
||||
.get("output")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("")
|
||||
.to_string(),
|
||||
}),
|
||||
"status" => Some(SseEvent::JobStatus {
|
||||
job_id: job_id_str,
|
||||
message: data
|
||||
.get("message")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("")
|
||||
.to_string(),
|
||||
}),
|
||||
"result" => Some(SseEvent::JobResult {
|
||||
job_id: job_id_str,
|
||||
status: data
|
||||
.get("status")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("completed")
|
||||
.to_string(),
|
||||
session_id: data
|
||||
.get("session_id")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string()),
|
||||
}),
|
||||
_ => None,
|
||||
};
|
||||
if let Some(event) = event {
|
||||
let _ = tx.send(event);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Run the worker until the job is complete or stopped.
|
||||
@@ -123,7 +198,7 @@ impl Worker {
|
||||
tracing::debug!("Worker for job {} stopped before starting", self.job_id);
|
||||
return Ok(());
|
||||
}
|
||||
Some(WorkerMessage::Ping) => {}
|
||||
Some(WorkerMessage::Ping) | Some(WorkerMessage::UserMessage(_)) => {}
|
||||
}
|
||||
|
||||
// Get job context
|
||||
@@ -219,6 +294,8 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
.unwrap_or(50) as usize;
|
||||
let max_iterations = max_iterations.min(MAX_WORKER_ITERATIONS);
|
||||
let mut iteration = 0;
|
||||
const MAX_CONSECUTIVE_RATE_LIMITS: usize = 10;
|
||||
let mut consecutive_rate_limits = 0usize;
|
||||
|
||||
// Initial tool definitions for planning (will be refreshed in loop)
|
||||
reason_ctx.available_tools = self.tools().tool_definitions().await;
|
||||
@@ -269,15 +346,27 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
None
|
||||
};
|
||||
|
||||
// If we have a plan, execute it
|
||||
// If we have a plan, execute it. Two exit paths:
|
||||
// 1. Plan ran to completion → job is Completed or needs continuation
|
||||
// (check state and only fall through if not terminal)
|
||||
// 2. Plan was interrupted by UserMessage → fall through to direct loop
|
||||
if let Some(ref plan) = plan {
|
||||
return self.execute_plan(rx, reasoning, reason_ctx, plan).await;
|
||||
self.execute_plan(rx, reasoning, reason_ctx, plan).await?;
|
||||
|
||||
// If the plan marked the job terminal, we're done. Only fall
|
||||
// through to the direct selection loop if the plan was
|
||||
// interrupted or explicitly left the job in-progress.
|
||||
if let Ok(ctx) = self.context_manager().get_context(self.job_id).await
|
||||
&& (ctx.state.is_terminal() || ctx.state == JobState::Stuck)
|
||||
{
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
|
||||
// Otherwise, use direct tool selection loop
|
||||
// Direct tool selection loop (also used as fallback after plan interruption)
|
||||
loop {
|
||||
// Check for stop signal
|
||||
if let Ok(msg) = rx.try_recv() {
|
||||
// Check for stop signal and injected user messages
|
||||
while let Ok(msg) = rx.try_recv() {
|
||||
match msg {
|
||||
WorkerMessage::Stop => {
|
||||
tracing::debug!("Worker for job {} received stop signal", self.job_id);
|
||||
@@ -287,6 +376,20 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
tracing::trace!("Worker for job {} received ping", self.job_id);
|
||||
}
|
||||
WorkerMessage::Start => {}
|
||||
WorkerMessage::UserMessage(content) => {
|
||||
tracing::info!(
|
||||
job_id = %self.job_id,
|
||||
"Worker received follow-up user message"
|
||||
);
|
||||
reason_ctx.messages.push(ChatMessage::user(&content));
|
||||
self.log_event(
|
||||
"message",
|
||||
serde_json::json!({
|
||||
"role": "user",
|
||||
"content": content,
|
||||
}),
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -307,12 +410,64 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
// Refresh tool definitions so newly built tools become visible
|
||||
reason_ctx.available_tools = self.tools().tool_definitions().await;
|
||||
|
||||
// Select next tool(s) to use
|
||||
let selections = reasoning.select_tools(reason_ctx).await?;
|
||||
// Select next tool(s) to use, with rate-limit retry.
|
||||
let selections = match reasoning.select_tools(reason_ctx).await {
|
||||
Ok(s) => s,
|
||||
Err(crate::error::LlmError::RateLimited { retry_after, .. }) => {
|
||||
consecutive_rate_limits += 1;
|
||||
let wait = retry_after.unwrap_or(Duration::from_secs(5));
|
||||
tracing::warn!(
|
||||
job_id = %self.job_id,
|
||||
wait_secs = wait.as_secs(),
|
||||
attempt = consecutive_rate_limits,
|
||||
"LLM rate limited during tool selection, backing off"
|
||||
);
|
||||
if consecutive_rate_limits >= MAX_CONSECUTIVE_RATE_LIMITS {
|
||||
self.mark_stuck("Persistent rate limiting").await?;
|
||||
return Ok(());
|
||||
}
|
||||
self.log_event(
|
||||
"status",
|
||||
serde_json::json!({
|
||||
"message": format!("Rate limited, retrying in {}s ({}/{})...",
|
||||
wait.as_secs(), consecutive_rate_limits, MAX_CONSECUTIVE_RATE_LIMITS),
|
||||
}),
|
||||
);
|
||||
tokio::time::sleep(wait).await;
|
||||
continue;
|
||||
}
|
||||
Err(e) => return Err(e.into()),
|
||||
};
|
||||
|
||||
if selections.is_empty() {
|
||||
// No tools from select_tools, ask LLM directly (may still return tool calls)
|
||||
let respond_output = reasoning.respond_with_tools(reason_ctx).await?;
|
||||
let respond_output = match reasoning.respond_with_tools(reason_ctx).await {
|
||||
Ok(o) => o,
|
||||
Err(crate::error::LlmError::RateLimited { retry_after, .. }) => {
|
||||
consecutive_rate_limits += 1;
|
||||
let wait = retry_after.unwrap_or(Duration::from_secs(5));
|
||||
tracing::warn!(
|
||||
job_id = %self.job_id,
|
||||
wait_secs = wait.as_secs(),
|
||||
attempt = consecutive_rate_limits,
|
||||
"LLM rate limited during respond_with_tools, backing off"
|
||||
);
|
||||
if consecutive_rate_limits >= MAX_CONSECUTIVE_RATE_LIMITS {
|
||||
self.mark_stuck("Persistent rate limiting").await?;
|
||||
return Ok(());
|
||||
}
|
||||
self.log_event(
|
||||
"status",
|
||||
serde_json::json!({
|
||||
"message": format!("Rate limited, retrying in {}s ({}/{})...",
|
||||
wait.as_secs(), consecutive_rate_limits, MAX_CONSECUTIVE_RATE_LIMITS),
|
||||
}),
|
||||
);
|
||||
tokio::time::sleep(wait).await;
|
||||
continue;
|
||||
}
|
||||
Err(e) => return Err(e.into()),
|
||||
};
|
||||
|
||||
match respond_output.result {
|
||||
RespondResult::Text(response) => {
|
||||
@@ -424,6 +579,11 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
}
|
||||
}
|
||||
|
||||
// Reset rate-limit counter after a successful iteration (all LLM
|
||||
// calls succeeded). Placed here so alternating success/fail between
|
||||
// select_tools and respond_with_tools cannot bypass the cap.
|
||||
consecutive_rate_limits = 0;
|
||||
|
||||
// Small delay between iterations
|
||||
tokio::time::sleep(Duration::from_millis(100)).await;
|
||||
}
|
||||
@@ -540,9 +700,10 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
// Run BeforeToolCall hook
|
||||
let params = {
|
||||
use crate::hooks::{HookError, HookEvent, HookOutcome};
|
||||
let hook_params = redact_params(params, tool.sensitive_params());
|
||||
let event = HookEvent::ToolCall {
|
||||
tool_name: tool_name.to_string(),
|
||||
parameters: params.clone(),
|
||||
parameters: hook_params,
|
||||
user_id: job_ctx.user_id.clone(),
|
||||
context: format!("job:{}", job_id),
|
||||
};
|
||||
@@ -598,9 +759,12 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
.into());
|
||||
}
|
||||
|
||||
// Redact sensitive parameter values (e.g. secret_save's "value") before
|
||||
// they touch any observability or audit path.
|
||||
let safe_params = redact_params(¶ms, tool.sensitive_params());
|
||||
tracing::debug!(
|
||||
tool = %tool_name,
|
||||
params = %params,
|
||||
params = %safe_params,
|
||||
job = %job_id,
|
||||
"Tool call started"
|
||||
);
|
||||
@@ -652,7 +816,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
match deps
|
||||
.context_manager
|
||||
.update_memory(job_id, |mem| {
|
||||
let rec = mem.create_action(tool_name, params.clone()).succeed(
|
||||
let rec = mem.create_action(tool_name, safe_params.clone()).succeed(
|
||||
output_str.clone(),
|
||||
output.result.clone(),
|
||||
elapsed,
|
||||
@@ -674,7 +838,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
.context_manager
|
||||
.update_memory(job_id, |mem| {
|
||||
let rec = mem
|
||||
.create_action(tool_name, params.clone())
|
||||
.create_action(tool_name, safe_params.clone())
|
||||
.fail(e.to_string(), elapsed);
|
||||
mem.record_action(rec.clone());
|
||||
rec
|
||||
@@ -693,7 +857,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
.context_manager
|
||||
.update_memory(job_id, |mem| {
|
||||
let rec = mem
|
||||
.create_action(tool_name, params.clone())
|
||||
.create_action(tool_name, safe_params.clone())
|
||||
.fail("Execution timeout", elapsed);
|
||||
mem.record_action(rec.clone());
|
||||
rec
|
||||
@@ -836,8 +1000,8 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
plan: &ActionPlan,
|
||||
) -> Result<(), Error> {
|
||||
for (i, action) in plan.actions.iter().enumerate() {
|
||||
// Check for stop signal
|
||||
if let Ok(msg) = rx.try_recv() {
|
||||
// Check for stop signal and injected user messages
|
||||
while let Ok(msg) = rx.try_recv() {
|
||||
match msg {
|
||||
WorkerMessage::Stop => {
|
||||
tracing::debug!(
|
||||
@@ -850,6 +1014,29 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
tracing::trace!("Worker for job {} received ping", self.job_id);
|
||||
}
|
||||
WorkerMessage::Start => {}
|
||||
WorkerMessage::UserMessage(content) => {
|
||||
tracing::info!(
|
||||
job_id = %self.job_id,
|
||||
"User message received during plan execution, abandoning plan"
|
||||
);
|
||||
reason_ctx.messages.push(ChatMessage::user(&content));
|
||||
self.log_event(
|
||||
"message",
|
||||
serde_json::json!({
|
||||
"role": "user",
|
||||
"content": content,
|
||||
}),
|
||||
);
|
||||
self.log_event(
|
||||
"status",
|
||||
serde_json::json!({
|
||||
"message": "Plan interrupted by user message, re-evaluating...",
|
||||
}),
|
||||
);
|
||||
// Return Ok to break out of plan; caller falls through to
|
||||
// the direct selection loop for LLM re-evaluation.
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -902,14 +1089,18 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
if crate::util::llm_signals_completion(&response) {
|
||||
self.mark_completed().await?;
|
||||
} else {
|
||||
// Job not complete, could re-plan or fall back to direct selection
|
||||
// Job not complete — return Ok without marking terminal so the
|
||||
// caller falls through to the direct selection loop for continuation.
|
||||
tracing::info!(
|
||||
"Job {} plan completed but work remains, falling back to direct selection",
|
||||
self.job_id
|
||||
);
|
||||
// Continue with standard execution loop by returning (will be picked up by main loop)
|
||||
self.mark_stuck("Plan completed but job incomplete - needs re-planning")
|
||||
.await?;
|
||||
self.log_event(
|
||||
"status",
|
||||
serde_json::json!({
|
||||
"message": "Plan completed but job needs more work, continuing...",
|
||||
}),
|
||||
);
|
||||
}
|
||||
|
||||
Ok(())
|
||||
@@ -940,6 +1131,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
self.log_event(
|
||||
"result",
|
||||
serde_json::json!({
|
||||
"status": "completed",
|
||||
"success": true,
|
||||
"message": "Job completed successfully",
|
||||
}),
|
||||
@@ -965,6 +1157,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
self.log_event(
|
||||
"result",
|
||||
serde_json::json!({
|
||||
"status": "failed",
|
||||
"success": false,
|
||||
"message": format!("Execution failed: {}", reason),
|
||||
}),
|
||||
@@ -985,6 +1178,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
self.log_event(
|
||||
"result",
|
||||
serde_json::json!({
|
||||
"status": "stuck",
|
||||
"success": false,
|
||||
"message": format!("Job stuck: {}", reason),
|
||||
}),
|
||||
@@ -1103,6 +1297,7 @@ mod tests {
|
||||
hooks: Arc::new(crate::hooks::HookRegistry::new()),
|
||||
timeout: Duration::from_secs(30),
|
||||
use_planning: false,
|
||||
sse_tx: None,
|
||||
};
|
||||
|
||||
Worker::new(job_id, deps)
|
||||
|
||||
+66
-5
@@ -15,7 +15,7 @@ use crate::context::ContextManager;
|
||||
use crate::db::Database;
|
||||
use crate::extensions::ExtensionManager;
|
||||
use crate::hooks::HookRegistry;
|
||||
use crate::llm::{LlmProvider, SessionManager};
|
||||
use crate::llm::{LlmProvider, RecordingLlm, SessionManager};
|
||||
use crate::safety::SafetyLayer;
|
||||
use crate::secrets::SecretsStore;
|
||||
use crate::skills::SkillRegistry;
|
||||
@@ -48,6 +48,7 @@ pub struct AppComponents {
|
||||
pub skill_registry: Option<Arc<std::sync::RwLock<SkillRegistry>>>,
|
||||
pub skill_catalog: Option<Arc<SkillCatalog>>,
|
||||
pub cost_guard: Arc<crate::agent::cost_guard::CostGuard>,
|
||||
pub recording_handle: Option<Arc<RecordingLlm>>,
|
||||
pub session: Arc<SessionManager>,
|
||||
pub catalog_entries: Vec<crate::extensions::RegistryEntry>,
|
||||
pub dev_loaded_tool_names: Vec<String>,
|
||||
@@ -71,6 +72,9 @@ pub struct AppBuilder {
|
||||
db: Option<Arc<dyn Database>>,
|
||||
secrets_store: Option<Arc<dyn SecretsStore + Send + Sync>>,
|
||||
|
||||
// Test overrides
|
||||
llm_override: Option<Arc<dyn LlmProvider>>,
|
||||
|
||||
// Backend-specific handles needed by secrets store
|
||||
#[cfg(feature = "postgres")]
|
||||
pg_pool: Option<deadpool_postgres::Pool>,
|
||||
@@ -99,6 +103,7 @@ impl AppBuilder {
|
||||
log_broadcaster,
|
||||
db: None,
|
||||
secrets_store: None,
|
||||
llm_override: None,
|
||||
#[cfg(feature = "postgres")]
|
||||
pg_pool: None,
|
||||
#[cfg(feature = "libsql")]
|
||||
@@ -106,11 +111,26 @@ impl AppBuilder {
|
||||
}
|
||||
}
|
||||
|
||||
/// Inject a pre-created database, skipping `init_database()`.
|
||||
pub fn with_database(&mut self, db: Arc<dyn Database>) {
|
||||
self.db = Some(db);
|
||||
}
|
||||
|
||||
/// Inject a pre-created LLM provider, skipping `init_llm()`.
|
||||
pub fn with_llm(&mut self, llm: Arc<dyn LlmProvider>) {
|
||||
self.llm_override = Some(llm);
|
||||
}
|
||||
|
||||
/// Phase 1: Initialize database backend.
|
||||
///
|
||||
/// Creates the database connection, runs migrations, reloads config
|
||||
/// from DB, attaches DB to session manager, and cleans up stale jobs.
|
||||
pub async fn init_database(&mut self) -> Result<(), anyhow::Error> {
|
||||
if self.db.is_some() {
|
||||
tracing::debug!("Database already provided, skipping init_database()");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
if self.flags.no_db {
|
||||
tracing::warn!("Running without database connection");
|
||||
return Ok(());
|
||||
@@ -297,10 +317,17 @@ impl AppBuilder {
|
||||
#[allow(clippy::type_complexity)]
|
||||
pub fn init_llm(
|
||||
&self,
|
||||
) -> Result<(Arc<dyn LlmProvider>, Option<Arc<dyn LlmProvider>>), anyhow::Error> {
|
||||
let (llm, cheap_llm) =
|
||||
) -> Result<
|
||||
(
|
||||
Arc<dyn LlmProvider>,
|
||||
Option<Arc<dyn LlmProvider>>,
|
||||
Option<Arc<RecordingLlm>>,
|
||||
),
|
||||
anyhow::Error,
|
||||
> {
|
||||
let (llm, cheap_llm, recording_handle) =
|
||||
crate::llm::build_provider_chain(&self.config.llm, self.session.clone())?;
|
||||
Ok((llm, cheap_llm))
|
||||
Ok((llm, cheap_llm, recording_handle))
|
||||
}
|
||||
|
||||
/// Phase 4: Initialize safety, tools, embeddings, and workspace.
|
||||
@@ -331,6 +358,10 @@ impl AppBuilder {
|
||||
};
|
||||
tools.register_builtin_tools();
|
||||
|
||||
if let Some(ref ss) = self.secrets_store {
|
||||
tools.register_secrets_tools(Arc::clone(ss));
|
||||
}
|
||||
|
||||
// Create embeddings provider using the unified method
|
||||
let embeddings = self
|
||||
.config
|
||||
@@ -649,7 +680,11 @@ impl AppBuilder {
|
||||
self.init_database().await?;
|
||||
self.init_secrets().await?;
|
||||
|
||||
let (llm, cheap_llm) = self.init_llm()?;
|
||||
let (llm, cheap_llm, recording_handle) = if let Some(llm) = self.llm_override.take() {
|
||||
(llm, None, None)
|
||||
} else {
|
||||
self.init_llm()?
|
||||
};
|
||||
let (safety, tools, embeddings, workspace) = self.init_tools(&llm).await?;
|
||||
|
||||
// Create hook registry early so runtime extension activation can register hooks.
|
||||
@@ -665,6 +700,31 @@ impl AppBuilder {
|
||||
|
||||
// Seed workspace and backfill embeddings
|
||||
if let Some(ref ws) = workspace {
|
||||
// Import workspace files from disk FIRST if WORKSPACE_IMPORT_DIR is set.
|
||||
// This lets Docker images / deployment scripts ship customized
|
||||
// workspace templates (e.g., AGENTS.md, TOOLS.md) that override
|
||||
// the generic seeds. Only imports files that don't already exist
|
||||
// in the database — never overwrites user edits.
|
||||
//
|
||||
// Runs before seed_if_empty() so that custom templates take priority
|
||||
// over generic seeds. seed_if_empty() then fills any remaining gaps.
|
||||
if let Ok(import_dir) = std::env::var("WORKSPACE_IMPORT_DIR") {
|
||||
let import_path = std::path::Path::new(&import_dir);
|
||||
match ws.import_from_directory(import_path).await {
|
||||
Ok(count) if count > 0 => {
|
||||
tracing::info!("Imported {} workspace file(s) from {}", count, import_dir);
|
||||
}
|
||||
Ok(_) => {}
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
"Failed to import workspace files from {}: {}",
|
||||
import_dir,
|
||||
e
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
match ws.seed_if_empty().await {
|
||||
Ok(_) => {}
|
||||
Err(e) => {
|
||||
@@ -736,6 +796,7 @@ impl AppBuilder {
|
||||
skill_registry,
|
||||
skill_catalog,
|
||||
cost_guard,
|
||||
recording_handle,
|
||||
session: self.session,
|
||||
catalog_entries,
|
||||
dev_loaded_tool_names,
|
||||
|
||||
+174
-1
@@ -117,7 +117,20 @@ pub enum StatusUpdate {
|
||||
/// Tool execution started.
|
||||
ToolStarted { name: String },
|
||||
/// Tool execution completed.
|
||||
ToolCompleted { name: String, success: bool },
|
||||
///
|
||||
/// Use [`StatusUpdate::tool_completed`] to construct this variant — it
|
||||
/// handles redaction of sensitive parameters and keeps the 9-line pattern
|
||||
/// in one place.
|
||||
ToolCompleted {
|
||||
name: String,
|
||||
success: bool,
|
||||
/// Error message when success is false.
|
||||
error: Option<String>,
|
||||
/// Tool input parameters (JSON string) for display on failure.
|
||||
/// Only populated when `success` is `false`. Values listed in the
|
||||
/// tool's `sensitive_params()` are replaced with `"[REDACTED]"`.
|
||||
parameters: Option<String>,
|
||||
},
|
||||
/// Brief preview of tool execution output.
|
||||
ToolResult { name: String, preview: String },
|
||||
/// Streaming text chunk.
|
||||
@@ -152,6 +165,38 @@ pub enum StatusUpdate {
|
||||
},
|
||||
}
|
||||
|
||||
impl StatusUpdate {
|
||||
/// Build a `ToolCompleted` status with redacted parameters.
|
||||
///
|
||||
/// On failure, serializes the tool's input parameters as pretty JSON after
|
||||
/// replacing any keys listed in the tool's `sensitive_params()` with
|
||||
/// `"[REDACTED]"`. On success, no parameters or error are included.
|
||||
///
|
||||
/// Pass the resolved `Tool` reference (if available) so this method can
|
||||
/// query `sensitive_params()` directly — callers don't need to manage the
|
||||
/// borrow lifetime of the sensitive slice.
|
||||
pub fn tool_completed(
|
||||
name: String,
|
||||
result: &Result<String, crate::error::Error>,
|
||||
params: &serde_json::Value,
|
||||
tool: Option<&dyn crate::tools::Tool>,
|
||||
) -> Self {
|
||||
let success = result.is_ok();
|
||||
let sensitive = tool.map(|t| t.sensitive_params()).unwrap_or(&[]);
|
||||
Self::ToolCompleted {
|
||||
name,
|
||||
success,
|
||||
error: result.as_ref().err().map(|e| e.to_string()),
|
||||
parameters: if !success {
|
||||
let safe = crate::tools::redact_params(params, sensitive);
|
||||
Some(serde_json::to_string_pretty(&safe).unwrap_or_else(|_| safe.to_string()))
|
||||
} else {
|
||||
None
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Trait for message channels.
|
||||
///
|
||||
/// Channels receive messages from external sources and convert them to
|
||||
@@ -223,3 +268,131 @@ pub trait Channel: Send + Sync {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
/// Stub tool that marks `"value"` as sensitive.
|
||||
struct SecretTool;
|
||||
|
||||
#[async_trait]
|
||||
impl crate::tools::Tool for SecretTool {
|
||||
fn name(&self) -> &str {
|
||||
"secret_save"
|
||||
}
|
||||
fn description(&self) -> &str {
|
||||
"stub"
|
||||
}
|
||||
fn parameters_schema(&self) -> serde_json::Value {
|
||||
serde_json::json!({"type": "object", "properties": {}})
|
||||
}
|
||||
async fn execute(
|
||||
&self,
|
||||
_params: serde_json::Value,
|
||||
_ctx: &crate::context::JobContext,
|
||||
) -> Result<crate::tools::ToolOutput, crate::tools::ToolError> {
|
||||
unreachable!()
|
||||
}
|
||||
fn sensitive_params(&self) -> &[&str] {
|
||||
&["value"]
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tool_completed_redacts_sensitive_params_on_failure() {
|
||||
let params = serde_json::json!({"name": "api_key", "value": "sk-secret-123"});
|
||||
let err: Result<String, crate::error::Error> =
|
||||
Err(crate::error::ToolError::ExecutionFailed {
|
||||
name: "secret_save".into(),
|
||||
reason: "db error".into(),
|
||||
}
|
||||
.into());
|
||||
let tool = SecretTool;
|
||||
|
||||
let status = StatusUpdate::tool_completed(
|
||||
"secret_save".into(),
|
||||
&err,
|
||||
¶ms,
|
||||
Some(&tool as &dyn crate::tools::Tool),
|
||||
);
|
||||
|
||||
if let StatusUpdate::ToolCompleted {
|
||||
success,
|
||||
error,
|
||||
parameters,
|
||||
..
|
||||
} = &status
|
||||
{
|
||||
assert!(!success);
|
||||
let err_msg = error.as_deref().expect("should have error");
|
||||
assert!(err_msg.contains("db error"), "error: {}", err_msg);
|
||||
let param_str = parameters
|
||||
.as_ref()
|
||||
.expect("should have parameters on failure");
|
||||
assert!(
|
||||
param_str.contains("[REDACTED]"),
|
||||
"sensitive value should be redacted: {}",
|
||||
param_str
|
||||
);
|
||||
assert!(
|
||||
!param_str.contains("sk-secret-123"),
|
||||
"raw secret should not appear: {}",
|
||||
param_str
|
||||
);
|
||||
assert!(
|
||||
param_str.contains("api_key"),
|
||||
"non-sensitive params should be preserved: {}",
|
||||
param_str
|
||||
);
|
||||
} else {
|
||||
panic!("expected ToolCompleted variant");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tool_completed_no_params_on_success() {
|
||||
let params = serde_json::json!({"name": "key", "value": "secret"});
|
||||
let ok: Result<String, crate::error::Error> = Ok("done".into());
|
||||
|
||||
let status = StatusUpdate::tool_completed("secret_save".into(), &ok, ¶ms, None);
|
||||
|
||||
if let StatusUpdate::ToolCompleted {
|
||||
success,
|
||||
error,
|
||||
parameters,
|
||||
..
|
||||
} = &status
|
||||
{
|
||||
assert!(success);
|
||||
assert!(error.is_none());
|
||||
assert!(parameters.is_none(), "no params should be sent on success");
|
||||
} else {
|
||||
panic!("expected ToolCompleted variant");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tool_completed_no_tool_passes_params_unredacted() {
|
||||
let params = serde_json::json!({"cmd": "ls -la"});
|
||||
let err: Result<String, crate::error::Error> =
|
||||
Err(crate::error::ToolError::ExecutionFailed {
|
||||
name: "shell".into(),
|
||||
reason: "timeout".into(),
|
||||
}
|
||||
.into());
|
||||
|
||||
let status = StatusUpdate::tool_completed("shell".into(), &err, ¶ms, None);
|
||||
|
||||
if let StatusUpdate::ToolCompleted { parameters, .. } = &status {
|
||||
let param_str = parameters.as_ref().expect("should have parameters");
|
||||
assert!(
|
||||
param_str.contains("ls -la"),
|
||||
"non-sensitive params should pass through: {}",
|
||||
param_str
|
||||
);
|
||||
} else {
|
||||
panic!("expected ToolCompleted variant");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -466,7 +466,7 @@ impl Channel for ReplChannel {
|
||||
StatusUpdate::ToolStarted { name } => {
|
||||
eprintln!(" \x1b[33m\u{25CB} {name}\x1b[0m");
|
||||
}
|
||||
StatusUpdate::ToolCompleted { name, success } => {
|
||||
StatusUpdate::ToolCompleted { name, success, .. } => {
|
||||
if success {
|
||||
eprintln!(" \x1b[32m\u{25CF} {name}\x1b[0m");
|
||||
} else {
|
||||
|
||||
@@ -974,7 +974,7 @@ impl Channel for SignalChannel {
|
||||
|
||||
// Send tool completed notification (debug mode only)
|
||||
if self.is_debug()
|
||||
&& let StatusUpdate::ToolCompleted { name, success } = &status
|
||||
&& let StatusUpdate::ToolCompleted { name, success, .. } = &status
|
||||
&& let Some(target_str) = metadata.get("signal_target").and_then(|v| v.as_str())
|
||||
{
|
||||
let (icon, color) = if *success {
|
||||
|
||||
@@ -80,6 +80,9 @@ pub enum WasmChannelError {
|
||||
|
||||
#[error("HTTP request error: {0}")]
|
||||
HttpRequest(String),
|
||||
|
||||
#[error("WIT version mismatch: {0}")]
|
||||
IncompatibleWitVersion(String),
|
||||
}
|
||||
|
||||
impl From<crate::tools::wasm::WasmError> for WasmChannelError {
|
||||
|
||||
@@ -19,12 +19,14 @@ use crate::channels::wasm::schema::ChannelCapabilitiesFile;
|
||||
use crate::channels::wasm::wrapper::WasmChannel;
|
||||
use crate::db::SettingsStore;
|
||||
use crate::pairing::PairingStore;
|
||||
use crate::secrets::SecretsStore;
|
||||
|
||||
/// Loads WASM channels from the filesystem.
|
||||
pub struct WasmChannelLoader {
|
||||
runtime: Arc<WasmChannelRuntime>,
|
||||
pairing_store: Arc<PairingStore>,
|
||||
settings_store: Option<Arc<dyn SettingsStore>>,
|
||||
secrets_store: Option<Arc<dyn SecretsStore + Send + Sync>>,
|
||||
}
|
||||
|
||||
impl WasmChannelLoader {
|
||||
@@ -38,9 +40,16 @@ impl WasmChannelLoader {
|
||||
runtime,
|
||||
pairing_store,
|
||||
settings_store,
|
||||
secrets_store: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Set the secrets store for host-based credential injection in WASM channels.
|
||||
pub fn with_secrets_store(mut self, store: Arc<dyn SecretsStore + Send + Sync>) -> Self {
|
||||
self.secrets_store = Some(store);
|
||||
self
|
||||
}
|
||||
|
||||
/// Load a single WASM channel from a file pair.
|
||||
///
|
||||
/// Expects:
|
||||
@@ -72,6 +81,7 @@ impl WasmChannelLoader {
|
||||
let cap_bytes = fs::read(cap_path).await?;
|
||||
let cap_file = ChannelCapabilitiesFile::from_bytes(&cap_bytes)
|
||||
.map_err(|e| WasmChannelError::InvalidCapabilities(e.to_string()))?;
|
||||
cap_file.validate();
|
||||
|
||||
// Debug: log raw capabilities
|
||||
tracing::debug!(
|
||||
@@ -80,6 +90,14 @@ impl WasmChannelLoader {
|
||||
"Parsed capabilities file"
|
||||
);
|
||||
|
||||
// Check WIT version compatibility
|
||||
crate::tools::wasm::loader::check_wit_version_compat(
|
||||
name,
|
||||
cap_file.wit_version.as_deref(),
|
||||
crate::tools::wasm::WIT_CHANNEL_VERSION,
|
||||
)
|
||||
.map_err(|e| WasmChannelError::IncompatibleWitVersion(e.to_string()))?;
|
||||
|
||||
let caps = cap_file.to_capabilities();
|
||||
|
||||
// Debug: log resulting capabilities
|
||||
@@ -127,7 +145,7 @@ impl WasmChannelLoader {
|
||||
.await?;
|
||||
|
||||
// Create the channel
|
||||
let channel = WasmChannel::new(
|
||||
let mut channel = WasmChannel::new(
|
||||
self.runtime.clone(),
|
||||
prepared,
|
||||
capabilities,
|
||||
@@ -135,6 +153,9 @@ impl WasmChannelLoader {
|
||||
self.pairing_store.clone(),
|
||||
self.settings_store.clone(),
|
||||
);
|
||||
if let Some(ref secrets) = self.secrets_store {
|
||||
channel = channel.with_secrets_store(Arc::clone(secrets));
|
||||
}
|
||||
|
||||
tracing::info!(
|
||||
name = name,
|
||||
@@ -264,6 +285,13 @@ impl LoadedChannel {
|
||||
.and_then(|f| f.signature_key_secret_name().map(|s| s.to_string()))
|
||||
}
|
||||
|
||||
/// Get the HMAC-SHA256 signing secret name from capabilities.
|
||||
pub fn hmac_secret_name(&self) -> Option<String> {
|
||||
self.capabilities_file
|
||||
.as_ref()
|
||||
.and_then(|f| f.hmac_secret_name().map(|s| s.to_string()))
|
||||
}
|
||||
|
||||
/// Get the webhook secret name from capabilities.
|
||||
pub fn webhook_secret_name(&self) -> String {
|
||||
self.capabilities_file
|
||||
|
||||
@@ -87,6 +87,8 @@ mod router;
|
||||
mod runtime;
|
||||
mod schema;
|
||||
pub(crate) mod signature;
|
||||
#[allow(dead_code)]
|
||||
pub(crate) mod storage;
|
||||
mod wrapper;
|
||||
|
||||
// Core types
|
||||
|
||||
+337
-1
@@ -44,6 +44,8 @@ pub struct WasmChannelRouter {
|
||||
secret_headers: RwLock<HashMap<String, String>>,
|
||||
/// Ed25519 public keys for signature verification by channel name (hex-encoded).
|
||||
signature_keys: RwLock<HashMap<String, String>>,
|
||||
/// HMAC-SHA256 signing secrets for signature verification by channel name (Slack-style).
|
||||
hmac_secrets: RwLock<HashMap<String, String>>,
|
||||
}
|
||||
|
||||
impl WasmChannelRouter {
|
||||
@@ -55,6 +57,7 @@ impl WasmChannelRouter {
|
||||
secrets: RwLock::new(HashMap::new()),
|
||||
secret_headers: RwLock::new(HashMap::new()),
|
||||
signature_keys: RwLock::new(HashMap::new()),
|
||||
hmac_secrets: RwLock::new(HashMap::new()),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -134,6 +137,7 @@ impl WasmChannelRouter {
|
||||
self.secrets.write().await.remove(channel_name);
|
||||
self.secret_headers.write().await.remove(channel_name);
|
||||
self.signature_keys.write().await.remove(channel_name);
|
||||
self.hmac_secrets.write().await.remove(channel_name);
|
||||
|
||||
// Remove all paths for this channel
|
||||
self.path_to_channel
|
||||
@@ -208,6 +212,24 @@ impl WasmChannelRouter {
|
||||
pub async fn get_signature_key(&self, channel_name: &str) -> Option<String> {
|
||||
self.signature_keys.read().await.get(channel_name).cloned()
|
||||
}
|
||||
|
||||
/// Register an HMAC-SHA256 signing secret for signature verification.
|
||||
///
|
||||
/// Channels with a registered secret will have Slack-style HMAC-SHA256
|
||||
/// signature validation performed before forwarding to WASM.
|
||||
pub async fn register_hmac_secret(&self, channel_name: &str, secret: &str) {
|
||||
self.hmac_secrets
|
||||
.write()
|
||||
.await
|
||||
.insert(channel_name.to_string(), secret.to_string());
|
||||
}
|
||||
|
||||
/// Get the HMAC signing secret for a channel.
|
||||
///
|
||||
/// Returns `None` if no secret is registered (no HMAC check needed).
|
||||
pub async fn get_hmac_secret(&self, channel_name: &str) -> Option<String> {
|
||||
self.hmac_secrets.read().await.get(channel_name).cloned()
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for WasmChannelRouter {
|
||||
@@ -427,6 +449,57 @@ async fn webhook_handler(
|
||||
}
|
||||
}
|
||||
|
||||
// HMAC-SHA256 signature verification (Slack-style)
|
||||
if let Some(hmac_secret) = state.router.get_hmac_secret(channel_name).await {
|
||||
let timestamp = headers
|
||||
.get("x-slack-request-timestamp")
|
||||
.and_then(|v| v.to_str().ok());
|
||||
let sig_header = headers
|
||||
.get("x-slack-signature")
|
||||
.and_then(|v| v.to_str().ok());
|
||||
|
||||
match (timestamp, sig_header) {
|
||||
(Some(ts), Some(sig)) => {
|
||||
let now_secs = std::time::SystemTime::now()
|
||||
.duration_since(std::time::UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_secs() as i64;
|
||||
|
||||
if !crate::channels::wasm::signature::verify_slack_signature(
|
||||
&hmac_secret,
|
||||
ts,
|
||||
&body,
|
||||
sig,
|
||||
now_secs,
|
||||
) {
|
||||
tracing::warn!(
|
||||
channel = %channel_name,
|
||||
"HMAC-SHA256 signature verification failed"
|
||||
);
|
||||
return (
|
||||
StatusCode::UNAUTHORIZED,
|
||||
Json(serde_json::json!({
|
||||
"error": "Invalid Slack signature"
|
||||
})),
|
||||
);
|
||||
}
|
||||
tracing::debug!(channel = %channel_name, "HMAC-SHA256 signature verified");
|
||||
}
|
||||
_ => {
|
||||
tracing::warn!(
|
||||
channel = %channel_name,
|
||||
"Slack signature headers missing but secret is registered"
|
||||
);
|
||||
return (
|
||||
StatusCode::UNAUTHORIZED,
|
||||
Json(serde_json::json!({
|
||||
"error": "Missing Slack signature headers"
|
||||
})),
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Convert headers to HashMap
|
||||
let headers_map: HashMap<String, String> = headers
|
||||
.iter()
|
||||
@@ -731,7 +804,59 @@ mod tests {
|
||||
assert_eq!(router.get_secret_header("slack").await, "X-Webhook-Secret");
|
||||
}
|
||||
|
||||
// ── Category 3: Router Signature Key Management ─────────────────────
|
||||
// ── Category 3: Router HMAC Secret Management ───────────────────────
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_register_and_get_hmac_secret() {
|
||||
let router = WasmChannelRouter::new();
|
||||
let channel = create_test_channel("slack");
|
||||
|
||||
router.register(channel, vec![], None, None).await;
|
||||
|
||||
let hmac_secret = "my-slack-signing-secret";
|
||||
router.register_hmac_secret("slack", hmac_secret).await;
|
||||
|
||||
let retrieved = router.get_hmac_secret("slack").await;
|
||||
assert_eq!(retrieved, Some(hmac_secret.to_string()));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_no_hmac_secret_returns_none() {
|
||||
let router = WasmChannelRouter::new();
|
||||
let channel = create_test_channel("slack");
|
||||
router.register(channel, vec![], None, None).await;
|
||||
|
||||
// Slack has no HMAC secret registered
|
||||
let secret = router.get_hmac_secret("slack").await;
|
||||
assert!(secret.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_unregister_removes_hmac_secret() {
|
||||
let router = WasmChannelRouter::new();
|
||||
let channel = create_test_channel("slack");
|
||||
|
||||
let endpoints = vec![RegisteredEndpoint {
|
||||
channel_name: "slack".to_string(),
|
||||
path: "/webhook/slack".to_string(),
|
||||
methods: vec!["POST".to_string()],
|
||||
require_secret: false,
|
||||
}];
|
||||
|
||||
router.register(channel, endpoints, None, None).await;
|
||||
router.register_hmac_secret("slack", "signing-secret").await;
|
||||
|
||||
// Secret should exist
|
||||
assert!(router.get_hmac_secret("slack").await.is_some());
|
||||
|
||||
// Unregister
|
||||
router.unregister("slack").await;
|
||||
|
||||
// Secret should be gone
|
||||
assert!(router.get_hmac_secret("slack").await.is_none());
|
||||
}
|
||||
|
||||
// ── Category 4: Router Signature Key Management ─────────────────────
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_register_and_get_signature_key() {
|
||||
@@ -1163,4 +1288,215 @@ mod tests {
|
||||
"Valid secret + valid signature should not return 401"
|
||||
);
|
||||
}
|
||||
|
||||
// ── HMAC-SHA256 Webhook Signature Tests ────────────────────────────
|
||||
|
||||
/// Helper to create a router with a registered channel at /webhook/slack.
|
||||
async fn setup_slack_router() -> (Arc<WasmChannelRouter>, AxumRouter) {
|
||||
let wasm_router = Arc::new(WasmChannelRouter::new());
|
||||
let channel = create_test_channel("slack");
|
||||
|
||||
let endpoints = vec![RegisteredEndpoint {
|
||||
channel_name: "slack".to_string(),
|
||||
path: "/webhook/slack".to_string(),
|
||||
methods: vec!["POST".to_string()],
|
||||
require_secret: false,
|
||||
}];
|
||||
|
||||
wasm_router.register(channel, endpoints, None, None).await;
|
||||
|
||||
let app = create_wasm_channel_router(wasm_router.clone(), None);
|
||||
(wasm_router, app)
|
||||
}
|
||||
|
||||
/// Helper: compute expected Slack signature for testing.
|
||||
fn slack_signature(signing_secret: &str, timestamp: &str, body: &[u8]) -> String {
|
||||
use hmac::{Hmac, Mac};
|
||||
use sha2::Sha256;
|
||||
|
||||
let mut basestring = Vec::new();
|
||||
basestring.extend_from_slice(b"v0:");
|
||||
basestring.extend_from_slice(timestamp.as_bytes());
|
||||
basestring.push(b':');
|
||||
basestring.extend_from_slice(body);
|
||||
|
||||
let mut mac = Hmac::<Sha256>::new_from_slice(signing_secret.as_bytes()).unwrap();
|
||||
mac.update(&basestring);
|
||||
let computed = mac.finalize().into_bytes();
|
||||
format!("v0={}", hex::encode(computed))
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_webhook_hmac_rejects_missing_sig_headers() {
|
||||
let (wasm_router, app) = setup_slack_router().await;
|
||||
|
||||
wasm_router
|
||||
.register_hmac_secret("slack", "my-signing-secret")
|
||||
.await;
|
||||
|
||||
// Send request without HMAC signature headers
|
||||
let req = Request::builder()
|
||||
.method("POST")
|
||||
.uri("/webhook/slack")
|
||||
.header("content-type", "application/json")
|
||||
.body(Body::from("token=xyzz0WbapA4vBCDEFasx0q6G"))
|
||||
.unwrap();
|
||||
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
assert_eq!(
|
||||
resp.status(),
|
||||
StatusCode::UNAUTHORIZED,
|
||||
"Missing HMAC signature headers should return 401"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_webhook_hmac_rejects_invalid_signature() {
|
||||
let (wasm_router, app) = setup_slack_router().await;
|
||||
|
||||
wasm_router
|
||||
.register_hmac_secret("slack", "my-signing-secret")
|
||||
.await;
|
||||
|
||||
let req = Request::builder()
|
||||
.method("POST")
|
||||
.uri("/webhook/slack")
|
||||
.header("content-type", "application/json")
|
||||
.header("x-slack-request-timestamp", "1234567890")
|
||||
.header("x-slack-signature", "v0=deadbeefdeadbeef")
|
||||
.body(Body::from("token=xyzz0WbapA4vBCDEFasx0q6G"))
|
||||
.unwrap();
|
||||
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
assert_eq!(
|
||||
resp.status(),
|
||||
StatusCode::UNAUTHORIZED,
|
||||
"Invalid HMAC signature should return 401"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_webhook_hmac_accepts_valid_signature() {
|
||||
let (wasm_router, app) = setup_slack_router().await;
|
||||
|
||||
let signing_secret = "my-signing-secret";
|
||||
wasm_router
|
||||
.register_hmac_secret("slack", signing_secret)
|
||||
.await;
|
||||
|
||||
let now_secs = std::time::SystemTime::now()
|
||||
.duration_since(std::time::UNIX_EPOCH)
|
||||
.unwrap()
|
||||
.as_secs();
|
||||
let timestamp = now_secs.to_string();
|
||||
let body = b"token=xyzz0WbapA4vBCDEFasx0q6G";
|
||||
|
||||
let signature = slack_signature(signing_secret, ×tamp, body);
|
||||
|
||||
let req = Request::builder()
|
||||
.method("POST")
|
||||
.uri("/webhook/slack")
|
||||
.header("content-type", "application/json")
|
||||
.header("x-slack-request-timestamp", ×tamp)
|
||||
.header("x-slack-signature", &signature)
|
||||
.body(Body::from(&body[..]))
|
||||
.unwrap();
|
||||
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
// Should NOT be 401 — signature is valid (may be 500 since no WASM module)
|
||||
assert_ne!(
|
||||
resp.status(),
|
||||
StatusCode::UNAUTHORIZED,
|
||||
"Valid HMAC signature should not return 401"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_webhook_hmac_skips_check_for_no_secret() {
|
||||
let (_wasm_router, app) = setup_slack_router().await;
|
||||
|
||||
// No HMAC secret registered — should not require signature
|
||||
let req = Request::builder()
|
||||
.method("POST")
|
||||
.uri("/webhook/slack")
|
||||
.header("content-type", "application/json")
|
||||
.body(Body::from("token=xyzz0WbapA4vBCDEFasx0q6G"))
|
||||
.unwrap();
|
||||
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
// Should NOT be 401 (may be 500 since no WASM module, but not auth failure)
|
||||
assert_ne!(
|
||||
resp.status(),
|
||||
StatusCode::UNAUTHORIZED,
|
||||
"No HMAC secret registered — should skip check"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_webhook_hmac_uses_correct_body() {
|
||||
let (wasm_router, app) = setup_slack_router().await;
|
||||
|
||||
let signing_secret = "my-signing-secret";
|
||||
wasm_router
|
||||
.register_hmac_secret("slack", signing_secret)
|
||||
.await;
|
||||
|
||||
let timestamp = "1234567890";
|
||||
let body_a = b"token=xyzz0WbapA4vBCDEFasx0q6G";
|
||||
let body_b = b"token=MODIFIED";
|
||||
|
||||
// Sign body A
|
||||
let signature = slack_signature(signing_secret, timestamp, body_a);
|
||||
|
||||
// But send body B
|
||||
let req = Request::builder()
|
||||
.method("POST")
|
||||
.uri("/webhook/slack")
|
||||
.header("content-type", "application/json")
|
||||
.header("x-slack-request-timestamp", timestamp)
|
||||
.header("x-slack-signature", &signature)
|
||||
.body(Body::from(&body_b[..]))
|
||||
.unwrap();
|
||||
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
assert_eq!(
|
||||
resp.status(),
|
||||
StatusCode::UNAUTHORIZED,
|
||||
"Signature for different body should return 401"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_webhook_hmac_uses_correct_timestamp() {
|
||||
let (wasm_router, app) = setup_slack_router().await;
|
||||
|
||||
let signing_secret = "my-signing-secret";
|
||||
wasm_router
|
||||
.register_hmac_secret("slack", signing_secret)
|
||||
.await;
|
||||
|
||||
let timestamp_a = "1234567890";
|
||||
let timestamp_b = "9999999999";
|
||||
let body = b"token=xyzz0WbapA4vBCDEFasx0q6G";
|
||||
|
||||
// Sign with timestamp A
|
||||
let signature = slack_signature(signing_secret, timestamp_a, body);
|
||||
|
||||
// But send timestamp B in the header
|
||||
let req = Request::builder()
|
||||
.method("POST")
|
||||
.uri("/webhook/slack")
|
||||
.header("content-type", "application/json")
|
||||
.header("x-slack-request-timestamp", timestamp_b)
|
||||
.header("x-slack-signature", &signature)
|
||||
.body(Body::from(&body[..]))
|
||||
.unwrap();
|
||||
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
assert_eq!(
|
||||
resp.status(),
|
||||
StatusCode::UNAUTHORIZED,
|
||||
"Signature with mismatched timestamp should return 401"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -51,6 +51,14 @@ use crate::tools::wasm::{CapabilitiesFile as ToolCapabilitiesFile, RateLimitSche
|
||||
/// Root schema for a channel capabilities JSON file.
|
||||
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
|
||||
pub struct ChannelCapabilitiesFile {
|
||||
/// Extension version (semver).
|
||||
#[serde(default)]
|
||||
pub version: Option<String>,
|
||||
|
||||
/// WIT interface version this channel was compiled against (semver).
|
||||
#[serde(default)]
|
||||
pub wit_version: Option<String>,
|
||||
|
||||
/// File type, must be "channel".
|
||||
#[serde(default = "default_type")]
|
||||
pub r#type: String,
|
||||
@@ -90,6 +98,37 @@ impl ChannelCapabilitiesFile {
|
||||
serde_json::from_slice(bytes)
|
||||
}
|
||||
|
||||
/// Validate the capabilities file and emit warnings for common misconfigurations.
|
||||
///
|
||||
/// Called once at load time to catch issues early. Warnings are emitted via
|
||||
/// `tracing::warn` so they show up in startup logs without blocking loading.
|
||||
pub fn validate(&self) {
|
||||
const MIN_PROMPT_LENGTH: usize = 30;
|
||||
|
||||
// Check for short prompts in required_secrets
|
||||
for secret in &self.setup.required_secrets {
|
||||
if secret.prompt.len() < MIN_PROMPT_LENGTH {
|
||||
tracing::warn!(
|
||||
channel = self.name,
|
||||
secret = secret.name,
|
||||
prompt = secret.prompt,
|
||||
"setup.required_secrets prompt is shorter than {} chars — \
|
||||
consider a more descriptive prompt that tells the user where to find this value",
|
||||
MIN_PROMPT_LENGTH
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// Has required_secrets but no setup_url
|
||||
if !self.setup.required_secrets.is_empty() && self.setup.setup_url.is_none() {
|
||||
tracing::warn!(
|
||||
channel = self.name,
|
||||
"setup.required_secrets defined but no setup.setup_url — \
|
||||
user has no link to obtain credentials"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
/// Convert to runtime ChannelCapabilities.
|
||||
pub fn to_capabilities(&self) -> ChannelCapabilities {
|
||||
self.capabilities.to_channel_capabilities(&self.name)
|
||||
@@ -123,6 +162,18 @@ impl ChannelCapabilitiesFile {
|
||||
.and_then(|w| w.signature_key_secret_name.as_deref())
|
||||
}
|
||||
|
||||
/// Get the HMAC-SHA256 signing secret name for this channel.
|
||||
///
|
||||
/// Returns the secret name declared in `webhook.hmac_secret_name`,
|
||||
/// used to look up the HMAC signing secret in the secrets store (Slack-style).
|
||||
pub fn hmac_secret_name(&self) -> Option<&str> {
|
||||
self.capabilities
|
||||
.channel
|
||||
.as_ref()
|
||||
.and_then(|c| c.webhook.as_ref())
|
||||
.and_then(|w| w.hmac_secret_name.as_deref())
|
||||
}
|
||||
|
||||
/// Get the webhook secret name for this channel.
|
||||
///
|
||||
/// Returns the configured secret name or defaults to "{channel_name}_webhook_secret".
|
||||
@@ -247,6 +298,10 @@ pub struct WebhookSchema {
|
||||
/// for signature verification (e.g., Discord interaction verification).
|
||||
#[serde(default)]
|
||||
pub signature_key_secret_name: Option<String>,
|
||||
|
||||
/// Secret name in secrets store for HMAC-SHA256 signing (Slack-style).
|
||||
#[serde(default)]
|
||||
pub hmac_secret_name: Option<String>,
|
||||
}
|
||||
|
||||
/// Setup configuration schema.
|
||||
@@ -262,6 +317,10 @@ pub struct SetupSchema {
|
||||
/// Placeholders like {secret_name} are replaced with actual values.
|
||||
#[serde(default)]
|
||||
pub validation_endpoint: Option<String>,
|
||||
|
||||
/// User-facing URL where they can create/manage credentials.
|
||||
#[serde(default)]
|
||||
pub setup_url: Option<String>,
|
||||
}
|
||||
|
||||
/// Configuration for a secret required during setup.
|
||||
@@ -605,6 +664,65 @@ mod tests {
|
||||
|
||||
// ── Category 5: Discord Capabilities Setup & Configuration ──────────
|
||||
|
||||
#[test]
|
||||
fn test_validate_channel_short_prompt() {
|
||||
// prompt < 30 chars — should not panic
|
||||
let json = r#"{
|
||||
"name": "test-channel",
|
||||
"setup": {
|
||||
"required_secrets": [
|
||||
{ "name": "bot_token", "prompt": "Bot token" }
|
||||
],
|
||||
"setup_url": "https://example.com"
|
||||
}
|
||||
}"#;
|
||||
|
||||
let file = ChannelCapabilitiesFile::from_json(json).unwrap();
|
||||
// Should not panic; warning emitted for short prompt
|
||||
file.validate();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate_channel_missing_setup_url() {
|
||||
// required_secrets without setup_url — should not panic
|
||||
let json = r#"{
|
||||
"name": "test-channel",
|
||||
"setup": {
|
||||
"required_secrets": [
|
||||
{
|
||||
"name": "bot_token",
|
||||
"prompt": "Enter your bot token from the developer portal settings"
|
||||
}
|
||||
]
|
||||
}
|
||||
}"#;
|
||||
|
||||
let file = ChannelCapabilitiesFile::from_json(json).unwrap();
|
||||
// Should not panic; warning emitted for missing setup_url
|
||||
file.validate();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate_clean_channel() {
|
||||
// Well-configured channel — should not panic or warn
|
||||
let json = r#"{
|
||||
"name": "good-channel",
|
||||
"setup": {
|
||||
"required_secrets": [
|
||||
{
|
||||
"name": "bot_token",
|
||||
"prompt": "Enter your bot token from https://example.com/bot-settings"
|
||||
}
|
||||
],
|
||||
"setup_url": "https://example.com/bot-settings"
|
||||
}
|
||||
}"#;
|
||||
|
||||
let file = ChannelCapabilitiesFile::from_json(json).unwrap();
|
||||
// Should not panic and emits no warnings
|
||||
file.validate();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_discord_capabilities_has_public_key_secret() {
|
||||
let json = include_str!("../../../channels-src/discord/discord.capabilities.json");
|
||||
|
||||
@@ -1,9 +1,11 @@
|
||||
//! Discord Ed25519 signature verification.
|
||||
//! Webhook signature verification (Discord Ed25519 and Slack HMAC-SHA256).
|
||||
//!
|
||||
//! Validates `X-Signature-Ed25519` and `X-Signature-Timestamp` headers
|
||||
//! on incoming Discord interaction webhooks, per Discord's security requirements.
|
||||
//! Validates request signatures for incoming webhooks:
|
||||
//! - Discord: `X-Signature-Ed25519` and `X-Signature-Timestamp` headers
|
||||
//! - Slack: `X-Slack-Signature` and `X-Slack-Request-Timestamp` headers
|
||||
//!
|
||||
//! See: <https://discord.com/developers/docs/interactions/overview#validating-security-request-headers>
|
||||
//! See: <https://api.slack.com/authentication/verifying-requests-from-slack>
|
||||
|
||||
/// Verify a Discord interaction signature.
|
||||
///
|
||||
@@ -50,6 +52,60 @@ pub fn verify_discord_signature(
|
||||
verifying_key.verify_strict(&message, &signature).is_ok()
|
||||
}
|
||||
|
||||
/// Verify a Slack webhook signature using HMAC-SHA256.
|
||||
///
|
||||
/// Slack signs each webhook request with HMAC-SHA256 using:
|
||||
/// - basestring = `"v0:" + timestamp + ":" + body`
|
||||
/// - signature = hex-encoded HMAC-SHA256(signing_secret, basestring)
|
||||
/// - header = `"v0=" + signature` (in `X-Slack-Signature` header)
|
||||
///
|
||||
/// Includes staleness check: rejects requests with timestamps older than 5 minutes.
|
||||
/// Returns `true` if the signature is valid, `false` on any error
|
||||
/// (bad timing, mismatched signature, invalid format, etc.).
|
||||
pub fn verify_slack_signature(
|
||||
signing_secret: &str,
|
||||
timestamp: &str,
|
||||
body: &[u8],
|
||||
signature_header: &str,
|
||||
now_secs: i64,
|
||||
) -> bool {
|
||||
use hmac::{Hmac, Mac};
|
||||
use sha2::Sha256;
|
||||
|
||||
// 1. Parse and check staleness (5-minute window)
|
||||
let ts: i64 = match timestamp.parse() {
|
||||
Ok(v) => v,
|
||||
Err(_) => return false,
|
||||
};
|
||||
if (now_secs - ts).abs() > 300 {
|
||||
return false;
|
||||
}
|
||||
|
||||
// 2. Build the basestring: "v0:{timestamp}:{body}"
|
||||
let mut basestring = Vec::with_capacity(3 + timestamp.len() + 1 + body.len());
|
||||
basestring.extend_from_slice(b"v0:");
|
||||
basestring.extend_from_slice(timestamp.as_bytes());
|
||||
basestring.push(b':');
|
||||
basestring.extend_from_slice(body);
|
||||
|
||||
// 3. Compute HMAC-SHA256
|
||||
let mut mac = match Hmac::<Sha256>::new_from_slice(signing_secret.as_bytes()) {
|
||||
Ok(m) => m,
|
||||
Err(_) => return false,
|
||||
};
|
||||
mac.update(&basestring);
|
||||
let computed = mac.finalize().into_bytes();
|
||||
let computed_hex = hex::encode(computed);
|
||||
let expected = format!("v0={}", computed_hex);
|
||||
|
||||
// 4. Constant-time compare (avoids timing side-channels)
|
||||
use subtle::ConstantTimeEq;
|
||||
expected
|
||||
.as_bytes()
|
||||
.ct_eq(signature_header.as_bytes())
|
||||
.into()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
@@ -338,4 +394,264 @@ mod tests {
|
||||
"Negative timestamp should be rejected"
|
||||
);
|
||||
}
|
||||
|
||||
// ── Category: HMAC-SHA256 Signature Verification (Slack) ────────────
|
||||
|
||||
/// Helper: compute expected Slack signature for a given secret, timestamp, and body.
|
||||
fn sign_slack_message(signing_secret: &str, timestamp: &str, body: &[u8]) -> String {
|
||||
use hmac::{Hmac, Mac};
|
||||
use sha2::Sha256;
|
||||
|
||||
let mut basestring = Vec::new();
|
||||
basestring.extend_from_slice(b"v0:");
|
||||
basestring.extend_from_slice(timestamp.as_bytes());
|
||||
basestring.push(b':');
|
||||
basestring.extend_from_slice(body);
|
||||
|
||||
let mut mac = Hmac::<Sha256>::new_from_slice(signing_secret.as_bytes()).unwrap();
|
||||
mac.update(&basestring);
|
||||
let computed = mac.finalize().into_bytes();
|
||||
format!("v0={}", hex::encode(computed))
|
||||
}
|
||||
|
||||
const SLACK_TEST_TS: i64 = 1234567890;
|
||||
|
||||
#[test]
|
||||
fn test_slack_valid_signature_succeeds() {
|
||||
let signing_secret = "my-signing-secret";
|
||||
let timestamp = "1234567890";
|
||||
let body = b"token=xyzz0WbapA4vBCDEFasx0q6G&team_id=T1DC2JH3J";
|
||||
|
||||
let signature = sign_slack_message(signing_secret, timestamp, body);
|
||||
assert!(verify_slack_signature(
|
||||
signing_secret,
|
||||
timestamp,
|
||||
body,
|
||||
&signature,
|
||||
SLACK_TEST_TS
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_slack_tampered_body_fails() {
|
||||
let signing_secret = "my-signing-secret";
|
||||
let timestamp = "1234567890";
|
||||
let original_body = b"token=xyzz0WbapA4vBCDEFasx0q6G&team_id=T1DC2JH3J";
|
||||
let tampered_body = b"token=MODIFIED&team_id=T1DC2JH3J";
|
||||
|
||||
let signature = sign_slack_message(signing_secret, timestamp, original_body);
|
||||
assert!(
|
||||
!verify_slack_signature(
|
||||
signing_secret,
|
||||
timestamp,
|
||||
tampered_body,
|
||||
&signature,
|
||||
SLACK_TEST_TS
|
||||
),
|
||||
"Signature for different body should fail"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_slack_tampered_timestamp_fails() {
|
||||
let signing_secret = "my-signing-secret";
|
||||
let timestamp = "1234567890";
|
||||
let body = b"token=xyzz0WbapA4vBCDEFasx0q6G&team_id=T1DC2JH3J";
|
||||
|
||||
let signature = sign_slack_message(signing_secret, timestamp, body);
|
||||
assert!(
|
||||
!verify_slack_signature(
|
||||
signing_secret,
|
||||
"9999999999", // Different timestamp in signature
|
||||
body,
|
||||
&signature,
|
||||
SLACK_TEST_TS
|
||||
),
|
||||
"Signature with wrong timestamp should fail"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_slack_tampered_signature_fails() {
|
||||
let signing_secret = "my-signing-secret";
|
||||
let timestamp = "1234567890";
|
||||
let body = b"token=xyzz0WbapA4vBCDEFasx0q6G&team_id=T1DC2JH3J";
|
||||
|
||||
let signature = sign_slack_message(signing_secret, timestamp, body);
|
||||
// Flip a byte in the signature hex (change first char after "v0=")
|
||||
let chars: Vec<char> = signature.chars().collect();
|
||||
let mut new_chars = chars.clone();
|
||||
if chars.len() > 3 {
|
||||
new_chars[3] = if chars[3] == 'a' { 'b' } else { 'a' };
|
||||
}
|
||||
let modified_sig: String = new_chars.iter().collect();
|
||||
|
||||
assert!(
|
||||
!verify_slack_signature(
|
||||
signing_secret,
|
||||
timestamp,
|
||||
body,
|
||||
&modified_sig,
|
||||
SLACK_TEST_TS
|
||||
),
|
||||
"Tampered signature should fail"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_slack_stale_timestamp_rejected() {
|
||||
let signing_secret = "my-signing-secret";
|
||||
let timestamp = "1234567890";
|
||||
let body = b"token=xyzz0WbapA4vBCDEFasx0q6G";
|
||||
|
||||
let signature = sign_slack_message(signing_secret, timestamp, body);
|
||||
// now_secs is 400 seconds after timestamp — too stale
|
||||
assert!(
|
||||
!verify_slack_signature(
|
||||
signing_secret,
|
||||
timestamp,
|
||||
body,
|
||||
&signature,
|
||||
SLACK_TEST_TS + 400
|
||||
),
|
||||
"Stale timestamp (400s old) should be rejected"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_slack_future_timestamp_rejected() {
|
||||
let signing_secret = "my-signing-secret";
|
||||
let timestamp = "1234567890";
|
||||
let body = b"token=xyzz0WbapA4vBCDEFasx0q6G";
|
||||
|
||||
let signature = sign_slack_message(signing_secret, timestamp, body);
|
||||
// now_secs is 400 seconds before timestamp — future
|
||||
assert!(
|
||||
!verify_slack_signature(
|
||||
signing_secret,
|
||||
timestamp,
|
||||
body,
|
||||
&signature,
|
||||
SLACK_TEST_TS - 400
|
||||
),
|
||||
"Future timestamp (400s ahead) should be rejected"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_slack_boundary_300s_accepted() {
|
||||
let signing_secret = "my-signing-secret";
|
||||
let timestamp = "1234567890";
|
||||
let body = b"token=xyzz0WbapA4vBCDEFasx0q6G";
|
||||
|
||||
let signature = sign_slack_message(signing_secret, timestamp, body);
|
||||
// Exactly 300 seconds difference — should be accepted
|
||||
assert!(
|
||||
verify_slack_signature(
|
||||
signing_secret,
|
||||
timestamp,
|
||||
body,
|
||||
&signature,
|
||||
SLACK_TEST_TS + 300
|
||||
),
|
||||
"Timestamp exactly 300s old should be accepted"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_slack_boundary_301s_rejected() {
|
||||
let signing_secret = "my-signing-secret";
|
||||
let timestamp = "1234567890";
|
||||
let body = b"token=xyzz0WbapA4vBCDEFasx0q6G";
|
||||
|
||||
let signature = sign_slack_message(signing_secret, timestamp, body);
|
||||
// 301 seconds difference — should be rejected
|
||||
assert!(
|
||||
!verify_slack_signature(
|
||||
signing_secret,
|
||||
timestamp,
|
||||
body,
|
||||
&signature,
|
||||
SLACK_TEST_TS + 301
|
||||
),
|
||||
"Timestamp 301s old should be rejected"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_slack_non_numeric_timestamp_rejected() {
|
||||
let signing_secret = "my-signing-secret";
|
||||
let body = b"token=xyzz0WbapA4vBCDEFasx0q6G";
|
||||
|
||||
assert!(
|
||||
!verify_slack_signature(signing_secret, "not-a-number", body, "v0=abc123", 0),
|
||||
"Non-numeric timestamp should be rejected"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_slack_missing_v0_prefix_fails() {
|
||||
let signing_secret = "my-signing-secret";
|
||||
let timestamp = "1234567890";
|
||||
let body = b"token=xyzz0WbapA4vBCDEFasx0q6G";
|
||||
|
||||
let signature = sign_slack_message(signing_secret, timestamp, body);
|
||||
// Remove the "v0=" prefix
|
||||
let bad_sig = signature.strip_prefix("v0=").unwrap_or(&signature);
|
||||
|
||||
assert!(
|
||||
!verify_slack_signature(signing_secret, timestamp, body, bad_sig, SLACK_TEST_TS),
|
||||
"Missing v0= prefix should fail"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_slack_wrong_signing_secret_fails() {
|
||||
let secret_a = "secret-a";
|
||||
let secret_b = "secret-b";
|
||||
let timestamp = "1234567890";
|
||||
let body = b"token=xyzz0WbapA4vBCDEFasx0q6G";
|
||||
|
||||
let signature = sign_slack_message(secret_a, timestamp, body);
|
||||
// Try to verify with a different secret
|
||||
assert!(
|
||||
!verify_slack_signature(secret_b, timestamp, body, &signature, SLACK_TEST_TS),
|
||||
"Signature from different secret should fail"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_slack_empty_body_valid() {
|
||||
let signing_secret = "my-signing-secret";
|
||||
let timestamp = "1234567890";
|
||||
let body = b"";
|
||||
|
||||
let signature = sign_slack_message(signing_secret, timestamp, body);
|
||||
assert!(
|
||||
verify_slack_signature(signing_secret, timestamp, body, &signature, SLACK_TEST_TS),
|
||||
"Empty body with valid signature should succeed"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_slack_negative_timestamp_rejected() {
|
||||
let signing_secret = "my-signing-secret";
|
||||
let body = b"token=xyzz0WbapA4vBCDEFasx0q6G";
|
||||
|
||||
assert!(
|
||||
!verify_slack_signature(signing_secret, "-1", body, "v0=abc123", 0),
|
||||
"Negative timestamp should be rejected"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_slack_empty_timestamp_rejected() {
|
||||
let signing_secret = "my-signing-secret";
|
||||
let body = b"token=xyzz0WbapA4vBCDEFasx0q6G";
|
||||
|
||||
assert!(
|
||||
!verify_slack_signature(signing_secret, "", body, "v0=abc123", 0),
|
||||
"Empty timestamp should be rejected"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,690 @@
|
||||
//! WASM channel binary storage with integrity verification.
|
||||
//!
|
||||
//! Stores compiled WASM channels in the database with BLAKE3 hash verification.
|
||||
//! Mirrors the pattern in `crate::tools::wasm::storage` but without capabilities table.
|
||||
//!
|
||||
//! # Storage Flow
|
||||
//!
|
||||
//! ```text
|
||||
//! WASM bytes ──► BLAKE3 hash ──► Store in database
|
||||
//! │ (binary + hash)
|
||||
//! │
|
||||
//! └──► Later: Load ──► Verify hash ──► Return bytes
|
||||
//! ```
|
||||
|
||||
use async_trait::async_trait;
|
||||
use chrono::{DateTime, Utc};
|
||||
#[cfg(feature = "postgres")]
|
||||
use deadpool_postgres::Pool;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::tools::wasm::storage::{compute_binary_hash, verify_binary_integrity};
|
||||
|
||||
/// A stored WASM channel (metadata only, no binary).
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct StoredWasmChannel {
|
||||
pub id: Uuid,
|
||||
pub user_id: String,
|
||||
pub name: String,
|
||||
pub version: String,
|
||||
pub wit_version: String,
|
||||
pub description: String,
|
||||
pub capabilities_json: String,
|
||||
pub status: String,
|
||||
pub created_at: DateTime<Utc>,
|
||||
pub updated_at: DateTime<Utc>,
|
||||
}
|
||||
|
||||
/// Full channel data including binary.
|
||||
#[derive(Debug)]
|
||||
pub struct StoredWasmChannelWithBinary {
|
||||
pub channel: StoredWasmChannel,
|
||||
pub wasm_binary: Vec<u8>,
|
||||
pub binary_hash: Vec<u8>,
|
||||
}
|
||||
|
||||
/// Parameters for storing a new WASM channel.
|
||||
pub struct StoreChannelParams {
|
||||
pub user_id: String,
|
||||
pub name: String,
|
||||
pub version: String,
|
||||
pub wit_version: String,
|
||||
pub description: String,
|
||||
pub wasm_binary: Vec<u8>,
|
||||
pub capabilities_json: String,
|
||||
}
|
||||
|
||||
/// Error from WASM channel storage operations.
|
||||
#[derive(Debug, Clone, thiserror::Error)]
|
||||
pub enum WasmChannelStoreError {
|
||||
#[error("Channel not found: {0}")]
|
||||
NotFound(String),
|
||||
|
||||
#[error("Binary integrity check failed: hash mismatch")]
|
||||
IntegrityCheckFailed,
|
||||
|
||||
#[error("Database error: {0}")]
|
||||
Database(String),
|
||||
|
||||
#[error("Invalid data: {0}")]
|
||||
InvalidData(String),
|
||||
}
|
||||
|
||||
/// Trait for WASM channel storage.
|
||||
#[async_trait]
|
||||
pub trait WasmChannelStore: Send + Sync {
|
||||
/// Store a new WASM channel.
|
||||
async fn store(
|
||||
&self,
|
||||
params: StoreChannelParams,
|
||||
) -> Result<StoredWasmChannel, WasmChannelStoreError>;
|
||||
|
||||
/// Get channel metadata (without binary).
|
||||
async fn get(
|
||||
&self,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
) -> Result<StoredWasmChannel, WasmChannelStoreError>;
|
||||
|
||||
/// Get channel with binary (verifies integrity).
|
||||
async fn get_with_binary(
|
||||
&self,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
) -> Result<StoredWasmChannelWithBinary, WasmChannelStoreError>;
|
||||
|
||||
/// List all channels for a user.
|
||||
async fn list(&self, user_id: &str) -> Result<Vec<StoredWasmChannel>, WasmChannelStoreError>;
|
||||
|
||||
/// Delete a channel.
|
||||
async fn delete(&self, user_id: &str, name: &str) -> Result<bool, WasmChannelStoreError>;
|
||||
}
|
||||
|
||||
// ==================== PostgreSQL implementation ====================
|
||||
|
||||
/// PostgreSQL implementation of WasmChannelStore.
|
||||
#[cfg(feature = "postgres")]
|
||||
pub struct PostgresWasmChannelStore {
|
||||
pool: Pool,
|
||||
}
|
||||
|
||||
#[cfg(feature = "postgres")]
|
||||
impl PostgresWasmChannelStore {
|
||||
pub fn new(pool: Pool) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "postgres")]
|
||||
#[async_trait]
|
||||
impl WasmChannelStore for PostgresWasmChannelStore {
|
||||
async fn store(
|
||||
&self,
|
||||
params: StoreChannelParams,
|
||||
) -> Result<StoredWasmChannel, WasmChannelStoreError> {
|
||||
let mut client = self
|
||||
.pool
|
||||
.get()
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
let binary_hash = compute_binary_hash(¶ms.wasm_binary);
|
||||
let id = Uuid::new_v4();
|
||||
let now = Utc::now();
|
||||
|
||||
// Wrap delete + insert in a transaction for atomicity
|
||||
let tx = client
|
||||
.transaction()
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
// Delete any existing version for this (user_id, name) — upgrade-in-place
|
||||
tx.execute(
|
||||
"DELETE FROM wasm_channels WHERE user_id = $1 AND name = $2",
|
||||
&[¶ms.user_id, ¶ms.name],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
let row = tx
|
||||
.query_one(
|
||||
r#"
|
||||
INSERT INTO wasm_channels (
|
||||
id, user_id, name, version, wit_version, description, wasm_binary, binary_hash,
|
||||
capabilities_json, status, created_at, updated_at
|
||||
)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, 'active', $10, $10)
|
||||
RETURNING id, user_id, name, version, wit_version, description,
|
||||
capabilities_json, status, created_at, updated_at
|
||||
"#,
|
||||
&[
|
||||
&id,
|
||||
¶ms.user_id,
|
||||
¶ms.name,
|
||||
¶ms.version,
|
||||
¶ms.wit_version,
|
||||
¶ms.description,
|
||||
¶ms.wasm_binary,
|
||||
&binary_hash,
|
||||
¶ms.capabilities_json,
|
||||
&now,
|
||||
],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
let channel = pg_row_to_channel(&row)?;
|
||||
|
||||
tx.commit()
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
Ok(channel)
|
||||
}
|
||||
|
||||
async fn get(
|
||||
&self,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
) -> Result<StoredWasmChannel, WasmChannelStoreError> {
|
||||
let client = self
|
||||
.pool
|
||||
.get()
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
let row = client
|
||||
.query_opt(
|
||||
r#"
|
||||
SELECT id, user_id, name, version, wit_version, description,
|
||||
capabilities_json, status, created_at, updated_at
|
||||
FROM wasm_channels
|
||||
WHERE user_id = $1 AND name = $2
|
||||
"#,
|
||||
&[&user_id, &name],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
match row {
|
||||
Some(r) => pg_row_to_channel(&r),
|
||||
None => Err(WasmChannelStoreError::NotFound(name.to_string())),
|
||||
}
|
||||
}
|
||||
|
||||
async fn get_with_binary(
|
||||
&self,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
) -> Result<StoredWasmChannelWithBinary, WasmChannelStoreError> {
|
||||
let client = self
|
||||
.pool
|
||||
.get()
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
let row = client
|
||||
.query_opt(
|
||||
r#"
|
||||
SELECT id, user_id, name, version, wit_version, description,
|
||||
wasm_binary, binary_hash,
|
||||
capabilities_json, status, created_at, updated_at
|
||||
FROM wasm_channels
|
||||
WHERE user_id = $1 AND name = $2
|
||||
"#,
|
||||
&[&user_id, &name],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
match row {
|
||||
Some(r) => {
|
||||
let wasm_binary: Vec<u8> = r.get("wasm_binary");
|
||||
let binary_hash: Vec<u8> = r.get("binary_hash");
|
||||
|
||||
if !verify_binary_integrity(&wasm_binary, &binary_hash) {
|
||||
tracing::error!(
|
||||
user_id = user_id,
|
||||
name = name,
|
||||
"WASM channel binary integrity check failed"
|
||||
);
|
||||
return Err(WasmChannelStoreError::IntegrityCheckFailed);
|
||||
}
|
||||
|
||||
let channel = StoredWasmChannel {
|
||||
id: r.get("id"),
|
||||
user_id: r.get("user_id"),
|
||||
name: r.get("name"),
|
||||
version: r.get("version"),
|
||||
wit_version: r.get("wit_version"),
|
||||
description: r.get("description"),
|
||||
capabilities_json: r.get("capabilities_json"),
|
||||
status: r.get("status"),
|
||||
created_at: r.get("created_at"),
|
||||
updated_at: r.get("updated_at"),
|
||||
};
|
||||
|
||||
Ok(StoredWasmChannelWithBinary {
|
||||
channel,
|
||||
wasm_binary,
|
||||
binary_hash,
|
||||
})
|
||||
}
|
||||
None => Err(WasmChannelStoreError::NotFound(name.to_string())),
|
||||
}
|
||||
}
|
||||
|
||||
async fn list(&self, user_id: &str) -> Result<Vec<StoredWasmChannel>, WasmChannelStoreError> {
|
||||
let client = self
|
||||
.pool
|
||||
.get()
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
let rows = client
|
||||
.query(
|
||||
r#"
|
||||
SELECT id, user_id, name, version, wit_version, description,
|
||||
capabilities_json, status, created_at, updated_at
|
||||
FROM wasm_channels
|
||||
WHERE user_id = $1
|
||||
ORDER BY name
|
||||
"#,
|
||||
&[&user_id],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
rows.into_iter().map(|r| pg_row_to_channel(&r)).collect()
|
||||
}
|
||||
|
||||
async fn delete(&self, user_id: &str, name: &str) -> Result<bool, WasmChannelStoreError> {
|
||||
let client = self
|
||||
.pool
|
||||
.get()
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
let result = client
|
||||
.execute(
|
||||
"DELETE FROM wasm_channels WHERE user_id = $1 AND name = $2",
|
||||
&[&user_id, &name],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
Ok(result > 0)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "postgres")]
|
||||
fn pg_row_to_channel(
|
||||
row: &tokio_postgres::Row,
|
||||
) -> Result<StoredWasmChannel, WasmChannelStoreError> {
|
||||
Ok(StoredWasmChannel {
|
||||
id: row.get("id"),
|
||||
user_id: row.get("user_id"),
|
||||
name: row.get("name"),
|
||||
version: row.get("version"),
|
||||
wit_version: row.get("wit_version"),
|
||||
description: row.get("description"),
|
||||
capabilities_json: row.get("capabilities_json"),
|
||||
status: row.get("status"),
|
||||
created_at: row.get("created_at"),
|
||||
updated_at: row.get("updated_at"),
|
||||
})
|
||||
}
|
||||
|
||||
// ==================== libSQL implementation ====================
|
||||
|
||||
/// libSQL/Turso implementation of WasmChannelStore.
|
||||
///
|
||||
/// Holds an `Arc<Database>` handle and creates a fresh connection per operation,
|
||||
/// matching the connection-per-request pattern used by the main `LibSqlBackend`.
|
||||
#[cfg(feature = "libsql")]
|
||||
pub struct LibSqlWasmChannelStore {
|
||||
db: std::sync::Arc<libsql::Database>,
|
||||
}
|
||||
|
||||
#[cfg(feature = "libsql")]
|
||||
impl LibSqlWasmChannelStore {
|
||||
pub fn new(db: std::sync::Arc<libsql::Database>) -> Self {
|
||||
Self { db }
|
||||
}
|
||||
|
||||
async fn connect(&self) -> Result<libsql::Connection, WasmChannelStoreError> {
|
||||
let conn = self
|
||||
.db
|
||||
.connect()
|
||||
.map_err(|e| WasmChannelStoreError::Database(format!("Connection failed: {}", e)))?;
|
||||
conn.query("PRAGMA busy_timeout = 5000", ())
|
||||
.await
|
||||
.map_err(|e| {
|
||||
WasmChannelStoreError::Database(format!("Failed to set busy_timeout: {}", e))
|
||||
})?;
|
||||
Ok(conn)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "libsql")]
|
||||
#[async_trait]
|
||||
impl WasmChannelStore for LibSqlWasmChannelStore {
|
||||
async fn store(
|
||||
&self,
|
||||
params: StoreChannelParams,
|
||||
) -> Result<StoredWasmChannel, WasmChannelStoreError> {
|
||||
let binary_hash = compute_binary_hash(¶ms.wasm_binary);
|
||||
let id = Uuid::new_v4();
|
||||
let now = Utc::now().to_rfc3339_opts(chrono::SecondsFormat::Millis, true);
|
||||
|
||||
let conn = self.connect().await?;
|
||||
let tx = conn
|
||||
.transaction()
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
// Delete any existing version for this (user_id, name) — upgrade-in-place
|
||||
tx.execute(
|
||||
"DELETE FROM wasm_channels WHERE user_id = ?1 AND name = ?2",
|
||||
libsql::params![params.user_id.as_str(), params.name.as_str()],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
tx.execute(
|
||||
r#"
|
||||
INSERT INTO wasm_channels (
|
||||
id, user_id, name, version, wit_version, description, wasm_binary, binary_hash,
|
||||
capabilities_json, status, created_at, updated_at
|
||||
)
|
||||
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, 'active', ?10, ?10)
|
||||
"#,
|
||||
libsql::params![
|
||||
id.to_string(),
|
||||
params.user_id.as_str(),
|
||||
params.name.as_str(),
|
||||
params.version.as_str(),
|
||||
params.wit_version.as_str(),
|
||||
params.description.as_str(),
|
||||
libsql::Value::Blob(params.wasm_binary),
|
||||
libsql::Value::Blob(binary_hash),
|
||||
params.capabilities_json.as_str(),
|
||||
now.as_str(),
|
||||
],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
// Read back the row within the same transaction
|
||||
let mut rows = tx
|
||||
.query(
|
||||
r#"
|
||||
SELECT id, user_id, name, version, wit_version, description,
|
||||
capabilities_json, status, created_at, updated_at
|
||||
FROM wasm_channels
|
||||
WHERE user_id = ?1 AND name = ?2
|
||||
"#,
|
||||
libsql::params![params.user_id.as_str(), params.name.as_str()],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
let row = rows
|
||||
.next()
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?
|
||||
.ok_or_else(|| {
|
||||
WasmChannelStoreError::Database("Insert succeeded but row not found".into())
|
||||
})?;
|
||||
|
||||
let channel = libsql_row_to_channel(&row)?;
|
||||
|
||||
tx.commit()
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
Ok(channel)
|
||||
}
|
||||
|
||||
async fn get(
|
||||
&self,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
) -> Result<StoredWasmChannel, WasmChannelStoreError> {
|
||||
let conn = self.connect().await?;
|
||||
let mut rows = conn
|
||||
.query(
|
||||
r#"
|
||||
SELECT id, user_id, name, version, wit_version, description,
|
||||
capabilities_json, status, created_at, updated_at
|
||||
FROM wasm_channels
|
||||
WHERE user_id = ?1 AND name = ?2
|
||||
"#,
|
||||
libsql::params![user_id, name],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
match rows
|
||||
.next()
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?
|
||||
{
|
||||
Some(row) => libsql_row_to_channel(&row),
|
||||
None => Err(WasmChannelStoreError::NotFound(name.to_string())),
|
||||
}
|
||||
}
|
||||
|
||||
async fn get_with_binary(
|
||||
&self,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
) -> Result<StoredWasmChannelWithBinary, WasmChannelStoreError> {
|
||||
let conn = self.connect().await?;
|
||||
let mut rows = conn
|
||||
.query(
|
||||
r#"
|
||||
SELECT id, user_id, name, version, wit_version, description,
|
||||
wasm_binary, binary_hash,
|
||||
capabilities_json, status, created_at, updated_at
|
||||
FROM wasm_channels
|
||||
WHERE user_id = ?1 AND name = ?2
|
||||
"#,
|
||||
libsql::params![user_id, name],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
match rows
|
||||
.next()
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?
|
||||
{
|
||||
Some(row) => {
|
||||
let wasm_binary: Vec<u8> = row
|
||||
.get(6)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
let binary_hash: Vec<u8> = row
|
||||
.get(7)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
if !verify_binary_integrity(&wasm_binary, &binary_hash) {
|
||||
tracing::error!(
|
||||
user_id = user_id,
|
||||
name = name,
|
||||
"WASM channel binary integrity check failed"
|
||||
);
|
||||
return Err(WasmChannelStoreError::IntegrityCheckFailed);
|
||||
}
|
||||
|
||||
let channel = libsql_row_to_channel_with_offset(&row)?;
|
||||
|
||||
Ok(StoredWasmChannelWithBinary {
|
||||
channel,
|
||||
wasm_binary,
|
||||
binary_hash,
|
||||
})
|
||||
}
|
||||
None => Err(WasmChannelStoreError::NotFound(name.to_string())),
|
||||
}
|
||||
}
|
||||
|
||||
async fn list(&self, user_id: &str) -> Result<Vec<StoredWasmChannel>, WasmChannelStoreError> {
|
||||
let conn = self.connect().await?;
|
||||
let mut rows = conn
|
||||
.query(
|
||||
r#"
|
||||
SELECT id, user_id, name, version, wit_version, description,
|
||||
capabilities_json, status, created_at, updated_at
|
||||
FROM wasm_channels
|
||||
WHERE user_id = ?1
|
||||
ORDER BY name
|
||||
"#,
|
||||
libsql::params![user_id],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
let mut channels = Vec::new();
|
||||
while let Some(row) = rows
|
||||
.next()
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?
|
||||
{
|
||||
channels.push(libsql_row_to_channel(&row)?);
|
||||
}
|
||||
Ok(channels)
|
||||
}
|
||||
|
||||
async fn delete(&self, user_id: &str, name: &str) -> Result<bool, WasmChannelStoreError> {
|
||||
let conn = self.connect().await?;
|
||||
let result = conn
|
||||
.execute(
|
||||
"DELETE FROM wasm_channels WHERE user_id = ?1 AND name = ?2",
|
||||
libsql::params![user_id, name],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
Ok(result > 0)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "libsql")]
|
||||
#[allow(dead_code)]
|
||||
fn libsql_channel_opt_text(s: Option<&str>) -> libsql::Value {
|
||||
match s {
|
||||
Some(s) => libsql::Value::Text(s.to_string()),
|
||||
None => libsql::Value::Null,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "libsql")]
|
||||
fn libsql_channel_parse_ts(s: &str) -> Result<DateTime<Utc>, WasmChannelStoreError> {
|
||||
if let Ok(dt) = chrono::DateTime::parse_from_rfc3339(s) {
|
||||
return Ok(dt.with_timezone(&Utc));
|
||||
}
|
||||
if let Ok(ndt) = chrono::NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S%.f") {
|
||||
return Ok(ndt.and_utc());
|
||||
}
|
||||
if let Ok(ndt) = chrono::NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S") {
|
||||
return Ok(ndt.and_utc());
|
||||
}
|
||||
Err(WasmChannelStoreError::InvalidData(format!(
|
||||
"unparseable timestamp: {:?}",
|
||||
s
|
||||
)))
|
||||
}
|
||||
|
||||
/// Parse a channel row with standard column order (no binary columns).
|
||||
/// Columns: id(0), user_id(1), name(2), version(3), wit_version(4), description(5),
|
||||
/// capabilities_json(6), status(7), created_at(8), updated_at(9)
|
||||
#[cfg(feature = "libsql")]
|
||||
fn libsql_row_to_channel(row: &libsql::Row) -> Result<StoredWasmChannel, WasmChannelStoreError> {
|
||||
let id_str: String = row
|
||||
.get(0)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
let created_at_str: String = row
|
||||
.get(8)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
let updated_at_str: String = row
|
||||
.get(9)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
Ok(StoredWasmChannel {
|
||||
id: id_str
|
||||
.parse()
|
||||
.map_err(|e: uuid::Error| WasmChannelStoreError::InvalidData(e.to_string()))?,
|
||||
user_id: row
|
||||
.get(1)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?,
|
||||
name: row
|
||||
.get(2)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?,
|
||||
version: row
|
||||
.get(3)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?,
|
||||
wit_version: row
|
||||
.get(4)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?,
|
||||
description: row
|
||||
.get(5)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?,
|
||||
capabilities_json: row
|
||||
.get(6)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?,
|
||||
status: row
|
||||
.get(7)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?,
|
||||
created_at: libsql_channel_parse_ts(&created_at_str)?,
|
||||
updated_at: libsql_channel_parse_ts(&updated_at_str)?,
|
||||
})
|
||||
}
|
||||
|
||||
/// Parse a channel row when binary columns are present (get_with_binary query).
|
||||
/// Columns: id(0), user_id(1), name(2), version(3), wit_version(4), description(5),
|
||||
/// wasm_binary(6), binary_hash(7),
|
||||
/// capabilities_json(8), status(9), created_at(10), updated_at(11)
|
||||
#[cfg(feature = "libsql")]
|
||||
fn libsql_row_to_channel_with_offset(
|
||||
row: &libsql::Row,
|
||||
) -> Result<StoredWasmChannel, WasmChannelStoreError> {
|
||||
let id_str: String = row
|
||||
.get(0)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
let created_at_str: String = row
|
||||
.get(10)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
let updated_at_str: String = row
|
||||
.get(11)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?;
|
||||
|
||||
Ok(StoredWasmChannel {
|
||||
id: id_str
|
||||
.parse()
|
||||
.map_err(|e: uuid::Error| WasmChannelStoreError::InvalidData(e.to_string()))?,
|
||||
user_id: row
|
||||
.get(1)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?,
|
||||
name: row
|
||||
.get(2)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?,
|
||||
version: row
|
||||
.get(3)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?,
|
||||
wit_version: row
|
||||
.get(4)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?,
|
||||
description: row
|
||||
.get(5)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?,
|
||||
capabilities_json: row
|
||||
.get(8)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?,
|
||||
status: row
|
||||
.get(9)
|
||||
.map_err(|e| WasmChannelStoreError::Database(e.to_string()))?,
|
||||
created_at: libsql_channel_parse_ts(&created_at_str)?,
|
||||
updated_at: libsql_channel_parse_ts(&updated_at_str)?,
|
||||
})
|
||||
}
|
||||
@@ -933,8 +933,19 @@ impl WasmChannel {
|
||||
Self::add_host_functions(&mut linker)?;
|
||||
|
||||
// Instantiate using the generated bindings
|
||||
let instance = SandboxedChannel::instantiate(store, &component, &linker)
|
||||
.map_err(|e| WasmChannelError::Instantiation(e.to_string()))?;
|
||||
let instance = SandboxedChannel::instantiate(store, &component, &linker).map_err(|e| {
|
||||
let msg = e.to_string();
|
||||
if msg.contains("near:agent") || msg.contains("import") {
|
||||
WasmChannelError::Instantiation(format!(
|
||||
"{msg}. This may indicate a WIT version mismatch — \
|
||||
the channel was compiled against a different WIT than the host supports \
|
||||
(host WIT: {}). Rebuild the channel against the current WIT.",
|
||||
crate::tools::wasm::WIT_CHANNEL_VERSION
|
||||
))
|
||||
} else {
|
||||
WasmChannelError::Instantiation(msg)
|
||||
}
|
||||
})?;
|
||||
|
||||
Ok(instance)
|
||||
}
|
||||
@@ -2479,7 +2490,7 @@ fn status_to_wit(status: &StatusUpdate, metadata: &serde_json::Value) -> wit_cha
|
||||
message: format!("Tool started: {}", name),
|
||||
metadata_json,
|
||||
},
|
||||
StatusUpdate::ToolCompleted { name, success } => wit_channel::StatusUpdate {
|
||||
StatusUpdate::ToolCompleted { name, success, .. } => wit_channel::StatusUpdate {
|
||||
status: wit_channel::StatusType::ToolCompleted,
|
||||
message: format!(
|
||||
"Tool completed: {} ({})",
|
||||
@@ -3387,6 +3398,8 @@ mod tests {
|
||||
&crate::channels::StatusUpdate::ToolCompleted {
|
||||
name: "http_request".to_string(),
|
||||
success: true,
|
||||
error: None,
|
||||
parameters: None,
|
||||
},
|
||||
&metadata,
|
||||
);
|
||||
@@ -3407,6 +3420,8 @@ mod tests {
|
||||
&crate::channels::StatusUpdate::ToolCompleted {
|
||||
name: "http_request".to_string(),
|
||||
success: false,
|
||||
error: Some("connection refused".to_string()),
|
||||
parameters: None,
|
||||
},
|
||||
&metadata,
|
||||
);
|
||||
|
||||
+148
-30
@@ -2,7 +2,7 @@
|
||||
|
||||
use axum::{
|
||||
extract::{Request, State},
|
||||
http::{HeaderMap, StatusCode},
|
||||
http::{HeaderMap, Method, StatusCode},
|
||||
middleware::Next,
|
||||
response::{IntoResponse, Response},
|
||||
};
|
||||
@@ -14,10 +14,44 @@ pub struct AuthState {
|
||||
pub token: String,
|
||||
}
|
||||
|
||||
/// Whether query-string token auth is allowed for this request.
|
||||
///
|
||||
/// Only GET requests to streaming endpoints may use `?token=xxx`. This
|
||||
/// minimizes token-in-URL exposure on state-changing routes, where the token
|
||||
/// would leak via server logs, Referer headers, and browser history.
|
||||
///
|
||||
/// Allowed endpoints:
|
||||
/// - SSE: `/api/chat/events`, `/api/logs/events` (EventSource can't set headers)
|
||||
/// - WebSocket: `/api/chat/ws` (WS upgrade can't set custom headers)
|
||||
///
|
||||
/// If you add a new SSE or WebSocket endpoint, add its path here.
|
||||
fn allows_query_token_auth(request: &Request) -> bool {
|
||||
if request.method() != Method::GET {
|
||||
return false;
|
||||
}
|
||||
|
||||
matches!(
|
||||
request.uri().path(),
|
||||
"/api/chat/events" | "/api/logs/events" | "/api/chat/ws"
|
||||
)
|
||||
}
|
||||
|
||||
/// Extract the `token` query parameter value, URL-decoded.
|
||||
fn query_token(request: &Request) -> Option<String> {
|
||||
let query = request.uri().query()?;
|
||||
url::form_urlencoded::parse(query.as_bytes()).find_map(|(k, v)| {
|
||||
if k == "token" {
|
||||
Some(v.into_owned())
|
||||
} else {
|
||||
None
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
/// Auth middleware that validates bearer token from header or query param.
|
||||
///
|
||||
/// SSE connections can't set headers from `EventSource`, so we also accept
|
||||
/// `?token=xxx` as a query parameter.
|
||||
/// `?token=xxx` as a query parameter, but only on SSE endpoints.
|
||||
pub async fn auth_middleware(
|
||||
State(auth): State<AuthState>,
|
||||
headers: HeaderMap,
|
||||
@@ -35,15 +69,12 @@ pub async fn auth_middleware(
|
||||
return next.run(request).await;
|
||||
}
|
||||
|
||||
// Fall back to query parameter for SSE EventSource (constant-time comparison)
|
||||
if let Some(query) = request.uri().query() {
|
||||
for pair in query.split('&') {
|
||||
if let Some(token) = pair.strip_prefix("token=")
|
||||
&& bool::from(token.as_bytes().ct_eq(auth.token.as_bytes()))
|
||||
{
|
||||
return next.run(request).await;
|
||||
}
|
||||
}
|
||||
// Fall back to query parameter, but only for SSE endpoints (constant-time comparison).
|
||||
if allows_query_token_auth(&request)
|
||||
&& let Some(token) = query_token(&request)
|
||||
&& bool::from(token.as_bytes().ct_eq(auth.token.as_bytes()))
|
||||
{
|
||||
return next.run(request).await;
|
||||
}
|
||||
|
||||
(StatusCode::UNAUTHORIZED, "Invalid or missing auth token").into_response()
|
||||
@@ -62,24 +93,28 @@ mod tests {
|
||||
assert_eq!(cloned.token, "test-token");
|
||||
}
|
||||
|
||||
// === QA Plan - Web gateway auth tests ===
|
||||
|
||||
use axum::Router;
|
||||
use axum::body::Body;
|
||||
use axum::middleware;
|
||||
use axum::routing::get;
|
||||
use axum::routing::{get, post};
|
||||
use tower::ServiceExt;
|
||||
|
||||
async fn dummy_handler() -> &'static str {
|
||||
"ok"
|
||||
}
|
||||
|
||||
/// Router with streaming endpoints (query auth allowed) and regular
|
||||
/// endpoints (query auth rejected).
|
||||
fn test_app(token: &str) -> Router {
|
||||
let state = AuthState {
|
||||
token: token.to_string(),
|
||||
};
|
||||
Router::new()
|
||||
.route("/test", get(dummy_handler))
|
||||
.route("/api/chat/events", get(dummy_handler))
|
||||
.route("/api/logs/events", get(dummy_handler))
|
||||
.route("/api/chat/ws", get(dummy_handler))
|
||||
.route("/api/chat/history", get(dummy_handler))
|
||||
.route("/api/chat/send", post(dummy_handler))
|
||||
.layer(middleware::from_fn_with_state(state, auth_middleware))
|
||||
}
|
||||
|
||||
@@ -87,7 +122,7 @@ mod tests {
|
||||
async fn test_valid_bearer_token_passes() {
|
||||
let app = test_app("secret-token");
|
||||
let req = Request::builder()
|
||||
.uri("/test")
|
||||
.uri("/api/chat/events")
|
||||
.header("Authorization", "Bearer secret-token")
|
||||
.body(Body::empty())
|
||||
.unwrap();
|
||||
@@ -99,7 +134,7 @@ mod tests {
|
||||
async fn test_invalid_bearer_token_rejected() {
|
||||
let app = test_app("secret-token");
|
||||
let req = Request::builder()
|
||||
.uri("/test")
|
||||
.uri("/api/chat/events")
|
||||
.header("Authorization", "Bearer wrong-token")
|
||||
.body(Body::empty())
|
||||
.unwrap();
|
||||
@@ -108,10 +143,10 @@ mod tests {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_missing_auth_header_falls_through_to_query() {
|
||||
async fn test_query_token_allowed_for_chat_events() {
|
||||
let app = test_app("secret-token");
|
||||
let req = Request::builder()
|
||||
.uri("/test?token=secret-token")
|
||||
.uri("/api/chat/events?token=secret-token")
|
||||
.body(Body::empty())
|
||||
.unwrap();
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
@@ -119,10 +154,80 @@ mod tests {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_query_param_invalid_token_rejected() {
|
||||
async fn test_query_token_allowed_for_logs_events() {
|
||||
let app = test_app("secret-token");
|
||||
let req = Request::builder()
|
||||
.uri("/test?token=wrong-token")
|
||||
.uri("/api/logs/events?token=secret-token")
|
||||
.body(Body::empty())
|
||||
.unwrap();
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
assert_eq!(resp.status(), StatusCode::OK);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_query_token_allowed_for_ws_upgrade() {
|
||||
let app = test_app("secret-token");
|
||||
let req = Request::builder()
|
||||
.uri("/api/chat/ws?token=secret-token")
|
||||
.body(Body::empty())
|
||||
.unwrap();
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
assert_eq!(resp.status(), StatusCode::OK);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_query_token_url_encoded() {
|
||||
// Token with characters that get percent-encoded in URLs.
|
||||
let raw_token = "tok+en/with spaces";
|
||||
let app = test_app(raw_token);
|
||||
let req = Request::builder()
|
||||
.uri("/api/chat/events?token=tok%2Ben%2Fwith%20spaces")
|
||||
.body(Body::empty())
|
||||
.unwrap();
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
assert_eq!(resp.status(), StatusCode::OK);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_query_token_url_encoded_mismatch() {
|
||||
let app = test_app("real-token");
|
||||
// Encoded value decodes to "wrong-token", not "real-token".
|
||||
let req = Request::builder()
|
||||
.uri("/api/chat/events?token=wrong%2Dtoken")
|
||||
.body(Body::empty())
|
||||
.unwrap();
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_query_token_rejected_for_non_sse_get() {
|
||||
let app = test_app("secret-token");
|
||||
let req = Request::builder()
|
||||
.uri("/api/chat/history?token=secret-token")
|
||||
.body(Body::empty())
|
||||
.unwrap();
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_query_token_rejected_for_post() {
|
||||
let app = test_app("secret-token");
|
||||
let req = Request::builder()
|
||||
.method(Method::POST)
|
||||
.uri("/api/chat/send?token=secret-token")
|
||||
.body(Body::empty())
|
||||
.unwrap();
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_query_token_invalid_rejected() {
|
||||
let app = test_app("secret-token");
|
||||
let req = Request::builder()
|
||||
.uri("/api/chat/events?token=wrong-token")
|
||||
.body(Body::empty())
|
||||
.unwrap();
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
@@ -132,17 +237,32 @@ mod tests {
|
||||
#[tokio::test]
|
||||
async fn test_no_auth_at_all_rejected() {
|
||||
let app = test_app("secret-token");
|
||||
let req = Request::builder().uri("/test").body(Body::empty()).unwrap();
|
||||
let req = Request::builder()
|
||||
.uri("/api/chat/events")
|
||||
.body(Body::empty())
|
||||
.unwrap();
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_bearer_prefix_case_insensitive() {
|
||||
// RFC 6750 Section 2.1: auth-scheme comparison must be case-insensitive.
|
||||
async fn test_bearer_header_works_for_post() {
|
||||
let app = test_app("secret-token");
|
||||
let req = Request::builder()
|
||||
.uri("/test")
|
||||
.method(Method::POST)
|
||||
.uri("/api/chat/send")
|
||||
.header("Authorization", "Bearer secret-token")
|
||||
.body(Body::empty())
|
||||
.unwrap();
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
assert_eq!(resp.status(), StatusCode::OK);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_bearer_prefix_case_insensitive() {
|
||||
let app = test_app("secret-token");
|
||||
let req = Request::builder()
|
||||
.uri("/api/chat/events")
|
||||
.header("Authorization", "bearer secret-token")
|
||||
.body(Body::empty())
|
||||
.unwrap();
|
||||
@@ -154,7 +274,7 @@ mod tests {
|
||||
async fn test_bearer_prefix_mixed_case() {
|
||||
let app = test_app("secret-token");
|
||||
let req = Request::builder()
|
||||
.uri("/test")
|
||||
.uri("/api/chat/events")
|
||||
.header("Authorization", "BEARER secret-token")
|
||||
.body(Body::empty())
|
||||
.unwrap();
|
||||
@@ -166,7 +286,7 @@ mod tests {
|
||||
async fn test_empty_bearer_token_rejected() {
|
||||
let app = test_app("secret-token");
|
||||
let req = Request::builder()
|
||||
.uri("/test")
|
||||
.uri("/api/chat/events")
|
||||
.header("Authorization", "Bearer ")
|
||||
.body(Body::empty())
|
||||
.unwrap();
|
||||
@@ -176,11 +296,9 @@ mod tests {
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_token_with_whitespace_rejected() {
|
||||
// Extra space after "Bearer " means the token value starts with a space,
|
||||
// which should not match the expected token.
|
||||
let app = test_app("secret-token");
|
||||
let req = Request::builder()
|
||||
.uri("/test")
|
||||
.uri("/api/chat/events")
|
||||
.header("Authorization", "Bearer secret-token")
|
||||
.body(Body::empty())
|
||||
.unwrap();
|
||||
|
||||
@@ -142,7 +142,7 @@ pub async fn chat_auth_token_handler(
|
||||
.await
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
|
||||
if result.status == "authenticated" {
|
||||
if result.is_authenticated() {
|
||||
// Auto-activate so tools are available immediately
|
||||
let msg = match ext_mgr.activate(&req.extension_name).await {
|
||||
Ok(r) => format!(
|
||||
@@ -170,13 +170,14 @@ pub async fn chat_auth_token_handler(
|
||||
// Re-emit auth_required for retry
|
||||
state.sse.broadcast(SseEvent::AuthRequired {
|
||||
extension_name: req.extension_name.clone(),
|
||||
instructions: result.instructions.clone(),
|
||||
auth_url: result.auth_url.clone(),
|
||||
setup_url: result.setup_url.clone(),
|
||||
instructions: result.instructions().map(String::from),
|
||||
auth_url: result.auth_url().map(String::from),
|
||||
setup_url: result.setup_url().map(String::from),
|
||||
});
|
||||
Ok(Json(ActionResponse::fail(
|
||||
result
|
||||
.instructions
|
||||
.instructions()
|
||||
.map(String::from)
|
||||
.unwrap_or_else(|| "Invalid token".to_string()),
|
||||
)))
|
||||
}
|
||||
|
||||
@@ -33,7 +33,7 @@ pub async fn extensions_list_handler(
|
||||
"failed".to_string()
|
||||
} else if !ext.authenticated {
|
||||
"installed".to_string()
|
||||
} else if ext.active && ext.name == "telegram" {
|
||||
} else if ext.active {
|
||||
let has_paired = pairing_store
|
||||
.read_allow_from(&ext.name)
|
||||
.map(|list| !list.is_empty())
|
||||
@@ -59,6 +59,7 @@ pub async fn extensions_list_handler(
|
||||
active: ext.active,
|
||||
tools: ext.tools,
|
||||
needs_setup: ext.needs_setup,
|
||||
has_auth: ext.has_auth,
|
||||
activation_status,
|
||||
activation_error: ext.activation_error,
|
||||
}
|
||||
@@ -123,7 +124,11 @@ pub async fn extensions_activate_handler(
|
||||
))?;
|
||||
|
||||
match ext_mgr.activate(&name).await {
|
||||
Ok(result) => Ok(Json(ActionResponse::ok(result.message))),
|
||||
Ok(result) => {
|
||||
// Activation just loads the WASM module. Auth (OAuth/manual) is
|
||||
// triggered separately via save_setup_secrets or the auth endpoint.
|
||||
Ok(Json(ActionResponse::ok(result.message)))
|
||||
}
|
||||
Err(activate_err) => {
|
||||
let err_str = activate_err.to_string();
|
||||
let needs_auth = err_str.contains("authentication")
|
||||
@@ -136,7 +141,7 @@ pub async fn extensions_activate_handler(
|
||||
|
||||
// Activation failed due to auth; try authenticating first.
|
||||
match ext_mgr.auth(&name, None).await {
|
||||
Ok(auth_result) if auth_result.status == "authenticated" => {
|
||||
Ok(auth_result) if auth_result.is_authenticated() => {
|
||||
// Auth succeeded, retry activation.
|
||||
match ext_mgr.activate(&name).await {
|
||||
Ok(result) => Ok(Json(ActionResponse::ok(result.message))),
|
||||
@@ -147,13 +152,13 @@ pub async fn extensions_activate_handler(
|
||||
// Auth in progress (OAuth URL or awaiting manual token).
|
||||
let mut resp = ActionResponse::fail(
|
||||
auth_result
|
||||
.instructions
|
||||
.clone()
|
||||
.instructions()
|
||||
.map(String::from)
|
||||
.unwrap_or_else(|| format!("'{}' requires authentication.", name)),
|
||||
);
|
||||
resp.auth_url = auth_result.auth_url;
|
||||
resp.awaiting_token = Some(auth_result.awaiting_token);
|
||||
resp.instructions = auth_result.instructions;
|
||||
resp.auth_url = auth_result.auth_url().map(String::from);
|
||||
resp.awaiting_token = Some(auth_result.is_awaiting_token());
|
||||
resp.instructions = auth_result.instructions().map(String::from);
|
||||
Ok(Json(resp))
|
||||
}
|
||||
Err(auth_err) => Ok(Json(ActionResponse::fail(format!(
|
||||
|
||||
@@ -181,6 +181,9 @@ pub async fn jobs_detail_handler(
|
||||
});
|
||||
}
|
||||
|
||||
let mode = store.get_sandbox_job_mode(job.id).await.ok().flatten();
|
||||
let is_claude_code = mode.as_deref() == Some("claude_code");
|
||||
|
||||
return Ok(Json(JobDetailResponse {
|
||||
id: job.id,
|
||||
title: job.task.clone(),
|
||||
@@ -193,11 +196,11 @@ pub async fn jobs_detail_handler(
|
||||
elapsed_secs,
|
||||
project_dir: Some(job.project_dir.clone()),
|
||||
browse_url: Some(format!("/projects/{}/", browse_id)),
|
||||
job_mode: {
|
||||
let mode = store.get_sandbox_job_mode(job.id).await.ok().flatten();
|
||||
mode.filter(|m| m != "worker")
|
||||
},
|
||||
job_mode: mode.filter(|m| m != "worker"),
|
||||
transitions,
|
||||
can_restart: state.job_manager.is_some(),
|
||||
can_prompt: is_claude_code && state.prompt_queue.is_some(),
|
||||
job_kind: Some("sandbox".to_string()),
|
||||
}));
|
||||
}
|
||||
|
||||
@@ -208,6 +211,12 @@ pub async fn jobs_detail_handler(
|
||||
(end - start).num_seconds().max(0) as u64
|
||||
});
|
||||
|
||||
// Only show prompt bar for jobs that have a running worker (Pending/InProgress).
|
||||
// Stuck jobs have no active worker loop, so messages would be silently dropped.
|
||||
let is_promptable = matches!(
|
||||
ctx.state,
|
||||
crate::context::JobState::Pending | crate::context::JobState::InProgress
|
||||
);
|
||||
return Ok(Json(JobDetailResponse {
|
||||
id: ctx.job_id,
|
||||
title: ctx.title.clone(),
|
||||
@@ -222,6 +231,9 @@ pub async fn jobs_detail_handler(
|
||||
browse_url: None,
|
||||
job_mode: None,
|
||||
transitions: Vec::new(),
|
||||
can_restart: state.scheduler.is_some(),
|
||||
can_prompt: is_promptable && state.scheduler.is_some(),
|
||||
job_kind: Some("agent".to_string()),
|
||||
}));
|
||||
}
|
||||
|
||||
@@ -295,108 +307,164 @@ pub async fn jobs_restart_handler(
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Database not available".to_string(),
|
||||
))?;
|
||||
let jm = state.job_manager.as_ref().ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Sandbox not enabled".to_string(),
|
||||
))?;
|
||||
|
||||
let old_job_id = Uuid::parse_str(&id)
|
||||
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
|
||||
|
||||
let old_job = store
|
||||
.get_sandbox_job(old_job_id)
|
||||
.await
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
|
||||
.ok_or((StatusCode::NOT_FOUND, "Job not found".to_string()))?;
|
||||
// Try sandbox job restart first.
|
||||
if let Ok(Some(old_job)) = store.get_sandbox_job(old_job_id).await {
|
||||
if old_job.status != "interrupted" && old_job.status != "failed" {
|
||||
return Err((
|
||||
StatusCode::CONFLICT,
|
||||
format!("Cannot restart job in state '{}'", old_job.status),
|
||||
));
|
||||
}
|
||||
|
||||
if old_job.status != "interrupted" && old_job.status != "failed" {
|
||||
return Err((
|
||||
StatusCode::CONFLICT,
|
||||
format!("Cannot restart job in state '{}'", old_job.status),
|
||||
));
|
||||
let jm = state.job_manager.as_ref().ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Sandbox not enabled".to_string(),
|
||||
))?;
|
||||
|
||||
// Enrich the task with failure context.
|
||||
let task = if let Some(ref reason) = old_job.failure_reason {
|
||||
format!(
|
||||
"Previous attempt failed: {}. Retry: {}",
|
||||
reason, old_job.task
|
||||
)
|
||||
} else {
|
||||
old_job.task.clone()
|
||||
};
|
||||
|
||||
let new_job_id = Uuid::new_v4();
|
||||
let now = chrono::Utc::now();
|
||||
|
||||
let record = crate::history::SandboxJobRecord {
|
||||
id: new_job_id,
|
||||
task: task.clone(),
|
||||
status: "creating".to_string(),
|
||||
user_id: old_job.user_id.clone(),
|
||||
project_dir: old_job.project_dir.clone(),
|
||||
success: None,
|
||||
failure_reason: None,
|
||||
created_at: now,
|
||||
started_at: None,
|
||||
completed_at: None,
|
||||
credential_grants_json: old_job.credential_grants_json.clone(),
|
||||
};
|
||||
store
|
||||
.save_sandbox_job(&record)
|
||||
.await
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
|
||||
let mode = match store.get_sandbox_job_mode(old_job_id).await {
|
||||
Ok(Some(m)) if m == "claude_code" => {
|
||||
crate::orchestrator::job_manager::JobMode::ClaudeCode
|
||||
}
|
||||
_ => crate::orchestrator::job_manager::JobMode::Worker,
|
||||
};
|
||||
|
||||
let credential_grants: Vec<crate::orchestrator::auth::CredentialGrant> =
|
||||
serde_json::from_str(&old_job.credential_grants_json).unwrap_or_else(|e| {
|
||||
tracing::warn!(
|
||||
job_id = %old_job.id,
|
||||
"Failed to deserialize credential grants from stored job: {}. \
|
||||
Restarted job will have no credentials.",
|
||||
e
|
||||
);
|
||||
vec![]
|
||||
});
|
||||
|
||||
let project_dir = std::path::PathBuf::from(&old_job.project_dir);
|
||||
let _token = jm
|
||||
.create_job(
|
||||
new_job_id,
|
||||
&task,
|
||||
Some(project_dir),
|
||||
mode,
|
||||
credential_grants,
|
||||
)
|
||||
.await
|
||||
.map_err(|e| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
format!("Failed to create container: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
store
|
||||
.update_sandbox_job_status(new_job_id, "running", None, None, Some(now), None)
|
||||
.await
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
|
||||
return Ok(Json(serde_json::json!({
|
||||
"status": "restarted",
|
||||
"old_job_id": old_job_id,
|
||||
"new_job_id": new_job_id,
|
||||
})));
|
||||
}
|
||||
|
||||
// Create a new job with the same task and project_dir.
|
||||
let new_job_id = Uuid::new_v4();
|
||||
let now = chrono::Utc::now();
|
||||
// Try agent job restart: dispatch a new job via the scheduler.
|
||||
if let Ok(Some(old_job)) = store.get_job(old_job_id).await {
|
||||
if old_job.state.is_active() {
|
||||
return Err((
|
||||
StatusCode::CONFLICT,
|
||||
format!("Cannot restart job in state '{}'", old_job.state),
|
||||
));
|
||||
}
|
||||
|
||||
let record = crate::history::SandboxJobRecord {
|
||||
id: new_job_id,
|
||||
task: old_job.task.clone(),
|
||||
status: "creating".to_string(),
|
||||
user_id: old_job.user_id.clone(),
|
||||
project_dir: old_job.project_dir.clone(),
|
||||
success: None,
|
||||
failure_reason: None,
|
||||
created_at: now,
|
||||
started_at: None,
|
||||
completed_at: None,
|
||||
credential_grants_json: old_job.credential_grants_json.clone(),
|
||||
};
|
||||
store
|
||||
.save_sandbox_job(&record)
|
||||
.await
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
let slot = state.scheduler.as_ref().ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Scheduler not available".to_string(),
|
||||
))?;
|
||||
let scheduler_guard = slot.read().await;
|
||||
let scheduler = scheduler_guard.as_ref().ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Agent not started yet".to_string(),
|
||||
))?;
|
||||
|
||||
// Look up the original job's mode so the restart uses the same mode.
|
||||
let mode = match store.get_sandbox_job_mode(old_job_id).await {
|
||||
Ok(Some(m)) if m == "claude_code" => crate::orchestrator::job_manager::JobMode::ClaudeCode,
|
||||
_ => crate::orchestrator::job_manager::JobMode::Worker,
|
||||
};
|
||||
// Look up failure reason (O(1) point lookup).
|
||||
let failure_reason = store
|
||||
.get_agent_job_failure_reason(old_job_id)
|
||||
.await
|
||||
.ok()
|
||||
.flatten()
|
||||
.unwrap_or_default();
|
||||
|
||||
// Restore credential grants from the original job so the restarted container
|
||||
// has access to the same secrets.
|
||||
let credential_grants: Vec<crate::orchestrator::auth::CredentialGrant> =
|
||||
serde_json::from_str(&old_job.credential_grants_json).unwrap_or_else(|e| {
|
||||
tracing::warn!(
|
||||
job_id = %old_job.id,
|
||||
"Failed to deserialize credential grants from stored job: {}. \
|
||||
Restarted job will have no credentials.",
|
||||
e
|
||||
);
|
||||
vec![]
|
||||
});
|
||||
|
||||
let project_dir = std::path::PathBuf::from(&old_job.project_dir);
|
||||
let _token = jm
|
||||
.create_job(
|
||||
new_job_id,
|
||||
&old_job.task,
|
||||
Some(project_dir),
|
||||
mode,
|
||||
credential_grants,
|
||||
)
|
||||
.await
|
||||
.map_err(|e| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
format!("Failed to create container: {}", e),
|
||||
let title = if !failure_reason.is_empty() {
|
||||
format!(
|
||||
"Previous attempt failed: {}. Retry: {}",
|
||||
failure_reason, old_job.title
|
||||
)
|
||||
})?;
|
||||
} else {
|
||||
old_job.title.clone()
|
||||
};
|
||||
|
||||
store
|
||||
.update_sandbox_job_status(new_job_id, "running", None, None, Some(now), None)
|
||||
.await
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
let new_job_id = scheduler
|
||||
.dispatch_job(&old_job.user_id, &title, &old_job.description, None)
|
||||
.await
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
|
||||
Ok(Json(serde_json::json!({
|
||||
"status": "restarted",
|
||||
"old_job_id": old_job_id,
|
||||
"new_job_id": new_job_id,
|
||||
})))
|
||||
return Ok(Json(serde_json::json!({
|
||||
"status": "restarted",
|
||||
"old_job_id": old_job_id,
|
||||
"new_job_id": new_job_id,
|
||||
})));
|
||||
}
|
||||
|
||||
Err((StatusCode::NOT_FOUND, "Job not found".to_string()))
|
||||
}
|
||||
|
||||
/// Submit a follow-up prompt to a running Claude Code sandbox job.
|
||||
/// Submit a follow-up prompt to a running job.
|
||||
///
|
||||
/// Routes to the appropriate backend:
|
||||
/// - Claude Code sandbox jobs → prompt queue (polled by the bridge)
|
||||
/// - Agent (non-sandbox) jobs → WorkerMessage injection via scheduler
|
||||
/// - Worker-mode sandbox jobs → not supported (no mechanism to inject)
|
||||
pub async fn jobs_prompt_handler(
|
||||
State(state): State<Arc<GatewayState>>,
|
||||
Path(id): Path<String>,
|
||||
Json(body): Json<serde_json::Value>,
|
||||
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
||||
let prompt_queue = state.prompt_queue.as_ref().ok_or((
|
||||
StatusCode::NOT_IMPLEMENTED,
|
||||
"Claude Code not configured".to_string(),
|
||||
))?;
|
||||
|
||||
let job_id: uuid::Uuid = id
|
||||
.parse()
|
||||
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
|
||||
@@ -412,17 +480,57 @@ pub async fn jobs_prompt_handler(
|
||||
|
||||
let done = body.get("done").and_then(|v| v.as_bool()).unwrap_or(false);
|
||||
|
||||
let prompt = crate::orchestrator::api::PendingPrompt { content, done };
|
||||
|
||||
// Try sandbox job path: check if we have a sandbox record for this ID.
|
||||
if let Some(ref s) = state.store
|
||||
&& let Ok(Some(_)) = s.get_sandbox_job(job_id).await
|
||||
{
|
||||
let mut queue = prompt_queue.lock().await;
|
||||
queue.entry(job_id).or_default().push_back(prompt);
|
||||
// It's a sandbox job. Check if Claude Code mode.
|
||||
let mode = s.get_sandbox_job_mode(job_id).await.ok().flatten();
|
||||
if mode.as_deref() == Some("claude_code") {
|
||||
let prompt_queue = state.prompt_queue.as_ref().ok_or((
|
||||
StatusCode::NOT_IMPLEMENTED,
|
||||
"Claude Code not configured".to_string(),
|
||||
))?;
|
||||
let prompt = crate::orchestrator::api::PendingPrompt { content, done };
|
||||
{
|
||||
let mut queue = prompt_queue.lock().await;
|
||||
queue.entry(job_id).or_default().push_back(prompt);
|
||||
}
|
||||
return Ok(Json(serde_json::json!({
|
||||
"status": "queued",
|
||||
"job_id": job_id.to_string(),
|
||||
})));
|
||||
} else {
|
||||
return Err((
|
||||
StatusCode::NOT_IMPLEMENTED,
|
||||
"Follow-up prompts are not supported for worker-mode sandbox jobs".to_string(),
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
Ok(Json(serde_json::json!({
|
||||
"status": "queued",
|
||||
"job_id": job_id.to_string(),
|
||||
})))
|
||||
// Try agent job path: send via scheduler.
|
||||
let slot = state.scheduler.as_ref().ok_or((
|
||||
StatusCode::NOT_IMPLEMENTED,
|
||||
"Agent job prompts require the scheduler to be configured".to_string(),
|
||||
))?;
|
||||
let scheduler_guard = slot.read().await;
|
||||
if let Some(ref scheduler) = *scheduler_guard
|
||||
&& scheduler.is_running(job_id).await
|
||||
{
|
||||
scheduler
|
||||
.send_message(job_id, content)
|
||||
.await
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
return Ok(Json(serde_json::json!({
|
||||
"status": "sent",
|
||||
"job_id": job_id.to_string(),
|
||||
})));
|
||||
}
|
||||
|
||||
Err((
|
||||
StatusCode::NOT_FOUND,
|
||||
"Job not found or not running".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
/// Load persisted job events for a job (for history replay on page open).
|
||||
|
||||
@@ -159,10 +159,10 @@ pub async fn memory_search_handler(
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
|
||||
let hits: Vec<SearchHit> = results
|
||||
.iter()
|
||||
.into_iter()
|
||||
.map(|r| SearchHit {
|
||||
path: r.document_id.to_string(),
|
||||
content: r.content.clone(),
|
||||
path: r.document_path,
|
||||
content: r.content,
|
||||
score: r.score as f64,
|
||||
})
|
||||
.collect();
|
||||
|
||||
@@ -147,6 +147,10 @@ pub async fn routines_trigger_handler(
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
|
||||
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
|
||||
|
||||
if routine.user_id != state.user_id {
|
||||
return Err((StatusCode::FORBIDDEN, "Access denied".to_string()));
|
||||
}
|
||||
|
||||
// Send the routine prompt through the message pipeline as a manual trigger.
|
||||
let prompt = match &routine.action {
|
||||
crate::agent::routine::RoutineAction::Lightweight { prompt, .. } => prompt.clone(),
|
||||
@@ -156,7 +160,12 @@ pub async fn routines_trigger_handler(
|
||||
};
|
||||
|
||||
let content = format!("[routine:{}] {}", routine.name, prompt);
|
||||
let msg = IncomingMessage::new("gateway", &state.user_id, content);
|
||||
let thread_id = format!(
|
||||
"routine-{}-{}",
|
||||
routine_id,
|
||||
chrono::Utc::now().timestamp_millis()
|
||||
);
|
||||
let msg = IncomingMessage::new("gateway", &state.user_id, content).with_thread(thread_id);
|
||||
|
||||
let tx_guard = state.msg_tx.read().await;
|
||||
let tx = tx_guard.as_ref().ok_or((
|
||||
|
||||
@@ -148,7 +148,14 @@ pub async fn skills_install_handler(
|
||||
.await
|
||||
.map_err(|e| (StatusCode::BAD_REQUEST, e.to_string()))?
|
||||
} else if let Some(ref catalog) = state.skill_catalog {
|
||||
let url = crate::skills::catalog::skill_download_url(catalog.registry_url(), &req.name);
|
||||
// Prefer slug (e.g. "owner/skill-name") over display name for the
|
||||
// download URL, since the registry endpoint expects a slug.
|
||||
let download_key = req
|
||||
.slug
|
||||
.as_deref()
|
||||
.filter(|s| !s.is_empty())
|
||||
.unwrap_or(&req.name);
|
||||
let url = crate::skills::catalog::skill_download_url(catalog.registry_url(), download_key);
|
||||
crate::tools::builtin::skill_tools::fetch_skill_content(&url)
|
||||
.await
|
||||
.map_err(|e| (StatusCode::BAD_GATEWAY, e.to_string()))?
|
||||
|
||||
+23
-11
@@ -63,13 +63,11 @@ impl GatewayChannel {
|
||||
/// If no auth token is configured, generates a random one and prints it.
|
||||
pub fn new(config: GatewayConfig) -> Self {
|
||||
let auth_token = config.auth_token.clone().unwrap_or_else(|| {
|
||||
use rand::Rng;
|
||||
let token: String = rand::thread_rng()
|
||||
.sample_iter(&rand::distributions::Alphanumeric)
|
||||
.take(32)
|
||||
.map(char::from)
|
||||
.collect();
|
||||
token
|
||||
use rand::RngCore;
|
||||
use rand::rngs::OsRng;
|
||||
let mut bytes = [0u8; 32];
|
||||
OsRng.fill_bytes(&mut bytes);
|
||||
bytes.iter().map(|b| format!("{b:02x}")).collect()
|
||||
});
|
||||
|
||||
let state = Arc::new(GatewayState {
|
||||
@@ -84,6 +82,7 @@ impl GatewayChannel {
|
||||
store: None,
|
||||
job_manager: None,
|
||||
prompt_queue: None,
|
||||
scheduler: None,
|
||||
user_id: config.user_id.clone(),
|
||||
shutdown_tx: tokio::sync::RwLock::new(None),
|
||||
ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())),
|
||||
@@ -94,7 +93,6 @@ impl GatewayChannel {
|
||||
registry_entries: Vec::new(),
|
||||
cost_guard: None,
|
||||
startup_time: std::time::Instant::now(),
|
||||
restart_requested: std::sync::atomic::AtomicBool::new(false),
|
||||
});
|
||||
|
||||
Self {
|
||||
@@ -108,7 +106,8 @@ impl GatewayChannel {
|
||||
fn rebuild_state(&mut self, mutate: impl FnOnce(&mut GatewayState)) {
|
||||
let mut new_state = GatewayState {
|
||||
msg_tx: tokio::sync::RwLock::new(None),
|
||||
sse: SseManager::new(),
|
||||
// Preserve the existing broadcast channel so sender handles remain valid.
|
||||
sse: SseManager::from_sender(self.state.sse.sender()),
|
||||
workspace: self.state.workspace.clone(),
|
||||
session_manager: self.state.session_manager.clone(),
|
||||
log_broadcaster: self.state.log_broadcaster.clone(),
|
||||
@@ -118,6 +117,7 @@ impl GatewayChannel {
|
||||
store: self.state.store.clone(),
|
||||
job_manager: self.state.job_manager.clone(),
|
||||
prompt_queue: self.state.prompt_queue.clone(),
|
||||
scheduler: self.state.scheduler.clone(),
|
||||
user_id: self.state.user_id.clone(),
|
||||
shutdown_tx: tokio::sync::RwLock::new(None),
|
||||
ws_tracker: self.state.ws_tracker.clone(),
|
||||
@@ -128,7 +128,6 @@ impl GatewayChannel {
|
||||
registry_entries: self.state.registry_entries.clone(),
|
||||
cost_guard: self.state.cost_guard.clone(),
|
||||
startup_time: self.state.startup_time,
|
||||
restart_requested: std::sync::atomic::AtomicBool::new(false),
|
||||
};
|
||||
mutate(&mut new_state);
|
||||
self.state = Arc::new(new_state);
|
||||
@@ -198,6 +197,12 @@ impl GatewayChannel {
|
||||
self
|
||||
}
|
||||
|
||||
/// Inject the scheduler for sending follow-up messages to agent jobs.
|
||||
pub fn with_scheduler(mut self, slot: crate::tools::builtin::SchedulerSlot) -> Self {
|
||||
self.rebuild_state(|s| s.scheduler = Some(slot));
|
||||
self
|
||||
}
|
||||
|
||||
/// Inject the skill registry for skill management API.
|
||||
pub fn with_skill_registry(mut self, sr: Arc<std::sync::RwLock<SkillRegistry>>) -> Self {
|
||||
self.rebuild_state(|s| s.skill_registry = Some(sr));
|
||||
@@ -297,9 +302,16 @@ impl Channel for GatewayChannel {
|
||||
name,
|
||||
thread_id: thread_id.clone(),
|
||||
},
|
||||
StatusUpdate::ToolCompleted { name, success } => SseEvent::ToolCompleted {
|
||||
StatusUpdate::ToolCompleted {
|
||||
name,
|
||||
success,
|
||||
error,
|
||||
parameters,
|
||||
} => SseEvent::ToolCompleted {
|
||||
name,
|
||||
success,
|
||||
error,
|
||||
parameters,
|
||||
thread_id: thread_id.clone(),
|
||||
},
|
||||
StatusUpdate::ToolResult { name, preview } => SseEvent::ToolResult {
|
||||
|
||||
+603
-50
@@ -156,6 +156,8 @@ pub struct GatewayState {
|
||||
pub skill_registry: Option<Arc<std::sync::RwLock<crate::skills::SkillRegistry>>>,
|
||||
/// Skill catalog for searching the ClawHub registry.
|
||||
pub skill_catalog: Option<Arc<crate::skills::catalog::SkillCatalog>>,
|
||||
/// Scheduler for sending follow-up messages to running agent jobs.
|
||||
pub scheduler: Option<crate::tools::builtin::SchedulerSlot>,
|
||||
/// Rate limiter for chat endpoints (30 messages per 60 seconds).
|
||||
pub chat_rate_limiter: RateLimiter,
|
||||
/// Registry catalog entries for the available extensions API.
|
||||
@@ -165,8 +167,6 @@ pub struct GatewayState {
|
||||
pub cost_guard: Option<Arc<crate::agent::cost_guard::CostGuard>>,
|
||||
/// Server startup time for uptime calculation.
|
||||
pub startup_time: std::time::Instant,
|
||||
/// Flag set when a restart has been requested via the API.
|
||||
pub restart_requested: std::sync::atomic::AtomicBool,
|
||||
}
|
||||
|
||||
/// Start the gateway HTTP server.
|
||||
@@ -192,7 +192,9 @@ pub async fn start_server(
|
||||
})?;
|
||||
|
||||
// Public routes (no auth)
|
||||
let public = Router::new().route("/api/health", get(health_handler));
|
||||
let public = Router::new()
|
||||
.route("/api/health", get(health_handler))
|
||||
.route("/oauth/callback", get(oauth_callback_handler));
|
||||
|
||||
// Protected routes (require auth)
|
||||
let auth_state = AuthState { token: auth_token };
|
||||
@@ -247,8 +249,6 @@ pub async fn start_server(
|
||||
"/api/extensions/{name}/setup",
|
||||
get(extensions_setup_handler).post(extensions_setup_submit_handler),
|
||||
)
|
||||
// Gateway management
|
||||
.route("/api/gateway/restart", post(gateway_restart_handler))
|
||||
// Pairing
|
||||
.route("/api/pairing/{channel}", get(pairing_list_handler))
|
||||
.route(
|
||||
@@ -426,12 +426,192 @@ async fn health_handler() -> Json<HealthResponse> {
|
||||
})
|
||||
}
|
||||
|
||||
/// Return an OAuth error landing page response.
|
||||
fn oauth_error_page(label: &str) -> axum::response::Response {
|
||||
let html = crate::cli::oauth_defaults::landing_html(label, false);
|
||||
axum::response::Html(html).into_response()
|
||||
}
|
||||
|
||||
/// OAuth callback handler for the web gateway.
|
||||
///
|
||||
/// This is a PUBLIC route (no Bearer token required) because OAuth providers
|
||||
/// redirect the user's browser here. The `state` query parameter correlates
|
||||
/// the callback with a pending OAuth flow registered by `start_wasm_oauth()`.
|
||||
///
|
||||
/// Used on hosted instances where `IRONCLAW_OAUTH_CALLBACK_URL` points to
|
||||
/// the gateway (e.g., `https://kind-deer.agent1.near.ai/oauth/callback`).
|
||||
/// Local/desktop mode continues to use the TCP listener on port 9876.
|
||||
async fn oauth_callback_handler(
|
||||
State(state): State<Arc<GatewayState>>,
|
||||
Query(params): Query<std::collections::HashMap<String, String>>,
|
||||
) -> impl IntoResponse {
|
||||
use crate::cli::oauth_defaults;
|
||||
|
||||
// Check for error from OAuth provider (e.g., user denied consent)
|
||||
if let Some(error) = params.get("error") {
|
||||
let description = params
|
||||
.get("error_description")
|
||||
.cloned()
|
||||
.unwrap_or_else(|| error.clone());
|
||||
return oauth_error_page(&description);
|
||||
}
|
||||
|
||||
let state_param = match params.get("state") {
|
||||
Some(s) if !s.is_empty() => s.clone(),
|
||||
_ => return oauth_error_page("IronClaw"),
|
||||
};
|
||||
|
||||
let code = match params.get("code") {
|
||||
Some(c) if !c.is_empty() => c.clone(),
|
||||
_ => return oauth_error_page("IronClaw"),
|
||||
};
|
||||
|
||||
// Look up the pending flow by CSRF state (atomic remove prevents replay)
|
||||
let ext_mgr = match state.extension_manager.as_ref() {
|
||||
Some(mgr) => mgr,
|
||||
None => return oauth_error_page("IronClaw"),
|
||||
};
|
||||
|
||||
// Strip instance prefix from state for registry lookup.
|
||||
// Platform nginx sends `state=instance:nonce` but flows are keyed by nonce only.
|
||||
let lookup_key = oauth_defaults::strip_instance_prefix(&state_param);
|
||||
|
||||
let flow = ext_mgr
|
||||
.pending_oauth_flows()
|
||||
.write()
|
||||
.await
|
||||
.remove(lookup_key);
|
||||
|
||||
let flow = match flow {
|
||||
Some(f) => f,
|
||||
None => {
|
||||
tracing::warn!(
|
||||
state = %state_param,
|
||||
lookup_key = %lookup_key,
|
||||
"OAuth callback received with unknown or expired state"
|
||||
);
|
||||
return oauth_error_page("IronClaw");
|
||||
}
|
||||
};
|
||||
|
||||
// Check flow expiry (5 minutes, matching TCP listener timeout)
|
||||
if flow.created_at.elapsed() > oauth_defaults::OAUTH_FLOW_EXPIRY {
|
||||
tracing::warn!(
|
||||
extension = %flow.extension_name,
|
||||
"OAuth flow expired"
|
||||
);
|
||||
return oauth_error_page(&flow.display_name);
|
||||
}
|
||||
|
||||
// Exchange the authorization code for tokens.
|
||||
// Use the platform exchange proxy when configured (keeps client_secret off container),
|
||||
// otherwise call the provider's token URL directly.
|
||||
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 gateway_token = flow.gateway_token.as_deref().unwrap_or_default();
|
||||
oauth_defaults::exchange_via_proxy(
|
||||
proxy_url,
|
||||
gateway_token,
|
||||
&code,
|
||||
&flow.redirect_uri,
|
||||
flow.code_verifier.as_deref(),
|
||||
&flow.access_token_field,
|
||||
)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?
|
||||
} else {
|
||||
oauth_defaults::exchange_oauth_code(
|
||||
&flow.token_url,
|
||||
&flow.client_id,
|
||||
flow.client_secret.as_deref(),
|
||||
&code,
|
||||
&flow.redirect_uri,
|
||||
flow.code_verifier.as_deref(),
|
||||
&flow.access_token_field,
|
||||
)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?
|
||||
};
|
||||
|
||||
// Validate the token before storing (catches wrong account, etc.)
|
||||
if let Some(ref validation) = flow.validation_endpoint {
|
||||
oauth_defaults::validate_oauth_token(&token_response.access_token, validation)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
}
|
||||
|
||||
// Store tokens encrypted in the secrets store
|
||||
oauth_defaults::store_oauth_tokens(
|
||||
flow.secrets.as_ref(),
|
||||
&flow.user_id,
|
||||
&flow.secret_name,
|
||||
flow.provider.as_deref(),
|
||||
&token_response.access_token,
|
||||
token_response.refresh_token.as_deref(),
|
||||
token_response.expires_in,
|
||||
&flow.scopes,
|
||||
)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
.await;
|
||||
|
||||
let (success, message) = match &result {
|
||||
Ok(()) => (
|
||||
true,
|
||||
format!("{} authenticated successfully", flow.display_name),
|
||||
),
|
||||
Err(e) => (
|
||||
false,
|
||||
format!("{} authentication failed: {}", flow.display_name, e),
|
||||
),
|
||||
};
|
||||
|
||||
match &result {
|
||||
Ok(()) => {
|
||||
tracing::info!(
|
||||
extension = %flow.extension_name,
|
||||
"OAuth completed successfully via gateway callback"
|
||||
);
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
extension = %flow.extension_name,
|
||||
error = %e,
|
||||
"OAuth failed via gateway callback"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// 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,
|
||||
});
|
||||
}
|
||||
|
||||
let html = oauth_defaults::landing_html(&flow.display_name, success);
|
||||
axum::response::Html(html).into_response()
|
||||
}
|
||||
|
||||
// --- Chat handlers ---
|
||||
|
||||
async fn chat_send_handler(
|
||||
State(state): State<Arc<GatewayState>>,
|
||||
Json(req): Json<SendMessageRequest>,
|
||||
) -> Result<(StatusCode, Json<SendMessageResponse>), (StatusCode, String)> {
|
||||
tracing::debug!(
|
||||
"[chat_send_handler] Received message: content={:?}, thread_id={:?}",
|
||||
req.content,
|
||||
req.thread_id
|
||||
);
|
||||
|
||||
if !state.chat_rate_limiter.check() {
|
||||
return Err((
|
||||
StatusCode::TOO_MANY_REQUESTS,
|
||||
@@ -447,6 +627,11 @@ async fn chat_send_handler(
|
||||
}
|
||||
|
||||
let msg_id = msg.id;
|
||||
tracing::debug!(
|
||||
"[chat_send_handler] Created message id={}, content={:?}",
|
||||
msg_id,
|
||||
req.content
|
||||
);
|
||||
|
||||
let tx_guard = state.msg_tx.read().await;
|
||||
let tx = tx_guard.as_ref().ok_or((
|
||||
@@ -454,6 +639,7 @@ async fn chat_send_handler(
|
||||
"Channel not started".to_string(),
|
||||
))?;
|
||||
|
||||
tracing::debug!("[chat_send_handler] Sending message through channel");
|
||||
tx.send(msg).await.map_err(|_| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
@@ -461,6 +647,8 @@ async fn chat_send_handler(
|
||||
)
|
||||
})?;
|
||||
|
||||
tracing::debug!("[chat_send_handler] Message sent successfully, returning 202 ACCEPTED");
|
||||
|
||||
Ok((
|
||||
StatusCode::ACCEPTED,
|
||||
Json(SendMessageResponse {
|
||||
@@ -554,7 +742,7 @@ async fn chat_auth_token_handler(
|
||||
.await
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
|
||||
if result.status == "authenticated" {
|
||||
if result.is_authenticated() {
|
||||
// Auto-activate so tools are available immediately
|
||||
let msg = match ext_mgr.activate(&req.extension_name).await {
|
||||
Ok(r) => format!(
|
||||
@@ -582,13 +770,14 @@ async fn chat_auth_token_handler(
|
||||
// Re-emit auth_required for retry
|
||||
state.sse.broadcast(SseEvent::AuthRequired {
|
||||
extension_name: req.extension_name.clone(),
|
||||
instructions: result.instructions.clone(),
|
||||
auth_url: result.auth_url.clone(),
|
||||
setup_url: result.setup_url.clone(),
|
||||
instructions: result.instructions().map(String::from),
|
||||
auth_url: result.auth_url().map(String::from),
|
||||
setup_url: result.setup_url().map(String::from),
|
||||
});
|
||||
Ok(Json(ActionResponse::fail(
|
||||
result
|
||||
.instructions
|
||||
.instructions()
|
||||
.map(String::from)
|
||||
.unwrap_or_else(|| "Invalid token".to_string()),
|
||||
)))
|
||||
}
|
||||
@@ -1218,8 +1407,8 @@ async fn extensions_list_handler(
|
||||
} else if !ext.authenticated {
|
||||
// No credentials configured yet.
|
||||
"installed".to_string()
|
||||
} else if ext.active && ext.name == "telegram" {
|
||||
// Telegram: check pairing status (end-to-end setup via web UI).
|
||||
} else if ext.active {
|
||||
// Check pairing status for active channels.
|
||||
let has_paired = pairing_store
|
||||
.read_allow_from(&ext.name)
|
||||
.map(|list| !list.is_empty())
|
||||
@@ -1230,7 +1419,7 @@ async fn extensions_list_handler(
|
||||
"pairing".to_string()
|
||||
}
|
||||
} else {
|
||||
// Authenticated but not fully active (or non-Telegram).
|
||||
// Authenticated but not yet active.
|
||||
"configured".to_string()
|
||||
})
|
||||
} else {
|
||||
@@ -1246,6 +1435,7 @@ async fn extensions_list_handler(
|
||||
active: ext.active,
|
||||
tools: ext.tools,
|
||||
needs_setup: ext.needs_setup,
|
||||
has_auth: ext.has_auth,
|
||||
activation_status,
|
||||
activation_error: ext.activation_error,
|
||||
}
|
||||
@@ -1315,7 +1505,34 @@ async fn extensions_install_handler(
|
||||
.install(&req.name, req.url.as_deref(), kind_hint)
|
||||
.await
|
||||
{
|
||||
Ok(result) => Ok(Json(ActionResponse::ok(result.message))),
|
||||
Ok(result) => {
|
||||
let mut resp = ActionResponse::ok(result.message);
|
||||
|
||||
// Auto-activate WASM tools after install (install = active).
|
||||
if result.kind == crate::extensions::ExtensionKind::WasmTool {
|
||||
if let Err(e) = ext_mgr.activate(&req.name).await {
|
||||
tracing::debug!(
|
||||
extension = %req.name,
|
||||
error = %e,
|
||||
"Auto-activation after install failed"
|
||||
);
|
||||
}
|
||||
|
||||
// Check auth after activation. This may initiate OAuth both for scope
|
||||
// expansion and for first-time auth when credentials are already
|
||||
// configured (e.g., built-in providers). We only surface an auth_url
|
||||
// when the extension reports it is awaiting authorization.
|
||||
match ext_mgr.auth(&req.name, None).await {
|
||||
Ok(auth_result) if auth_result.auth_url().is_some() => {
|
||||
// Scope expansion or initial OAuth: user needs to authorize
|
||||
resp.auth_url = auth_result.auth_url().map(String::from);
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(Json(resp))
|
||||
}
|
||||
Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))),
|
||||
}
|
||||
}
|
||||
@@ -1330,7 +1547,19 @@ async fn extensions_activate_handler(
|
||||
))?;
|
||||
|
||||
match ext_mgr.activate(&name).await {
|
||||
Ok(result) => Ok(Json(ActionResponse::ok(result.message))),
|
||||
Ok(result) => {
|
||||
// Activation loaded the WASM module. Check if the tool needs
|
||||
// OAuth scope expansion (e.g., adding google-docs when gmail
|
||||
// already has a token but missing the documents scope).
|
||||
// Initial OAuth setup is triggered via save_setup_secrets.
|
||||
let mut resp = ActionResponse::ok(result.message);
|
||||
if let Ok(auth_result) = ext_mgr.auth(&name, None).await
|
||||
&& auth_result.auth_url().is_some()
|
||||
{
|
||||
resp.auth_url = auth_result.auth_url().map(String::from);
|
||||
}
|
||||
Ok(Json(resp))
|
||||
}
|
||||
Err(activate_err) => {
|
||||
let err_str = activate_err.to_string();
|
||||
let needs_auth = err_str.contains("authentication")
|
||||
@@ -1343,7 +1572,7 @@ async fn extensions_activate_handler(
|
||||
|
||||
// Activation failed due to auth; try authenticating first.
|
||||
match ext_mgr.auth(&name, None).await {
|
||||
Ok(auth_result) if auth_result.status == "authenticated" => {
|
||||
Ok(auth_result) if auth_result.is_authenticated() => {
|
||||
// Auth succeeded, retry activation.
|
||||
match ext_mgr.activate(&name).await {
|
||||
Ok(result) => Ok(Json(ActionResponse::ok(result.message))),
|
||||
@@ -1354,13 +1583,13 @@ async fn extensions_activate_handler(
|
||||
// Auth in progress (OAuth URL or awaiting manual token).
|
||||
let mut resp = ActionResponse::fail(
|
||||
auth_result
|
||||
.instructions
|
||||
.clone()
|
||||
.instructions()
|
||||
.map(String::from)
|
||||
.unwrap_or_else(|| format!("'{}' requires authentication.", name)),
|
||||
);
|
||||
resp.auth_url = auth_result.auth_url;
|
||||
resp.awaiting_token = Some(auth_result.awaiting_token);
|
||||
resp.instructions = auth_result.instructions;
|
||||
resp.auth_url = auth_result.auth_url().map(String::from);
|
||||
resp.awaiting_token = Some(auth_result.is_awaiting_token());
|
||||
resp.instructions = auth_result.instructions().map(String::from);
|
||||
Ok(Json(resp))
|
||||
}
|
||||
Err(auth_err) => Ok(Json(ActionResponse::fail(format!(
|
||||
@@ -1550,43 +1779,22 @@ async fn extensions_setup_submit_handler(
|
||||
|
||||
match ext_mgr.save_setup_secrets(&name, &req.secrets).await {
|
||||
Ok(result) => {
|
||||
// Broadcast auth_completed so the chat UI can dismiss any in-progress
|
||||
// auth card or setup modal that was triggered by tool_auth/tool_activate.
|
||||
state.sse.broadcast(SseEvent::AuthCompleted {
|
||||
extension_name: name.clone(),
|
||||
success: true,
|
||||
message: result.message.clone(),
|
||||
});
|
||||
let mut resp = ActionResponse::ok(result.message);
|
||||
resp.activated = Some(result.activated);
|
||||
if !result.activated {
|
||||
resp.needs_restart = Some(true);
|
||||
}
|
||||
resp.auth_url = result.auth_url;
|
||||
Ok(Json(resp))
|
||||
}
|
||||
Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))),
|
||||
}
|
||||
}
|
||||
|
||||
// --- Gateway management handlers ---
|
||||
|
||||
async fn gateway_restart_handler(State(state): State<Arc<GatewayState>>) -> Json<ActionResponse> {
|
||||
// Idempotency guard: only allow one restart at a time.
|
||||
if state
|
||||
.restart_requested
|
||||
.compare_exchange(
|
||||
false,
|
||||
true,
|
||||
std::sync::atomic::Ordering::SeqCst,
|
||||
std::sync::atomic::Ordering::SeqCst,
|
||||
)
|
||||
.is_err()
|
||||
{
|
||||
return Json(ActionResponse::ok("Restart already in progress"));
|
||||
}
|
||||
|
||||
// Take the shutdown sender and trigger graceful shutdown.
|
||||
if let Some(tx) = state.shutdown_tx.write().await.take() {
|
||||
let _ = tx.send(());
|
||||
tracing::info!("Gateway restart requested via API");
|
||||
}
|
||||
|
||||
Json(ActionResponse::ok("Restarting..."))
|
||||
}
|
||||
|
||||
// --- Pairing handlers ---
|
||||
|
||||
async fn pairing_list_handler(
|
||||
@@ -2106,11 +2314,16 @@ async fn gateway_status_handler(
|
||||
(None, None, None)
|
||||
};
|
||||
|
||||
let restart_enabled = std::env::var("IRONCLAW_IN_DOCKER")
|
||||
.map(|v| v.to_lowercase() == "true")
|
||||
.unwrap_or(false);
|
||||
|
||||
Json(GatewayStatusResponse {
|
||||
sse_connections,
|
||||
ws_connections,
|
||||
total_connections: sse_connections + ws_connections,
|
||||
uptime_secs,
|
||||
restart_enabled,
|
||||
daily_cost,
|
||||
actions_this_hour,
|
||||
model_usage,
|
||||
@@ -2131,6 +2344,7 @@ struct GatewayStatusResponse {
|
||||
ws_connections: u64,
|
||||
total_connections: u64,
|
||||
uptime_secs: u64,
|
||||
restart_enabled: bool,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
daily_cost: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
@@ -2218,4 +2432,343 @@ mod tests {
|
||||
let turns = build_turns_from_db_messages(&[]);
|
||||
assert!(turns.is_empty());
|
||||
}
|
||||
|
||||
// --- OAuth callback handler tests ---
|
||||
|
||||
/// Build a minimal `GatewayState` for testing the OAuth callback handler.
|
||||
fn test_gateway_state(ext_mgr: Option<Arc<ExtensionManager>>) -> Arc<GatewayState> {
|
||||
Arc::new(GatewayState {
|
||||
msg_tx: tokio::sync::RwLock::new(None),
|
||||
sse: SseManager::new(),
|
||||
workspace: None,
|
||||
session_manager: None,
|
||||
log_broadcaster: None,
|
||||
log_level_handle: None,
|
||||
extension_manager: ext_mgr,
|
||||
tool_registry: None,
|
||||
store: None,
|
||||
job_manager: None,
|
||||
prompt_queue: None,
|
||||
user_id: "test".to_string(),
|
||||
shutdown_tx: tokio::sync::RwLock::new(None),
|
||||
ws_tracker: None,
|
||||
llm_provider: None,
|
||||
skill_registry: None,
|
||||
skill_catalog: None,
|
||||
scheduler: None,
|
||||
chat_rate_limiter: RateLimiter::new(30, 60),
|
||||
registry_entries: vec![],
|
||||
cost_guard: None,
|
||||
startup_time: std::time::Instant::now(),
|
||||
})
|
||||
}
|
||||
|
||||
/// Build a test router with just the OAuth callback route.
|
||||
fn test_oauth_router(state: Arc<GatewayState>) -> Router {
|
||||
Router::new()
|
||||
.route("/oauth/callback", get(oauth_callback_handler))
|
||||
.with_state(state)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_oauth_callback_missing_params() {
|
||||
use axum::body::Body;
|
||||
use tower::ServiceExt;
|
||||
|
||||
let state = test_gateway_state(None);
|
||||
let app = test_oauth_router(state);
|
||||
|
||||
let req = axum::http::Request::builder()
|
||||
.uri("/oauth/callback")
|
||||
.body(Body::empty())
|
||||
.expect("request");
|
||||
|
||||
let resp = ServiceExt::<axum::http::Request<Body>>::oneshot(app, req)
|
||||
.await
|
||||
.expect("response");
|
||||
assert_eq!(resp.status(), StatusCode::OK);
|
||||
|
||||
let body = axum::body::to_bytes(resp.into_body(), 1024 * 64)
|
||||
.await
|
||||
.expect("body");
|
||||
let html = String::from_utf8_lossy(&body);
|
||||
assert!(html.contains("Authorization Failed"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_oauth_callback_error_from_provider() {
|
||||
use axum::body::Body;
|
||||
use tower::ServiceExt;
|
||||
|
||||
let state = test_gateway_state(None);
|
||||
let app = test_oauth_router(state);
|
||||
|
||||
let req = axum::http::Request::builder()
|
||||
.uri("/oauth/callback?error=access_denied&error_description=access_denied")
|
||||
.body(Body::empty())
|
||||
.expect("request");
|
||||
|
||||
let resp = ServiceExt::<axum::http::Request<Body>>::oneshot(app, req)
|
||||
.await
|
||||
.expect("response");
|
||||
assert_eq!(resp.status(), StatusCode::OK);
|
||||
|
||||
let body = axum::body::to_bytes(resp.into_body(), 1024 * 64)
|
||||
.await
|
||||
.expect("body");
|
||||
let html = String::from_utf8_lossy(&body);
|
||||
assert!(html.contains("Authorization Failed"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_oauth_callback_unknown_state() {
|
||||
use axum::body::Body;
|
||||
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-key-at-least-32-chars-long!!".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,
|
||||
secrets,
|
||||
tool_registry,
|
||||
None,
|
||||
None,
|
||||
std::path::PathBuf::from("/tmp/wasm_tools"),
|
||||
std::path::PathBuf::from("/tmp/wasm_channels"),
|
||||
None,
|
||||
"test".to_string(),
|
||||
None,
|
||||
vec![],
|
||||
));
|
||||
|
||||
let state = test_gateway_state(Some(ext_mgr));
|
||||
let app = test_oauth_router(state);
|
||||
|
||||
let req = axum::http::Request::builder()
|
||||
.uri("/oauth/callback?code=test_code&state=unknown_state_value")
|
||||
.body(Body::empty())
|
||||
.expect("request");
|
||||
|
||||
let resp = ServiceExt::<axum::http::Request<Body>>::oneshot(app, req)
|
||||
.await
|
||||
.expect("response");
|
||||
assert_eq!(resp.status(), StatusCode::OK);
|
||||
|
||||
let body = axum::body::to_bytes(resp.into_body(), 1024 * 64)
|
||||
.await
|
||||
.expect("body");
|
||||
let html = String::from_utf8_lossy(&body);
|
||||
assert!(html.contains("Authorization Failed"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_oauth_callback_expired_flow() {
|
||||
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-key-at-least-32-chars-long!!".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,
|
||||
secrets.clone(),
|
||||
tool_registry,
|
||||
None,
|
||||
None,
|
||||
std::path::PathBuf::from("/tmp/wasm_tools"),
|
||||
std::path::PathBuf::from("/tmp/wasm_channels"),
|
||||
None,
|
||||
"test".to_string(),
|
||||
None,
|
||||
vec![],
|
||||
));
|
||||
|
||||
// Insert an expired flow (created 10 minutes ago)
|
||||
let flow = crate::cli::oauth_defaults::PendingOAuthFlow {
|
||||
extension_name: "test_tool".to_string(),
|
||||
display_name: "Test Tool".to_string(),
|
||||
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: None,
|
||||
gateway_token: None,
|
||||
created_at: std::time::Instant::now() - std::time::Duration::from_secs(600),
|
||||
};
|
||||
|
||||
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);
|
||||
|
||||
let body = axum::body::to_bytes(resp.into_body(), 1024 * 64)
|
||||
.await
|
||||
.expect("body");
|
||||
let html = String::from_utf8_lossy(&body);
|
||||
// Expired flow → error landing page
|
||||
assert!(html.contains("Authorization Failed"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_oauth_callback_no_extension_manager() {
|
||||
use axum::body::Body;
|
||||
use tower::ServiceExt;
|
||||
|
||||
// No extension manager set → graceful error
|
||||
let state = test_gateway_state(None);
|
||||
let app = test_oauth_router(state);
|
||||
|
||||
let req = axum::http::Request::builder()
|
||||
.uri("/oauth/callback?code=test_code&state=some_state")
|
||||
.body(Body::empty())
|
||||
.expect("request");
|
||||
|
||||
let resp = ServiceExt::<axum::http::Request<Body>>::oneshot(app, req)
|
||||
.await
|
||||
.expect("response");
|
||||
assert_eq!(resp.status(), StatusCode::OK);
|
||||
|
||||
let body = axum::body::to_bytes(resp.into_body(), 1024 * 64)
|
||||
.await
|
||||
.expect("body");
|
||||
let html = String::from_utf8_lossy(&body);
|
||||
assert!(html.contains("Authorization Failed"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_oauth_callback_strips_instance_prefix() {
|
||||
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-key-at-least-32-chars-long!!".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,
|
||||
secrets.clone(),
|
||||
tool_registry,
|
||||
None,
|
||||
None,
|
||||
std::path::PathBuf::from("/tmp/wasm_tools"),
|
||||
std::path::PathBuf::from("/tmp/wasm_channels"),
|
||||
None,
|
||||
"test".to_string(),
|
||||
None,
|
||||
vec![],
|
||||
));
|
||||
|
||||
// Insert a flow keyed by raw nonce "test_nonce" (without instance prefix).
|
||||
// Use an expired flow so the handler exits before attempting a real HTTP
|
||||
// token exchange — we only need to verify that the instance prefix was
|
||||
// stripped and the flow was found by the raw nonce.
|
||||
let 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: None,
|
||||
gateway_token: None,
|
||||
// Expired — handler will reject after lookup (no network I/O)
|
||||
created_at: std::time::Instant::now() - std::time::Duration::from_secs(600),
|
||||
};
|
||||
|
||||
ext_mgr
|
||||
.pending_oauth_flows()
|
||||
.write()
|
||||
.await
|
||||
.insert("test_nonce".to_string(), flow);
|
||||
|
||||
let state = test_gateway_state(Some(ext_mgr.clone()));
|
||||
let app = test_oauth_router(state);
|
||||
|
||||
// Send callback with instance prefix: "myinstance:test_nonce"
|
||||
// The handler should strip "myinstance:" and find the flow keyed by "test_nonce"
|
||||
let req = axum::http::Request::builder()
|
||||
.uri("/oauth/callback?code=fake_code&state=myinstance:test_nonce")
|
||||
.body(Body::empty())
|
||||
.expect("request");
|
||||
|
||||
let resp = ServiceExt::<axum::http::Request<Body>>::oneshot(app, req)
|
||||
.await
|
||||
.expect("response");
|
||||
assert_eq!(resp.status(), StatusCode::OK);
|
||||
|
||||
let body = axum::body::to_bytes(resp.into_body(), 1024 * 64)
|
||||
.await
|
||||
.expect("body");
|
||||
let html = String::from_utf8_lossy(&body);
|
||||
|
||||
// The flow was found (stripped prefix matched) but is expired, so the
|
||||
// handler returns an error landing page. The flow being consumed from
|
||||
// the registry (checked below) proves the prefix was stripped correctly.
|
||||
assert!(
|
||||
html.contains("Authorization Failed"),
|
||||
"Expected error page, html was: {}",
|
||||
&html[..html.len().min(500)]
|
||||
);
|
||||
|
||||
// Verify the flow was consumed (removed from registry)
|
||||
assert!(
|
||||
ext_mgr
|
||||
.pending_oauth_flows()
|
||||
.read()
|
||||
.await
|
||||
.get("test_nonce")
|
||||
.is_none()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -36,6 +36,23 @@ impl SseManager {
|
||||
}
|
||||
}
|
||||
|
||||
/// Create an SSE manager that reuses an existing broadcast sender.
|
||||
///
|
||||
/// This preserves the broadcast channel across `rebuild_state` calls so
|
||||
/// that sender handles captured by other components remain valid.
|
||||
///
|
||||
/// **Important:** The connection counter is reset to zero. This method must
|
||||
/// only be called before the server starts accepting connections (i.e.,
|
||||
/// during startup wiring). Calling it after connections are established
|
||||
/// will break connection tracking and allow exceeding `MAX_CONNECTIONS`.
|
||||
pub fn from_sender(tx: broadcast::Sender<SseEvent>) -> Self {
|
||||
Self {
|
||||
tx,
|
||||
connection_count: Arc::new(AtomicU64::new(0)),
|
||||
max_connections: MAX_CONNECTIONS,
|
||||
}
|
||||
}
|
||||
|
||||
/// Broadcast an event to all connected clients.
|
||||
pub fn broadcast(&self, event: SseEvent) {
|
||||
// Ignore send errors (no receivers is fine)
|
||||
|
||||
+255
-108
@@ -133,6 +133,110 @@ function apiFetch(path, options) {
|
||||
});
|
||||
}
|
||||
|
||||
// --- Restart Feature ---
|
||||
|
||||
let isRestarting = false; // Track if we're currently restarting
|
||||
let restartEnabled = false; // Track if restart is available in this deployment
|
||||
|
||||
function triggerRestart() {
|
||||
if (!currentThreadId) {
|
||||
alert('Please start a conversation first');
|
||||
return;
|
||||
}
|
||||
|
||||
// Show the confirmation modal
|
||||
const confirmModal = document.getElementById('restart-confirm-modal');
|
||||
confirmModal.style.display = 'flex';
|
||||
}
|
||||
|
||||
function confirmRestart() {
|
||||
if (!currentThreadId) {
|
||||
alert('Please start a conversation first');
|
||||
return;
|
||||
}
|
||||
|
||||
// Hide confirmation modal
|
||||
const confirmModal = document.getElementById('restart-confirm-modal');
|
||||
confirmModal.style.display = 'none';
|
||||
|
||||
const restartBtn = document.getElementById('restart-btn');
|
||||
const restartIcon = document.getElementById('restart-icon');
|
||||
|
||||
// Mark as restarting
|
||||
isRestarting = true;
|
||||
restartBtn.disabled = true;
|
||||
if (restartIcon) restartIcon.classList.add('spinning');
|
||||
|
||||
// Show progress modal
|
||||
const loaderEl = document.getElementById('restart-loader');
|
||||
loaderEl.style.display = 'flex';
|
||||
|
||||
// Send restart command via chat
|
||||
console.log('[confirmRestart] Sending /restart command to server');
|
||||
apiFetch('/api/chat/send', {
|
||||
method: 'POST',
|
||||
body: {
|
||||
content: '/restart',
|
||||
thread_id: currentThreadId,
|
||||
},
|
||||
})
|
||||
.then((response) => {
|
||||
console.log('[confirmRestart] API call succeeded, response:', response);
|
||||
})
|
||||
.catch((err) => {
|
||||
console.error('[confirmRestart] Restart request failed:', err);
|
||||
addMessage('system', 'Restart failed: ' + err.message);
|
||||
isRestarting = false;
|
||||
restartBtn.disabled = false;
|
||||
if (restartIcon) restartIcon.classList.remove('spinning');
|
||||
loaderEl.style.display = 'none';
|
||||
});
|
||||
}
|
||||
|
||||
function cancelRestart() {
|
||||
const confirmModal = document.getElementById('restart-confirm-modal');
|
||||
confirmModal.style.display = 'none';
|
||||
}
|
||||
|
||||
function tryShowRestartModal() {
|
||||
// Defensive callback for when restart is detected in messages.
|
||||
if (!isRestarting) {
|
||||
isRestarting = true;
|
||||
const restartBtn = document.getElementById('restart-btn');
|
||||
const restartIcon = document.getElementById('restart-icon');
|
||||
restartBtn.disabled = true;
|
||||
if (restartIcon) restartIcon.classList.add('spinning');
|
||||
|
||||
// Show progress modal
|
||||
const loaderEl = document.getElementById('restart-loader');
|
||||
loaderEl.style.display = 'flex';
|
||||
}
|
||||
}
|
||||
|
||||
function updateRestartButtonVisibility() {
|
||||
const restartBtn = document.getElementById('restart-btn');
|
||||
if (restartBtn) {
|
||||
restartBtn.style.display = restartEnabled ? 'block' : 'none';
|
||||
}
|
||||
}
|
||||
|
||||
function startGatewayStatusPolling() {
|
||||
fetchGatewayStatus();
|
||||
// Poll every 5 seconds
|
||||
setInterval(fetchGatewayStatus, 5000);
|
||||
}
|
||||
|
||||
function fetchGatewayStatus() {
|
||||
apiFetch('/api/gateway/status')
|
||||
.then((data) => {
|
||||
restartEnabled = data.restart_enabled || false;
|
||||
updateRestartButtonVisibility();
|
||||
})
|
||||
.catch((err) => {
|
||||
console.warn('[gateway status] Failed to fetch:', err);
|
||||
});
|
||||
}
|
||||
|
||||
// --- SSE ---
|
||||
|
||||
function connectSSE() {
|
||||
@@ -143,6 +247,18 @@ function connectSSE() {
|
||||
eventSource.onopen = () => {
|
||||
document.getElementById('sse-dot').classList.remove('disconnected');
|
||||
document.getElementById('sse-status').textContent = 'Connected';
|
||||
|
||||
// If we were restarting, close the modal and reset button now that server is back
|
||||
if (isRestarting) {
|
||||
const loaderEl = document.getElementById('restart-loader');
|
||||
if (loaderEl) loaderEl.style.display = 'none';
|
||||
const restartBtn = document.getElementById('restart-btn');
|
||||
const restartIcon = document.getElementById('restart-icon');
|
||||
if (restartBtn) restartBtn.disabled = false;
|
||||
if (restartIcon) restartIcon.classList.remove('spinning');
|
||||
isRestarting = false;
|
||||
}
|
||||
|
||||
if (sseHasConnectedBefore && currentThreadId) {
|
||||
finalizeActivityGroup();
|
||||
loadHistory();
|
||||
@@ -163,6 +279,11 @@ function connectSSE() {
|
||||
enableChatInput();
|
||||
// Refresh thread list so new titles appear after first message
|
||||
loadThreads();
|
||||
|
||||
// Show restart modal if the response indicates restart was initiated
|
||||
if (data.content && data.content.toLowerCase().includes('restart initiated')) {
|
||||
setTimeout(() => tryShowRestartModal(), 500);
|
||||
}
|
||||
});
|
||||
|
||||
eventSource.addEventListener('thinking', (e) => {
|
||||
@@ -180,7 +301,12 @@ function connectSSE() {
|
||||
eventSource.addEventListener('tool_completed', (e) => {
|
||||
const data = JSON.parse(e.data);
|
||||
if (!isCurrentThread(data.thread_id)) return;
|
||||
completeToolCard(data.name, data.success);
|
||||
completeToolCard(data.name, data.success, data.error, data.parameters);
|
||||
|
||||
// Show restart modal only when the restart tool succeeds
|
||||
if (data.name.toLowerCase() === 'restart' && data.success) {
|
||||
setTimeout(() => tryShowRestartModal(), 500);
|
||||
}
|
||||
});
|
||||
|
||||
eventSource.addEventListener('tool_result', (e) => {
|
||||
@@ -222,13 +348,24 @@ function connectSSE() {
|
||||
|
||||
eventSource.addEventListener('auth_required', (e) => {
|
||||
const data = JSON.parse(e.data);
|
||||
showAuthCard(data);
|
||||
if (data.auth_url) {
|
||||
// OAuth flow: show the auth card with an OAuth button + optional token paste field.
|
||||
showAuthCard(data);
|
||||
} else {
|
||||
// Setup flow: fetch the extension's credential schema and show the multi-field
|
||||
// configure modal (the same UI used by the Extensions tab "Setup" button).
|
||||
showConfigureModal(data.extension_name);
|
||||
}
|
||||
});
|
||||
|
||||
eventSource.addEventListener('auth_completed', (e) => {
|
||||
const data = JSON.parse(e.data);
|
||||
// Dismiss whichever UI path was active: auth card (OAuth) or configure modal (setup).
|
||||
removeAuthCard(data.extension_name);
|
||||
showToast(data.message, 'success');
|
||||
closeConfigureModal();
|
||||
showToast(data.message, data.success ? 'success' : 'error');
|
||||
// Refresh extensions list so status indicators update
|
||||
if (currentTab === 'extensions') loadExtensions();
|
||||
enableChatInput();
|
||||
});
|
||||
|
||||
@@ -359,6 +496,9 @@ function selectSlashItem(cmd) {
|
||||
function updateSlashHighlight() {
|
||||
const items = document.querySelectorAll('#slash-autocomplete .slash-ac-item');
|
||||
items.forEach((el, i) => el.classList.toggle('selected', i === _slashSelected));
|
||||
if (_slashSelected >= 0 && items[_slashSelected]) {
|
||||
items[_slashSelected].scrollIntoView({ block: 'nearest' });
|
||||
}
|
||||
}
|
||||
|
||||
function filterSlashCommands(value) {
|
||||
@@ -581,7 +721,7 @@ function addToolCard(name) {
|
||||
container.scrollTop = container.scrollHeight;
|
||||
}
|
||||
|
||||
function completeToolCard(name, success) {
|
||||
function completeToolCard(name, success, error, parameters) {
|
||||
const entries = _activeToolCards[name];
|
||||
if (!entries || entries.length === 0) return;
|
||||
// Find first running card
|
||||
@@ -602,6 +742,27 @@ function completeToolCard(name, success) {
|
||||
? '<span class="activity-icon-success">✓</span>'
|
||||
: '<span class="activity-icon-fail">✗</span>';
|
||||
entry.card.setAttribute('data-status', success ? 'success' : 'fail');
|
||||
|
||||
// For failed tools, populate the body with error details and auto-expand
|
||||
if (!success && (error || parameters)) {
|
||||
const output = entry.card.querySelector('.activity-tool-output');
|
||||
if (output) {
|
||||
let detail = '';
|
||||
if (parameters) {
|
||||
detail += 'Input:\n' + parameters + '\n\n';
|
||||
}
|
||||
if (error) {
|
||||
detail += 'Error:\n' + error;
|
||||
}
|
||||
output.textContent = detail;
|
||||
|
||||
// Auto-expand so the error is immediately visible
|
||||
const body = entry.card.querySelector('.activity-tool-body');
|
||||
const chevron = entry.card.querySelector('.activity-tool-chevron');
|
||||
if (body) body.style.display = 'block';
|
||||
if (chevron) chevron.classList.add('expanded');
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
function setToolCardOutput(name, preview) {
|
||||
@@ -842,7 +1003,7 @@ function showAuthCard(data) {
|
||||
oauthBtn.className = 'auth-oauth';
|
||||
oauthBtn.textContent = 'Authenticate with ' + data.extension_name;
|
||||
oauthBtn.addEventListener('click', () => {
|
||||
window.open(data.auth_url, '_blank', 'width=600,height=700');
|
||||
openOAuthUrl(data.auth_url);
|
||||
});
|
||||
links.appendChild(oauthBtn);
|
||||
}
|
||||
@@ -865,7 +1026,7 @@ function showAuthCard(data) {
|
||||
|
||||
const tokenInput = document.createElement('input');
|
||||
tokenInput.type = 'password';
|
||||
tokenInput.placeholder = 'Paste your API key or token';
|
||||
tokenInput.placeholder = data.instructions || 'Paste your API key or token';
|
||||
tokenInput.addEventListener('keydown', (e) => {
|
||||
if (e.key === 'Enter') submitAuthToken(data.extension_name, tokenInput.value);
|
||||
});
|
||||
@@ -1196,7 +1357,7 @@ chatInput.addEventListener('keydown', (e) => {
|
||||
updateSlashHighlight();
|
||||
return;
|
||||
}
|
||||
if (e.key === 'Tab' || (e.key === 'Enter' && _slashSelected >= 0)) {
|
||||
if (e.key === 'Tab' || e.key === 'Enter') {
|
||||
e.preventDefault();
|
||||
const pick = _slashSelected >= 0 ? _slashMatches[_slashSelected] : _slashMatches[0];
|
||||
if (pick) selectSlashItem(pick.cmd);
|
||||
@@ -1757,6 +1918,11 @@ function renderAvailableExtensionCard(entry) {
|
||||
}).then(function(res) {
|
||||
if (res.success) {
|
||||
showToast('Installed ' + entry.display_name, 'success');
|
||||
// OAuth popup if auth started during install (builtin creds)
|
||||
if (res.auth_url) {
|
||||
showToast('Opening authentication for ' + entry.display_name, 'info');
|
||||
openOAuthUrl(res.auth_url);
|
||||
}
|
||||
loadExtensions();
|
||||
// Auto-open configure for WASM channels
|
||||
if (entry.kind === 'wasm_channel') {
|
||||
@@ -1913,7 +2079,7 @@ function renderExtensionCard(ext) {
|
||||
card.appendChild(url);
|
||||
}
|
||||
|
||||
if (ext.tools.length > 0) {
|
||||
if (ext.tools && ext.tools.length > 0) {
|
||||
const tools = document.createElement('div');
|
||||
tools.className = 'ext-tools';
|
||||
tools.textContent = 'Tools: ' + ext.tools.join(', ');
|
||||
@@ -1928,14 +2094,6 @@ function renderExtensionCard(ext) {
|
||||
card.appendChild(errorDiv);
|
||||
}
|
||||
|
||||
// Show "coming soon" note for non-Telegram channels that are configured but not fully supported yet
|
||||
if (ext.kind === 'wasm_channel' && ext.name !== 'telegram'
|
||||
&& (ext.activation_status === 'configured' || ext.active)) {
|
||||
const noteDiv = document.createElement('div');
|
||||
noteDiv.className = 'ext-note';
|
||||
noteDiv.textContent = 'Full integration coming soon. Use the CLI to complete setup.';
|
||||
card.appendChild(noteDiv);
|
||||
}
|
||||
|
||||
const actions = document.createElement('div');
|
||||
actions.className = 'ext-actions';
|
||||
@@ -1966,24 +2124,29 @@ function renderExtensionCard(ext) {
|
||||
actions.appendChild(setupBtn);
|
||||
}
|
||||
} else {
|
||||
// Non-WASM-channel extensions: original behavior
|
||||
if (!ext.active) {
|
||||
// WASM tools / MCP servers
|
||||
const activeLabel = document.createElement('span');
|
||||
activeLabel.className = 'ext-active-label';
|
||||
activeLabel.textContent = ext.active ? 'Active' : 'Installed';
|
||||
actions.appendChild(activeLabel);
|
||||
|
||||
// MCP servers may be installed but inactive — show Activate button
|
||||
if (ext.kind === 'mcp_server' && !ext.active) {
|
||||
const activateBtn = document.createElement('button');
|
||||
activateBtn.className = 'btn-ext activate';
|
||||
activateBtn.textContent = 'Activate';
|
||||
activateBtn.addEventListener('click', () => activateExtension(ext.name));
|
||||
actions.appendChild(activateBtn);
|
||||
} else {
|
||||
const activeLabel = document.createElement('span');
|
||||
activeLabel.className = 'ext-active-label';
|
||||
activeLabel.textContent = 'Active';
|
||||
actions.appendChild(activeLabel);
|
||||
}
|
||||
|
||||
if (ext.needs_setup) {
|
||||
// Show Configure/Reconfigure button when there are secrets to enter.
|
||||
// Skip when has_auth is true but needs_setup is false and not yet authenticated —
|
||||
// this means OAuth credentials resolve automatically (builtin/env) and the user
|
||||
// just needs to complete the OAuth flow, not fill in a config form.
|
||||
if (ext.needs_setup || (ext.has_auth && ext.authenticated)) {
|
||||
const configBtn = document.createElement('button');
|
||||
configBtn.className = 'btn-ext configure';
|
||||
configBtn.textContent = ext.authenticated ? 'Reconfigure' : 'Setup';
|
||||
configBtn.textContent = ext.authenticated ? 'Reconfigure' : 'Configure';
|
||||
configBtn.addEventListener('click', () => showConfigureModal(ext.name));
|
||||
actions.appendChild(configBtn);
|
||||
}
|
||||
@@ -2013,13 +2176,18 @@ function activateExtension(name) {
|
||||
apiFetch('/api/extensions/' + encodeURIComponent(name) + '/activate', { method: 'POST' })
|
||||
.then((res) => {
|
||||
if (res.success) {
|
||||
// Even on success, the tool may need OAuth (e.g., WASM loaded but no token yet)
|
||||
if (res.auth_url) {
|
||||
showToast('Opening authentication for ' + name, 'info');
|
||||
openOAuthUrl(res.auth_url);
|
||||
}
|
||||
loadExtensions();
|
||||
return;
|
||||
}
|
||||
|
||||
if (res.auth_url) {
|
||||
showToast('Opening authentication for ' + name, 'info');
|
||||
window.open(res.auth_url, '_blank');
|
||||
openOAuthUrl(res.auth_url);
|
||||
} else if (res.awaiting_token) {
|
||||
showConfigureModal(name);
|
||||
} else {
|
||||
@@ -2161,19 +2329,22 @@ function submitConfigureModal(name, fields) {
|
||||
body: { secrets },
|
||||
})
|
||||
.then((res) => {
|
||||
closeConfigureModal();
|
||||
if (res.success) {
|
||||
if (res.activated) {
|
||||
showToast('Configured and activated ' + name, 'success');
|
||||
} else if (res.needs_restart) {
|
||||
showToast('Configured ' + name + '. Use Reconfigure to re-enter credentials and activate.', 'info');
|
||||
} else {
|
||||
showToast(res.message, '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.
|
||||
showToast('Opening OAuth authorization for ' + name, 'info');
|
||||
openOAuthUrl(res.auth_url);
|
||||
loadExtensions();
|
||||
}
|
||||
// For non-OAuth success: the server always broadcasts auth_completed SSE,
|
||||
// which will show the toast and refresh extensions — no need to do it here too.
|
||||
} else {
|
||||
// Keep modal open so the user can correct their input and retry.
|
||||
btns.forEach(function(b) { b.disabled = false; });
|
||||
showToast(res.message || 'Configuration failed', 'error');
|
||||
}
|
||||
loadExtensions();
|
||||
})
|
||||
.catch((err) => {
|
||||
btns.forEach(function(b) { b.disabled = false; });
|
||||
@@ -2186,6 +2357,25 @@ function closeConfigureModal() {
|
||||
if (existing) existing.remove();
|
||||
}
|
||||
|
||||
// Validate that a server-supplied OAuth URL is HTTPS before opening a popup.
|
||||
// Rejects javascript:, data:, and other non-HTTPS schemes to prevent URL-injection.
|
||||
// Uses the URL constructor to safely parse and validate the scheme, which also
|
||||
// handles non-string values (objects, null, etc.) that would throw on .startsWith().
|
||||
function openOAuthUrl(url) {
|
||||
let parsed;
|
||||
try {
|
||||
parsed = new URL(url);
|
||||
if (parsed.protocol !== 'https:') {
|
||||
throw new Error('non-HTTPS protocol: ' + parsed.protocol);
|
||||
}
|
||||
} catch (e) {
|
||||
console.warn('Blocked invalid/non-HTTPS OAuth URL:', url, e.message);
|
||||
showToast('Invalid OAuth URL returned by server', 'error');
|
||||
return;
|
||||
}
|
||||
window.open(parsed.href, '_blank', 'width=600,height=700');
|
||||
}
|
||||
|
||||
// --- Pairing ---
|
||||
|
||||
function loadPairingRequests(channel, container) {
|
||||
@@ -2232,7 +2422,7 @@ function approvePairing(channel, code, container) {
|
||||
}).then(res => {
|
||||
if (res.success) {
|
||||
showToast('Pairing approved', 'success');
|
||||
loadPairingRequests(channel, container);
|
||||
loadExtensions();
|
||||
} else {
|
||||
showToast(res.message || 'Approve failed', 'error');
|
||||
}
|
||||
@@ -2255,53 +2445,6 @@ function stopPairingPoll() {
|
||||
}
|
||||
}
|
||||
|
||||
// --- Gateway restart ---
|
||||
|
||||
function restartGateway() {
|
||||
if (!confirm('Restart IronClaw gateway? Active connections will be dropped.')) return;
|
||||
|
||||
apiFetch('/api/gateway/restart', { method: 'POST' })
|
||||
.then(function() {
|
||||
showRestartOverlay();
|
||||
})
|
||||
.catch(function() {
|
||||
showRestartOverlay();
|
||||
});
|
||||
}
|
||||
|
||||
function showRestartOverlay() {
|
||||
var overlay = document.createElement('div');
|
||||
overlay.className = 'restart-overlay';
|
||||
overlay.innerHTML = '<div class="restart-message">'
|
||||
+ '<div class="restart-spinner"></div>'
|
||||
+ '<h2>Restarting IronClaw...</h2>'
|
||||
+ '<p>Waiting for server to come back online</p>'
|
||||
+ '</div>';
|
||||
document.body.appendChild(overlay);
|
||||
|
||||
var pollCount = 0;
|
||||
var pollTimer = setInterval(function() {
|
||||
pollCount++;
|
||||
if (pollCount > 30) { // 60 seconds
|
||||
clearInterval(pollTimer);
|
||||
overlay.querySelector('h2').textContent = 'Restart timed out';
|
||||
overlay.querySelector('p').textContent = 'Server did not come back within 60 seconds. Check logs.';
|
||||
overlay.querySelector('.restart-spinner').style.display = 'none';
|
||||
return;
|
||||
}
|
||||
fetch('/api/gateway/status', {
|
||||
headers: { 'Authorization': 'Bearer ' + token },
|
||||
})
|
||||
.then(function(r) {
|
||||
if (r.ok) {
|
||||
clearInterval(pollTimer);
|
||||
window.location.reload();
|
||||
}
|
||||
})
|
||||
.catch(function() { /* still restarting */ });
|
||||
}, 2000);
|
||||
}
|
||||
|
||||
// --- WASM channel stepper ---
|
||||
|
||||
function renderWasmChannelStepper(ext) {
|
||||
@@ -2309,23 +2452,17 @@ function renderWasmChannelStepper(ext) {
|
||||
stepper.className = 'ext-stepper';
|
||||
|
||||
var status = ext.activation_status || 'installed';
|
||||
var isTelegram = ext.name === 'telegram';
|
||||
|
||||
// Telegram gets a 3-step stepper (Installed → Configured → Active/Pairing).
|
||||
// Other channels only get 2 steps (Installed → Configured) since full
|
||||
// integration isn't available in the web UI yet.
|
||||
var steps = [
|
||||
{ label: 'Installed', key: 'installed' },
|
||||
{ label: 'Configured', key: 'configured' },
|
||||
{ label: status === 'pairing' ? 'Awaiting Pairing' : 'Active', key: 'active' },
|
||||
];
|
||||
if (isTelegram) {
|
||||
steps.push({ label: status === 'pairing' ? 'Awaiting Pairing' : 'Active', key: 'active' });
|
||||
}
|
||||
|
||||
var reachedIdx;
|
||||
if (status === 'active') reachedIdx = isTelegram ? 2 : 1;
|
||||
if (status === 'active') reachedIdx = 2;
|
||||
else if (status === 'pairing') reachedIdx = 2;
|
||||
else if (status === 'failed') reachedIdx = isTelegram ? 2 : 1;
|
||||
else if (status === 'failed') reachedIdx = 2;
|
||||
else if (status === 'configured') reachedIdx = 1;
|
||||
else reachedIdx = 0;
|
||||
|
||||
@@ -2436,9 +2573,8 @@ 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>';
|
||||
} else if (job.state === 'failed' || job.state === 'interrupted') {
|
||||
actionBtns = '<button class="btn-restart" onclick="event.stopPropagation(); restartJob(\'' + job.id + '\')">Restart</button>';
|
||||
}
|
||||
// Retry is only shown in the detail view where can_restart is available.
|
||||
|
||||
return '<tr class="job-row" onclick="openJobDetail(\'' + job.id + '\')">'
|
||||
+ '<td title="' + escapeHtml(job.id) + '">' + shortId + '</td>'
|
||||
@@ -2467,10 +2603,12 @@ function restartJob(jobId) {
|
||||
apiFetch('/api/jobs/' + jobId + '/restart', { method: 'POST' })
|
||||
.then((res) => {
|
||||
showToast('Job restarted as ' + (res.new_job_id || '').substring(0, 8), 'success');
|
||||
loadJobs();
|
||||
})
|
||||
.catch((err) => {
|
||||
showToast('Failed to restart job: ' + err.message, 'error');
|
||||
})
|
||||
.finally(() => {
|
||||
loadJobs();
|
||||
});
|
||||
}
|
||||
|
||||
@@ -2505,8 +2643,8 @@ function renderJobDetail(job) {
|
||||
+ '<h2>' + escapeHtml(job.title) + '</h2>'
|
||||
+ '<span class="badge ' + stateClass + '">' + escapeHtml(job.state) + '</span>';
|
||||
|
||||
if (job.state === 'failed' || job.state === 'interrupted') {
|
||||
headerHtml += '<button class="btn-restart" onclick="restartJob(\'' + job.id + '\')">Restart</button>';
|
||||
if ((job.state === 'failed' || job.state === 'interrupted') && job.can_restart === true) {
|
||||
headerHtml += '<button class="btn-restart" onclick="restartJob(\'' + job.id + '\')">Retry</button>';
|
||||
}
|
||||
if (job.browse_url) {
|
||||
headerHtml += '<a class="btn-browse" href="' + escapeHtml(job.browse_url) + '" target="_blank">Browse Files</a>';
|
||||
@@ -2753,7 +2891,7 @@ function renderJobActivity(container, job) {
|
||||
activityCurrentJobId = job ? job.id : null;
|
||||
activityRenderedLiveIndex = 0;
|
||||
|
||||
container.innerHTML = '<div class="activity-toolbar">'
|
||||
let html = '<div class="activity-toolbar">'
|
||||
+ '<select id="activity-type-filter">'
|
||||
+ '<option value="all">All Events</option>'
|
||||
+ '<option value="message">Messages</option>'
|
||||
@@ -2762,12 +2900,17 @@ function renderJobActivity(container, job) {
|
||||
+ '</select>'
|
||||
+ '<label class="logs-checkbox"><input type="checkbox" id="activity-autoscroll" checked> Auto-scroll</label>'
|
||||
+ '</div>'
|
||||
+ '<div class="activity-terminal" id="activity-terminal"></div>'
|
||||
+ '<div class="activity-input-bar" id="activity-input-bar">'
|
||||
+ '<input type="text" id="activity-prompt-input" placeholder="Send follow-up prompt..." />'
|
||||
+ '<button id="activity-send-btn">Send</button>'
|
||||
+ '<button id="activity-done-btn" title="Signal done">Done</button>'
|
||||
+ '</div>';
|
||||
+ '<div class="activity-terminal" id="activity-terminal"></div>';
|
||||
|
||||
if (job && job.can_prompt === true) {
|
||||
html += '<div class="activity-input-bar" id="activity-input-bar">'
|
||||
+ '<input type="text" id="activity-prompt-input" placeholder="Send follow-up prompt..." />'
|
||||
+ '<button id="activity-send-btn">Send</button>'
|
||||
+ '<button id="activity-done-btn" title="Signal done">Done</button>'
|
||||
+ '</div>';
|
||||
}
|
||||
|
||||
container.innerHTML = html;
|
||||
|
||||
document.getElementById('activity-type-filter').addEventListener('change', applyActivityFilter);
|
||||
|
||||
@@ -2776,9 +2919,9 @@ function renderJobActivity(container, job) {
|
||||
const sendBtn = document.getElementById('activity-send-btn');
|
||||
const doneBtn = document.getElementById('activity-done-btn');
|
||||
|
||||
sendBtn.addEventListener('click', () => sendJobPrompt(job.id, false));
|
||||
doneBtn.addEventListener('click', () => sendJobPrompt(job.id, true));
|
||||
input.addEventListener('keydown', (e) => {
|
||||
if (sendBtn) sendBtn.addEventListener('click', () => sendJobPrompt(job.id, false));
|
||||
if (doneBtn) doneBtn.addEventListener('click', () => sendJobPrompt(job.id, true));
|
||||
if (input) input.addEventListener('keydown', (e) => {
|
||||
if (e.key === 'Enter') sendJobPrompt(job.id, false);
|
||||
});
|
||||
|
||||
@@ -3065,7 +3208,11 @@ function renderRoutineDetail(routine) {
|
||||
|
||||
function triggerRoutine(id) {
|
||||
apiFetch('/api/routines/' + id + '/trigger', { method: 'POST' })
|
||||
.then(() => showToast('Routine triggered', 'success'))
|
||||
.then(() => {
|
||||
showToast('Routine triggered', 'success');
|
||||
if (currentRoutineId === id) openRoutineDetail(id);
|
||||
else loadRoutines();
|
||||
})
|
||||
.catch((err) => showToast('Trigger failed: ' + err.message, 'error'));
|
||||
}
|
||||
|
||||
@@ -3615,7 +3762,7 @@ function formatTimeAgo(epochMs) {
|
||||
}
|
||||
|
||||
function installSkill(nameOrSlug, url, btn) {
|
||||
var body = { name: nameOrSlug };
|
||||
var body = { name: nameOrSlug, slug: nameOrSlug };
|
||||
if (url) body.url = url;
|
||||
|
||||
apiFetch('/api/skills/install', {
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
<html lang="en">
|
||||
<head>
|
||||
<meta charset="UTF-8">
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0">
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0, viewport-fit=cover">
|
||||
<title>IronClaw</title>
|
||||
<link rel="icon" href="/favicon.ico" type="image/x-icon">
|
||||
<link rel="preconnect" href="https://fonts.googleapis.com">
|
||||
@@ -33,6 +33,48 @@
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- 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-content">
|
||||
<div class="restart-modal-header">
|
||||
<h2>Restart IronClaw Instance</h2>
|
||||
<button class="restart-modal-close" onclick="cancelRestart()" title="Close">×</button>
|
||||
</div>
|
||||
<div class="restart-modal-body">
|
||||
<p class="restart-modal-description">
|
||||
Are you sure you want to restart the IronClaw instance? This will gracefully restart the process.
|
||||
</p>
|
||||
<div class="restart-modal-warning">
|
||||
<span class="restart-modal-warning-icon">⚠️</span>
|
||||
<p>Any in-progress jobs may be interrupted. The restart will complete within a few seconds.</p>
|
||||
</div>
|
||||
</div>
|
||||
<div class="restart-modal-footer">
|
||||
<button class="restart-modal-btn cancel" onclick="cancelRestart()">Cancel</button>
|
||||
<button class="restart-modal-btn confirm" onclick="confirmRestart()">Confirm Restart</button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Restart Progress Modal -->
|
||||
<div id="restart-loader" class="restart-loader" style="display: none;">
|
||||
<div class="restart-loader-overlay"></div>
|
||||
<div class="restart-loader-content">
|
||||
<div class="restart-spinner"></div>
|
||||
<div class="restart-loader-text">
|
||||
<p class="restart-title">Restarting IronClaw</p>
|
||||
<p class="restart-subtitle">Please wait while the process restarts...</p>
|
||||
</div>
|
||||
<div class="restart-progress-bar">
|
||||
<div class="restart-progress-fill"></div>
|
||||
</div>
|
||||
<p class="restart-modal-info">
|
||||
Check the Logs tab for details after the restart completes.
|
||||
</p>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Main App (hidden until authenticated) -->
|
||||
<div id="app">
|
||||
<!-- Tab Bar -->
|
||||
@@ -57,6 +99,14 @@
|
||||
<span id="sse-status">Connected</span>
|
||||
<div class="gateway-popover" id="gateway-popover"></div>
|
||||
</div>
|
||||
<button class="restart-btn" id="restart-btn" onclick="triggerRestart()" title="Gracefully restart the process">
|
||||
<svg id="restart-icon" width="13" height="13" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round">
|
||||
<path d="M23 4v6h-6"></path>
|
||||
<path d="M1 20v-6h6"></path>
|
||||
<path d="M3.51 9a9 9 0 0114.85-3.36M20.49 15a9 9 0 01-14.85 3.36"></path>
|
||||
</svg>
|
||||
<span>Restart</span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
<!-- Chat Tab -->
|
||||
|
||||
@@ -30,6 +30,7 @@ body {
|
||||
background: var(--bg);
|
||||
color: var(--text);
|
||||
height: 100vh;
|
||||
height: 100dvh;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
overflow: hidden;
|
||||
@@ -41,6 +42,7 @@ body {
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
height: 100vh;
|
||||
height: 100dvh;
|
||||
}
|
||||
|
||||
.auth-card-login {
|
||||
@@ -141,6 +143,7 @@ body {
|
||||
display: none;
|
||||
flex-direction: column;
|
||||
height: 100vh;
|
||||
height: 100dvh;
|
||||
}
|
||||
|
||||
/* Tab Bar */
|
||||
@@ -256,6 +259,284 @@ body {
|
||||
white-space: nowrap;
|
||||
}
|
||||
|
||||
/* Restart Button */
|
||||
.restart-btn {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 0.375rem;
|
||||
padding: 0.25rem 0.75rem;
|
||||
border-radius: 0.5rem;
|
||||
font-size: 0.8rem;
|
||||
border: 1px solid;
|
||||
border-color: #00d894;
|
||||
color: #00d894;
|
||||
background-color: transparent;
|
||||
cursor: pointer;
|
||||
transition: color 150ms, background-color 150ms, border-color 150ms;
|
||||
}
|
||||
|
||||
.restart-btn:hover:not(:disabled) {
|
||||
background-color: rgba(0, 216, 148, 0.1);
|
||||
}
|
||||
|
||||
.restart-btn:disabled {
|
||||
border-color: #333;
|
||||
color: #666;
|
||||
cursor: not-allowed;
|
||||
}
|
||||
|
||||
.restart-btn:disabled:hover {
|
||||
background-color: transparent;
|
||||
}
|
||||
|
||||
.restart-btn svg {
|
||||
flex-shrink: 0;
|
||||
width: 13px;
|
||||
height: 13px;
|
||||
}
|
||||
|
||||
.restart-btn svg.spinning {
|
||||
animation: spin-icon 1s linear infinite;
|
||||
}
|
||||
|
||||
@keyframes spin-icon {
|
||||
from { transform: rotate(0deg); }
|
||||
to { transform: rotate(360deg); }
|
||||
}
|
||||
|
||||
/* Restart Loader Overlay */
|
||||
.restart-loader {
|
||||
position: fixed;
|
||||
top: 0;
|
||||
left: 0;
|
||||
right: 0;
|
||||
bottom: 0;
|
||||
z-index: 9999;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
}
|
||||
|
||||
.restart-loader-overlay {
|
||||
position: absolute;
|
||||
top: 0;
|
||||
left: 0;
|
||||
right: 0;
|
||||
bottom: 0;
|
||||
background: rgba(0, 0, 0, 0.5);
|
||||
backdrop-filter: blur(4px);
|
||||
z-index: -1;
|
||||
}
|
||||
|
||||
.restart-loader-content {
|
||||
position: relative;
|
||||
z-index: 10000;
|
||||
background-color: #1a1a1a;
|
||||
border: 1px solid #333;
|
||||
border-radius: 0.75rem;
|
||||
box-shadow: 0 25px 50px -12px rgba(0, 0, 0, 0.25);
|
||||
width: 100%;
|
||||
max-width: 28rem;
|
||||
margin: 0 1rem;
|
||||
overflow: hidden;
|
||||
padding: 1.25rem;
|
||||
}
|
||||
|
||||
.restart-spinner {
|
||||
display: none;
|
||||
}
|
||||
|
||||
.restart-loader-text {
|
||||
padding: 0;
|
||||
}
|
||||
|
||||
.restart-title {
|
||||
color: #e0e0e0;
|
||||
font-size: 0.85rem;
|
||||
margin-bottom: 1rem;
|
||||
margin-top: 0;
|
||||
}
|
||||
|
||||
.restart-subtitle {
|
||||
display: none;
|
||||
}
|
||||
|
||||
/* Restart Modal (Confirmation) */
|
||||
.restart-modal {
|
||||
position: fixed;
|
||||
top: 0;
|
||||
left: 0;
|
||||
right: 0;
|
||||
bottom: 0;
|
||||
z-index: 9999;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
}
|
||||
|
||||
.restart-modal-overlay {
|
||||
position: absolute;
|
||||
top: 0;
|
||||
left: 0;
|
||||
right: 0;
|
||||
bottom: 0;
|
||||
background: rgba(0, 0, 0, 0.5);
|
||||
backdrop-filter: blur(4px);
|
||||
}
|
||||
|
||||
.restart-modal-content {
|
||||
position: relative;
|
||||
z-index: 10000;
|
||||
background-color: #1a1a1a;
|
||||
border: 1px solid #333;
|
||||
border-radius: 0.75rem;
|
||||
box-shadow: 0 25px 50px -12px rgba(0, 0, 0, 0.25);
|
||||
width: 100%;
|
||||
max-width: 28rem;
|
||||
margin: 0 1rem;
|
||||
overflow: hidden;
|
||||
}
|
||||
|
||||
.restart-modal-header {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
padding: 1rem 1.25rem;
|
||||
border-bottom: 1px solid #2a2a2a;
|
||||
}
|
||||
|
||||
.restart-modal-header h2 {
|
||||
color: #e0e0e0;
|
||||
font-size: 0.95rem;
|
||||
margin: 0;
|
||||
}
|
||||
|
||||
.restart-modal-close {
|
||||
color: #888;
|
||||
padding: 0.25rem;
|
||||
border-radius: 0.25rem;
|
||||
background-color: transparent;
|
||||
border: none;
|
||||
cursor: pointer;
|
||||
transition: color 150ms, background-color 150ms;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
}
|
||||
|
||||
.restart-modal-close:hover {
|
||||
color: #ccc;
|
||||
background-color: #2a2a2a;
|
||||
}
|
||||
|
||||
.restart-modal-body {
|
||||
padding: 1.25rem;
|
||||
}
|
||||
|
||||
.restart-modal-description {
|
||||
color: #aaa;
|
||||
font-size: 0.85rem;
|
||||
margin: 0;
|
||||
}
|
||||
|
||||
.restart-modal-warning {
|
||||
margin-top: 1rem;
|
||||
background-color: #1e1400;
|
||||
border: 1px solid #3a2a00;
|
||||
border-radius: 0.5rem;
|
||||
padding: 0.75rem 1rem;
|
||||
}
|
||||
|
||||
.restart-modal-warning p {
|
||||
color: #facc15;
|
||||
font-size: 0.8rem;
|
||||
margin: 0;
|
||||
}
|
||||
|
||||
.restart-modal-footer {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: flex-end;
|
||||
gap: 0.75rem;
|
||||
padding: 1rem 1.25rem;
|
||||
border-top: 1px solid #2a2a2a;
|
||||
}
|
||||
|
||||
.restart-modal-btn {
|
||||
padding: 0.5rem 1rem;
|
||||
border-radius: 0.5rem;
|
||||
font-size: 0.85rem;
|
||||
border: none;
|
||||
cursor: pointer;
|
||||
transition: background-color 150ms;
|
||||
}
|
||||
|
||||
.restart-modal-btn.cancel {
|
||||
color: #ccc;
|
||||
background-color: transparent;
|
||||
}
|
||||
|
||||
.restart-modal-btn.cancel:hover {
|
||||
background-color: #2a2a2a;
|
||||
}
|
||||
|
||||
.restart-modal-btn.confirm {
|
||||
background-color: #00D894;
|
||||
color: #111;
|
||||
}
|
||||
|
||||
.restart-modal-btn.confirm:hover {
|
||||
background-color: #00be82;
|
||||
}
|
||||
|
||||
/* Progress Bar for Restart */
|
||||
.restart-progress-bar {
|
||||
width: 100%;
|
||||
height: 0.375rem;
|
||||
background-color: #2a2a2a;
|
||||
border-radius: 9999px;
|
||||
overflow: hidden;
|
||||
}
|
||||
|
||||
.restart-progress-fill {
|
||||
height: 100%;
|
||||
border-radius: 9999px;
|
||||
background-color: #00D894;
|
||||
width: 40%;
|
||||
animation: indeterminate 1.5s ease-in-out infinite;
|
||||
}
|
||||
|
||||
@keyframes indeterminate {
|
||||
0% {
|
||||
margin-left: 0;
|
||||
width: 40%;
|
||||
}
|
||||
50% {
|
||||
margin-left: 60%;
|
||||
width: 40%;
|
||||
}
|
||||
100% {
|
||||
margin-left: 0;
|
||||
width: 40%;
|
||||
}
|
||||
}
|
||||
|
||||
.restart-modal-info {
|
||||
color: #666;
|
||||
font-size: 0.8rem;
|
||||
margin-top: 1.25rem;
|
||||
margin-bottom: 0;
|
||||
}
|
||||
|
||||
.restart-modal-info a {
|
||||
color: #00D894;
|
||||
text-decoration: none;
|
||||
}
|
||||
|
||||
.restart-modal-info a:hover {
|
||||
text-decoration: underline;
|
||||
}
|
||||
|
||||
.tee-popover {
|
||||
display: none;
|
||||
position: absolute;
|
||||
@@ -550,6 +831,10 @@ body {
|
||||
border-color: rgba(230, 76, 76, 0.3);
|
||||
}
|
||||
|
||||
.activity-tool-card[data-status="fail"] .activity-tool-name {
|
||||
color: var(--danger);
|
||||
}
|
||||
|
||||
.activity-tool-header {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
@@ -987,7 +1272,7 @@ body {
|
||||
/* Chat input */
|
||||
.chat-input {
|
||||
display: flex;
|
||||
padding: 12px 16px;
|
||||
padding: 12px 16px max(12px, env(safe-area-inset-bottom)) 16px;
|
||||
gap: 8px;
|
||||
background: var(--bg-secondary);
|
||||
border-top: 1px solid var(--border);
|
||||
@@ -1808,6 +2093,7 @@ body {
|
||||
.job-files {
|
||||
display: flex;
|
||||
height: calc(100vh - 280px);
|
||||
height: calc(100dvh - 280px);
|
||||
border: 1px solid var(--border);
|
||||
border-radius: var(--radius);
|
||||
overflow: hidden;
|
||||
@@ -2312,43 +2598,6 @@ body {
|
||||
margin-top: 6px;
|
||||
}
|
||||
|
||||
/* Restart overlay */
|
||||
.restart-overlay {
|
||||
position: fixed;
|
||||
top: 0;
|
||||
left: 0;
|
||||
width: 100%;
|
||||
height: 100%;
|
||||
background: rgba(0, 0, 0, 0.8);
|
||||
z-index: 2000;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
}
|
||||
|
||||
.restart-message {
|
||||
text-align: center;
|
||||
color: var(--text);
|
||||
}
|
||||
|
||||
.restart-message h2 {
|
||||
margin: 16px 0 8px;
|
||||
}
|
||||
|
||||
.restart-message p {
|
||||
color: var(--text-secondary);
|
||||
}
|
||||
|
||||
.restart-spinner {
|
||||
width: 40px;
|
||||
height: 40px;
|
||||
border: 3px solid var(--border);
|
||||
border-top-color: var(--accent);
|
||||
border-radius: 50%;
|
||||
animation: spin 0.8s linear infinite;
|
||||
margin: 0 auto;
|
||||
}
|
||||
|
||||
@keyframes spin {
|
||||
to { transform: rotate(360deg); }
|
||||
}
|
||||
|
||||
@@ -123,6 +123,10 @@ pub enum SseEvent {
|
||||
name: String,
|
||||
success: bool,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
error: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
parameters: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
thread_id: Option<String>,
|
||||
},
|
||||
#[serde(rename = "tool_result")]
|
||||
@@ -332,6 +336,15 @@ pub struct JobDetailResponse {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub job_mode: Option<String>,
|
||||
pub transitions: Vec<TransitionInfo>,
|
||||
/// Whether this job can be restarted from the UI.
|
||||
#[serde(default)]
|
||||
pub can_restart: bool,
|
||||
/// Whether follow-up prompts can be sent to this job.
|
||||
#[serde(default)]
|
||||
pub can_prompt: bool,
|
||||
/// The kind of job: "sandbox" or "agent".
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub job_kind: Option<String>,
|
||||
}
|
||||
|
||||
// --- Project Files ---
|
||||
@@ -379,6 +392,9 @@ pub struct ExtensionInfo {
|
||||
/// Whether this extension has configurable secrets (setup schema).
|
||||
#[serde(default)]
|
||||
pub needs_setup: bool,
|
||||
/// Whether this extension has an auth configuration (OAuth or manual token).
|
||||
#[serde(default)]
|
||||
pub has_auth: bool,
|
||||
/// WASM channel activation status: "installed", "configured", "active", "failed".
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub activation_status: Option<String>,
|
||||
@@ -451,9 +467,6 @@ pub struct ActionResponse {
|
||||
/// Whether the channel was successfully activated after setup.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub activated: Option<bool>,
|
||||
/// Whether a gateway restart is needed (activation failed).
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub needs_restart: Option<bool>,
|
||||
}
|
||||
|
||||
impl ActionResponse {
|
||||
@@ -465,7 +478,6 @@ impl ActionResponse {
|
||||
awaiting_token: None,
|
||||
instructions: None,
|
||||
activated: None,
|
||||
needs_restart: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -477,7 +489,6 @@ impl ActionResponse {
|
||||
awaiting_token: None,
|
||||
instructions: None,
|
||||
activated: None,
|
||||
needs_restart: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -562,6 +573,9 @@ pub struct SkillSearchResponse {
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct SkillInstallRequest {
|
||||
pub name: String,
|
||||
/// Registry slug (e.g. "owner/skill-name"). Preferred over `name` for
|
||||
/// constructing the download URL when fetching from ClawHub.
|
||||
pub slug: Option<String>,
|
||||
pub url: Option<String>,
|
||||
pub content: Option<String>,
|
||||
}
|
||||
|
||||
@@ -242,7 +242,7 @@ async fn handle_client_message(
|
||||
} => {
|
||||
if let Some(ref ext_mgr) = state.extension_manager {
|
||||
match ext_mgr.auth(&extension_name, Some(&token)).await {
|
||||
Ok(result) if result.status == "authenticated" => {
|
||||
Ok(result) if result.is_authenticated() => {
|
||||
let msg = match ext_mgr.activate(&extension_name).await {
|
||||
Ok(r) => format!(
|
||||
"{} authenticated ({} tools loaded)",
|
||||
@@ -268,9 +268,9 @@ async fn handle_client_message(
|
||||
.sse
|
||||
.broadcast(crate::channels::web::types::SseEvent::AuthRequired {
|
||||
extension_name,
|
||||
instructions: result.instructions,
|
||||
auth_url: result.auth_url,
|
||||
setup_url: result.setup_url,
|
||||
instructions: result.instructions().map(String::from),
|
||||
auth_url: result.auth_url().map(String::from),
|
||||
setup_url: result.setup_url().map(String::from),
|
||||
});
|
||||
}
|
||||
Err(e) => {
|
||||
@@ -483,6 +483,7 @@ mod tests {
|
||||
store: None,
|
||||
job_manager: None,
|
||||
prompt_queue: None,
|
||||
scheduler: None,
|
||||
user_id: "test".to_string(),
|
||||
shutdown_tx: tokio::sync::RwLock::new(None),
|
||||
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
|
||||
@@ -493,7 +494,6 @@ mod tests {
|
||||
registry_entries: Vec::new(),
|
||||
cost_guard: None,
|
||||
startup_time: std::time::Instant::now(),
|
||||
restart_requested: std::sync::atomic::AtomicBool::new(false),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+812
-13
@@ -17,10 +17,18 @@
|
||||
//! - **Runtime**: Users can set GOOGLE_OAUTH_CLIENT_ID / GOOGLE_OAUTH_CLIENT_SECRET
|
||||
//! env vars, which take priority over built-in defaults.
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD};
|
||||
use rand::RngCore;
|
||||
use sha2::{Digest, Sha256};
|
||||
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
|
||||
use tokio::net::TcpListener;
|
||||
use tokio::sync::RwLock;
|
||||
|
||||
use crate::secrets::{CreateSecretParams, SecretsStore};
|
||||
|
||||
// ── Built-in credentials ────────────────────────────────────────────────
|
||||
|
||||
@@ -121,6 +129,9 @@ pub enum OAuthCallbackError {
|
||||
#[error("Timed out waiting for authorization")]
|
||||
Timeout,
|
||||
|
||||
#[error("CSRF state mismatch: expected {expected}, got {actual}")]
|
||||
StateMismatch { expected: String, actual: String },
|
||||
|
||||
#[error("IO error: {0}")]
|
||||
Io(String),
|
||||
}
|
||||
@@ -177,16 +188,22 @@ pub async fn bind_callback_listener() -> Result<TcpListener, OAuthCallbackError>
|
||||
/// extracts the value of `param_name` (e.g., "code" or "token"), and shows a branded
|
||||
/// landing page using `display_name` (e.g., "Google", "Notion", "NEAR AI").
|
||||
///
|
||||
/// When `expected_state` is `Some`, the callback's `state` query parameter is validated
|
||||
/// against it to prevent CSRF attacks. If the state doesn't match, the callback is
|
||||
/// rejected with an error page.
|
||||
///
|
||||
/// Times out after 5 minutes.
|
||||
pub async fn wait_for_callback(
|
||||
listener: TcpListener,
|
||||
path_prefix: &str,
|
||||
param_name: &str,
|
||||
display_name: &str,
|
||||
expected_state: Option<&str>,
|
||||
) -> Result<String, OAuthCallbackError> {
|
||||
let path_prefix = path_prefix.to_string();
|
||||
let param_name = param_name.to_string();
|
||||
let display_name = display_name.to_string();
|
||||
let expected_state = expected_state.map(String::from);
|
||||
|
||||
tokio::time::timeout(Duration::from_secs(300), async move {
|
||||
loop {
|
||||
@@ -221,17 +238,29 @@ pub async fn wait_for_callback(
|
||||
return Err(OAuthCallbackError::Denied);
|
||||
}
|
||||
|
||||
// Look for the target parameter
|
||||
for param in query.split('&') {
|
||||
let parts: Vec<&str> = param.splitn(2, '=').collect();
|
||||
if parts.len() == 2 && parts[0] == param_name {
|
||||
let value = urlencoding::decode(parts[1])
|
||||
.unwrap_or_else(|_| parts[1].into())
|
||||
.into_owned();
|
||||
// Parse all query params into a map for validation
|
||||
let params: HashMap<&str, String> = query
|
||||
.split('&')
|
||||
.filter_map(|p| {
|
||||
let mut parts = p.splitn(2, '=');
|
||||
let key = parts.next()?;
|
||||
let val = parts.next().unwrap_or("");
|
||||
Some((
|
||||
key,
|
||||
urlencoding::decode(val)
|
||||
.unwrap_or_else(|_| val.into())
|
||||
.into_owned(),
|
||||
))
|
||||
})
|
||||
.collect();
|
||||
|
||||
let html = landing_html(&display_name, true);
|
||||
// Validate CSRF state parameter
|
||||
if let Some(ref expected) = expected_state {
|
||||
let actual = params.get("state").cloned().unwrap_or_default();
|
||||
if actual != *expected {
|
||||
let html = landing_html(&display_name, false);
|
||||
let response = format!(
|
||||
"HTTP/1.1 200 OK\r\n\
|
||||
"HTTP/1.1 403 Forbidden\r\n\
|
||||
Content-Type: text/html; charset=utf-8\r\n\
|
||||
Connection: close\r\n\
|
||||
\r\n\
|
||||
@@ -239,11 +268,29 @@ pub async fn wait_for_callback(
|
||||
html
|
||||
);
|
||||
let _ = socket.write_all(response.as_bytes()).await;
|
||||
let _ = socket.shutdown().await;
|
||||
|
||||
return Ok(value);
|
||||
return Err(OAuthCallbackError::StateMismatch {
|
||||
expected: expected.clone(),
|
||||
actual,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
// Look for the target parameter
|
||||
if let Some(value) = params.get(param_name.as_str()) {
|
||||
let html = landing_html(&display_name, true);
|
||||
let response = format!(
|
||||
"HTTP/1.1 200 OK\r\n\
|
||||
Content-Type: text/html; charset=utf-8\r\n\
|
||||
Connection: close\r\n\
|
||||
\r\n\
|
||||
{}",
|
||||
html
|
||||
);
|
||||
let _ = socket.write_all(response.as_bytes()).await;
|
||||
let _ = socket.shutdown().await;
|
||||
|
||||
return Ok(value.clone());
|
||||
}
|
||||
}
|
||||
|
||||
// Not the callback we're looking for
|
||||
@@ -271,7 +318,288 @@ fn html_escape(s: &str) -> String {
|
||||
out
|
||||
}
|
||||
|
||||
/// HTML landing page shown in the browser after an OAuth redirect.
|
||||
// ── Shared OAuth flow steps ─────────────────────────────────────────
|
||||
|
||||
/// Response from the OAuth token exchange.
|
||||
pub struct OAuthTokenResponse {
|
||||
pub access_token: String,
|
||||
pub refresh_token: Option<String>,
|
||||
pub expires_in: Option<u64>,
|
||||
}
|
||||
|
||||
/// Result of building an OAuth 2.0 authorization URL.
|
||||
pub struct OAuthUrlResult {
|
||||
/// The full authorization URL to redirect the user to.
|
||||
pub url: String,
|
||||
/// PKCE code verifier (must be sent with the token exchange request).
|
||||
pub code_verifier: Option<String>,
|
||||
/// Random state parameter for CSRF protection (must be validated in callback).
|
||||
pub state: String,
|
||||
}
|
||||
|
||||
/// Build an OAuth 2.0 authorization URL with optional PKCE and CSRF state.
|
||||
///
|
||||
/// Returns an `OAuthUrlResult` containing the authorization URL, optional PKCE
|
||||
/// code verifier, and a random `state` parameter for CSRF protection. The caller
|
||||
/// must validate the `state` value in the callback before exchanging the code.
|
||||
pub fn build_oauth_url(
|
||||
authorization_url: &str,
|
||||
client_id: &str,
|
||||
redirect_uri: &str,
|
||||
scopes: &[String],
|
||||
use_pkce: bool,
|
||||
extra_params: &HashMap<String, String>,
|
||||
) -> OAuthUrlResult {
|
||||
// Generate PKCE verifier and challenge
|
||||
let (code_verifier, code_challenge) = if use_pkce {
|
||||
let mut verifier_bytes = [0u8; 32];
|
||||
rand::rngs::OsRng.fill_bytes(&mut verifier_bytes);
|
||||
let verifier = URL_SAFE_NO_PAD.encode(verifier_bytes);
|
||||
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(verifier.as_bytes());
|
||||
let challenge = URL_SAFE_NO_PAD.encode(hasher.finalize());
|
||||
|
||||
(Some(verifier), Some(challenge))
|
||||
} else {
|
||||
(None, None)
|
||||
};
|
||||
|
||||
// Generate random state for CSRF protection
|
||||
let mut state_bytes = [0u8; 32];
|
||||
rand::rngs::OsRng.fill_bytes(&mut state_bytes);
|
||||
let state = URL_SAFE_NO_PAD.encode(state_bytes);
|
||||
|
||||
// Build authorization URL
|
||||
let mut auth_url = format!(
|
||||
"{}?client_id={}&response_type=code&redirect_uri={}&state={}",
|
||||
authorization_url,
|
||||
urlencoding::encode(client_id),
|
||||
urlencoding::encode(redirect_uri),
|
||||
urlencoding::encode(&state),
|
||||
);
|
||||
|
||||
if !scopes.is_empty() {
|
||||
auth_url.push_str(&format!(
|
||||
"&scope={}",
|
||||
urlencoding::encode(&scopes.join(" "))
|
||||
));
|
||||
}
|
||||
|
||||
if let Some(ref challenge) = code_challenge {
|
||||
auth_url.push_str(&format!(
|
||||
"&code_challenge={}&code_challenge_method=S256",
|
||||
challenge
|
||||
));
|
||||
}
|
||||
|
||||
for (key, value) in extra_params {
|
||||
auth_url.push_str(&format!(
|
||||
"&{}={}",
|
||||
urlencoding::encode(key),
|
||||
urlencoding::encode(value)
|
||||
));
|
||||
}
|
||||
|
||||
OAuthUrlResult {
|
||||
url: auth_url,
|
||||
code_verifier,
|
||||
state,
|
||||
}
|
||||
}
|
||||
|
||||
/// Exchange an OAuth authorization code for tokens.
|
||||
///
|
||||
/// POSTs to `token_url` with the authorization code and optional PKCE verifier.
|
||||
/// If `client_secret` is provided, uses HTTP Basic auth; otherwise includes
|
||||
/// `client_id` in the form body (for public clients).
|
||||
pub async fn exchange_oauth_code(
|
||||
token_url: &str,
|
||||
client_id: &str,
|
||||
client_secret: Option<&str>,
|
||||
code: &str,
|
||||
redirect_uri: &str,
|
||||
code_verifier: Option<&str>,
|
||||
access_token_field: &str,
|
||||
) -> Result<OAuthTokenResponse, OAuthCallbackError> {
|
||||
let client = reqwest::Client::new();
|
||||
let mut token_params = vec![
|
||||
("grant_type", "authorization_code".to_string()),
|
||||
("code", code.to_string()),
|
||||
("redirect_uri", redirect_uri.to_string()),
|
||||
];
|
||||
|
||||
if let Some(verifier) = code_verifier {
|
||||
token_params.push(("code_verifier", verifier.to_string()));
|
||||
}
|
||||
|
||||
let mut request = client.post(token_url);
|
||||
|
||||
if let Some(secret) = client_secret {
|
||||
request = request.basic_auth(client_id, Some(secret));
|
||||
} else {
|
||||
token_params.push(("client_id", client_id.to_string()));
|
||||
}
|
||||
|
||||
let token_response = request
|
||||
.form(&token_params)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| OAuthCallbackError::Io(format!("Token exchange request failed: {}", e)))?;
|
||||
|
||||
if !token_response.status().is_success() {
|
||||
let status = token_response.status();
|
||||
let body = token_response.text().await.unwrap_or_default();
|
||||
return Err(OAuthCallbackError::Io(format!(
|
||||
"Token exchange failed: {} - {}",
|
||||
status, body
|
||||
)));
|
||||
}
|
||||
|
||||
let token_data: serde_json::Value = token_response
|
||||
.json()
|
||||
.await
|
||||
.map_err(|e| OAuthCallbackError::Io(format!("Failed to parse token response: {}", e)))?;
|
||||
|
||||
let access_token = token_data
|
||||
.get(access_token_field)
|
||||
.and_then(|v| v.as_str())
|
||||
.ok_or_else(|| {
|
||||
// Log only the field names present, not values (which may contain tokens)
|
||||
let fields: Vec<&str> = token_data
|
||||
.as_object()
|
||||
.map(|o| o.keys().map(|k| k.as_str()).collect())
|
||||
.unwrap_or_default();
|
||||
OAuthCallbackError::Io(format!(
|
||||
"No '{}' field in token response (fields present: {:?})",
|
||||
access_token_field, fields
|
||||
))
|
||||
})?
|
||||
.to_string();
|
||||
|
||||
let refresh_token = token_data
|
||||
.get("refresh_token")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(String::from);
|
||||
let expires_in = token_data.get("expires_in").and_then(|v| v.as_u64());
|
||||
|
||||
Ok(OAuthTokenResponse {
|
||||
access_token,
|
||||
refresh_token,
|
||||
expires_in,
|
||||
})
|
||||
}
|
||||
|
||||
/// Store OAuth tokens (access + refresh) in the secrets store.
|
||||
///
|
||||
/// Also stores the granted scopes as `{secret_name}_scopes` so that scope
|
||||
/// expansion can be detected on subsequent activations.
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub async fn store_oauth_tokens(
|
||||
store: &(dyn SecretsStore + Send + Sync),
|
||||
user_id: &str,
|
||||
secret_name: &str,
|
||||
provider: Option<&str>,
|
||||
access_token: &str,
|
||||
refresh_token: Option<&str>,
|
||||
expires_in: Option<u64>,
|
||||
scopes: &[String],
|
||||
) -> Result<(), OAuthCallbackError> {
|
||||
let mut params = CreateSecretParams::new(secret_name, access_token);
|
||||
|
||||
if let Some(prov) = provider {
|
||||
params = params.with_provider(prov);
|
||||
}
|
||||
|
||||
if let Some(secs) = expires_in {
|
||||
let expires_at = chrono::Utc::now() + chrono::Duration::seconds(secs as i64);
|
||||
params = params.with_expiry(expires_at);
|
||||
}
|
||||
|
||||
store
|
||||
.create(user_id, params)
|
||||
.await
|
||||
.map_err(|e| OAuthCallbackError::Io(format!("Failed to save token: {}", e)))?;
|
||||
|
||||
// Store refresh token separately (no expiry, it's long-lived)
|
||||
if let Some(rt) = refresh_token {
|
||||
let refresh_name = format!("{}_refresh_token", secret_name);
|
||||
let mut refresh_params = CreateSecretParams::new(&refresh_name, rt);
|
||||
if let Some(prov) = provider {
|
||||
refresh_params = refresh_params.with_provider(prov);
|
||||
}
|
||||
store
|
||||
.create(user_id, refresh_params)
|
||||
.await
|
||||
.map_err(|e| OAuthCallbackError::Io(format!("Failed to save refresh token: {}", e)))?;
|
||||
}
|
||||
|
||||
// Store granted scopes for scope expansion detection
|
||||
if !scopes.is_empty() {
|
||||
let scopes_name = format!("{}_scopes", secret_name);
|
||||
let scopes_value = scopes.join(" ");
|
||||
let scopes_params = CreateSecretParams::new(&scopes_name, &scopes_value);
|
||||
// Best-effort: scope tracking failure shouldn't block auth
|
||||
let _ = store.create(user_id, scopes_params).await;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Validate an OAuth token against a tool's validation endpoint.
|
||||
///
|
||||
/// Sends a request to the configured endpoint with the token as a Bearer header.
|
||||
/// Returns `Ok(())` if the response status matches the expected success status,
|
||||
/// or an error with details if validation fails (wrong account, expired token, etc.).
|
||||
pub async fn validate_oauth_token(
|
||||
token: &str,
|
||||
validation: &crate::tools::wasm::ValidationEndpointSchema,
|
||||
) -> Result<(), OAuthCallbackError> {
|
||||
let client = reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(10))
|
||||
.build()
|
||||
.map_err(|e| OAuthCallbackError::Io(format!("Failed to build HTTP client: {}", e)))?;
|
||||
|
||||
let request = match validation.method.to_uppercase().as_str() {
|
||||
"POST" => client.post(&validation.url),
|
||||
_ => client.get(&validation.url),
|
||||
};
|
||||
|
||||
let mut request = request.header("Authorization", format!("Bearer {}", token));
|
||||
|
||||
// Add custom headers from the validation schema (e.g., Notion-Version)
|
||||
for (key, value) in &validation.headers {
|
||||
request = request.header(key, value);
|
||||
}
|
||||
|
||||
let response = request
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| OAuthCallbackError::Io(format!("Validation request failed: {}", e)))?;
|
||||
|
||||
if response.status().as_u16() == validation.success_status {
|
||||
Ok(())
|
||||
} else {
|
||||
let status = response.status();
|
||||
let body = response.text().await.unwrap_or_default();
|
||||
let truncated: String = if body.len() > 200 {
|
||||
let mut end = 200;
|
||||
while end > 0 && !body.is_char_boundary(end) {
|
||||
end -= 1;
|
||||
}
|
||||
format!("{}...", &body[..end])
|
||||
} else {
|
||||
body
|
||||
};
|
||||
Err(OAuthCallbackError::Io(format!(
|
||||
"Token validation failed: HTTP {} (expected {}): {}",
|
||||
status, validation.success_status, truncated
|
||||
)))
|
||||
}
|
||||
}
|
||||
|
||||
// ── Landing pages ───────────────────────────────────────────────────
|
||||
|
||||
pub fn landing_html(provider_name: &str, success: bool) -> String {
|
||||
let safe_name = html_escape(provider_name);
|
||||
let (icon, heading, subtitle, accent) = if success {
|
||||
@@ -357,6 +685,219 @@ pub fn landing_html(provider_name: &str, success: bool) -> String {
|
||||
)
|
||||
}
|
||||
|
||||
// ── Gateway callback support ─────────────────────────────────────────
|
||||
|
||||
/// State for an in-progress OAuth flow, keyed by CSRF `state` parameter.
|
||||
///
|
||||
/// Created by `start_wasm_oauth()` and consumed by the web gateway's
|
||||
/// `/oauth/callback` handler when running in hosted mode.
|
||||
pub struct PendingOAuthFlow {
|
||||
/// Extension name (e.g., "google_calendar").
|
||||
pub extension_name: String,
|
||||
/// Human-readable display name (e.g., "Google Calendar").
|
||||
pub display_name: String,
|
||||
/// OAuth token exchange URL.
|
||||
pub token_url: String,
|
||||
/// OAuth client ID.
|
||||
pub client_id: String,
|
||||
/// OAuth client secret (optional for PKCE-only flows).
|
||||
pub client_secret: Option<String>,
|
||||
/// The redirect_uri used in the authorization request.
|
||||
pub redirect_uri: String,
|
||||
/// PKCE code verifier (must match the code_challenge sent in the auth URL).
|
||||
pub code_verifier: Option<String>,
|
||||
/// Field name in token response containing the access token.
|
||||
pub access_token_field: String,
|
||||
/// Secret name for storage (e.g., "google_oauth_token").
|
||||
pub secret_name: String,
|
||||
/// Provider hint (e.g., "google").
|
||||
pub provider: Option<String>,
|
||||
/// Token validation endpoint (optional).
|
||||
pub validation_endpoint: Option<crate::tools::wasm::ValidationEndpointSchema>,
|
||||
/// Scopes that were requested.
|
||||
pub scopes: Vec<String>,
|
||||
/// User ID for secret storage.
|
||||
pub user_id: String,
|
||||
/// Secrets store reference for token persistence.
|
||||
pub secrets: Arc<dyn SecretsStore + Send + Sync>,
|
||||
/// SSE broadcast sender for notifying the web UI.
|
||||
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>,
|
||||
/// When this flow was created (for expiry).
|
||||
pub created_at: std::time::Instant,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for PendingOAuthFlow {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("PendingOAuthFlow")
|
||||
.field("extension_name", &self.extension_name)
|
||||
.field("display_name", &self.display_name)
|
||||
.field("secret_name", &self.secret_name)
|
||||
.field("created_at", &self.created_at)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
/// Thread-safe registry of pending OAuth flows, keyed by CSRF `state` parameter.
|
||||
pub type PendingOAuthRegistry = Arc<RwLock<HashMap<String, PendingOAuthFlow>>>;
|
||||
|
||||
/// Create a new empty pending OAuth flow registry.
|
||||
pub fn new_pending_oauth_registry() -> PendingOAuthRegistry {
|
||||
Arc::new(RwLock::new(HashMap::new()))
|
||||
}
|
||||
|
||||
/// Returns `true` if OAuth callbacks should be routed through the web gateway
|
||||
/// instead of the local TCP listener.
|
||||
///
|
||||
/// This is the case when `IRONCLAW_OAUTH_CALLBACK_URL` is set to a non-loopback
|
||||
/// URL, meaning the user's browser will redirect to a hosted gateway rather than
|
||||
/// localhost.
|
||||
pub fn use_gateway_callback() -> bool {
|
||||
std::env::var("IRONCLAW_OAUTH_CALLBACK_URL")
|
||||
.ok()
|
||||
.filter(|v| !v.is_empty())
|
||||
.map(|raw| {
|
||||
url::Url::parse(&raw)
|
||||
.ok()
|
||||
.and_then(|u| u.host_str().map(String::from))
|
||||
.map(|host| !is_loopback_host(&host))
|
||||
.unwrap_or(false)
|
||||
})
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
/// Maximum age for pending OAuth flows (5 minutes, matching TCP listener timeout).
|
||||
pub const OAUTH_FLOW_EXPIRY: Duration = Duration::from_secs(300);
|
||||
|
||||
/// Remove expired flows from the registry.
|
||||
///
|
||||
/// Called when inserting new flows to prevent accumulation from abandoned
|
||||
/// OAuth attempts.
|
||||
pub async fn sweep_expired_flows(registry: &PendingOAuthRegistry) {
|
||||
let mut flows = registry.write().await;
|
||||
flows.retain(|_, flow| flow.created_at.elapsed() < OAUTH_FLOW_EXPIRY);
|
||||
}
|
||||
|
||||
// ── Platform routing helpers ────────────────────────────────────────
|
||||
|
||||
/// Prepend instance name to CSRF state for platform routing.
|
||||
///
|
||||
/// The NEAR AI platform nginx proxy at `auth.DOMAIN` parses the instance name
|
||||
/// from the `state` query parameter (format: `instance:nonce`) to route the
|
||||
/// OAuth callback to the correct container.
|
||||
///
|
||||
/// Returns the nonce unchanged when `IRONCLAW_INSTANCE_NAME` is not set
|
||||
/// (local/non-platform mode).
|
||||
pub fn build_platform_state(nonce: &str) -> String {
|
||||
let instance = std::env::var("IRONCLAW_INSTANCE_NAME")
|
||||
.or_else(|_| std::env::var("OPENCLAW_INSTANCE_NAME"))
|
||||
.ok()
|
||||
.filter(|v| !v.is_empty());
|
||||
match instance {
|
||||
Some(name) => format!("{}:{}", name, nonce),
|
||||
None => nonce.to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Strip the instance prefix from a state parameter to recover the lookup nonce.
|
||||
///
|
||||
/// `"myinstance:abc123"` → `"abc123"`, `"abc123"` → `"abc123"` (no prefix).
|
||||
///
|
||||
/// Safe because nonces are base64url-encoded (`[A-Za-z0-9_-]`, no colons).
|
||||
pub fn strip_instance_prefix(state: &str) -> &str {
|
||||
state
|
||||
.split_once(':')
|
||||
.map(|(_, nonce)| nonce)
|
||||
.unwrap_or(state)
|
||||
}
|
||||
|
||||
/// Exchange an OAuth authorization code via the platform's token exchange proxy.
|
||||
///
|
||||
/// The proxy holds `client_secret` server-side so the container never sees it.
|
||||
/// Authenticated via the gateway auth token (Bearer header).
|
||||
///
|
||||
/// The proxy expects form params `{code, redirect_uri, code_verifier}` and
|
||||
/// returns a standard Google token response `{access_token, refresh_token, expires_in}`.
|
||||
pub async fn exchange_via_proxy(
|
||||
proxy_url: &str,
|
||||
gateway_token: &str,
|
||||
code: &str,
|
||||
redirect_uri: &str,
|
||||
code_verifier: Option<&str>,
|
||||
access_token_field: &str,
|
||||
) -> Result<OAuthTokenResponse, OAuthCallbackError> {
|
||||
if gateway_token.is_empty() {
|
||||
return Err(OAuthCallbackError::Io(
|
||||
"Gateway auth token is required for proxy token exchange".to_string(),
|
||||
));
|
||||
}
|
||||
let exchange_url = format!("{}/oauth/exchange", proxy_url.trim_end_matches('/'));
|
||||
|
||||
let client = reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(60))
|
||||
.build()
|
||||
.map_err(|e| OAuthCallbackError::Io(format!("Failed to build HTTP client: {}", e)))?;
|
||||
let mut params = vec![
|
||||
("code", code.to_string()),
|
||||
("redirect_uri", redirect_uri.to_string()),
|
||||
];
|
||||
if let Some(verifier) = code_verifier {
|
||||
params.push(("code_verifier", verifier.to_string()));
|
||||
}
|
||||
|
||||
let response = client
|
||||
.post(&exchange_url)
|
||||
.bearer_auth(gateway_token)
|
||||
.form(¶ms)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| {
|
||||
OAuthCallbackError::Io(format!("Token exchange proxy request failed: {}", e))
|
||||
})?;
|
||||
|
||||
if !response.status().is_success() {
|
||||
let status = response.status();
|
||||
let body = response.text().await.unwrap_or_default();
|
||||
return Err(OAuthCallbackError::Io(format!(
|
||||
"Token exchange proxy failed: {} - {}",
|
||||
status, body
|
||||
)));
|
||||
}
|
||||
|
||||
let token_data: serde_json::Value = response
|
||||
.json()
|
||||
.await
|
||||
.map_err(|e| OAuthCallbackError::Io(format!("Failed to parse proxy response: {}", e)))?;
|
||||
|
||||
let access_token = token_data
|
||||
.get(access_token_field)
|
||||
.and_then(|v| v.as_str())
|
||||
.ok_or_else(|| {
|
||||
let fields: Vec<&str> = token_data
|
||||
.as_object()
|
||||
.map(|o| o.keys().map(|k| k.as_str()).collect())
|
||||
.unwrap_or_default();
|
||||
OAuthCallbackError::Io(format!(
|
||||
"No '{}' field in proxy response (fields present: {:?})",
|
||||
access_token_field, fields
|
||||
))
|
||||
})?
|
||||
.to_string();
|
||||
|
||||
let refresh_token = token_data
|
||||
.get("refresh_token")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(String::from);
|
||||
let expires_in = token_data.get("expires_in").and_then(|v| v.as_u64());
|
||||
|
||||
Ok(OAuthTokenResponse {
|
||||
access_token,
|
||||
refresh_token,
|
||||
expires_in,
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Mutex;
|
||||
@@ -512,4 +1053,262 @@ mod tests {
|
||||
assert!(html.contains("#ef4444")); // red accent
|
||||
assert!(!html.contains("Connected"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_oauth_url_basic() {
|
||||
use std::collections::HashMap;
|
||||
|
||||
use crate::cli::oauth_defaults::build_oauth_url;
|
||||
|
||||
let result = build_oauth_url(
|
||||
"https://accounts.google.com/o/oauth2/auth",
|
||||
"my-client-id",
|
||||
"http://localhost:9876/callback",
|
||||
&["openid".to_string(), "email".to_string()],
|
||||
false,
|
||||
&HashMap::new(),
|
||||
);
|
||||
|
||||
assert!(
|
||||
result
|
||||
.url
|
||||
.starts_with("https://accounts.google.com/o/oauth2/auth?")
|
||||
);
|
||||
assert!(result.url.contains("client_id=my-client-id"));
|
||||
assert!(result.url.contains("response_type=code"));
|
||||
assert!(result.url.contains("redirect_uri="));
|
||||
assert!(result.url.contains("scope=openid%20email"));
|
||||
assert!(result.url.contains("state="));
|
||||
assert!(result.code_verifier.is_none());
|
||||
assert!(!result.state.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_oauth_url_with_pkce() {
|
||||
use std::collections::HashMap;
|
||||
|
||||
use crate::cli::oauth_defaults::build_oauth_url;
|
||||
|
||||
let result = build_oauth_url(
|
||||
"https://auth.example.com/authorize",
|
||||
"client-123",
|
||||
"http://localhost:9876/callback",
|
||||
&[],
|
||||
true,
|
||||
&HashMap::new(),
|
||||
);
|
||||
|
||||
assert!(result.url.contains("code_challenge="));
|
||||
assert!(result.url.contains("code_challenge_method=S256"));
|
||||
assert!(result.code_verifier.is_some());
|
||||
let verifier = result.code_verifier.unwrap();
|
||||
assert!(!verifier.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_oauth_url_with_extra_params() {
|
||||
use std::collections::HashMap;
|
||||
|
||||
use crate::cli::oauth_defaults::build_oauth_url;
|
||||
|
||||
let mut extra = HashMap::new();
|
||||
extra.insert("access_type".to_string(), "offline".to_string());
|
||||
extra.insert("prompt".to_string(), "consent".to_string());
|
||||
|
||||
let result = build_oauth_url(
|
||||
"https://auth.example.com/authorize",
|
||||
"client-123",
|
||||
"http://localhost:9876/callback",
|
||||
&["read".to_string()],
|
||||
false,
|
||||
&extra,
|
||||
);
|
||||
|
||||
assert!(result.url.contains("access_type=offline"));
|
||||
assert!(result.url.contains("prompt=consent"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_oauth_url_state_is_unique() {
|
||||
use std::collections::HashMap;
|
||||
|
||||
use crate::cli::oauth_defaults::build_oauth_url;
|
||||
|
||||
let result1 = build_oauth_url(
|
||||
"https://auth.example.com/authorize",
|
||||
"client",
|
||||
"http://localhost:9876/callback",
|
||||
&[],
|
||||
false,
|
||||
&HashMap::new(),
|
||||
);
|
||||
let result2 = build_oauth_url(
|
||||
"https://auth.example.com/authorize",
|
||||
"client",
|
||||
"http://localhost:9876/callback",
|
||||
&[],
|
||||
false,
|
||||
&HashMap::new(),
|
||||
);
|
||||
|
||||
// State should be different each time (random)
|
||||
assert_ne!(result1.state, result2.state);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_use_gateway_callback_false_by_default() {
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok();
|
||||
// SAFETY: Under ENV_MUTEX, no concurrent env access.
|
||||
unsafe {
|
||||
std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL");
|
||||
}
|
||||
assert!(!crate::cli::oauth_defaults::use_gateway_callback());
|
||||
unsafe {
|
||||
if let Some(val) = original {
|
||||
std::env::set_var("IRONCLAW_OAUTH_CALLBACK_URL", val);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_use_gateway_callback_true_for_hosted() {
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok();
|
||||
// SAFETY: Under ENV_MUTEX, no concurrent env access.
|
||||
unsafe {
|
||||
std::env::set_var(
|
||||
"IRONCLAW_OAUTH_CALLBACK_URL",
|
||||
"https://kind-deer.agent1.near.ai",
|
||||
);
|
||||
}
|
||||
assert!(crate::cli::oauth_defaults::use_gateway_callback());
|
||||
unsafe {
|
||||
if let Some(val) = original {
|
||||
std::env::set_var("IRONCLAW_OAUTH_CALLBACK_URL", val);
|
||||
} else {
|
||||
std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_use_gateway_callback_false_for_localhost() {
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok();
|
||||
// SAFETY: Under ENV_MUTEX, no concurrent env access.
|
||||
unsafe {
|
||||
std::env::set_var("IRONCLAW_OAUTH_CALLBACK_URL", "http://127.0.0.1:3001");
|
||||
}
|
||||
assert!(!crate::cli::oauth_defaults::use_gateway_callback());
|
||||
unsafe {
|
||||
if let Some(val) = original {
|
||||
std::env::set_var("IRONCLAW_OAUTH_CALLBACK_URL", val);
|
||||
} else {
|
||||
std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_use_gateway_callback_false_for_empty() {
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
let original = std::env::var("IRONCLAW_OAUTH_CALLBACK_URL").ok();
|
||||
// SAFETY: Under ENV_MUTEX, no concurrent env access.
|
||||
unsafe {
|
||||
std::env::set_var("IRONCLAW_OAUTH_CALLBACK_URL", "");
|
||||
}
|
||||
assert!(!crate::cli::oauth_defaults::use_gateway_callback());
|
||||
unsafe {
|
||||
if let Some(val) = original {
|
||||
std::env::set_var("IRONCLAW_OAUTH_CALLBACK_URL", val);
|
||||
} else {
|
||||
std::env::remove_var("IRONCLAW_OAUTH_CALLBACK_URL");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_platform_state_with_instance() {
|
||||
use crate::cli::oauth_defaults::build_platform_state;
|
||||
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
let original = std::env::var("IRONCLAW_INSTANCE_NAME").ok();
|
||||
// SAFETY: Under ENV_MUTEX, no concurrent env access.
|
||||
unsafe {
|
||||
std::env::set_var("IRONCLAW_INSTANCE_NAME", "kind-deer");
|
||||
}
|
||||
assert_eq!(build_platform_state("abc123"), "kind-deer:abc123");
|
||||
unsafe {
|
||||
if let Some(val) = original {
|
||||
std::env::set_var("IRONCLAW_INSTANCE_NAME", val);
|
||||
} else {
|
||||
std::env::remove_var("IRONCLAW_INSTANCE_NAME");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_platform_state_without_instance() {
|
||||
use crate::cli::oauth_defaults::build_platform_state;
|
||||
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
let original = std::env::var("IRONCLAW_INSTANCE_NAME").ok();
|
||||
let original_oc = std::env::var("OPENCLAW_INSTANCE_NAME").ok();
|
||||
// SAFETY: Under ENV_MUTEX, no concurrent env access.
|
||||
unsafe {
|
||||
std::env::remove_var("IRONCLAW_INSTANCE_NAME");
|
||||
std::env::remove_var("OPENCLAW_INSTANCE_NAME");
|
||||
}
|
||||
assert_eq!(build_platform_state("abc123"), "abc123");
|
||||
unsafe {
|
||||
if let Some(val) = original {
|
||||
std::env::set_var("IRONCLAW_INSTANCE_NAME", val);
|
||||
}
|
||||
if let Some(val) = original_oc {
|
||||
std::env::set_var("OPENCLAW_INSTANCE_NAME", val);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_platform_state_with_openclaw_instance() {
|
||||
use crate::cli::oauth_defaults::build_platform_state;
|
||||
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
let original_ic = std::env::var("IRONCLAW_INSTANCE_NAME").ok();
|
||||
let original_oc = std::env::var("OPENCLAW_INSTANCE_NAME").ok();
|
||||
// SAFETY: Under ENV_MUTEX, no concurrent env access.
|
||||
unsafe {
|
||||
std::env::remove_var("IRONCLAW_INSTANCE_NAME");
|
||||
std::env::set_var("OPENCLAW_INSTANCE_NAME", "quiet-lion");
|
||||
}
|
||||
assert_eq!(build_platform_state("xyz789"), "quiet-lion:xyz789");
|
||||
unsafe {
|
||||
if let Some(val) = original_ic {
|
||||
std::env::set_var("IRONCLAW_INSTANCE_NAME", val);
|
||||
}
|
||||
if let Some(val) = original_oc {
|
||||
std::env::set_var("OPENCLAW_INSTANCE_NAME", val);
|
||||
} else {
|
||||
std::env::remove_var("OPENCLAW_INSTANCE_NAME");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_strip_instance_prefix_with_colon() {
|
||||
use crate::cli::oauth_defaults::strip_instance_prefix;
|
||||
|
||||
assert_eq!(strip_instance_prefix("kind-deer:abc123"), "abc123");
|
||||
assert_eq!(strip_instance_prefix("my-instance:xyz"), "xyz");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_strip_instance_prefix_without_colon() {
|
||||
use crate::cli::oauth_defaults::strip_instance_prefix;
|
||||
|
||||
assert_eq!(strip_instance_prefix("abc123"), "abc123");
|
||||
assert_eq!(strip_instance_prefix(""), "");
|
||||
}
|
||||
}
|
||||
|
||||
+58
-184
@@ -782,11 +782,7 @@ async fn auth_tool_oauth(
|
||||
auth: &crate::tools::wasm::AuthCapabilitySchema,
|
||||
oauth: &crate::tools::wasm::OAuthConfigSchema,
|
||||
) -> anyhow::Result<()> {
|
||||
use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD};
|
||||
use rand::RngCore;
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
use crate::cli::oauth_defaults::{self, OAUTH_CALLBACK_PORT};
|
||||
use crate::cli::oauth_defaults;
|
||||
|
||||
let display_name = auth.display_name.as_deref().unwrap_or(&auth.secret_name);
|
||||
|
||||
@@ -827,142 +823,69 @@ async fn auth_tool_oauth(
|
||||
println!();
|
||||
|
||||
let listener = oauth_defaults::bind_callback_listener().await?;
|
||||
let redirect_uri = format!("http://localhost:{}/callback", OAUTH_CALLBACK_PORT);
|
||||
let redirect_uri = format!("{}/callback", oauth_defaults::callback_url());
|
||||
|
||||
// Generate PKCE verifier and challenge
|
||||
let (code_verifier, code_challenge) = if oauth.use_pkce {
|
||||
let mut verifier_bytes = [0u8; 32];
|
||||
rand::thread_rng().fill_bytes(&mut verifier_bytes);
|
||||
let verifier = URL_SAFE_NO_PAD.encode(verifier_bytes);
|
||||
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(verifier.as_bytes());
|
||||
let challenge = URL_SAFE_NO_PAD.encode(hasher.finalize());
|
||||
|
||||
(Some(verifier), Some(challenge))
|
||||
} else {
|
||||
(None, None)
|
||||
};
|
||||
|
||||
// Build authorization URL
|
||||
let mut auth_url = format!(
|
||||
"{}?client_id={}&response_type=code&redirect_uri={}",
|
||||
oauth.authorization_url,
|
||||
urlencoding::encode(&client_id),
|
||||
urlencoding::encode(&redirect_uri)
|
||||
// Build authorization URL with PKCE and CSRF state
|
||||
let oauth_result = oauth_defaults::build_oauth_url(
|
||||
&oauth.authorization_url,
|
||||
&client_id,
|
||||
&redirect_uri,
|
||||
&oauth.scopes,
|
||||
oauth.use_pkce,
|
||||
&oauth.extra_params,
|
||||
);
|
||||
|
||||
if !oauth.scopes.is_empty() {
|
||||
auth_url.push_str(&format!(
|
||||
"&scope={}",
|
||||
urlencoding::encode(&oauth.scopes.join(" "))
|
||||
));
|
||||
}
|
||||
|
||||
if let Some(ref challenge) = code_challenge {
|
||||
auth_url.push_str(&format!(
|
||||
"&code_challenge={}&code_challenge_method=S256",
|
||||
challenge
|
||||
));
|
||||
}
|
||||
|
||||
// Add extra params
|
||||
for (key, value) in &oauth.extra_params {
|
||||
auth_url.push_str(&format!(
|
||||
"&{}={}",
|
||||
urlencoding::encode(key),
|
||||
urlencoding::encode(value)
|
||||
));
|
||||
}
|
||||
let code_verifier = oauth_result.code_verifier;
|
||||
|
||||
println!(" Opening browser for {} login...", display_name);
|
||||
println!();
|
||||
|
||||
if let Err(e) = open::that(&auth_url) {
|
||||
if let Err(e) = open::that(&oauth_result.url) {
|
||||
println!(" Could not open browser: {}", e);
|
||||
println!(" Please open this URL manually:");
|
||||
println!(" {}", auth_url);
|
||||
println!(" {}", oauth_result.url);
|
||||
}
|
||||
|
||||
println!(" Waiting for authorization...");
|
||||
|
||||
let code =
|
||||
oauth_defaults::wait_for_callback(listener, "/callback", "code", display_name).await?;
|
||||
let code = oauth_defaults::wait_for_callback(
|
||||
listener,
|
||||
"/callback",
|
||||
"code",
|
||||
display_name,
|
||||
Some(&oauth_result.state),
|
||||
)
|
||||
.await?;
|
||||
|
||||
println!();
|
||||
println!(" Exchanging code for token...");
|
||||
|
||||
// Exchange code for token
|
||||
let client = reqwest::Client::new();
|
||||
let mut token_params = vec![
|
||||
("grant_type", "authorization_code".to_string()),
|
||||
("code", code),
|
||||
("redirect_uri", redirect_uri),
|
||||
];
|
||||
|
||||
if let Some(ref verifier) = code_verifier {
|
||||
token_params.push(("code_verifier", verifier.to_string()));
|
||||
}
|
||||
|
||||
// Build token request
|
||||
let mut request = client.post(&oauth.token_url);
|
||||
|
||||
// Use Basic auth if client_secret is provided, otherwise include client_id in body
|
||||
if let Some(ref secret) = client_secret {
|
||||
request = request.basic_auth(&client_id, Some(secret));
|
||||
} else {
|
||||
token_params.push(("client_id", client_id));
|
||||
}
|
||||
|
||||
let token_response = request.form(&token_params).send().await?;
|
||||
|
||||
if !token_response.status().is_success() {
|
||||
let status = token_response.status();
|
||||
let body = token_response.text().await.unwrap_or_default();
|
||||
return Err(anyhow::anyhow!(
|
||||
"Token exchange failed: {} - {}",
|
||||
status,
|
||||
body
|
||||
));
|
||||
}
|
||||
|
||||
let token_data: serde_json::Value = token_response.json().await?;
|
||||
let access_token = token_data
|
||||
.get(&oauth.access_token_field)
|
||||
.and_then(|v| v.as_str())
|
||||
.ok_or_else(|| {
|
||||
anyhow::anyhow!(
|
||||
"No {} in token response: {:?}",
|
||||
oauth.access_token_field,
|
||||
token_data
|
||||
)
|
||||
})?;
|
||||
|
||||
let refresh_token = token_data.get("refresh_token").and_then(|v| v.as_str());
|
||||
let expires_in = token_data.get("expires_in").and_then(|v| v.as_u64());
|
||||
|
||||
// Save the token (with refresh token and expiry if provided)
|
||||
save_token(
|
||||
store,
|
||||
user_id,
|
||||
auth,
|
||||
access_token,
|
||||
refresh_token,
|
||||
expires_in,
|
||||
let token_response = oauth_defaults::exchange_oauth_code(
|
||||
&oauth.token_url,
|
||||
&client_id,
|
||||
client_secret.as_deref(),
|
||||
&code,
|
||||
&redirect_uri,
|
||||
code_verifier.as_deref(),
|
||||
&oauth.access_token_field,
|
||||
)
|
||||
.await?;
|
||||
|
||||
// Extract any additional info for display
|
||||
let workspace_name = token_data
|
||||
.get("workspace_name")
|
||||
.and_then(|v| v.as_str())
|
||||
.or_else(|| token_data.get("team_name").and_then(|v| v.as_str()));
|
||||
// Save tokens (access + refresh + scopes)
|
||||
oauth_defaults::store_oauth_tokens(
|
||||
store,
|
||||
user_id,
|
||||
&auth.secret_name,
|
||||
auth.provider.as_deref(),
|
||||
&token_response.access_token,
|
||||
token_response.refresh_token.as_deref(),
|
||||
token_response.expires_in,
|
||||
&oauth.scopes,
|
||||
)
|
||||
.await?;
|
||||
|
||||
println!();
|
||||
println!(" ✓ {} connected!", display_name);
|
||||
if let Some(workspace) = workspace_name {
|
||||
println!(" Workspace: {}", workspace);
|
||||
}
|
||||
println!();
|
||||
println!(" The tool can now access the API.");
|
||||
println!();
|
||||
@@ -1107,46 +1030,15 @@ async fn validate_token(
|
||||
validation: &crate::tools::wasm::ValidationEndpointSchema,
|
||||
_secret_name: &str,
|
||||
) -> anyhow::Result<()> {
|
||||
let client = reqwest::Client::builder()
|
||||
.timeout(std::time::Duration::from_secs(10))
|
||||
.build()?;
|
||||
|
||||
// Build request based on method
|
||||
let request = match validation.method.to_uppercase().as_str() {
|
||||
"GET" => client.get(&validation.url),
|
||||
"POST" => client.post(&validation.url),
|
||||
_ => client.get(&validation.url),
|
||||
};
|
||||
|
||||
// Add authorization header (assume Bearer for now, could be extended)
|
||||
let response = request
|
||||
.header("Authorization", format!("Bearer {}", token))
|
||||
.header("Notion-Version", "2022-06-28") // Notion-specific, but harmless for others
|
||||
.send()
|
||||
.await?;
|
||||
|
||||
if response.status().as_u16() == validation.success_status {
|
||||
Ok(())
|
||||
} else {
|
||||
let status = response.status();
|
||||
let body = response.text().await.unwrap_or_default();
|
||||
Err(anyhow::anyhow!(
|
||||
"HTTP {} (expected {}): {}",
|
||||
status,
|
||||
validation.success_status,
|
||||
if body.len() > 100 {
|
||||
format!("{}...", &body[..100])
|
||||
} else {
|
||||
body
|
||||
}
|
||||
))
|
||||
}
|
||||
crate::cli::oauth_defaults::validate_oauth_token(token, validation)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("{}", e))
|
||||
}
|
||||
|
||||
/// Save token to secrets store.
|
||||
///
|
||||
/// Optionally stores a refresh token (as `{secret_name}_refresh_token`) and
|
||||
/// sets `expires_at` on the access token so the runtime can auto-refresh.
|
||||
/// Delegates to the shared `store_oauth_tokens` for OAuth tokens, or stores
|
||||
/// directly for manual/env-var tokens (no scopes or refresh token).
|
||||
async fn save_token(
|
||||
store: &(dyn SecretsStore + Send + Sync),
|
||||
user_id: &str,
|
||||
@@ -1155,36 +1047,18 @@ async fn save_token(
|
||||
refresh_token: Option<&str>,
|
||||
expires_in: Option<u64>,
|
||||
) -> anyhow::Result<()> {
|
||||
let mut params = CreateSecretParams::new(&auth.secret_name, token);
|
||||
|
||||
if let Some(ref provider) = auth.provider {
|
||||
params = params.with_provider(provider);
|
||||
}
|
||||
|
||||
if let Some(secs) = expires_in {
|
||||
let expires_at = chrono::Utc::now() + chrono::Duration::seconds(secs as i64);
|
||||
params = params.with_expiry(expires_at);
|
||||
}
|
||||
|
||||
store
|
||||
.create(user_id, params)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("Failed to save token: {}", e))?;
|
||||
|
||||
// Store refresh token separately (no expiry, it's long-lived)
|
||||
if let Some(rt) = refresh_token {
|
||||
let refresh_name = format!("{}_refresh_token", auth.secret_name);
|
||||
let mut refresh_params = CreateSecretParams::new(&refresh_name, rt);
|
||||
if let Some(ref provider) = auth.provider {
|
||||
refresh_params = refresh_params.with_provider(provider);
|
||||
}
|
||||
store
|
||||
.create(user_id, refresh_params)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("Failed to save refresh token: {}", e))?;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
crate::cli::oauth_defaults::store_oauth_tokens(
|
||||
store,
|
||||
user_id,
|
||||
&auth.secret_name,
|
||||
auth.provider.as_deref(),
|
||||
token,
|
||||
refresh_token,
|
||||
expires_in,
|
||||
&[], // No scopes for manual/env-var tokens
|
||||
)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("{}", e))
|
||||
}
|
||||
|
||||
/// Print success message.
|
||||
|
||||
@@ -30,6 +30,26 @@ pub struct AgentConfig {
|
||||
}
|
||||
|
||||
impl AgentConfig {
|
||||
/// Create a test-friendly config without reading env vars.
|
||||
#[cfg(feature = "libsql")]
|
||||
pub fn for_testing() -> Self {
|
||||
Self {
|
||||
name: "test-rig".to_string(),
|
||||
max_parallel_jobs: 1,
|
||||
job_timeout: Duration::from_secs(30),
|
||||
stuck_threshold: Duration::from_secs(300),
|
||||
repair_check_interval: Duration::from_secs(3600),
|
||||
max_repair_attempts: 0,
|
||||
use_planning: false,
|
||||
session_idle_timeout: Duration::from_secs(3600),
|
||||
allow_local_tools: true,
|
||||
max_cost_per_day_cents: None,
|
||||
max_actions_per_hour: None,
|
||||
max_tool_iterations: 10,
|
||||
auto_approve_tools: true,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn resolve(settings: &Settings) -> Result<Self, ConfigError> {
|
||||
Ok(Self {
|
||||
name: parse_optional_env("AGENT_NAME", settings.agent.name.clone())?,
|
||||
|
||||
+18
-10
@@ -1,3 +1,4 @@
|
||||
use std::collections::HashMap;
|
||||
use std::path::PathBuf;
|
||||
|
||||
use secrecy::SecretString;
|
||||
@@ -18,8 +19,9 @@ pub struct ChannelsConfig {
|
||||
pub wasm_channels_dir: std::path::PathBuf,
|
||||
/// Whether WASM channels are enabled.
|
||||
pub wasm_channels_enabled: bool,
|
||||
/// Telegram owner user ID. When set, the bot only responds to this user.
|
||||
pub telegram_owner_id: Option<i64>,
|
||||
/// Per-channel owner user IDs. When set, the channel only responds to this user.
|
||||
/// Key: channel name (e.g., "telegram"), Value: owner user ID.
|
||||
pub wasm_channel_owner_ids: HashMap<String, i64>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
@@ -180,14 +182,20 @@ impl ChannelsConfig {
|
||||
.map(PathBuf::from)
|
||||
.unwrap_or_else(default_channels_dir),
|
||||
wasm_channels_enabled: parse_bool_env("WASM_CHANNELS_ENABLED", true)?,
|
||||
telegram_owner_id: optional_env("TELEGRAM_OWNER_ID")?
|
||||
.map(|s| s.parse())
|
||||
.transpose()
|
||||
.map_err(|e: std::num::ParseIntError| ConfigError::InvalidValue {
|
||||
key: "TELEGRAM_OWNER_ID".to_string(),
|
||||
message: format!("must be an integer: {e}"),
|
||||
})?
|
||||
.or(settings.channels.telegram_owner_id),
|
||||
wasm_channel_owner_ids: {
|
||||
let mut ids = settings.channels.wasm_channel_owner_ids.clone();
|
||||
// Backwards compat: TELEGRAM_OWNER_ID env var
|
||||
if let Some(id_str) = optional_env("TELEGRAM_OWNER_ID")? {
|
||||
let id: i64 = id_str.parse().map_err(|e: std::num::ParseIntError| {
|
||||
ConfigError::InvalidValue {
|
||||
key: "TELEGRAM_OWNER_ID".to_string(),
|
||||
message: format!("must be an integer: {e}"),
|
||||
}
|
||||
})?;
|
||||
ids.insert("telegram".to_string(), id);
|
||||
}
|
||||
ids
|
||||
},
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -195,6 +195,40 @@ pub struct NearAiConfig {
|
||||
}
|
||||
|
||||
impl LlmConfig {
|
||||
/// Create a test-friendly config without reading env vars.
|
||||
///
|
||||
/// Uses NearAi backend with dummy values. The LLM provider is replaced
|
||||
/// by `TraceLlm` via `AppBuilder::with_llm()`, so these values are unused.
|
||||
#[cfg(feature = "libsql")]
|
||||
pub fn for_testing() -> Self {
|
||||
Self {
|
||||
backend: LlmBackend::NearAi,
|
||||
nearai: NearAiConfig {
|
||||
model: "test-model".to_string(),
|
||||
cheap_model: None,
|
||||
base_url: "http://localhost:0".to_string(),
|
||||
auth_base_url: "http://localhost:0".to_string(),
|
||||
session_path: PathBuf::from("/tmp/ironclaw-test-session.json"),
|
||||
api_key: None,
|
||||
fallback_model: None,
|
||||
max_retries: 0,
|
||||
circuit_breaker_threshold: None,
|
||||
circuit_breaker_recovery_secs: 30,
|
||||
response_cache_enabled: false,
|
||||
response_cache_ttl_secs: 3600,
|
||||
response_cache_max_entries: 100,
|
||||
failover_cooldown_secs: 300,
|
||||
failover_cooldown_threshold: 3,
|
||||
smart_routing_cascade: false,
|
||||
},
|
||||
openai: None,
|
||||
anthropic: None,
|
||||
ollama: None,
|
||||
openai_compatible: None,
|
||||
tinfoil: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Resolve a model name from env var → settings.selected_model → hardcoded default.
|
||||
fn resolve_model(
|
||||
env_var: &str,
|
||||
|
||||
@@ -78,6 +78,77 @@ pub struct Config {
|
||||
}
|
||||
|
||||
impl Config {
|
||||
/// Create a full Config for integration tests without reading env vars.
|
||||
///
|
||||
/// Requires the `libsql` feature. Sets up:
|
||||
/// - libSQL database at the given path
|
||||
/// - WASM and embeddings disabled
|
||||
/// - Skills enabled with the given directories
|
||||
/// - Heartbeat, routines, sandbox, builder all disabled
|
||||
/// - Safety with injection check off, 100k output limit
|
||||
#[cfg(feature = "libsql")]
|
||||
pub fn for_testing(
|
||||
libsql_path: std::path::PathBuf,
|
||||
skills_dir: std::path::PathBuf,
|
||||
installed_skills_dir: std::path::PathBuf,
|
||||
) -> Self {
|
||||
Self {
|
||||
database: DatabaseConfig {
|
||||
backend: DatabaseBackend::LibSql,
|
||||
url: secrecy::SecretString::from("unused://test".to_string()),
|
||||
pool_size: 1,
|
||||
ssl_mode: SslMode::Disable,
|
||||
libsql_path: Some(libsql_path),
|
||||
libsql_url: None,
|
||||
libsql_auth_token: None,
|
||||
},
|
||||
llm: LlmConfig::for_testing(),
|
||||
embeddings: EmbeddingsConfig::default(),
|
||||
tunnel: TunnelConfig::default(),
|
||||
channels: ChannelsConfig {
|
||||
cli: CliConfig { enabled: false },
|
||||
http: None,
|
||||
gateway: None,
|
||||
signal: None,
|
||||
wasm_channels_dir: std::path::PathBuf::from("/tmp/ironclaw-test-channels"),
|
||||
wasm_channels_enabled: false,
|
||||
wasm_channel_owner_ids: HashMap::new(),
|
||||
},
|
||||
agent: AgentConfig::for_testing(),
|
||||
safety: SafetyConfig {
|
||||
max_output_length: 100_000,
|
||||
injection_check_enabled: false,
|
||||
},
|
||||
wasm: WasmConfig {
|
||||
enabled: false,
|
||||
..WasmConfig::default()
|
||||
},
|
||||
secrets: SecretsConfig::default(),
|
||||
builder: BuilderModeConfig {
|
||||
enabled: false,
|
||||
..BuilderModeConfig::default()
|
||||
},
|
||||
heartbeat: HeartbeatConfig::default(),
|
||||
hygiene: HygieneConfig::default(),
|
||||
routines: RoutineConfig {
|
||||
enabled: false,
|
||||
..RoutineConfig::default()
|
||||
},
|
||||
sandbox: SandboxModeConfig {
|
||||
enabled: false,
|
||||
..SandboxModeConfig::default()
|
||||
},
|
||||
claude_code: ClaudeCodeConfig::default(),
|
||||
skills: SkillsConfig {
|
||||
enabled: true,
|
||||
local_dir: skills_dir,
|
||||
installed_dir: installed_skills_dir,
|
||||
..SkillsConfig::default()
|
||||
},
|
||||
observability: crate::observability::ObservabilityConfig::default(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Load configuration from environment variables and the database.
|
||||
///
|
||||
/// Priority: env var > TOML config file > DB settings > default.
|
||||
|
||||
@@ -9,6 +9,8 @@ use rust_decimal::Decimal;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::llm::recording::HttpInterceptor;
|
||||
|
||||
/// State of a job.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
@@ -146,6 +148,22 @@ pub struct JobContext {
|
||||
/// Wrapped in `Arc` for cheap cloning on every tool invocation.
|
||||
#[serde(skip)]
|
||||
pub extra_env: Arc<HashMap<String, String>>,
|
||||
/// Optional HTTP interceptor for trace recording/replay.
|
||||
///
|
||||
/// When set, tools that make outgoing HTTP requests should check this
|
||||
/// interceptor before sending real requests. During recording, the
|
||||
/// interceptor captures request/response pairs. During replay, it
|
||||
/// returns pre-recorded responses.
|
||||
#[serde(skip)]
|
||||
pub http_interceptor: Option<Arc<dyn HttpInterceptor>>,
|
||||
/// Stash of full tool outputs keyed by tool_call_id.
|
||||
///
|
||||
/// Tool outputs may be truncated before reaching the LLM context window,
|
||||
/// but subsequent tools (e.g., `json`) may need the full output. This
|
||||
/// stash stores the complete, unsanitized output so tools can reference
|
||||
/// previous results by ID via `$tool_call_id` parameter syntax.
|
||||
#[serde(skip)]
|
||||
pub tool_output_stash: Arc<tokio::sync::RwLock<HashMap<String, String>>>,
|
||||
}
|
||||
|
||||
impl JobContext {
|
||||
@@ -182,7 +200,9 @@ impl JobContext {
|
||||
repair_attempts: 0,
|
||||
transitions: Vec::new(),
|
||||
extra_env: Arc::new(HashMap::new()),
|
||||
http_interceptor: None,
|
||||
metadata: serde_json::Value::Null,
|
||||
tool_output_stash: Arc::new(tokio::sync::RwLock::new(HashMap::new())),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -117,6 +117,10 @@ impl JobStore for LibSqlBackend {
|
||||
transitions: Vec::new(),
|
||||
metadata: serde_json::Value::Null,
|
||||
extra_env: std::sync::Arc::new(std::collections::HashMap::new()),
|
||||
http_interceptor: None,
|
||||
tool_output_stash: std::sync::Arc::new(tokio::sync::RwLock::new(
|
||||
std::collections::HashMap::new(),
|
||||
)),
|
||||
}))
|
||||
}
|
||||
None => Ok(None),
|
||||
@@ -213,6 +217,30 @@ impl JobStore for LibSqlBackend {
|
||||
Ok(jobs)
|
||||
}
|
||||
|
||||
async fn get_agent_job_failure_reason(
|
||||
&self,
|
||||
id: Uuid,
|
||||
) -> Result<Option<String>, DatabaseError> {
|
||||
let conn = self.connect().await?;
|
||||
let mut rows = conn
|
||||
.query(
|
||||
"SELECT failure_reason FROM agent_jobs WHERE id = ?1",
|
||||
[id.to_string()],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?;
|
||||
|
||||
if let Some(row) = rows
|
||||
.next()
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?
|
||||
{
|
||||
Ok(get_opt_text(&row, 0))
|
||||
} else {
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
|
||||
async fn agent_job_summary(&self) -> Result<AgentJobSummary, DatabaseError> {
|
||||
let conn = self.connect().await?;
|
||||
let mut rows = conn
|
||||
|
||||
@@ -515,7 +515,7 @@ impl WorkspaceStore for LibSqlBackend {
|
||||
let mut rows = conn
|
||||
.query(
|
||||
r#"
|
||||
SELECT c.id, c.document_id, c.content
|
||||
SELECT c.id, c.document_id, d.path, c.content
|
||||
FROM memory_chunks_fts fts
|
||||
JOIN memory_chunks c ON c._rowid = fts.rowid
|
||||
JOIN memory_documents d ON d.id = c.document_id
|
||||
@@ -542,7 +542,8 @@ impl WorkspaceStore for LibSqlBackend {
|
||||
results.push(RankedResult {
|
||||
chunk_id: get_text(&row, 0).parse().unwrap_or_default(),
|
||||
document_id: get_text(&row, 1).parse().unwrap_or_default(),
|
||||
content: get_text(&row, 2),
|
||||
document_path: get_text(&row, 2),
|
||||
content: get_text(&row, 3),
|
||||
rank: results.len() as u32 + 1,
|
||||
});
|
||||
}
|
||||
@@ -563,7 +564,7 @@ impl WorkspaceStore for LibSqlBackend {
|
||||
let mut rows = conn
|
||||
.query(
|
||||
r#"
|
||||
SELECT c.id, c.document_id, c.content
|
||||
SELECT c.id, c.document_id, d.path, c.content
|
||||
FROM vector_top_k('idx_memory_chunks_embedding', vector(?1), ?2) AS top_k
|
||||
JOIN memory_chunks c ON c._rowid = top_k.id
|
||||
JOIN memory_documents d ON d.id = c.document_id
|
||||
@@ -587,7 +588,8 @@ impl WorkspaceStore for LibSqlBackend {
|
||||
results.push(RankedResult {
|
||||
chunk_id: get_text(&row, 0).parse().unwrap_or_default(),
|
||||
document_id: get_text(&row, 1).parse().unwrap_or_default(),
|
||||
content: get_text(&row, 2),
|
||||
document_path: get_text(&row, 2),
|
||||
content: get_text(&row, 3),
|
||||
rank: results.len() as u32 + 1,
|
||||
});
|
||||
}
|
||||
|
||||
@@ -298,6 +298,7 @@ CREATE TABLE IF NOT EXISTS wasm_tools (
|
||||
user_id TEXT NOT NULL,
|
||||
name TEXT NOT NULL,
|
||||
version TEXT NOT NULL DEFAULT '1.0.0',
|
||||
wit_version TEXT NOT NULL DEFAULT '0.1.0',
|
||||
description TEXT NOT NULL,
|
||||
wasm_binary BLOB NOT NULL,
|
||||
binary_hash BLOB NOT NULL,
|
||||
@@ -314,6 +315,24 @@ CREATE INDEX IF NOT EXISTS idx_wasm_tools_user ON wasm_tools(user_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_wasm_tools_name ON wasm_tools(user_id, name);
|
||||
CREATE INDEX IF NOT EXISTS idx_wasm_tools_status ON wasm_tools(status);
|
||||
|
||||
-- ==================== WASM Channel Extensions ====================
|
||||
|
||||
CREATE TABLE IF NOT EXISTS wasm_channels (
|
||||
id TEXT PRIMARY KEY,
|
||||
user_id TEXT NOT NULL,
|
||||
name TEXT NOT NULL,
|
||||
version TEXT NOT NULL DEFAULT '0.1.0',
|
||||
wit_version TEXT NOT NULL DEFAULT '0.1.0',
|
||||
description TEXT NOT NULL DEFAULT '',
|
||||
wasm_binary BLOB NOT NULL,
|
||||
binary_hash BLOB NOT NULL,
|
||||
capabilities_json TEXT NOT NULL DEFAULT '{}',
|
||||
status TEXT NOT NULL DEFAULT 'active',
|
||||
created_at TEXT NOT NULL DEFAULT (datetime('now')),
|
||||
updated_at TEXT NOT NULL DEFAULT (datetime('now')),
|
||||
UNIQUE (user_id, name)
|
||||
);
|
||||
|
||||
-- ==================== Tool Capabilities ====================
|
||||
|
||||
CREATE TABLE IF NOT EXISTS tool_capabilities (
|
||||
|
||||
@@ -177,6 +177,9 @@ pub trait JobStore: Send + Sync {
|
||||
async fn get_stuck_jobs(&self) -> Result<Vec<Uuid>, DatabaseError>;
|
||||
async fn list_agent_jobs(&self) -> Result<Vec<AgentJobRecord>, DatabaseError>;
|
||||
async fn agent_job_summary(&self) -> Result<AgentJobSummary, DatabaseError>;
|
||||
/// Get the failure reason for a single agent job (O(1) lookup).
|
||||
async fn get_agent_job_failure_reason(&self, id: Uuid)
|
||||
-> Result<Option<String>, DatabaseError>;
|
||||
async fn save_action(&self, job_id: Uuid, action: &ActionRecord) -> Result<(), DatabaseError>;
|
||||
async fn get_job_actions(&self, job_id: Uuid) -> Result<Vec<ActionRecord>, DatabaseError>;
|
||||
async fn record_llm_call(&self, record: &LlmCallRecord<'_>) -> Result<Uuid, DatabaseError>;
|
||||
|
||||
@@ -223,6 +223,13 @@ impl JobStore for PgBackend {
|
||||
self.store.agent_job_summary().await
|
||||
}
|
||||
|
||||
async fn get_agent_job_failure_reason(
|
||||
&self,
|
||||
id: Uuid,
|
||||
) -> Result<Option<String>, DatabaseError> {
|
||||
self.store.get_agent_job_failure_reason(id).await
|
||||
}
|
||||
|
||||
async fn save_action(&self, job_id: Uuid, action: &ActionRecord) -> Result<(), DatabaseError> {
|
||||
self.store.save_action(job_id, action).await
|
||||
}
|
||||
|
||||
+147
@@ -331,6 +331,9 @@ pub enum WorkspaceError {
|
||||
|
||||
#[error("Heartbeat error: {reason}")]
|
||||
HeartbeatError { reason: String },
|
||||
|
||||
#[error("I/O error: {reason}")]
|
||||
IoError { reason: String },
|
||||
}
|
||||
|
||||
/// Orchestrator errors (internal API, container management).
|
||||
@@ -419,3 +422,147 @@ pub enum RoutineError {
|
||||
|
||||
/// Result type alias for the agent.
|
||||
pub type Result<T> = std::result::Result<T, Error>;
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn config_error_display() {
|
||||
let err = ConfigError::MissingEnvVar("DATABASE_URL".to_string());
|
||||
let msg = err.to_string();
|
||||
assert!(
|
||||
msg.contains("DATABASE_URL"),
|
||||
"Should mention the variable name: {msg}"
|
||||
);
|
||||
|
||||
let err = ConfigError::MissingRequired {
|
||||
key: "llm.model".to_string(),
|
||||
hint: "Set LLM_MODEL env var".to_string(),
|
||||
};
|
||||
let msg = err.to_string();
|
||||
assert!(msg.contains("llm.model"), "Should mention the key: {msg}");
|
||||
assert!(
|
||||
msg.contains("Set LLM_MODEL"),
|
||||
"Should include the hint: {msg}"
|
||||
);
|
||||
|
||||
let err = ConfigError::InvalidValue {
|
||||
key: "port".to_string(),
|
||||
message: "must be a number".to_string(),
|
||||
};
|
||||
let msg = err.to_string();
|
||||
assert!(msg.contains("port"), "Should mention the key: {msg}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn database_error_display() {
|
||||
let err = DatabaseError::NotFound {
|
||||
entity: "conversation".to_string(),
|
||||
id: "abc-123".to_string(),
|
||||
};
|
||||
let msg = err.to_string();
|
||||
assert!(msg.contains("conversation"), "Should mention entity: {msg}");
|
||||
assert!(msg.contains("abc-123"), "Should mention id: {msg}");
|
||||
|
||||
let err = DatabaseError::Query("syntax error near SELECT".to_string());
|
||||
assert!(err.to_string().contains("syntax error"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn channel_error_display() {
|
||||
let err = ChannelError::StartupFailed {
|
||||
name: "telegram".to_string(),
|
||||
reason: "invalid token".to_string(),
|
||||
};
|
||||
let msg = err.to_string();
|
||||
assert!(msg.contains("telegram"), "Should mention channel: {msg}");
|
||||
assert!(
|
||||
msg.contains("invalid token"),
|
||||
"Should mention reason: {msg}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn llm_error_display() {
|
||||
let err = LlmError::ContextLengthExceeded {
|
||||
used: 100_000,
|
||||
limit: 50_000,
|
||||
};
|
||||
let msg = err.to_string();
|
||||
assert!(msg.contains("100000"), "Should mention used tokens: {msg}");
|
||||
assert!(msg.contains("50000"), "Should mention limit: {msg}");
|
||||
|
||||
let err = LlmError::RateLimited {
|
||||
provider: "openai".to_string(),
|
||||
retry_after: Some(Duration::from_secs(30)),
|
||||
};
|
||||
let msg = err.to_string();
|
||||
assert!(msg.contains("openai"), "Should mention provider: {msg}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn job_error_display() {
|
||||
let err = JobError::MaxJobsExceeded { max: 5 };
|
||||
let msg = err.to_string();
|
||||
assert!(msg.contains("5"), "Should mention max: {msg}");
|
||||
|
||||
let id = Uuid::new_v4();
|
||||
let err = JobError::NotFound { id };
|
||||
let msg = err.to_string();
|
||||
assert!(
|
||||
msg.contains(&id.to_string()),
|
||||
"Should mention job id: {msg}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn safety_error_display() {
|
||||
let err = SafetyError::InjectionDetected {
|
||||
pattern: "SYSTEM:".to_string(),
|
||||
};
|
||||
let msg = err.to_string();
|
||||
assert!(msg.contains("SYSTEM:"), "Should mention pattern: {msg}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn workspace_error_display() {
|
||||
let err = WorkspaceError::DocumentNotFound {
|
||||
doc_type: "notes".to_string(),
|
||||
user_id: "user1".to_string(),
|
||||
};
|
||||
let msg = err.to_string();
|
||||
assert!(msg.contains("notes"), "Should mention doc_type: {msg}");
|
||||
assert!(msg.contains("user1"), "Should mention user_id: {msg}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn routine_error_display() {
|
||||
let err = RoutineError::InvalidCron {
|
||||
reason: "bad format".to_string(),
|
||||
};
|
||||
let msg = err.to_string();
|
||||
assert!(msg.contains("bad format"), "Should mention reason: {msg}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn top_level_error_from_conversions() {
|
||||
let config_err = ConfigError::MissingEnvVar("TEST".to_string());
|
||||
let err: Error = config_err.into();
|
||||
assert!(matches!(err, Error::Config(_)));
|
||||
|
||||
let db_err = DatabaseError::Query("test".to_string());
|
||||
let err: Error = db_err.into();
|
||||
assert!(matches!(err, Error::Database(_)));
|
||||
|
||||
let job_err = JobError::MaxJobsExceeded { max: 1 };
|
||||
let err: Error = job_err.into();
|
||||
assert!(matches!(err, Error::Job(_)));
|
||||
|
||||
let safety_err = SafetyError::ValidationFailed {
|
||||
reason: "test".to_string(),
|
||||
};
|
||||
let err: Error = safety_err.into();
|
||||
assert!(matches!(err, Error::Safety(_)));
|
||||
}
|
||||
}
|
||||
|
||||
+920
-290
File diff suppressed because it is too large
Load Diff
+382
-18
@@ -24,6 +24,7 @@ pub use discovery::OnlineDiscovery;
|
||||
pub use manager::ExtensionManager;
|
||||
pub use registry::ExtensionRegistry;
|
||||
|
||||
use serde::ser::SerializeMap;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
/// The kind of extension, determining how it's installed, authenticated, and activated.
|
||||
@@ -145,28 +146,267 @@ pub struct InstallResult {
|
||||
pub message: String,
|
||||
}
|
||||
|
||||
/// Auth readiness state for the extensions list UI.
|
||||
///
|
||||
/// Used by `check_tool_auth_status` and `check_channel_auth_status` to
|
||||
/// communicate a tool's credential state to the list handler without
|
||||
/// ambiguous `(bool, bool)` tuples.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum ToolAuthState {
|
||||
/// Token/credentials are present — ready to use.
|
||||
Ready,
|
||||
/// Auth section exists but the access token is missing (OAuth not completed).
|
||||
NeedsAuth,
|
||||
/// Setup credentials (client_id/secret) must be configured before OAuth can start.
|
||||
NeedsSetup,
|
||||
/// No auth configuration at all (no capabilities or auth section).
|
||||
NoAuth,
|
||||
}
|
||||
|
||||
/// The typed auth status, carrying only the data relevant to each state.
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum AuthStatus {
|
||||
/// Authentication is complete; no further action needed.
|
||||
Authenticated,
|
||||
/// No authentication is required for this extension.
|
||||
NoAuthRequired,
|
||||
/// OAuth flow started — user must open `auth_url` in their browser.
|
||||
AwaitingAuthorization {
|
||||
auth_url: String,
|
||||
callback_type: String,
|
||||
},
|
||||
/// Waiting for user to provide a token/key manually.
|
||||
AwaitingToken {
|
||||
instructions: String,
|
||||
setup_url: Option<String>,
|
||||
},
|
||||
/// OAuth client credentials need to be configured before auth can proceed.
|
||||
NeedsSetup {
|
||||
instructions: String,
|
||||
setup_url: Option<String>,
|
||||
},
|
||||
}
|
||||
|
||||
impl AuthStatus {
|
||||
/// The wire-format status string (backward-compatible with JS consumers).
|
||||
pub fn as_str(&self) -> &'static str {
|
||||
match self {
|
||||
AuthStatus::Authenticated => "authenticated",
|
||||
AuthStatus::NoAuthRequired => "no_auth_required",
|
||||
AuthStatus::AwaitingAuthorization { .. } => "awaiting_authorization",
|
||||
AuthStatus::AwaitingToken { .. } => "awaiting_token",
|
||||
AuthStatus::NeedsSetup { .. } => "needs_setup",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Result of authenticating an extension.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct AuthResult {
|
||||
pub name: String,
|
||||
pub kind: ExtensionKind,
|
||||
/// OAuth URL to open (for OAuth flows).
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub auth_url: Option<String>,
|
||||
/// Whether using local or remote callback.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub callback_type: Option<String>,
|
||||
/// Instructions for manual token entry (for WASM tools).
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub instructions: Option<String>,
|
||||
/// URL for manual token setup.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub setup_url: Option<String>,
|
||||
/// Whether the tool is waiting for a token from the user.
|
||||
#[serde(default)]
|
||||
pub awaiting_token: bool,
|
||||
/// Current auth status.
|
||||
pub status: String,
|
||||
pub status: AuthStatus,
|
||||
}
|
||||
|
||||
impl AuthResult {
|
||||
// ── Constructors ──────────────────────────────────────────────────
|
||||
|
||||
pub fn authenticated(name: impl Into<String>, kind: ExtensionKind) -> Self {
|
||||
Self {
|
||||
name: name.into(),
|
||||
kind,
|
||||
status: AuthStatus::Authenticated,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn no_auth_required(name: impl Into<String>, kind: ExtensionKind) -> Self {
|
||||
Self {
|
||||
name: name.into(),
|
||||
kind,
|
||||
status: AuthStatus::NoAuthRequired,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn awaiting_authorization(
|
||||
name: impl Into<String>,
|
||||
kind: ExtensionKind,
|
||||
auth_url: String,
|
||||
callback_type: String,
|
||||
) -> Self {
|
||||
Self {
|
||||
name: name.into(),
|
||||
kind,
|
||||
status: AuthStatus::AwaitingAuthorization {
|
||||
auth_url,
|
||||
callback_type,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
pub fn awaiting_token(
|
||||
name: impl Into<String>,
|
||||
kind: ExtensionKind,
|
||||
instructions: String,
|
||||
setup_url: Option<String>,
|
||||
) -> Self {
|
||||
Self {
|
||||
name: name.into(),
|
||||
kind,
|
||||
status: AuthStatus::AwaitingToken {
|
||||
instructions,
|
||||
setup_url,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
pub fn needs_setup(
|
||||
name: impl Into<String>,
|
||||
kind: ExtensionKind,
|
||||
instructions: String,
|
||||
setup_url: Option<String>,
|
||||
) -> Self {
|
||||
Self {
|
||||
name: name.into(),
|
||||
kind,
|
||||
status: AuthStatus::NeedsSetup {
|
||||
instructions,
|
||||
setup_url,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// ── Accessors ─────────────────────────────────────────────────────
|
||||
|
||||
pub fn is_authenticated(&self) -> bool {
|
||||
matches!(self.status, AuthStatus::Authenticated)
|
||||
}
|
||||
|
||||
pub fn auth_url(&self) -> Option<&str> {
|
||||
match &self.status {
|
||||
AuthStatus::AwaitingAuthorization { auth_url, .. } => Some(auth_url),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn callback_type(&self) -> Option<&str> {
|
||||
match &self.status {
|
||||
AuthStatus::AwaitingAuthorization { callback_type, .. } => Some(callback_type),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn instructions(&self) -> Option<&str> {
|
||||
match &self.status {
|
||||
AuthStatus::AwaitingToken { instructions, .. }
|
||||
| AuthStatus::NeedsSetup { instructions, .. } => Some(instructions),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn setup_url(&self) -> Option<&str> {
|
||||
match &self.status {
|
||||
AuthStatus::AwaitingToken { setup_url, .. }
|
||||
| AuthStatus::NeedsSetup { setup_url, .. } => setup_url.as_deref(),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn is_awaiting_token(&self) -> bool {
|
||||
matches!(self.status, AuthStatus::AwaitingToken { .. })
|
||||
}
|
||||
|
||||
pub fn status_str(&self) -> &'static str {
|
||||
self.status.as_str()
|
||||
}
|
||||
}
|
||||
|
||||
/// Serialize `AuthResult` to the same flat JSON shape the JS frontend expects.
|
||||
impl Serialize for AuthResult {
|
||||
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
|
||||
// Count fields: name + kind + status + optional fields
|
||||
let optional_count = self.auth_url().is_some() as usize
|
||||
+ self.callback_type().is_some() as usize
|
||||
+ self.instructions().is_some() as usize
|
||||
+ self.setup_url().is_some() as usize;
|
||||
let mut map = serializer.serialize_map(Some(4 + optional_count))?;
|
||||
|
||||
map.serialize_entry("name", &self.name)?;
|
||||
map.serialize_entry("kind", &self.kind)?;
|
||||
if let Some(url) = self.auth_url() {
|
||||
map.serialize_entry("auth_url", url)?;
|
||||
}
|
||||
if let Some(cb) = self.callback_type() {
|
||||
map.serialize_entry("callback_type", cb)?;
|
||||
}
|
||||
if let Some(inst) = self.instructions() {
|
||||
map.serialize_entry("instructions", inst)?;
|
||||
}
|
||||
if let Some(url) = self.setup_url() {
|
||||
map.serialize_entry("setup_url", url)?;
|
||||
}
|
||||
map.serialize_entry("awaiting_token", &self.is_awaiting_token())?;
|
||||
map.serialize_entry("status", self.status_str())?;
|
||||
map.end()
|
||||
}
|
||||
}
|
||||
|
||||
/// Deserialize from the flat JSON shape back into the typed enum.
|
||||
impl<'de> Deserialize<'de> for AuthResult {
|
||||
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
|
||||
/// Flat helper matching the old JSON shape.
|
||||
#[derive(Deserialize)]
|
||||
#[allow(dead_code)]
|
||||
struct Raw {
|
||||
name: String,
|
||||
kind: ExtensionKind,
|
||||
#[serde(default)]
|
||||
auth_url: Option<String>,
|
||||
#[serde(default)]
|
||||
callback_type: Option<String>,
|
||||
#[serde(default)]
|
||||
instructions: Option<String>,
|
||||
#[serde(default)]
|
||||
setup_url: Option<String>,
|
||||
#[serde(default)]
|
||||
awaiting_token: bool,
|
||||
status: String,
|
||||
}
|
||||
|
||||
let raw = Raw::deserialize(deserializer)?;
|
||||
let status = match raw.status.as_str() {
|
||||
"authenticated" => AuthStatus::Authenticated,
|
||||
"no_auth_required" => AuthStatus::NoAuthRequired,
|
||||
"awaiting_authorization" => AuthStatus::AwaitingAuthorization {
|
||||
auth_url: raw.auth_url.unwrap_or_default(),
|
||||
callback_type: raw.callback_type.unwrap_or_default(),
|
||||
},
|
||||
"awaiting_token" => AuthStatus::AwaitingToken {
|
||||
instructions: raw.instructions.unwrap_or_default(),
|
||||
setup_url: raw.setup_url,
|
||||
},
|
||||
"needs_setup" => AuthStatus::NeedsSetup {
|
||||
instructions: raw.instructions.unwrap_or_default(),
|
||||
setup_url: raw.setup_url,
|
||||
},
|
||||
other => {
|
||||
return Err(serde::de::Error::unknown_variant(
|
||||
other,
|
||||
&[
|
||||
"authenticated",
|
||||
"no_auth_required",
|
||||
"awaiting_authorization",
|
||||
"awaiting_token",
|
||||
"needs_setup",
|
||||
],
|
||||
));
|
||||
}
|
||||
};
|
||||
Ok(AuthResult {
|
||||
name: raw.name,
|
||||
kind: raw.kind,
|
||||
status,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// Result of activating an extension.
|
||||
@@ -204,6 +444,9 @@ pub struct InstalledExtension {
|
||||
/// Whether this extension has a setup schema (required_secrets) that can be configured.
|
||||
#[serde(default)]
|
||||
pub needs_setup: bool,
|
||||
/// Whether this extension has an auth configuration (OAuth or manual token).
|
||||
#[serde(default)]
|
||||
pub has_auth: bool,
|
||||
/// Whether this extension is installed locally (false = available in registry but not installed).
|
||||
#[serde(default = "default_true")]
|
||||
pub installed: bool,
|
||||
@@ -254,3 +497,124 @@ pub enum ExtensionError {
|
||||
#[error("{0}")]
|
||||
Other(String),
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn auth_result_authenticated_round_trip() {
|
||||
let result = AuthResult::authenticated("gmail", ExtensionKind::WasmTool);
|
||||
let json = serde_json::to_value(&result).unwrap();
|
||||
|
||||
assert_eq!(json["status"], "authenticated");
|
||||
assert_eq!(json["name"], "gmail");
|
||||
assert_eq!(json["kind"], "wasm_tool");
|
||||
assert_eq!(json["awaiting_token"], false);
|
||||
assert!(json.get("auth_url").is_none());
|
||||
assert!(json.get("instructions").is_none());
|
||||
|
||||
let back: AuthResult = serde_json::from_value(json).unwrap();
|
||||
assert!(back.is_authenticated());
|
||||
assert!(back.auth_url().is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn auth_result_awaiting_authorization_round_trip() {
|
||||
let result = AuthResult::awaiting_authorization(
|
||||
"google-drive",
|
||||
ExtensionKind::WasmTool,
|
||||
"https://accounts.google.com/o/oauth2/v2/auth?state=abc".to_string(),
|
||||
"local".to_string(),
|
||||
);
|
||||
let json = serde_json::to_value(&result).unwrap();
|
||||
|
||||
assert_eq!(json["status"], "awaiting_authorization");
|
||||
assert_eq!(
|
||||
json["auth_url"],
|
||||
"https://accounts.google.com/o/oauth2/v2/auth?state=abc"
|
||||
);
|
||||
assert_eq!(json["callback_type"], "local");
|
||||
assert_eq!(json["awaiting_token"], false);
|
||||
|
||||
let back: AuthResult = serde_json::from_value(json).unwrap();
|
||||
assert_eq!(
|
||||
back.auth_url(),
|
||||
Some("https://accounts.google.com/o/oauth2/v2/auth?state=abc")
|
||||
);
|
||||
assert_eq!(back.callback_type(), Some("local"));
|
||||
assert!(!back.is_authenticated());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn auth_result_awaiting_token_round_trip() {
|
||||
let result = AuthResult::awaiting_token(
|
||||
"telegram",
|
||||
ExtensionKind::WasmChannel,
|
||||
"Enter your bot token".to_string(),
|
||||
None,
|
||||
);
|
||||
let json = serde_json::to_value(&result).unwrap();
|
||||
|
||||
assert_eq!(json["status"], "awaiting_token");
|
||||
assert_eq!(json["instructions"], "Enter your bot token");
|
||||
assert_eq!(json["awaiting_token"], true);
|
||||
assert!(json.get("auth_url").is_none());
|
||||
|
||||
let back: AuthResult = serde_json::from_value(json).unwrap();
|
||||
assert!(back.is_awaiting_token());
|
||||
assert_eq!(back.instructions(), Some("Enter your bot token"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn auth_result_needs_setup_round_trip() {
|
||||
let result = AuthResult::needs_setup(
|
||||
"custom-tool",
|
||||
ExtensionKind::WasmTool,
|
||||
"Configure OAuth credentials in the Setup tab.".to_string(),
|
||||
Some("https://console.cloud.google.com".to_string()),
|
||||
);
|
||||
let json = serde_json::to_value(&result).unwrap();
|
||||
|
||||
assert_eq!(json["status"], "needs_setup");
|
||||
assert_eq!(json["setup_url"], "https://console.cloud.google.com");
|
||||
assert_eq!(json["awaiting_token"], false);
|
||||
|
||||
let back: AuthResult = serde_json::from_value(json).unwrap();
|
||||
assert!(!back.is_authenticated());
|
||||
assert!(!back.is_awaiting_token());
|
||||
assert_eq!(back.setup_url(), Some("https://console.cloud.google.com"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn auth_result_no_auth_required_round_trip() {
|
||||
let result = AuthResult::no_auth_required("echo", ExtensionKind::WasmTool);
|
||||
let json = serde_json::to_value(&result).unwrap();
|
||||
|
||||
assert_eq!(json["status"], "no_auth_required");
|
||||
assert_eq!(json["awaiting_token"], false);
|
||||
|
||||
let back: AuthResult = serde_json::from_value(json).unwrap();
|
||||
assert!(!back.is_authenticated());
|
||||
assert_eq!(back.status, AuthStatus::NoAuthRequired);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn auth_status_type_safety() {
|
||||
// AwaitingAuthorization always has auth_url
|
||||
let result = AuthResult::awaiting_authorization(
|
||||
"test",
|
||||
ExtensionKind::WasmTool,
|
||||
"https://example.com".to_string(),
|
||||
"local".to_string(),
|
||||
);
|
||||
assert!(result.auth_url().is_some());
|
||||
assert!(!result.is_awaiting_token());
|
||||
|
||||
// Authenticated never has auth_url
|
||||
let result = AuthResult::authenticated("test", ExtensionKind::WasmTool);
|
||||
assert!(result.auth_url().is_none());
|
||||
assert!(result.instructions().is_none());
|
||||
assert!(result.setup_url().is_none());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -237,6 +237,10 @@ impl Store {
|
||||
total_tokens_used: 0,
|
||||
max_tokens: 0,
|
||||
extra_env: std::sync::Arc::new(std::collections::HashMap::new()),
|
||||
http_interceptor: None,
|
||||
tool_output_stash: std::sync::Arc::new(tokio::sync::RwLock::new(
|
||||
std::collections::HashMap::new(),
|
||||
)),
|
||||
}))
|
||||
}
|
||||
None => Ok(None),
|
||||
@@ -821,6 +825,21 @@ impl Store {
|
||||
.collect())
|
||||
}
|
||||
|
||||
/// Get the failure reason for a single agent job.
|
||||
pub async fn get_agent_job_failure_reason(
|
||||
&self,
|
||||
id: Uuid,
|
||||
) -> Result<Option<String>, DatabaseError> {
|
||||
let conn = self.conn().await?;
|
||||
let row = conn
|
||||
.query_opt(
|
||||
"SELECT failure_reason FROM agent_jobs WHERE id = $1",
|
||||
&[&id],
|
||||
)
|
||||
.await?;
|
||||
Ok(row.and_then(|r| r.get::<_, Option<String>>("failure_reason")))
|
||||
}
|
||||
|
||||
/// Summary counts for agent (non-sandbox) jobs.
|
||||
pub async fn agent_job_summary(&self) -> Result<AgentJobSummary, DatabaseError> {
|
||||
let conn = self.conn().await?;
|
||||
|
||||
+19
-2
@@ -13,6 +13,7 @@ pub mod failover;
|
||||
mod nearai_chat;
|
||||
mod provider;
|
||||
mod reasoning;
|
||||
pub mod recording;
|
||||
pub mod response_cache;
|
||||
pub mod retry;
|
||||
mod rig_adapter;
|
||||
@@ -30,6 +31,7 @@ pub use reasoning::{
|
||||
ActionPlan, Reasoning, ReasoningContext, RespondOutput, RespondResult, SILENT_REPLY_TOKEN,
|
||||
TokenUsage, ToolSelection, is_silent_reply,
|
||||
};
|
||||
pub use recording::RecordingLlm;
|
||||
pub use response_cache::{CachedProvider, ResponseCacheConfig};
|
||||
pub use retry::{RetryConfig, RetryProvider};
|
||||
pub use rig_adapter::RigAdapter;
|
||||
@@ -314,7 +316,14 @@ pub fn create_cheap_llm_provider(
|
||||
pub fn build_provider_chain(
|
||||
config: &LlmConfig,
|
||||
session: Arc<SessionManager>,
|
||||
) -> Result<(Arc<dyn LlmProvider>, Option<Arc<dyn LlmProvider>>), LlmError> {
|
||||
) -> Result<
|
||||
(
|
||||
Arc<dyn LlmProvider>,
|
||||
Option<Arc<dyn LlmProvider>>,
|
||||
Option<Arc<RecordingLlm>>,
|
||||
),
|
||||
LlmError,
|
||||
> {
|
||||
let llm = create_llm_provider(config, session.clone())?;
|
||||
tracing::info!("LLM provider initialized: {}", llm.model_name());
|
||||
|
||||
@@ -427,13 +436,21 @@ pub fn build_provider_chain(
|
||||
llm
|
||||
};
|
||||
|
||||
// 6. Recording (trace capture for replay testing)
|
||||
let recording_handle = RecordingLlm::from_env(llm.clone());
|
||||
let llm: Arc<dyn LlmProvider> = if let Some(ref recorder) = recording_handle {
|
||||
Arc::clone(recorder) as Arc<dyn LlmProvider>
|
||||
} else {
|
||||
llm
|
||||
};
|
||||
|
||||
// Standalone cheap LLM for heartbeat/evaluation (not part of the chain)
|
||||
let cheap_llm = create_cheap_llm_provider(config, session)?;
|
||||
if let Some(ref cheap) = cheap_llm {
|
||||
tracing::info!("Cheap LLM provider initialized: {}", cheap.model_name());
|
||||
}
|
||||
|
||||
Ok((llm, cheap_llm))
|
||||
Ok((llm, cheap_llm, recording_handle))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
||||
+141
-4
@@ -199,6 +199,29 @@ impl NearAiChatProvider {
|
||||
})?;
|
||||
|
||||
let status = response.status();
|
||||
// Extract Retry-After header before consuming the response body.
|
||||
// Supports both delay-seconds (RFC 7231 §7.1.3) and HTTP-date formats.
|
||||
let retry_after_header = response
|
||||
.headers()
|
||||
.get("retry-after")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.and_then(|v| {
|
||||
// Try delay-seconds first (most common from API providers)
|
||||
if let Ok(secs) = v.trim().parse::<u64>() {
|
||||
return Some(std::time::Duration::from_secs(secs));
|
||||
}
|
||||
// Try HTTP-date (e.g. "Mon, 02 Mar 2026 18:00:00 GMT")
|
||||
if let Ok(dt) = chrono::DateTime::parse_from_rfc2822(v.trim()) {
|
||||
let now = chrono::Utc::now();
|
||||
let delta = dt.signed_duration_since(now);
|
||||
// Use max(0) so past/present dates yield Duration::ZERO
|
||||
// rather than None (which would cause an immediate retry).
|
||||
return Some(std::time::Duration::from_secs(
|
||||
delta.num_seconds().max(0) as u64
|
||||
));
|
||||
}
|
||||
None
|
||||
});
|
||||
let response_text = response.text().await.map_err(|e| LlmError::RequestFailed {
|
||||
provider: "nearai_chat".to_string(),
|
||||
reason: format!("Failed to read response body: {}", e),
|
||||
@@ -230,7 +253,7 @@ impl NearAiChatProvider {
|
||||
if status_code == 429 {
|
||||
return Err(LlmError::RateLimited {
|
||||
provider: "nearai_chat".to_string(),
|
||||
retry_after: None,
|
||||
retry_after: retry_after_header,
|
||||
});
|
||||
}
|
||||
|
||||
@@ -499,9 +522,6 @@ impl LlmProvider for NearAiChatProvider {
|
||||
reason: "No choices in response".to_string(),
|
||||
})?;
|
||||
|
||||
// Fall back to reasoning_content when content is null (e.g. GLM-5
|
||||
// returns its answer in reasoning_content instead of content).
|
||||
let content = choice.message.content.or(choice.message.reasoning_content);
|
||||
let tool_calls: Vec<ToolCall> = choice
|
||||
.message
|
||||
.tool_calls
|
||||
@@ -518,6 +538,18 @@ impl LlmProvider for NearAiChatProvider {
|
||||
})
|
||||
.collect();
|
||||
|
||||
// Fall back to reasoning_content when content is null (e.g. GLM-5
|
||||
// returns its answer in reasoning_content instead of content), but
|
||||
// only for final text responses. Tool-call responses often have
|
||||
// content: null + reasoning_content filled with chain-of-thought;
|
||||
// leaking that into conversation history inflates context and
|
||||
// confuses the model.
|
||||
let content = if tool_calls.is_empty() {
|
||||
choice.message.content.or(choice.message.reasoning_content)
|
||||
} else {
|
||||
choice.message.content
|
||||
};
|
||||
|
||||
let finish_reason = match choice.finish_reason.as_deref() {
|
||||
Some("stop") => FinishReason::Stop,
|
||||
Some("length") => FinishReason::Length,
|
||||
@@ -1262,4 +1294,109 @@ mod tests {
|
||||
assert_eq!(input, default_in);
|
||||
assert_eq!(output, default_out);
|
||||
}
|
||||
|
||||
/// Regression: reasoning_content must NOT leak into tool-call responses.
|
||||
#[test]
|
||||
fn test_reasoning_content_not_leaked_into_tool_call_response() {
|
||||
let response: ChatCompletionResponse = serde_json::from_value(serde_json::json!({
|
||||
"id": "chatcmpl-test",
|
||||
"choices": [{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": null,
|
||||
"reasoning_content": "Let me think about which tool to call...",
|
||||
"tool_calls": [{
|
||||
"id": "call_abc123",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "search",
|
||||
"arguments": "{\"query\":\"test\"}"
|
||||
}
|
||||
}]
|
||||
},
|
||||
"finish_reason": "tool_calls"
|
||||
}],
|
||||
"usage": { "prompt_tokens": 100, "completion_tokens": 50 }
|
||||
}))
|
||||
.unwrap();
|
||||
|
||||
let choice = response.choices.into_iter().next().unwrap();
|
||||
let tool_calls: Vec<ToolCall> = choice
|
||||
.message
|
||||
.tool_calls
|
||||
.unwrap_or_default()
|
||||
.into_iter()
|
||||
.map(|tc| {
|
||||
let arguments = serde_json::from_str(&tc.function.arguments)
|
||||
.unwrap_or(serde_json::Value::Object(Default::default()));
|
||||
ToolCall {
|
||||
id: tc.id,
|
||||
name: tc.function.name,
|
||||
arguments,
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
let content = if tool_calls.is_empty() {
|
||||
choice.message.content.or(choice.message.reasoning_content)
|
||||
} else {
|
||||
choice.message.content
|
||||
};
|
||||
|
||||
assert!(
|
||||
content.is_none(),
|
||||
"reasoning_content should NOT leak into tool-call responses, got: {:?}",
|
||||
content
|
||||
);
|
||||
assert_eq!(tool_calls.len(), 1);
|
||||
assert_eq!(tool_calls[0].name, "search");
|
||||
}
|
||||
|
||||
/// Regression: reasoning_content SHOULD be used as fallback for text responses.
|
||||
#[test]
|
||||
fn test_reasoning_content_used_for_text_response() {
|
||||
let response: ChatCompletionResponse = serde_json::from_value(serde_json::json!({
|
||||
"id": "chatcmpl-test",
|
||||
"choices": [{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": null,
|
||||
"reasoning_content": "The answer is 42."
|
||||
},
|
||||
"finish_reason": "stop"
|
||||
}],
|
||||
"usage": { "prompt_tokens": 50, "completion_tokens": 20 }
|
||||
}))
|
||||
.unwrap();
|
||||
|
||||
let choice = response.choices.into_iter().next().unwrap();
|
||||
let tool_calls: Vec<ToolCall> = choice
|
||||
.message
|
||||
.tool_calls
|
||||
.unwrap_or_default()
|
||||
.into_iter()
|
||||
.map(|tc| {
|
||||
let arguments = serde_json::from_str(&tc.function.arguments)
|
||||
.unwrap_or(serde_json::Value::Object(Default::default()));
|
||||
ToolCall {
|
||||
id: tc.id,
|
||||
name: tc.function.name,
|
||||
arguments,
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
let content = if tool_calls.is_empty() {
|
||||
choice.message.content.or(choice.message.reasoning_content)
|
||||
} else {
|
||||
choice.message.content
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
content,
|
||||
Some("The answer is 42.".to_string()),
|
||||
"reasoning_content should be used as fallback for text responses"
|
||||
);
|
||||
assert!(tool_calls.is_empty());
|
||||
}
|
||||
}
|
||||
|
||||
+170
-3
@@ -335,8 +335,9 @@ impl Reasoning {
|
||||
|
||||
let response = self.llm.complete(request).await?;
|
||||
|
||||
// Parse the plan from the response
|
||||
self.parse_plan(&response.content)
|
||||
// Clean reasoning model artifacts before parsing JSON
|
||||
let cleaned = clean_response(&response.content);
|
||||
self.parse_plan(&cleaned)
|
||||
}
|
||||
|
||||
/// Select the best tool for the current situation.
|
||||
@@ -429,7 +430,9 @@ Respond in JSON format:
|
||||
|
||||
let response = self.llm.complete(request).await?;
|
||||
|
||||
self.parse_evaluation(&response.content)
|
||||
// Clean reasoning model artifacts before parsing JSON
|
||||
let cleaned = clean_response(&response.content);
|
||||
self.parse_evaluation(&cleaned)
|
||||
}
|
||||
|
||||
/// Generate a response to a user message.
|
||||
@@ -689,6 +692,8 @@ Example:
|
||||
- If tools return empty or irrelevant results, answer with what you already know rather than retrying
|
||||
|
||||
## Tool Call Style
|
||||
- ALWAYS call tools via tool_calls — never just describe what you would do
|
||||
- If you say "let me fetch/check/look up X", you MUST include the actual tool call in the same response
|
||||
- Do not narrate routine, low-risk tool calls; just call the tool
|
||||
- Narrate only when it helps: multi-step work, sensitive actions, or when the user asks
|
||||
- For multi-step tasks, call independent tools in parallel when possible
|
||||
@@ -1131,6 +1136,51 @@ fn recover_tool_calls_from_content(
|
||||
}
|
||||
}
|
||||
|
||||
// Bracket format from flatten_tool_messages:
|
||||
// [Called tool `name` with arguments: {...}]
|
||||
{
|
||||
let mut remaining = content;
|
||||
while let Some(start) = remaining.find("[Called tool `") {
|
||||
let after_prefix = &remaining[start + "[Called tool `".len()..];
|
||||
let Some(backtick_end) = after_prefix.find('`') else {
|
||||
break;
|
||||
};
|
||||
let name = &after_prefix[..backtick_end];
|
||||
let after_name = &after_prefix[backtick_end + 1..];
|
||||
|
||||
if !tool_names.contains(name) {
|
||||
remaining = after_name;
|
||||
continue;
|
||||
}
|
||||
|
||||
// Look for " with arguments: " followed by JSON until "]"
|
||||
if let Some(args_start) = after_name.strip_prefix(" with arguments: ") {
|
||||
// Find the closing "]" — but the JSON itself may contain "]",
|
||||
// so find the last "]" on this logical line.
|
||||
if let Some(bracket_end) = args_start.rfind(']') {
|
||||
let args_str = &args_start[..bracket_end];
|
||||
let arguments = serde_json::from_str::<serde_json::Value>(args_str)
|
||||
.unwrap_or(serde_json::Value::Object(Default::default()));
|
||||
calls.push(ToolCall {
|
||||
id: format!("recovered_{}", calls.len()),
|
||||
name: name.to_string(),
|
||||
arguments,
|
||||
});
|
||||
remaining = &args_start[bracket_end + 1..];
|
||||
continue;
|
||||
}
|
||||
}
|
||||
|
||||
// No arguments or malformed — call with empty args
|
||||
calls.push(ToolCall {
|
||||
id: format!("recovered_{}", calls.len()),
|
||||
name: name.to_string(),
|
||||
arguments: serde_json::Value::Object(Default::default()),
|
||||
});
|
||||
remaining = after_name;
|
||||
}
|
||||
}
|
||||
|
||||
calls
|
||||
}
|
||||
|
||||
@@ -1174,10 +1224,39 @@ fn clean_response(text: &str) -> String {
|
||||
result = strip_pipe_tag(&result, tag);
|
||||
}
|
||||
|
||||
// 6b. Strip bracket-format inline tool calls: [Called tool `name` with arguments: {...}]
|
||||
result = strip_bracket_tool_calls(&result);
|
||||
|
||||
// 7. Collapse triple+ newlines, trim
|
||||
collapse_newlines(&result)
|
||||
}
|
||||
|
||||
/// Strip bracket-format inline tool calls produced by `flatten_tool_messages`.
|
||||
///
|
||||
/// Removes patterns like `[Called tool `name` with arguments: {...}]` from text
|
||||
/// so the user doesn't see raw tool call syntax when the model echoes it back.
|
||||
fn strip_bracket_tool_calls(text: &str) -> String {
|
||||
let mut result = String::with_capacity(text.len());
|
||||
let mut remaining = text;
|
||||
while let Some(start) = remaining.find("[Called tool `") {
|
||||
result.push_str(&remaining[..start]);
|
||||
let after = &remaining[start..];
|
||||
// Find the closing "]" for this bracket expression
|
||||
if let Some(end) = after.find("]\n").map(|i| i + 2).or_else(|| {
|
||||
// If it's at the end of the string, just find "]"
|
||||
after.rfind(']').map(|i| i + 1)
|
||||
}) {
|
||||
remaining = &after[end..];
|
||||
} else {
|
||||
// Malformed — keep the rest
|
||||
result.push_str(after);
|
||||
return result;
|
||||
}
|
||||
}
|
||||
result.push_str(remaining);
|
||||
result
|
||||
}
|
||||
|
||||
/// Tool-related tags stripped with simple string matching (no code-awareness needed).
|
||||
const TOOL_TAGS: &[&str] = &["tool_call", "function_call", "tool_calls"];
|
||||
|
||||
@@ -1216,8 +1295,15 @@ fn strip_thinking_tags_regex(text: &str, code_regions: &[CodeRegion]) -> String
|
||||
}
|
||||
|
||||
// Strict mode: if still inside an unclosed thinking tag, discard trailing text
|
||||
// BUT preserve any <final> block embedded in the discarded region
|
||||
if !in_thinking {
|
||||
result.push_str(&text[last_index..]);
|
||||
} else {
|
||||
let trailing = &text[last_index..];
|
||||
let trailing_regions = find_code_regions(trailing);
|
||||
if let Some(final_content) = extract_final_content(trailing, &trailing_regions) {
|
||||
result.push_str(&final_content);
|
||||
}
|
||||
}
|
||||
|
||||
result
|
||||
@@ -1841,4 +1927,85 @@ That's my plan."#;
|
||||
assert_eq!(calls.len(), 1);
|
||||
assert_eq!(calls[0].name, "tool_list");
|
||||
}
|
||||
|
||||
// ---- plan/evaluate bypass clean_response (Bug #564-2) ----
|
||||
|
||||
#[test]
|
||||
fn test_clean_response_strips_think_before_json_plan() {
|
||||
let raw = r#"<think>I need to plan the steps carefully...</think>{"steps": [{"description": "Step 1", "tool": "search", "expected_outcome": "results"}], "reasoning": "Simple plan"}"#;
|
||||
let cleaned = clean_response(raw);
|
||||
// After cleaning, the JSON should be parseable
|
||||
let json_str = extract_json(&cleaned).unwrap();
|
||||
let parsed: serde_json::Value = serde_json::from_str(json_str).unwrap();
|
||||
assert!(parsed.get("steps").is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_clean_response_strips_think_before_json_evaluation() {
|
||||
let raw = r#"<think>Let me evaluate whether this was successful...</think>{"success": true, "confidence": 0.95, "reasoning": "Task completed", "issues": [], "suggestions": []}"#;
|
||||
let cleaned = clean_response(raw);
|
||||
let json_str = extract_json(&cleaned).unwrap();
|
||||
let eval: SuccessEvaluation = serde_json::from_str(json_str).unwrap();
|
||||
assert!(eval.success);
|
||||
assert_eq!(eval.confidence, 0.95);
|
||||
}
|
||||
|
||||
// ---- Unclosed think before final (Bug #564-3) ----
|
||||
|
||||
#[test]
|
||||
fn test_unclosed_think_before_final() {
|
||||
assert_eq!(
|
||||
clean_response("<think>reasoning no close tag <final>actual answer</final>"),
|
||||
"actual answer"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_unclosed_thinking_before_final() {
|
||||
assert_eq!(
|
||||
clean_response("<thinking>long reasoning... <final>the real answer</final>"),
|
||||
"the real answer"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_unclosed_think_before_final_with_prefix() {
|
||||
assert_eq!(
|
||||
clean_response("Hello <think>reasoning <final>world</final>"),
|
||||
"Hello world"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_unclosed_think_no_final_still_discards() {
|
||||
assert_eq!(clean_response("Hello <thinking>this never closes"), "Hello");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_recover_bracket_format_tool_call() {
|
||||
let tools = make_tools(&["http"]);
|
||||
let content = "Let me try that. [Called tool `http` with arguments: {\"method\":\"GET\",\"url\":\"https://example.com\"}]";
|
||||
let calls = recover_tool_calls_from_content(content, &tools);
|
||||
assert_eq!(calls.len(), 1);
|
||||
assert_eq!(calls[0].name, "http");
|
||||
assert_eq!(calls[0].arguments["method"], "GET");
|
||||
assert_eq!(calls[0].arguments["url"], "https://example.com");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_recover_bracket_format_unknown_tool_ignored() {
|
||||
let tools = make_tools(&["http"]);
|
||||
let content = "[Called tool `unknown_tool` with arguments: {}]";
|
||||
let calls = recover_tool_calls_from_content(content, &tools);
|
||||
assert!(calls.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_clean_response_strips_bracket_tool_calls() {
|
||||
let input = "Let me fetch that.\n[Called tool `http` with arguments: {\"method\":\"GET\",\"url\":\"https://example.com\"}]\nHere are the results.";
|
||||
let cleaned = clean_response(input);
|
||||
assert!(!cleaned.contains("[Called tool"));
|
||||
assert!(cleaned.contains("Let me fetch that."));
|
||||
assert!(cleaned.contains("Here are the results."));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,917 @@
|
||||
//! Live trace recording mode.
|
||||
//!
|
||||
//! Wraps any [`LlmProvider`] and captures every LLM interaction into
|
||||
//! the trace fixture format used by `TraceLlm` for deterministic E2E
|
||||
//! testing. Recorded traces can be replayed later via `TraceLlm`.
|
||||
//!
|
||||
//! The trace includes:
|
||||
//! - **Memory snapshot**: workspace documents captured before the first LLM call
|
||||
//! - **HTTP exchanges**: all outgoing HTTP request/response pairs from tools
|
||||
//! - **Steps**: user inputs, LLM responses (text/tool_calls), and expected tool
|
||||
//! results for verifying tool output during replay
|
||||
//!
|
||||
//! Enable by setting `IRONCLAW_RECORD_TRACE=1` at runtime.
|
||||
|
||||
use std::collections::VecDeque;
|
||||
use std::path::PathBuf;
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use rust_decimal::Decimal;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tokio::sync::Mutex;
|
||||
|
||||
use crate::error::LlmError;
|
||||
use crate::llm::provider::{
|
||||
ChatMessage, CompletionRequest, CompletionResponse, LlmProvider, ModelMetadata, Role,
|
||||
ToolCompletionRequest, ToolCompletionResponse,
|
||||
};
|
||||
|
||||
// ── Trace format types ─────────────────────────────────────────────
|
||||
|
||||
/// Top-level trace file — extended format with memory snapshot and HTTP exchanges.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct TraceFile {
|
||||
pub model_name: String,
|
||||
/// Workspace memory documents captured before the recording session.
|
||||
/// Replay should restore these before running the trace.
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
pub memory_snapshot: Vec<MemorySnapshotEntry>,
|
||||
/// HTTP exchanges recorded during the session, in order.
|
||||
/// Replay should return these instead of making real HTTP requests.
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
pub http_exchanges: Vec<HttpExchange>,
|
||||
pub steps: Vec<TraceStep>,
|
||||
}
|
||||
|
||||
/// A memory document captured at recording start.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct MemorySnapshotEntry {
|
||||
pub path: String,
|
||||
pub content: String,
|
||||
}
|
||||
|
||||
/// A recorded HTTP request/response pair.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct HttpExchange {
|
||||
pub request: HttpExchangeRequest,
|
||||
pub response: HttpExchangeResponse,
|
||||
}
|
||||
|
||||
/// The request side of an HTTP exchange.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct HttpExchangeRequest {
|
||||
pub method: String,
|
||||
pub url: String,
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
pub headers: Vec<(String, String)>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub body: Option<String>,
|
||||
}
|
||||
|
||||
/// The response side of an HTTP exchange.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct HttpExchangeResponse {
|
||||
pub status: u16,
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
pub headers: Vec<(String, String)>,
|
||||
pub body: String,
|
||||
}
|
||||
|
||||
/// A single step in the trace.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct TraceStep {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub request_hint: Option<RequestHint>,
|
||||
pub response: TraceResponse,
|
||||
/// Tool results that appeared in the message context since the previous step.
|
||||
/// During replay, the test harness can compare actual tool results against
|
||||
/// these to verify tool output hasn't changed (regression detection).
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
pub expected_tool_results: Vec<ExpectedToolResult>,
|
||||
}
|
||||
|
||||
/// Soft validation hints for matching a step to a request.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct RequestHint {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub last_user_message_contains: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub min_message_count: Option<usize>,
|
||||
}
|
||||
|
||||
/// Tagged response enum — text, tool_calls, or user_input.
|
||||
///
|
||||
/// `user_input` steps are metadata markers — they record what the user said
|
||||
/// but do **not** correspond to an LLM call. During replay, `TraceLlm` must
|
||||
/// skip `user_input` steps and only consume `text`/`tool_calls` steps.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum TraceResponse {
|
||||
Text {
|
||||
content: String,
|
||||
input_tokens: u32,
|
||||
output_tokens: u32,
|
||||
},
|
||||
ToolCalls {
|
||||
tool_calls: Vec<TraceToolCall>,
|
||||
input_tokens: u32,
|
||||
output_tokens: u32,
|
||||
},
|
||||
/// Marker for a user message that triggered subsequent LLM calls.
|
||||
/// Not an LLM response — replay providers must skip these.
|
||||
UserInput { content: String },
|
||||
}
|
||||
|
||||
/// A tool call in a trace step.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct TraceToolCall {
|
||||
pub id: String,
|
||||
pub name: String,
|
||||
pub arguments: serde_json::Value,
|
||||
}
|
||||
|
||||
/// Recorded tool result for regression checking during replay.
|
||||
///
|
||||
/// During replay, after tools execute and before returning the canned LLM
|
||||
/// response, the test harness should compare actual `Role::Tool` messages
|
||||
/// against these entries. A content mismatch indicates a tool behavior change.
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ExpectedToolResult {
|
||||
pub tool_call_id: String,
|
||||
pub name: String,
|
||||
/// The full tool result content as it appeared in the message context.
|
||||
pub content: String,
|
||||
}
|
||||
|
||||
// ── HTTP interceptor ───────────────────────────────────────────────
|
||||
|
||||
/// Trait for intercepting HTTP requests from tools.
|
||||
///
|
||||
/// During recording, the interceptor captures exchanges after the real
|
||||
/// request completes. During replay, it short-circuits with a recorded response.
|
||||
#[async_trait]
|
||||
pub trait HttpInterceptor: Send + Sync + std::fmt::Debug {
|
||||
/// Called before making an HTTP request.
|
||||
///
|
||||
/// Return `Some(response)` to short-circuit (replay mode).
|
||||
/// Return `None` to let the real request proceed (recording mode).
|
||||
async fn before_request(&self, request: &HttpExchangeRequest) -> Option<HttpExchangeResponse>;
|
||||
|
||||
/// Called after a real HTTP request completes (recording mode only).
|
||||
async fn after_response(&self, request: &HttpExchangeRequest, response: &HttpExchangeResponse);
|
||||
}
|
||||
|
||||
/// Records HTTP exchanges during a live session.
|
||||
#[derive(Debug)]
|
||||
pub struct RecordingHttpInterceptor {
|
||||
exchanges: Mutex<Vec<HttpExchange>>,
|
||||
}
|
||||
|
||||
impl Default for RecordingHttpInterceptor {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
impl RecordingHttpInterceptor {
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
exchanges: Mutex::new(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
/// Return all recorded exchanges.
|
||||
pub async fn take_exchanges(&self) -> Vec<HttpExchange> {
|
||||
self.exchanges.lock().await.clone()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl HttpInterceptor for RecordingHttpInterceptor {
|
||||
async fn before_request(&self, _request: &HttpExchangeRequest) -> Option<HttpExchangeResponse> {
|
||||
// Recording mode: let the real request proceed
|
||||
None
|
||||
}
|
||||
|
||||
async fn after_response(&self, request: &HttpExchangeRequest, response: &HttpExchangeResponse) {
|
||||
self.exchanges.lock().await.push(HttpExchange {
|
||||
request: request.clone(),
|
||||
response: response.clone(),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
/// Replays recorded HTTP exchanges during test runs.
|
||||
///
|
||||
/// Returns responses in order. If more requests arrive than recorded
|
||||
/// exchanges, returns a 599 error response.
|
||||
#[derive(Debug)]
|
||||
pub struct ReplayingHttpInterceptor {
|
||||
exchanges: Mutex<VecDeque<HttpExchange>>,
|
||||
}
|
||||
|
||||
impl ReplayingHttpInterceptor {
|
||||
pub fn new(exchanges: Vec<HttpExchange>) -> Self {
|
||||
Self {
|
||||
exchanges: Mutex::new(VecDeque::from(exchanges)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl HttpInterceptor for ReplayingHttpInterceptor {
|
||||
async fn before_request(&self, request: &HttpExchangeRequest) -> Option<HttpExchangeResponse> {
|
||||
let mut queue = self.exchanges.lock().await;
|
||||
if let Some(exchange) = queue.pop_front() {
|
||||
// Soft-check: warn if the request doesn't match
|
||||
if exchange.request.url != request.url || exchange.request.method != request.method {
|
||||
tracing::warn!(
|
||||
expected_url = %exchange.request.url,
|
||||
actual_url = %request.url,
|
||||
expected_method = %exchange.request.method,
|
||||
actual_method = %request.method,
|
||||
"HTTP replay: request mismatch (returning recorded response anyway)"
|
||||
);
|
||||
}
|
||||
Some(exchange.response)
|
||||
} else {
|
||||
tracing::error!(
|
||||
url = %request.url,
|
||||
method = %request.method,
|
||||
"HTTP replay: no more recorded exchanges, returning error"
|
||||
);
|
||||
Some(HttpExchangeResponse {
|
||||
status: 599,
|
||||
headers: Vec::new(),
|
||||
body: "trace replay: no more recorded HTTP exchanges".to_string(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
async fn after_response(
|
||||
&self,
|
||||
_request: &HttpExchangeRequest,
|
||||
_response: &HttpExchangeResponse,
|
||||
) {
|
||||
// Replay mode: nothing to record
|
||||
}
|
||||
}
|
||||
|
||||
// ── RecordingLlm ───────────────────────────────────────────────────
|
||||
|
||||
/// LLM provider decorator that records interactions into a trace file.
|
||||
pub struct RecordingLlm {
|
||||
inner: Arc<dyn LlmProvider>,
|
||||
steps: Mutex<Vec<TraceStep>>,
|
||||
prev_message_count: Mutex<usize>,
|
||||
output_path: PathBuf,
|
||||
model_name: String,
|
||||
memory_snapshot: Mutex<Vec<MemorySnapshotEntry>>,
|
||||
http_interceptor: Arc<RecordingHttpInterceptor>,
|
||||
}
|
||||
|
||||
impl RecordingLlm {
|
||||
/// Wrap a provider for recording.
|
||||
pub fn new(inner: Arc<dyn LlmProvider>, output_path: PathBuf, model_name: String) -> Self {
|
||||
Self {
|
||||
inner,
|
||||
steps: Mutex::new(Vec::new()),
|
||||
prev_message_count: Mutex::new(0),
|
||||
output_path,
|
||||
model_name,
|
||||
memory_snapshot: Mutex::new(Vec::new()),
|
||||
http_interceptor: Arc::new(RecordingHttpInterceptor::new()),
|
||||
}
|
||||
}
|
||||
|
||||
/// Create from environment variables if recording is enabled.
|
||||
///
|
||||
/// - `IRONCLAW_RECORD_TRACE` — any non-empty value enables recording
|
||||
/// - `IRONCLAW_TRACE_OUTPUT` — file path (default: `./trace_{timestamp}.json`)
|
||||
/// - `IRONCLAW_TRACE_MODEL_NAME` — model_name field (default: `recorded-{inner.model_name()}`)
|
||||
pub fn from_env(inner: Arc<dyn LlmProvider>) -> Option<Arc<Self>> {
|
||||
let enabled = std::env::var("IRONCLAW_RECORD_TRACE")
|
||||
.ok()
|
||||
.filter(|v| !v.is_empty());
|
||||
enabled?;
|
||||
|
||||
let output_path = std::env::var("IRONCLAW_TRACE_OUTPUT")
|
||||
.ok()
|
||||
.filter(|v| !v.is_empty())
|
||||
.map(PathBuf::from)
|
||||
.unwrap_or_else(|| {
|
||||
let ts = chrono::Local::now().format("%Y%m%dT%H%M%S");
|
||||
PathBuf::from(format!("trace_{ts}.json"))
|
||||
});
|
||||
|
||||
let model_name = std::env::var("IRONCLAW_TRACE_MODEL_NAME")
|
||||
.ok()
|
||||
.filter(|v| !v.is_empty())
|
||||
.unwrap_or_else(|| format!("recorded-{}", inner.model_name()));
|
||||
|
||||
tracing::info!(
|
||||
output = %output_path.display(),
|
||||
model = %model_name,
|
||||
"LLM trace recording enabled"
|
||||
);
|
||||
|
||||
Some(Arc::new(Self::new(inner, output_path, model_name)))
|
||||
}
|
||||
|
||||
/// Get the HTTP interceptor for wiring into tools.
|
||||
///
|
||||
/// Pass this to `JobContext` or `HttpTool` so outgoing HTTP requests
|
||||
/// are recorded into the trace.
|
||||
pub fn http_interceptor(&self) -> Arc<dyn HttpInterceptor> {
|
||||
Arc::clone(&self.http_interceptor) as Arc<dyn HttpInterceptor>
|
||||
}
|
||||
|
||||
/// Snapshot all memory documents from a workspace.
|
||||
///
|
||||
/// Call this once after creation, before the agent starts processing.
|
||||
pub async fn snapshot_memory(&self, workspace: &crate::workspace::Workspace) {
|
||||
match workspace.list_all().await {
|
||||
Ok(paths) => {
|
||||
let mut snapshot = self.memory_snapshot.lock().await;
|
||||
for path in paths {
|
||||
match workspace.read(&path).await {
|
||||
Ok(doc) => {
|
||||
snapshot.push(MemorySnapshotEntry {
|
||||
path: doc.path,
|
||||
content: doc.content,
|
||||
});
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::debug!(path = %path, error = %e, "Skipped memory doc in snapshot");
|
||||
}
|
||||
}
|
||||
}
|
||||
tracing::info!(
|
||||
documents = snapshot.len(),
|
||||
"Captured memory snapshot for trace recording"
|
||||
);
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::warn!("Failed to snapshot memory for trace recording: {}", e);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Flush accumulated steps, memory snapshot, and HTTP exchanges to the output file.
|
||||
pub async fn flush(&self) -> Result<(), std::io::Error> {
|
||||
let steps = self.steps.lock().await;
|
||||
let memory_snapshot = self.memory_snapshot.lock().await;
|
||||
let http_exchanges = self.http_interceptor.take_exchanges().await;
|
||||
|
||||
let trace = TraceFile {
|
||||
model_name: self.model_name.clone(),
|
||||
memory_snapshot: memory_snapshot.clone(),
|
||||
http_exchanges,
|
||||
steps: steps.clone(),
|
||||
};
|
||||
let json = serde_json::to_string_pretty(&trace).map_err(std::io::Error::other)?;
|
||||
tokio::fs::write(&self.output_path, json).await?;
|
||||
tracing::info!(
|
||||
steps = steps.len(),
|
||||
memory_docs = memory_snapshot.len(),
|
||||
path = %self.output_path.display(),
|
||||
"Flushed LLM trace recording"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Extract new user messages, tool results, and build request hint.
|
||||
///
|
||||
/// Returns `(hint, tool_results)` where tool_results are new `Role::Tool`
|
||||
/// messages since the last call — these become `expected_tool_results` on
|
||||
/// the next step for replay verification.
|
||||
async fn capture_new_messages(
|
||||
&self,
|
||||
messages: &[ChatMessage],
|
||||
) -> (Option<RequestHint>, Vec<ExpectedToolResult>) {
|
||||
let mut prev_count = self.prev_message_count.lock().await;
|
||||
let current_count = messages.len();
|
||||
// After context compaction, the message list may shrink below
|
||||
// prev_count. Clamp to avoid an out-of-bounds slice.
|
||||
let start = (*prev_count).min(current_count);
|
||||
|
||||
let new_messages = &messages[start..];
|
||||
|
||||
// Emit UserInput steps for new user messages
|
||||
let new_user_messages: Vec<&ChatMessage> = new_messages
|
||||
.iter()
|
||||
.filter(|m| m.role == Role::User)
|
||||
.collect();
|
||||
|
||||
if !new_user_messages.is_empty() {
|
||||
let mut steps = self.steps.lock().await;
|
||||
for msg in &new_user_messages {
|
||||
steps.push(TraceStep {
|
||||
request_hint: None,
|
||||
response: TraceResponse::UserInput {
|
||||
content: msg.content.clone(),
|
||||
},
|
||||
expected_tool_results: Vec::new(),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
// Capture new tool result messages for expected_tool_results
|
||||
let tool_results: Vec<ExpectedToolResult> = new_messages
|
||||
.iter()
|
||||
.filter(|m| m.role == Role::Tool)
|
||||
.map(|m| ExpectedToolResult {
|
||||
tool_call_id: m.tool_call_id.clone().unwrap_or_default(),
|
||||
name: m.name.clone().unwrap_or_default(),
|
||||
content: m.content.clone(),
|
||||
})
|
||||
.collect();
|
||||
|
||||
*prev_count = current_count;
|
||||
|
||||
// Build request hint from last user message
|
||||
let hint = messages
|
||||
.iter()
|
||||
.rev()
|
||||
.find(|m| m.role == Role::User)
|
||||
.map(|msg| {
|
||||
let hint_text = if msg.content.len() > 80 {
|
||||
msg.content[..80].to_string()
|
||||
} else {
|
||||
msg.content.clone()
|
||||
};
|
||||
RequestHint {
|
||||
last_user_message_contains: Some(hint_text),
|
||||
min_message_count: Some(current_count),
|
||||
}
|
||||
});
|
||||
|
||||
(hint, tool_results)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LlmProvider for RecordingLlm {
|
||||
fn model_name(&self) -> &str {
|
||||
self.inner.model_name()
|
||||
}
|
||||
|
||||
fn cost_per_token(&self) -> (Decimal, Decimal) {
|
||||
self.inner.cost_per_token()
|
||||
}
|
||||
|
||||
async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse, LlmError> {
|
||||
let (hint, tool_results) = self.capture_new_messages(&request.messages).await;
|
||||
let response = self.inner.complete(request).await?;
|
||||
|
||||
self.steps.lock().await.push(TraceStep {
|
||||
request_hint: hint,
|
||||
response: TraceResponse::Text {
|
||||
content: response.content.clone(),
|
||||
input_tokens: response.input_tokens,
|
||||
output_tokens: response.output_tokens,
|
||||
},
|
||||
expected_tool_results: tool_results,
|
||||
});
|
||||
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
async fn complete_with_tools(
|
||||
&self,
|
||||
request: ToolCompletionRequest,
|
||||
) -> Result<ToolCompletionResponse, LlmError> {
|
||||
let (hint, tool_results) = self.capture_new_messages(&request.messages).await;
|
||||
let response = self.inner.complete_with_tools(request).await?;
|
||||
|
||||
let step = if response.tool_calls.is_empty() {
|
||||
TraceStep {
|
||||
request_hint: hint,
|
||||
response: TraceResponse::Text {
|
||||
content: response.content.clone().unwrap_or_default(),
|
||||
input_tokens: response.input_tokens,
|
||||
output_tokens: response.output_tokens,
|
||||
},
|
||||
expected_tool_results: tool_results,
|
||||
}
|
||||
} else {
|
||||
TraceStep {
|
||||
request_hint: hint,
|
||||
response: TraceResponse::ToolCalls {
|
||||
tool_calls: response
|
||||
.tool_calls
|
||||
.iter()
|
||||
.map(|tc| TraceToolCall {
|
||||
id: tc.id.clone(),
|
||||
name: tc.name.clone(),
|
||||
arguments: tc.arguments.clone(),
|
||||
})
|
||||
.collect(),
|
||||
input_tokens: response.input_tokens,
|
||||
output_tokens: response.output_tokens,
|
||||
},
|
||||
expected_tool_results: tool_results,
|
||||
}
|
||||
};
|
||||
|
||||
self.steps.lock().await.push(step);
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
async fn list_models(&self) -> Result<Vec<String>, LlmError> {
|
||||
self.inner.list_models().await
|
||||
}
|
||||
|
||||
async fn model_metadata(&self) -> Result<ModelMetadata, LlmError> {
|
||||
self.inner.model_metadata().await
|
||||
}
|
||||
|
||||
fn effective_model_name(&self, requested_model: Option<&str>) -> String {
|
||||
self.inner.effective_model_name(requested_model)
|
||||
}
|
||||
|
||||
fn active_model_name(&self) -> String {
|
||||
self.inner.active_model_name()
|
||||
}
|
||||
|
||||
fn set_model(&self, model: &str) -> Result<(), LlmError> {
|
||||
self.inner.set_model(model)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::testing::StubLlm;
|
||||
|
||||
fn make_recorder(stub: Arc<StubLlm>) -> RecordingLlm {
|
||||
RecordingLlm::new(
|
||||
stub,
|
||||
PathBuf::from("/tmp/test_recording.json"),
|
||||
"test-recording".to_string(),
|
||||
)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn captures_user_input_before_first_response() {
|
||||
let stub = Arc::new(StubLlm::new("hello back"));
|
||||
let recorder = make_recorder(stub);
|
||||
|
||||
let request = CompletionRequest::new(vec![
|
||||
ChatMessage::system("You are helpful."),
|
||||
ChatMessage::user("Hello!"),
|
||||
]);
|
||||
recorder.complete(request).await.unwrap();
|
||||
|
||||
let steps = recorder.steps.lock().await;
|
||||
assert_eq!(steps.len(), 2);
|
||||
|
||||
// First step: user_input
|
||||
assert!(
|
||||
matches!(&steps[0].response, TraceResponse::UserInput { content } if content == "Hello!")
|
||||
);
|
||||
|
||||
// Second step: text response
|
||||
assert!(
|
||||
matches!(&steps[1].response, TraceResponse::Text { content, .. } if content == "hello back")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn captures_text_response_correctly() {
|
||||
let stub = Arc::new(StubLlm::new("test response"));
|
||||
let recorder = make_recorder(stub);
|
||||
|
||||
let request = CompletionRequest::new(vec![ChatMessage::user("question")]);
|
||||
recorder.complete(request).await.unwrap();
|
||||
|
||||
let steps = recorder.steps.lock().await;
|
||||
// user_input + text
|
||||
assert_eq!(steps.len(), 2);
|
||||
match &steps[1].response {
|
||||
TraceResponse::Text {
|
||||
content,
|
||||
input_tokens,
|
||||
output_tokens,
|
||||
} => {
|
||||
assert_eq!(content, "test response");
|
||||
// StubLlm returns 0s for tokens, which is fine
|
||||
let _ = (*input_tokens, *output_tokens);
|
||||
}
|
||||
_ => panic!("Expected Text response"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn captures_tool_calls_response() {
|
||||
let stub = Arc::new(StubLlm::new("tool result"));
|
||||
let recorder = make_recorder(stub);
|
||||
|
||||
// complete_with_tools on StubLlm returns text, not tool_calls.
|
||||
// But we can still verify the recording captures it as text.
|
||||
let request = ToolCompletionRequest::new(vec![ChatMessage::user("use a tool")], vec![]);
|
||||
recorder.complete_with_tools(request).await.unwrap();
|
||||
|
||||
let steps = recorder.steps.lock().await;
|
||||
assert_eq!(steps.len(), 2); // user_input + text (StubLlm doesn't return tool_calls)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn no_spurious_user_input_for_tool_iterations() {
|
||||
let stub = Arc::new(StubLlm::new("response"));
|
||||
let recorder = make_recorder(stub);
|
||||
|
||||
// First call with user message
|
||||
let request = CompletionRequest::new(vec![
|
||||
ChatMessage::system("sys"),
|
||||
ChatMessage::user("Do something"),
|
||||
]);
|
||||
recorder.complete(request).await.unwrap();
|
||||
|
||||
// Second call: same messages plus tool result (no new user message)
|
||||
let request = CompletionRequest::new(vec![
|
||||
ChatMessage::system("sys"),
|
||||
ChatMessage::user("Do something"),
|
||||
ChatMessage::assistant("I'll use a tool"),
|
||||
ChatMessage::tool_result("call_1", "echo", "result"),
|
||||
]);
|
||||
recorder.complete(request).await.unwrap();
|
||||
|
||||
let steps = recorder.steps.lock().await;
|
||||
// Step 0: user_input "Do something"
|
||||
// Step 1: text response
|
||||
// Step 2: text response (no new user_input since no new user messages)
|
||||
assert_eq!(steps.len(), 3);
|
||||
assert!(matches!(
|
||||
&steps[0].response,
|
||||
TraceResponse::UserInput { .. }
|
||||
));
|
||||
assert!(matches!(&steps[1].response, TraceResponse::Text { .. }));
|
||||
assert!(matches!(&steps[2].response, TraceResponse::Text { .. }));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn captures_tool_results_for_verification() {
|
||||
let stub = Arc::new(StubLlm::new("response"));
|
||||
let recorder = make_recorder(stub);
|
||||
|
||||
// First call: user asks something
|
||||
let request = CompletionRequest::new(vec![
|
||||
ChatMessage::system("sys"),
|
||||
ChatMessage::user("Do something"),
|
||||
]);
|
||||
recorder.complete(request).await.unwrap();
|
||||
|
||||
// Second call: includes tool results from previous tool_calls
|
||||
let request = CompletionRequest::new(vec![
|
||||
ChatMessage::system("sys"),
|
||||
ChatMessage::user("Do something"),
|
||||
ChatMessage::assistant("I'll use a tool"),
|
||||
ChatMessage::tool_result("call_1", "echo", "echoed: hello"),
|
||||
ChatMessage::tool_result("call_2", "time", "2026-03-04T14:00:00Z"),
|
||||
]);
|
||||
recorder.complete(request).await.unwrap();
|
||||
|
||||
let steps = recorder.steps.lock().await;
|
||||
// Step 2 (the second LLM response) should have expected_tool_results
|
||||
let step = &steps[2];
|
||||
assert_eq!(step.expected_tool_results.len(), 2);
|
||||
assert_eq!(step.expected_tool_results[0].name, "echo");
|
||||
assert_eq!(step.expected_tool_results[0].content, "echoed: hello");
|
||||
assert_eq!(step.expected_tool_results[1].name, "time");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn request_hint_extraction() {
|
||||
let stub = Arc::new(StubLlm::new("response"));
|
||||
let recorder = make_recorder(stub);
|
||||
|
||||
let request = CompletionRequest::new(vec![
|
||||
ChatMessage::system("sys"),
|
||||
ChatMessage::user("What time is it?"),
|
||||
]);
|
||||
recorder.complete(request).await.unwrap();
|
||||
|
||||
let steps = recorder.steps.lock().await;
|
||||
let text_step = &steps[1];
|
||||
let hint = text_step.request_hint.as_ref().unwrap();
|
||||
assert_eq!(
|
||||
hint.last_user_message_contains.as_deref(),
|
||||
Some("What time is it?")
|
||||
);
|
||||
assert_eq!(hint.min_message_count, Some(2));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn flush_writes_valid_json_with_all_fields() {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let path = dir.path().join("trace.json");
|
||||
|
||||
let stub = Arc::new(StubLlm::new("response"));
|
||||
let recorder = RecordingLlm::new(stub, path.clone(), "flush-test".to_string());
|
||||
|
||||
// Simulate a memory snapshot
|
||||
recorder
|
||||
.memory_snapshot
|
||||
.lock()
|
||||
.await
|
||||
.push(MemorySnapshotEntry {
|
||||
path: "context/test.md".to_string(),
|
||||
content: "test content".to_string(),
|
||||
});
|
||||
|
||||
// Simulate an HTTP exchange
|
||||
recorder
|
||||
.http_interceptor
|
||||
.after_response(
|
||||
&HttpExchangeRequest {
|
||||
method: "GET".to_string(),
|
||||
url: "https://api.example.com/data".to_string(),
|
||||
headers: Vec::new(),
|
||||
body: None,
|
||||
},
|
||||
&HttpExchangeResponse {
|
||||
status: 200,
|
||||
headers: Vec::new(),
|
||||
body: r#"{"ok": true}"#.to_string(),
|
||||
},
|
||||
)
|
||||
.await;
|
||||
|
||||
let request = CompletionRequest::new(vec![ChatMessage::user("hello")]);
|
||||
recorder.complete(request).await.unwrap();
|
||||
recorder.flush().await.unwrap();
|
||||
|
||||
let content = tokio::fs::read_to_string(&path).await.unwrap();
|
||||
let trace: TraceFile = serde_json::from_str(&content).unwrap();
|
||||
assert_eq!(trace.model_name, "flush-test");
|
||||
assert_eq!(trace.memory_snapshot.len(), 1);
|
||||
assert_eq!(trace.memory_snapshot[0].path, "context/test.md");
|
||||
assert_eq!(trace.http_exchanges.len(), 1);
|
||||
assert_eq!(trace.http_exchanges[0].response.status, 200);
|
||||
assert_eq!(trace.steps.len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn from_env_returns_none_when_unset() {
|
||||
// SAFETY: This test is single-threaded and no other thread reads this var.
|
||||
unsafe { std::env::remove_var("IRONCLAW_RECORD_TRACE") };
|
||||
let stub = Arc::new(StubLlm::new("response"));
|
||||
let result = RecordingLlm::from_env(stub);
|
||||
assert!(result.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn recording_http_interceptor_passes_through_and_records() {
|
||||
let interceptor = RecordingHttpInterceptor::new();
|
||||
|
||||
let req = HttpExchangeRequest {
|
||||
method: "GET".to_string(),
|
||||
url: "https://example.com".to_string(),
|
||||
headers: Vec::new(),
|
||||
body: None,
|
||||
};
|
||||
|
||||
// before_request should return None (pass through)
|
||||
assert!(interceptor.before_request(&req).await.is_none());
|
||||
|
||||
// after_response records the exchange
|
||||
let resp = HttpExchangeResponse {
|
||||
status: 200,
|
||||
headers: Vec::new(),
|
||||
body: "ok".to_string(),
|
||||
};
|
||||
interceptor.after_response(&req, &resp).await;
|
||||
|
||||
let exchanges = interceptor.take_exchanges().await;
|
||||
assert_eq!(exchanges.len(), 1);
|
||||
assert_eq!(exchanges[0].request.url, "https://example.com");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn replaying_http_interceptor_returns_recorded_responses() {
|
||||
let exchanges = vec![HttpExchange {
|
||||
request: HttpExchangeRequest {
|
||||
method: "GET".to_string(),
|
||||
url: "https://api.example.com/data".to_string(),
|
||||
headers: Vec::new(),
|
||||
body: None,
|
||||
},
|
||||
response: HttpExchangeResponse {
|
||||
status: 200,
|
||||
headers: Vec::new(),
|
||||
body: r#"{"items": []}"#.to_string(),
|
||||
},
|
||||
}];
|
||||
let interceptor = ReplayingHttpInterceptor::new(exchanges);
|
||||
|
||||
// First request: returns recorded response
|
||||
let req = HttpExchangeRequest {
|
||||
method: "GET".to_string(),
|
||||
url: "https://api.example.com/data".to_string(),
|
||||
headers: Vec::new(),
|
||||
body: None,
|
||||
};
|
||||
let resp = interceptor.before_request(&req).await.unwrap();
|
||||
assert_eq!(resp.status, 200);
|
||||
assert_eq!(resp.body, r#"{"items": []}"#);
|
||||
|
||||
// Second request: no more exchanges → 599
|
||||
let resp = interceptor.before_request(&req).await.unwrap();
|
||||
assert_eq!(resp.status, 599);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn serde_roundtrip_extended_format() {
|
||||
let trace = TraceFile {
|
||||
model_name: "test".to_string(),
|
||||
memory_snapshot: vec![MemorySnapshotEntry {
|
||||
path: "context/vision.md".to_string(),
|
||||
content: "Be helpful.".to_string(),
|
||||
}],
|
||||
http_exchanges: vec![HttpExchange {
|
||||
request: HttpExchangeRequest {
|
||||
method: "GET".to_string(),
|
||||
url: "https://api.example.com".to_string(),
|
||||
headers: vec![("Accept".to_string(), "application/json".to_string())],
|
||||
body: None,
|
||||
},
|
||||
response: HttpExchangeResponse {
|
||||
status: 200,
|
||||
headers: Vec::new(),
|
||||
body: "{}".to_string(),
|
||||
},
|
||||
}],
|
||||
steps: vec![
|
||||
TraceStep {
|
||||
request_hint: None,
|
||||
response: TraceResponse::UserInput {
|
||||
content: "hello".to_string(),
|
||||
},
|
||||
expected_tool_results: Vec::new(),
|
||||
},
|
||||
TraceStep {
|
||||
request_hint: Some(RequestHint {
|
||||
last_user_message_contains: Some("hello".to_string()),
|
||||
min_message_count: Some(2),
|
||||
}),
|
||||
response: TraceResponse::ToolCalls {
|
||||
tool_calls: vec![TraceToolCall {
|
||||
id: "call_1".to_string(),
|
||||
name: "echo".to_string(),
|
||||
arguments: serde_json::json!({"message": "hi"}),
|
||||
}],
|
||||
input_tokens: 50,
|
||||
output_tokens: 20,
|
||||
},
|
||||
expected_tool_results: Vec::new(),
|
||||
},
|
||||
TraceStep {
|
||||
request_hint: None,
|
||||
response: TraceResponse::Text {
|
||||
content: "done".to_string(),
|
||||
input_tokens: 80,
|
||||
output_tokens: 10,
|
||||
},
|
||||
expected_tool_results: vec![ExpectedToolResult {
|
||||
tool_call_id: "call_1".to_string(),
|
||||
name: "echo".to_string(),
|
||||
content: "hi".to_string(),
|
||||
}],
|
||||
},
|
||||
],
|
||||
};
|
||||
|
||||
let json = serde_json::to_string_pretty(&trace).unwrap();
|
||||
let parsed: TraceFile = serde_json::from_str(&json).unwrap();
|
||||
assert_eq!(parsed.model_name, "test");
|
||||
assert_eq!(parsed.memory_snapshot.len(), 1);
|
||||
assert_eq!(parsed.http_exchanges.len(), 1);
|
||||
assert_eq!(parsed.steps.len(), 3);
|
||||
assert_eq!(parsed.steps[2].expected_tool_results.len(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backward_compatible_with_old_format() {
|
||||
// Old format without memory_snapshot, http_exchanges, expected_tool_results
|
||||
let json = r#"{
|
||||
"model_name": "old-trace",
|
||||
"steps": [
|
||||
{
|
||||
"response": {
|
||||
"type": "text",
|
||||
"content": "hello",
|
||||
"input_tokens": 10,
|
||||
"output_tokens": 5
|
||||
}
|
||||
}
|
||||
]
|
||||
}"#;
|
||||
let trace: TraceFile = serde_json::from_str(json).unwrap();
|
||||
assert_eq!(trace.model_name, "old-trace");
|
||||
assert!(trace.memory_snapshot.is_empty());
|
||||
assert!(trace.http_exchanges.is_empty());
|
||||
assert!(trace.steps[0].expected_tool_results.is_empty());
|
||||
}
|
||||
}
|
||||
+333
-33
@@ -16,13 +16,14 @@
|
||||
//! ```
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use rust_decimal::Decimal;
|
||||
use sha2::{Digest, Sha256};
|
||||
use tokio::sync::Mutex;
|
||||
|
||||
use crate::error::LlmError;
|
||||
use crate::llm::provider::{
|
||||
@@ -30,6 +31,9 @@ use crate::llm::provider::{
|
||||
ToolCompletionResponse,
|
||||
};
|
||||
|
||||
/// How often (in requests) to emit a cache statistics log line.
|
||||
const STATS_LOG_EVERY_N: u64 = 100;
|
||||
|
||||
/// Configuration for the response cache.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ResponseCacheConfig {
|
||||
@@ -61,8 +65,16 @@ struct CacheEntry {
|
||||
/// tool calls can have side effects that should not be replayed.
|
||||
pub struct CachedProvider {
|
||||
inner: Arc<dyn LlmProvider>,
|
||||
/// `std::sync::Mutex` (not tokio) — never held across an `.await` point,
|
||||
/// so blocking acquisition is safe and keeps `set_model()` synchronous.
|
||||
cache: Mutex<HashMap<String, CacheEntry>>,
|
||||
config: ResponseCacheConfig,
|
||||
/// Total `complete()` calls (hits + misses) for periodic stats logging.
|
||||
request_count: AtomicU64,
|
||||
/// Running total of cache hits, independent of entry lifecycle.
|
||||
/// Never decremented on eviction, so `hit_rate_pct` in stats doesn't
|
||||
/// drift down as entries expire or are LRU-evicted.
|
||||
total_hit_count: AtomicU64,
|
||||
}
|
||||
|
||||
impl CachedProvider {
|
||||
@@ -72,27 +84,53 @@ impl CachedProvider {
|
||||
inner,
|
||||
cache: Mutex::new(HashMap::new()),
|
||||
config,
|
||||
request_count: AtomicU64::new(0),
|
||||
total_hit_count: AtomicU64::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
/// Number of entries currently in the cache.
|
||||
pub async fn len(&self) -> usize {
|
||||
self.cache.lock().await.len()
|
||||
pub fn len(&self) -> usize {
|
||||
self.cache.lock().unwrap_or_else(|e| e.into_inner()).len()
|
||||
}
|
||||
|
||||
/// Whether the cache is empty.
|
||||
pub async fn is_empty(&self) -> bool {
|
||||
self.cache.lock().await.is_empty()
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.cache
|
||||
.lock()
|
||||
.unwrap_or_else(|e| e.into_inner())
|
||||
.is_empty()
|
||||
}
|
||||
|
||||
/// Total cache hits across all entries.
|
||||
pub async fn total_hits(&self) -> u64 {
|
||||
self.cache.lock().await.values().map(|e| e.hit_count).sum()
|
||||
/// Total cache hits since this provider was created.
|
||||
///
|
||||
/// Backed by an atomic counter that is never decremented on eviction,
|
||||
/// so the value is accurate even under high eviction pressure.
|
||||
pub fn total_hits(&self) -> u64 {
|
||||
self.total_hit_count.load(Ordering::Relaxed)
|
||||
}
|
||||
|
||||
/// Clear all cached entries.
|
||||
pub async fn clear(&self) {
|
||||
self.cache.lock().await.clear();
|
||||
pub fn clear(&self) {
|
||||
self.cache.lock().unwrap_or_else(|e| e.into_inner()).clear();
|
||||
}
|
||||
|
||||
/// Emit a cache statistics log line if `req_no` is a multiple of
|
||||
/// [`STATS_LOG_EVERY_N`]. `total_hits` must come from the `total_hit_count`
|
||||
/// atomic so it accurately reflects hits that occurred on since-evicted
|
||||
/// entries. Must be called while holding the cache lock so that
|
||||
/// `entry_count` is consistent with the snapshot.
|
||||
fn maybe_log_stats(guard: &HashMap<String, CacheEntry>, req_no: u64, total_hits: u64) {
|
||||
if req_no.is_multiple_of(STATS_LOG_EVERY_N) {
|
||||
let hit_rate = total_hits as f64 / req_no as f64 * 100.0;
|
||||
tracing::info!(
|
||||
total_requests = req_no,
|
||||
total_hits,
|
||||
hit_rate_pct = format!("{hit_rate:.1}"),
|
||||
entry_count = guard.len(),
|
||||
"LLM response cache statistics"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -147,28 +185,47 @@ impl LlmProvider for CachedProvider {
|
||||
let effective_model = self.inner.effective_model_name(request.model.as_deref());
|
||||
let key = cache_key(&effective_model, &request);
|
||||
let now = Instant::now();
|
||||
let req_no = self.request_count.fetch_add(1, Ordering::Relaxed) + 1;
|
||||
|
||||
// Check cache
|
||||
// Check cache — lock not held across the .await below.
|
||||
{
|
||||
let mut guard = self.cache.lock().await;
|
||||
let mut guard = self.cache.lock().unwrap_or_else(|e| e.into_inner());
|
||||
if let Some(entry) = guard.get_mut(&key) {
|
||||
if now.duration_since(entry.created_at) < self.config.ttl {
|
||||
entry.last_accessed = now;
|
||||
entry.hit_count += 1;
|
||||
tracing::debug!(hits = entry.hit_count, "response cache hit");
|
||||
return Ok(entry.response.clone());
|
||||
let hit_count = entry.hit_count;
|
||||
// Clone now so we can release the mutable borrow before stats.
|
||||
let cached_response = entry.response.clone();
|
||||
tracing::debug!(hits = hit_count, "response cache hit");
|
||||
// Drop the mutable borrow of `entry` before reading `guard` immutably.
|
||||
let _ = entry;
|
||||
let total_hits = self.total_hit_count.fetch_add(1, Ordering::Relaxed) + 1;
|
||||
Self::maybe_log_stats(&guard, req_no, total_hits);
|
||||
return Ok(cached_response);
|
||||
}
|
||||
// Expired, remove it
|
||||
guard.remove(&key);
|
||||
}
|
||||
}
|
||||
|
||||
// Cache miss, call the real provider
|
||||
let response = self.inner.complete(request).await?;
|
||||
// Cache miss — call the real provider.
|
||||
let result = self.inner.complete(request).await;
|
||||
|
||||
// Store in cache
|
||||
// Store result and maybe log stats, all within one lock acquisition.
|
||||
// Stats are logged even on provider error so milestone intervals are
|
||||
// not silently skipped.
|
||||
{
|
||||
let mut guard = self.cache.lock().await;
|
||||
let mut guard = self.cache.lock().unwrap_or_else(|e| e.into_inner());
|
||||
let total_hits = self.total_hit_count.load(Ordering::Relaxed);
|
||||
|
||||
let response = match result {
|
||||
Err(e) => {
|
||||
Self::maybe_log_stats(&guard, req_no, total_hits);
|
||||
return Err(e);
|
||||
}
|
||||
Ok(r) => r,
|
||||
};
|
||||
|
||||
// Evict expired entries
|
||||
guard.retain(|_, entry| now.duration_since(entry.created_at) < self.config.ttl);
|
||||
@@ -196,9 +253,10 @@ impl LlmProvider for CachedProvider {
|
||||
hit_count: 0,
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
Ok(response)
|
||||
Self::maybe_log_stats(&guard, req_no, total_hits);
|
||||
Ok(response)
|
||||
}
|
||||
}
|
||||
|
||||
async fn complete_with_tools(
|
||||
@@ -226,16 +284,91 @@ impl LlmProvider for CachedProvider {
|
||||
}
|
||||
|
||||
fn set_model(&self, model: &str) -> Result<(), LlmError> {
|
||||
// Cache keys embed the active model name via `effective_model_name()`, so
|
||||
// requests to the new model automatically land in a separate cache slot.
|
||||
// Entries for the old model remain valid: if we switch back, they will be
|
||||
// hit again rather than wasted. Natural TTL / LRU eviction cleans them up.
|
||||
self.inner.set_model(model)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::llm::provider::ChatMessage;
|
||||
use std::sync::atomic::{AtomicU32, Ordering};
|
||||
|
||||
use rust_decimal::Decimal;
|
||||
use tracing_test::traced_test;
|
||||
|
||||
use crate::error::LlmError;
|
||||
use crate::llm::provider::{
|
||||
ChatMessage, CompletionResponse, FinishReason, ToolCompletionRequest,
|
||||
ToolCompletionResponse,
|
||||
};
|
||||
use crate::llm::response_cache::*;
|
||||
use crate::testing::StubLlm;
|
||||
|
||||
/// Minimal provider stub that supports `set_model()` — used to test
|
||||
/// per-model cache key isolation.
|
||||
struct SwitchableStub {
|
||||
call_count: AtomicU32,
|
||||
active_model: std::sync::RwLock<String>,
|
||||
}
|
||||
|
||||
impl SwitchableStub {
|
||||
fn new() -> Self {
|
||||
Self {
|
||||
call_count: AtomicU32::new(0),
|
||||
active_model: std::sync::RwLock::new("stub-model".to_string()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LlmProvider for SwitchableStub {
|
||||
fn model_name(&self) -> &str {
|
||||
"stub-model"
|
||||
}
|
||||
|
||||
fn active_model_name(&self) -> String {
|
||||
self.active_model.read().unwrap().clone()
|
||||
}
|
||||
|
||||
fn cost_per_token(&self) -> (Decimal, Decimal) {
|
||||
(Decimal::ZERO, Decimal::ZERO)
|
||||
}
|
||||
|
||||
fn set_model(&self, model: &str) -> Result<(), LlmError> {
|
||||
*self.active_model.write().unwrap() = model.to_string();
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn complete(
|
||||
&self,
|
||||
_request: CompletionRequest,
|
||||
) -> Result<CompletionResponse, LlmError> {
|
||||
self.call_count.fetch_add(1, Ordering::Relaxed);
|
||||
Ok(CompletionResponse {
|
||||
content: "ok".into(),
|
||||
input_tokens: 1,
|
||||
output_tokens: 1,
|
||||
finish_reason: FinishReason::Stop,
|
||||
})
|
||||
}
|
||||
|
||||
async fn complete_with_tools(
|
||||
&self,
|
||||
_request: ToolCompletionRequest,
|
||||
) -> Result<ToolCompletionResponse, LlmError> {
|
||||
Ok(ToolCompletionResponse {
|
||||
content: Some("ok".into()),
|
||||
tool_calls: vec![],
|
||||
input_tokens: 1,
|
||||
output_tokens: 1,
|
||||
finish_reason: FinishReason::Stop,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn simple_request() -> CompletionRequest {
|
||||
CompletionRequest {
|
||||
messages: vec![ChatMessage::user("hello")],
|
||||
@@ -321,7 +454,7 @@ mod tests {
|
||||
assert_eq!(stub.calls(), 1); // still 1
|
||||
assert_eq!(r2.content, "cached response");
|
||||
|
||||
assert_eq!(cached.total_hits().await, 1);
|
||||
assert_eq!(cached.total_hits(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -333,7 +466,7 @@ mod tests {
|
||||
cached.complete(different_request()).await.unwrap();
|
||||
|
||||
assert_eq!(stub.calls(), 2);
|
||||
assert_eq!(cached.len().await, 2);
|
||||
assert_eq!(cached.len(), 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -372,7 +505,7 @@ mod tests {
|
||||
// Fill cache with 2 entries
|
||||
cached.complete(simple_request()).await.unwrap();
|
||||
cached.complete(different_request()).await.unwrap();
|
||||
assert_eq!(cached.len().await, 2);
|
||||
assert_eq!(cached.len(), 2);
|
||||
|
||||
// Add a third: should evict the oldest
|
||||
let third = CompletionRequest {
|
||||
@@ -384,7 +517,7 @@ mod tests {
|
||||
metadata: Default::default(),
|
||||
};
|
||||
cached.complete(third).await.unwrap();
|
||||
assert_eq!(cached.len().await, 2);
|
||||
assert_eq!(cached.len(), 2);
|
||||
assert_eq!(stub.calls(), 3);
|
||||
}
|
||||
|
||||
@@ -408,7 +541,7 @@ mod tests {
|
||||
|
||||
// Both should have called through
|
||||
assert_eq!(stub.calls(), 2);
|
||||
assert!(cached.is_empty().await);
|
||||
assert!(cached.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -425,12 +558,12 @@ mod tests {
|
||||
stub.set_failing(true);
|
||||
let result = cached.complete(simple_request()).await;
|
||||
assert!(result.is_err());
|
||||
assert!(cached.is_empty().await);
|
||||
assert!(cached.is_empty());
|
||||
|
||||
// After fixing the provider, should succeed and cache
|
||||
stub.set_failing(false);
|
||||
cached.complete(simple_request()).await.unwrap();
|
||||
assert_eq!(cached.len().await, 1);
|
||||
assert_eq!(cached.len(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -439,10 +572,10 @@ mod tests {
|
||||
let cached = CachedProvider::new(stub.clone(), ResponseCacheConfig::default());
|
||||
|
||||
cached.complete(simple_request()).await.unwrap();
|
||||
assert_eq!(cached.len().await, 1);
|
||||
assert_eq!(cached.len(), 1);
|
||||
|
||||
cached.clear().await;
|
||||
assert!(cached.is_empty().await);
|
||||
cached.clear();
|
||||
assert!(cached.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -459,7 +592,7 @@ mod tests {
|
||||
cached.complete(req_b).await.unwrap();
|
||||
|
||||
assert_eq!(stub.calls(), 2);
|
||||
assert_eq!(cached.len().await, 2);
|
||||
assert_eq!(cached.len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -475,4 +608,171 @@ mod tests {
|
||||
let cached = CachedProvider::new(stub.clone(), ResponseCacheConfig::default());
|
||||
assert_eq!(cached.model_name(), "stub-model");
|
||||
}
|
||||
|
||||
/// Switching models preserves existing cached entries and routes subsequent
|
||||
/// requests to a separate cache slot. Switching back replays the old slot.
|
||||
#[tokio::test]
|
||||
async fn set_model_isolates_per_model_via_key() {
|
||||
let stub = Arc::new(SwitchableStub::new());
|
||||
let cached = CachedProvider::new(stub.clone(), ResponseCacheConfig::default());
|
||||
|
||||
// Populate cache under the initial model ("stub-model").
|
||||
cached.complete(simple_request()).await.unwrap();
|
||||
assert_eq!(stub.call_count.load(Ordering::Relaxed), 1);
|
||||
assert_eq!(cached.len(), 1, "one entry cached for stub-model");
|
||||
|
||||
// Switch to a different model — old entries must survive.
|
||||
cached.set_model("model-b").unwrap();
|
||||
assert_eq!(cached.len(), 1, "old entries preserved after model switch");
|
||||
|
||||
// Same request under model-b is a cache miss (different key).
|
||||
cached.complete(simple_request()).await.unwrap();
|
||||
assert_eq!(
|
||||
stub.call_count.load(Ordering::Relaxed),
|
||||
2,
|
||||
"cache miss for model-b"
|
||||
);
|
||||
assert_eq!(cached.len(), 2, "separate slots for stub-model and model-b");
|
||||
|
||||
// Switch back — original slot is still valid (cache hit, no extra call).
|
||||
cached.set_model("stub-model").unwrap();
|
||||
cached.complete(simple_request()).await.unwrap();
|
||||
assert_eq!(
|
||||
stub.call_count.load(Ordering::Relaxed),
|
||||
2,
|
||||
"cache hit when switching back to stub-model"
|
||||
);
|
||||
}
|
||||
|
||||
/// When `set_model()` fails the error is propagated and the cache is unaffected.
|
||||
#[tokio::test]
|
||||
async fn set_model_error_leaves_cache_intact() {
|
||||
// StubLlm does not override set_model() — returns an error by default.
|
||||
let stub = Arc::new(StubLlm::default());
|
||||
let cached = CachedProvider::new(stub, ResponseCacheConfig::default());
|
||||
|
||||
cached.complete(simple_request()).await.unwrap();
|
||||
assert_eq!(cached.len(), 1);
|
||||
|
||||
let result = cached.set_model("new-model");
|
||||
assert!(result.is_err());
|
||||
assert_eq!(cached.len(), 1, "cache unaffected by failed set_model");
|
||||
}
|
||||
|
||||
/// `hit_rate_pct` stays accurate even after entries are evicted.
|
||||
/// The `total_hit_count` atomic is never decremented on eviction.
|
||||
#[tokio::test]
|
||||
async fn total_hits_survives_eviction() {
|
||||
let stub = Arc::new(StubLlm::new("response"));
|
||||
// max_entries = 1 so the first entry is LRU-evicted when a second arrives.
|
||||
let cached = CachedProvider::new(
|
||||
stub.clone(),
|
||||
ResponseCacheConfig {
|
||||
ttl: Duration::from_secs(60),
|
||||
max_entries: 1,
|
||||
},
|
||||
);
|
||||
|
||||
// Populate the cache and score a hit.
|
||||
cached.complete(simple_request()).await.unwrap();
|
||||
cached.complete(simple_request()).await.unwrap();
|
||||
assert_eq!(cached.total_hits(), 1);
|
||||
|
||||
// Add a different request — LRU evicts the first entry.
|
||||
cached.complete(different_request()).await.unwrap();
|
||||
assert_eq!(cached.len(), 1, "first entry was evicted");
|
||||
|
||||
// The hit from the evicted entry must still be counted.
|
||||
assert_eq!(cached.total_hits(), 1, "hit count survives eviction");
|
||||
}
|
||||
|
||||
/// A stats line is emitted exactly at the 100th request.
|
||||
#[tokio::test]
|
||||
#[traced_test]
|
||||
async fn stats_logged_at_request_100() {
|
||||
let stub = Arc::new(StubLlm::new("response"));
|
||||
let cached = CachedProvider::new(
|
||||
stub.clone(),
|
||||
ResponseCacheConfig {
|
||||
ttl: Duration::from_secs(60),
|
||||
max_entries: 2000,
|
||||
},
|
||||
);
|
||||
|
||||
// 99 distinct requests — no stats line yet.
|
||||
for i in 0..99u32 {
|
||||
let req = CompletionRequest {
|
||||
messages: vec![ChatMessage::user(format!("request {i}"))],
|
||||
model: None,
|
||||
max_tokens: None,
|
||||
temperature: None,
|
||||
stop_sequences: None,
|
||||
metadata: Default::default(),
|
||||
};
|
||||
cached.complete(req).await.unwrap();
|
||||
}
|
||||
assert!(
|
||||
!logs_contain("LLM response cache statistics"),
|
||||
"no stats before request 100"
|
||||
);
|
||||
|
||||
// 100th request triggers the first stats line.
|
||||
let req = CompletionRequest {
|
||||
messages: vec![ChatMessage::user("request 99")],
|
||||
model: None,
|
||||
max_tokens: None,
|
||||
temperature: None,
|
||||
stop_sequences: None,
|
||||
metadata: Default::default(),
|
||||
};
|
||||
cached.complete(req).await.unwrap();
|
||||
assert!(
|
||||
logs_contain("LLM response cache statistics"),
|
||||
"stats emitted at request 100"
|
||||
);
|
||||
}
|
||||
|
||||
/// Stats are emitted even when the inner provider returns an error.
|
||||
#[tokio::test]
|
||||
#[traced_test]
|
||||
async fn stats_logged_on_provider_error_at_interval() {
|
||||
let stub = Arc::new(StubLlm::new("response"));
|
||||
let cached = CachedProvider::new(
|
||||
stub.clone(),
|
||||
ResponseCacheConfig {
|
||||
ttl: Duration::from_secs(60),
|
||||
max_entries: 2000,
|
||||
},
|
||||
);
|
||||
|
||||
// 99 successful requests.
|
||||
for i in 0..99u32 {
|
||||
let req = CompletionRequest {
|
||||
messages: vec![ChatMessage::user(format!("req {i}"))],
|
||||
model: None,
|
||||
max_tokens: None,
|
||||
temperature: None,
|
||||
stop_sequences: None,
|
||||
metadata: Default::default(),
|
||||
};
|
||||
cached.complete(req).await.unwrap();
|
||||
}
|
||||
|
||||
// 100th request fails — stats must still be logged.
|
||||
stub.set_failing(true);
|
||||
let req = CompletionRequest {
|
||||
messages: vec![ChatMessage::user("req 99")],
|
||||
model: None,
|
||||
max_tokens: None,
|
||||
temperature: None,
|
||||
stop_sequences: None,
|
||||
metadata: Default::default(),
|
||||
};
|
||||
let result = cached.complete(req).await;
|
||||
assert!(result.is_err());
|
||||
assert!(
|
||||
logs_contain("LLM response cache statistics"),
|
||||
"stats emitted even when provider errors on request 100"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
+1
-1
@@ -347,7 +347,7 @@ impl SessionManager {
|
||||
|
||||
// The NEAR AI API redirects to: {frontend_callback}/auth/callback?token=X&...
|
||||
let session_token =
|
||||
oauth_defaults::wait_for_callback(listener, "/auth/callback", "token", "NEAR AI")
|
||||
oauth_defaults::wait_for_callback(listener, "/auth/callback", "token", "NEAR AI", None)
|
||||
.await
|
||||
.map_err(|e| LlmError::SessionRenewalFailed {
|
||||
provider: "nearai".to_string(),
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user