Compare commits

..
Author SHA1 Message Date
Henry ParkandClaude Opus 4.6 6e972863e7 fix(ci): prevent staging-ci tag failure and chained PR auto-close
- Use fetch-depth: 0 in update-tag to ensure current_head SHA is available
- Only merge promotion PRs targeting main; leave chained PRs open to
  prevent delete_branch_on_merge from auto-closing downstream PRs

Co-Authored-By: Claude Opus 4.6 <[email protected]>
2026-03-10 13:35:10 -07:00
165 changed files with 3628 additions and 12514 deletions
+1 -1
View File
@@ -64,7 +64,7 @@ If the event needs custom UI (cards, badges, etc.), add styles. Follow the exist
Identify where in the backend this event should be triggered. Common locations: Identify where in the backend this event should be triggered. Common locations:
- `src/agent/agent_loop.rs` - During message processing or tool execution - `src/agent/agent_loop.rs` - During message processing or tool execution
- `src/worker/job.rs` - During job execution - `src/agent/worker.rs` - During job execution
- `src/agent/heartbeat.rs` - During periodic execution - `src/agent/heartbeat.rs` - During periodic execution
Use the existing pattern: Use the existing pattern:
-50
View File
@@ -1,50 +0,0 @@
## Summary
<!-- 2-5 bullet points: what changed and why -->
-
## Change Type
<!-- Check one -->
- [ ] Bug fix
- [ ] New feature
- [ ] Refactor
- [ ] Documentation
- [ ] CI/Infrastructure
- [ ] Security
- [ ] Dependencies
## Linked Issue
<!-- Closes #N, or "None" -->
## Validation
<!-- How did you verify this works? -->
- [ ] `cargo fmt`
- [ ] `cargo clippy --all --benches --tests --examples --all-features`
- [ ] Relevant tests pass: <!-- list specific tests -->
- [ ] Manual testing: <!-- describe what you tested -->
## Security Impact
<!-- Does this change affect: permissions, network calls, secrets, file access, tool execution, sandbox policy? If yes, describe. If no, write "None". -->
## Database Impact
<!-- Does this add/modify migrations, change schema, or affect both PostgreSQL and libSQL? If yes, describe. If no, write "None". -->
## Blast Radius
<!-- What subsystems does this touch? What could break? -->
## Rollback Plan
<!-- How to revert if this causes problems? For Track C changes, this is mandatory. -->
---
**Review track**: <!-- A (docs/tests/chore) | B (feature/refactor) | C (security/runtime/DB/CI) -->
+23 -63
View File
@@ -144,8 +144,6 @@ jobs:
- name: Patch manifests with WASM checksums - name: Patch manifests with WASM checksums
if: ${{ needs.plan.outputs.publishing == 'true' }} if: ${{ needs.plan.outputs.publishing == 'true' }}
shell: bash shell: bash
env:
RELEASE_TAG: ${{ github.ref_name }}
run: | run: |
CHECKSUMS="target/distrib/checksums.txt" CHECKSUMS="target/distrib/checksums.txt"
if [ ! -f "$CHECKSUMS" ]; then if [ ! -f "$CHECKSUMS" ]; then
@@ -156,17 +154,12 @@ jobs:
while IFS= read -r line; do while IFS= read -r line; do
sha256=$(echo "$line" | awk '{print $1}') sha256=$(echo "$line" | awk '{print $1}')
filename=$(echo "$line" | awk '{print $2}') filename=$(echo "$line" | awk '{print $2}')
# Strip -{version}-wasm32-wasip2.tar.gz to get the extension name. name=$(echo "$filename" | sed 's/-wasm32-wasip2\.tar\.gz$//')
# Use '.*' (greedy) so pre-release suffixes like -alpha.1 are consumed too.
name=$(echo "$filename" | sed 's/-[0-9].*-wasm32-wasip2\.tar\.gz$//')
url="https://github.com/nearai/ironclaw/releases/download/${RELEASE_TAG}/${filename}"
for manifest in registry/tools/${name}.json registry/channels/${name}.json; do for manifest in registry/tools/${name}.json registry/channels/${name}.json; do
if [ -f "$manifest" ]; then if [ -f "$manifest" ]; then
jq --arg sha "$sha256" --arg url "$url" \ jq --arg sha "$sha256" '.artifacts["wasm32-wasip2"].sha256 = $sha' "$manifest" > "${manifest}.tmp" && mv "${manifest}.tmp" "$manifest"
'.artifacts["wasm32-wasip2"].sha256 = $sha | .artifacts["wasm32-wasip2"].url = $url' \ echo "Patched $manifest with sha256=$sha256"
"$manifest" > "${manifest}.tmp" && mv "${manifest}.tmp" "$manifest"
echo "Patched $manifest with sha256=$sha256 url=$url"
fi fi
done done
done < "$CHECKSUMS" done < "$CHECKSUMS"
@@ -275,41 +268,21 @@ jobs:
for manifest in registry/tools/*.json registry/channels/*.json; do for manifest in registry/tools/*.json registry/channels/*.json; do
[ -f "$manifest" ] || continue [ -f "$manifest" ] || continue
# file_stem: JSON filename without extension (e.g. "slack" for slack.json). name=$(jq -r '.name' "$manifest")
# Used for the bundle filename and CI manifest lookup, so patching always
# finds the right file regardless of whether manifest.name matches the filename.
file_stem=$(basename "$manifest" .json)
# ext_name: the manifest's .name field (e.g. "slack-tool").
# Used for file names *inside* the archive — the installer extracts by manifest.name.
ext_name=$(jq -r '.name' "$manifest")
source_dir=$(jq -r '.source.dir' "$manifest") source_dir=$(jq -r '.source.dir' "$manifest")
caps_file=$(jq -r '.source.capabilities' "$manifest") caps_file=$(jq -r '.source.capabilities' "$manifest")
crate_name=$(jq -r '.source.crate_name' "$manifest") crate_name=$(jq -r '.source.crate_name' "$manifest")
ext_version=$(jq -r '.version // ""' "$manifest")
if [ ! -d "$source_dir" ]; then if [ ! -d "$source_dir" ]; then
echo "::warning::Source dir '$source_dir' not found for '$file_stem', skipping" echo "::warning::Source dir '$source_dir' not found for '$name', skipping"
continue continue
fi fi
# Skip rebuild if this exact version was already built and checksummed. echo "=== Building $name from $source_dir ==="
# Checks that (1) the manifest already has a sha256, and (2) the version
# embedded in the existing artifact URL matches the current manifest version.
# This ensures stable checksums: only rebuild when the source version changes.
existing_sha=$(jq -r '.artifacts["wasm32-wasip2"].sha256 // ""' "$manifest")
existing_url=$(jq -r '.artifacts["wasm32-wasip2"].url // ""' "$manifest")
url_version=$(echo "$existing_url" | sed -n 's/.*-\([0-9].*\)-wasm32-wasip2\.tar\.gz$/\1/p')
if [[ -n "$ext_version" && "$url_version" == "$ext_version" && -n "$existing_sha" ]]; then
echo "=== Skipping $file_stem v$ext_version — already checksummed at $existing_url ==="
continue
fi
echo "=== Building $file_stem ($ext_name) v$ext_version from $source_dir ==="
# Build WASM component # Build WASM component
cargo component build --release --manifest-path "$source_dir/Cargo.toml" || { cargo component build --release --manifest-path "$source_dir/Cargo.toml" || {
echo "::warning::Build failed for '$file_stem', skipping" echo "::warning::Build failed for '$name', skipping"
continue continue
} }
@@ -325,36 +298,30 @@ jobs:
done done
if [ -z "$wasm_path" ]; then if [ -z "$wasm_path" ]; then
echo "::warning::No WASM output found for '$file_stem', skipping" echo "::warning::No WASM output found for '$name', skipping"
continue continue
fi fi
# Archive contents use ext_name (manifest .name) — the installer extracts # Copy files with standardized names for the archive
# files by manifest.name, so these must match even when file_stem differs. cp "$wasm_path" "target/wasm-bundles/${name}.wasm"
cp "$wasm_path" "target/wasm-bundles/${ext_name}.wasm"
caps_path="$source_dir/$caps_file" caps_path="$source_dir/$caps_file"
if [ -f "$caps_path" ]; then if [ -f "$caps_path" ]; then
cp "$caps_path" "target/wasm-bundles/${ext_name}.capabilities.json" cp "$caps_path" "target/wasm-bundles/${name}.capabilities.json"
else else
echo "::warning::No capabilities file at '$caps_path' for '$file_stem'" echo "::warning::No capabilities file at '$caps_path' for '$name'"
fi fi
# Bundle filename uses file_stem so CI patching can find the manifest by # Create tar.gz bundle
# filename (e.g. slack-0.1.0-wasm32-wasip2.tar.gz → registry/tools/slack.json). bundle="target/wasm-bundles/${name}-wasm32-wasip2.tar.gz"
bundle="target/wasm-bundles/${file_stem}-${ext_version}-wasm32-wasip2.tar.gz" (cd target/wasm-bundles && if [ -f "${name}.capabilities.json" ]; then tar czf "${name}-wasm32-wasip2.tar.gz" "${name}.wasm" "${name}.capabilities.json"; else tar czf "${name}-wasm32-wasip2.tar.gz" "${name}.wasm"; fi)
(cd target/wasm-bundles && if [ -f "${ext_name}.capabilities.json" ]; then
tar czf "${file_stem}-${ext_version}-wasm32-wasip2.tar.gz" "${ext_name}.wasm" "${ext_name}.capabilities.json"
else
tar czf "${file_stem}-${ext_version}-wasm32-wasip2.tar.gz" "${ext_name}.wasm"
fi)
# Compute SHA256 # Compute SHA256
sha256=$(sha256sum "$bundle" | cut -d' ' -f1) sha256=$(sha256sum "$bundle" | cut -d' ' -f1)
echo "$sha256 ${file_stem}-${ext_version}-wasm32-wasip2.tar.gz" >> target/wasm-bundles/checksums.txt echo "$sha256 ${name}-wasm32-wasip2.tar.gz" >> target/wasm-bundles/checksums.txt
# Clean up intermediate files # Clean up intermediate files
rm -f "target/wasm-bundles/${ext_name}.wasm" "target/wasm-bundles/${ext_name}.capabilities.json" rm -f "target/wasm-bundles/${name}.wasm" "target/wasm-bundles/${name}.capabilities.json"
echo " -> $bundle ($sha256)" echo " -> $bundle ($sha256)"
done done
@@ -460,10 +427,8 @@ jobs:
with: with:
name: artifacts-wasm-extensions name: artifacts-wasm-extensions
path: target/wasm-bundles/ path: target/wasm-bundles/
- name: Patch manifests with SHA256 and version-pinned URL - name: Patch manifests with SHA256
shell: bash shell: bash
env:
RELEASE_TAG: ${{ github.ref_name }}
run: | run: |
CHECKSUMS="target/wasm-bundles/checksums.txt" CHECKSUMS="target/wasm-bundles/checksums.txt"
if [ ! -f "$CHECKSUMS" ]; then if [ ! -f "$CHECKSUMS" ]; then
@@ -474,17 +439,12 @@ jobs:
while IFS= read -r line; do while IFS= read -r line; do
sha256=$(echo "$line" | awk '{print $1}') sha256=$(echo "$line" | awk '{print $1}')
filename=$(echo "$line" | awk '{print $2}') filename=$(echo "$line" | awk '{print $2}')
# Strip -{version}-wasm32-wasip2.tar.gz to get the extension name. name=$(echo "$filename" | sed 's/-wasm32-wasip2\.tar\.gz$//')
# Use '.*' (greedy) so pre-release suffixes like -alpha.1 are consumed too.
name=$(echo "$filename" | sed 's/-[0-9].*-wasm32-wasip2\.tar\.gz$//')
url="https://github.com/nearai/ironclaw/releases/download/${RELEASE_TAG}/${filename}"
for manifest in registry/tools/${name}.json registry/channels/${name}.json; do for manifest in registry/tools/${name}.json registry/channels/${name}.json; do
if [ -f "$manifest" ]; then if [ -f "$manifest" ]; then
jq --arg sha "$sha256" --arg url "$url" \ jq --arg sha "$sha256" '.artifacts["wasm32-wasip2"].sha256 = $sha' "$manifest" > "${manifest}.tmp" && mv "${manifest}.tmp" "$manifest"
'.artifacts["wasm32-wasip2"].sha256 = $sha | .artifacts["wasm32-wasip2"].url = $url' \ echo "Patched $manifest with sha256=$sha256"
"$manifest" > "${manifest}.tmp" && mv "${manifest}.tmp" "$manifest"
echo "Patched $manifest with sha256=$sha256 url=$url"
fi fi
done done
done < "$CHECKSUMS" done < "$CHECKSUMS"
@@ -501,8 +461,8 @@ jobs:
git commit -m "chore: update WASM artifact SHA256 checksums [skip ci]" git commit -m "chore: update WASM artifact SHA256 checksums [skip ci]"
git push origin "$BRANCH" git push origin "$BRANCH"
gh pr create \ gh pr create \
--title "chore: update WASM artifact checksums and version-pinned URLs" \ --title "chore: update WASM artifact SHA256 checksums" \
--body "Auto-generated by release CI. Updates SHA256 checksums and version-pinned artifact URLs in registry manifests to match the released WASM artifacts. Only extensions whose version changed since the last release are included." \ --body "Auto-generated by release CI. Updates SHA256 checksums in registry manifests to match the released WASM artifacts." \
--base main \ --base main \
--head "$BRANCH" --head "$BRANCH"
fi fi
+12 -5
View File
@@ -406,6 +406,10 @@ jobs:
echo "passed=true" >> "$GITHUB_OUTPUT" echo "passed=true" >> "$GITHUB_OUTPUT"
fi fi
# Only merge PRs targeting main. Chained PRs (targeting another
# promotion branch) stay open — when the base PR merges into main,
# GitHub auto-retargets the chained PR. Merging chained PRs would
# trigger delete_branch_on_merge, auto-closing downstream PRs.
- name: Merge promotion PR - name: Merge promotion PR
id: merge id: merge
if: steps.evaluate.outputs.passed == 'true' if: steps.evaluate.outputs.passed == 'true'
@@ -414,12 +418,15 @@ jobs:
PR_NUMBER: ${{ needs.create-promotion-pr.outputs.pr_number }} PR_NUMBER: ${{ needs.create-promotion-pr.outputs.pr_number }}
run: | run: |
if [ -n "$PR_NUMBER" ]; then if [ -n "$PR_NUMBER" ]; then
echo "Merging promotion PR #${PR_NUMBER}" BASE=$(gh pr view "$PR_NUMBER" --json baseRefName --jq '.baseRefName')
# Do NOT use --delete-branch: deleting a promotion branch closes if [ "$BASE" = "main" ]; then
# any chained PRs that use it as their base (verified in ironclaw-ci-test). echo "Merging promotion PR #${PR_NUMBER} (targets main)"
# Stale promotion branches are cleaned up separately.
gh pr merge "$PR_NUMBER" --merge gh pr merge "$PR_NUMBER" --merge
echo "merged=true" >> "$GITHUB_OUTPUT" echo "merged=true" >> "$GITHUB_OUTPUT"
else
echo "PR #${PR_NUMBER} targets '${BASE}' (not main) — leaving open for chain resolution"
echo "merged=false" >> "$GITHUB_OUTPUT"
fi
fi fi
# ── Update tested tag (always, so next batch covers only new commits) ── # ── Update tested tag (always, so next batch covers only new commits) ──
@@ -437,7 +444,7 @@ jobs:
- uses: actions/checkout@v6 - uses: actions/checkout@v6
with: with:
ref: staging ref: staging
fetch-depth: 1 fetch-depth: 0
- name: Update staging-tested tag - name: Update staging-tested tag
run: | run: |
-1
View File
@@ -28,4 +28,3 @@ trace_*.json
# Local Claude Code settings (machine-specific, should not be committed) # Local Claude Code settings (machine-specific, should not be committed)
.claude/settings.local.json .claude/settings.local.json
.worktrees/
-9
View File
@@ -7,15 +7,6 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
## [Unreleased] ## [Unreleased]
## [0.18.0](https://github.com/nearai/ironclaw/compare/v0.17.0...v0.18.0) - 2026-03-11
### Other
- Merge pull request #907 from nearai/staging-promote/b0214fef-22930316561
- promote staging to main (2026-03-10 15:19 UTC) ([#865](https://github.com/nearai/ironclaw/pull/865))
- Merge pull request #830 from nearai/staging-promote/3a2989d0-22888378864
- update WASM artifact SHA256 checksums [skip ci] ([#876](https://github.com/nearai/ironclaw/pull/876))
## [0.17.0](https://github.com/nearai/ironclaw/compare/v0.16.1...v0.17.0) - 2026-03-10 ## [0.17.0](https://github.com/nearai/ironclaw/compare/v0.16.1...v0.17.0) - 2026-03-10
### Added ### Added
+4 -38
View File
@@ -64,13 +64,6 @@ src/
│ ├── repl.rs # Simple REPL (for testing) │ ├── repl.rs # Simple REPL (for testing)
│ ├── web/ # Web gateway (browser UI) — see src/channels/web/CLAUDE.md │ ├── web/ # Web gateway (browser UI) — see src/channels/web/CLAUDE.md
│ └── wasm/ # WASM channel runtime │ └── wasm/ # WASM channel runtime
│ ├── mod.rs
│ ├── bundled.rs # Bundled channel discovery
│ ├── capabilities.rs # Channel-specific capabilities (HTTP endpoint, emit rate)
│ ├── error.rs # WASM channel error types
│ ├── runtime.rs # WASM channel execution runtime
│ ├── setup.rs # WasmChannelSetup, setup_wasm_channels(), inject_channel_credentials()
│ └── wrapper.rs # Channel trait wrapper for WASM modules
├── cli/ # CLI subcommands (clap) ├── cli/ # CLI subcommands (clap)
│ ├── mod.rs # Cli struct, Command enum (run/onboard/config/tool/registry/mcp/memory/pairing/service/doctor/status/completion) │ ├── mod.rs # Cli struct, Command enum (run/onboard/config/tool/registry/mcp/memory/pairing/service/doctor/status/completion)
@@ -83,13 +76,7 @@ src/
├── hooks/ # Lifecycle hooks (6 points: BeforeInbound, BeforeToolCall, BeforeOutbound, OnSessionStart, OnSessionEnd, TransformResponse) ├── hooks/ # Lifecycle hooks (6 points: BeforeInbound, BeforeToolCall, BeforeOutbound, OnSessionStart, OnSessionEnd, TransformResponse)
├── tunnel/ # Tunnel abstraction for public internet exposure ├── tunnel/ # Tunnel abstraction (cloudflare, ngrok, tailscale, custom, none)
│ ├── mod.rs # Tunnel trait, TunnelProviderConfig, create_tunnel(), start_managed_tunnel()
│ ├── cloudflare.rs # CloudflareTunnel (cloudflared binary)
│ ├── ngrok.rs # NgrokTunnel
│ ├── tailscale.rs # TailscaleTunnel (serve/funnel modes)
│ ├── custom.rs # CustomTunnel (arbitrary command with {host}/{port})
│ └── none.rs # NoneTunnel (local-only, no exposure)
├── observability/ # Pluggable event/metric recording (noop, log, multi) ├── observability/ # Pluggable event/metric recording (noop, log, multi)
@@ -99,8 +86,7 @@ src/
│ └── job_manager.rs # Container lifecycle (create, stop, cleanup) │ └── job_manager.rs # Container lifecycle (create, stop, cleanup)
├── worker/ # Runs inside Docker containers ├── worker/ # Runs inside Docker containers
│ ├── container.rs # Container worker runtime (ContainerDelegate + shared agentic loop) │ ├── runtime.rs # Worker execution loop (tool calls, LLM)
│ ├── job.rs # Background job worker (JobDelegate + shared agentic loop)
│ ├── claude_bridge.rs # Claude Code bridge (spawns claude CLI) │ ├── claude_bridge.rs # Claude Code bridge (spawns claude CLI)
│ └── proxy_llm.rs # LlmProvider that proxies through orchestrator │ └── proxy_llm.rs # LlmProvider that proxies through orchestrator
@@ -119,26 +105,8 @@ src/
│ ├── rate_limiter.rs # Shared sliding-window rate limiter │ ├── rate_limiter.rs # Shared sliding-window rate limiter
│ ├── builtin/ # Built-in tools (echo, time, json, http, web_fetch, file, shell, memory, message, job, routine, extension_tools, skill_tools, secrets_tools) │ ├── builtin/ # Built-in tools (echo, time, json, http, web_fetch, file, shell, memory, message, job, routine, extension_tools, skill_tools, secrets_tools)
│ ├── builder/ # Dynamic tool building │ ├── builder/ # Dynamic tool building
│ ├── core.rs # BuildRequirement, SoftwareType, Language │ ├── mcp/ # Model Context Protocol client
│ ├── templates.rs # Project scaffolding └── wasm/ # Full WASM sandbox (wasmtime) — runtime, host functions, fuel metering, allowlist, credential injection
│ │ ├── testing.rs # Test harness integration
│ │ └── validation.rs # WASM validation
│ ├── mcp/ # Model Context Protocol
│ │ ├── client.rs # MCP client over HTTP
│ │ ├── factory.rs # create_client_from_config() — transport dispatch factory
│ │ ├── protocol.rs # JSON-RPC types
│ │ └── session.rs # MCP session management (Mcp-Session-Id header, per-server state)
│ └── wasm/ # Full WASM sandbox (wasmtime)
│ ├── runtime.rs # Module compilation and caching
│ ├── wrapper.rs # Tool trait wrapper for WASM modules
│ ├── host.rs # Host functions (logging, time, workspace)
│ ├── limits.rs # Fuel metering and memory limiting
│ ├── allowlist.rs # Network endpoint allowlisting
│ ├── credential_injector.rs # Safe credential injection
│ ├── loader.rs # WASM tool discovery from filesystem
│ ├── rate_limiter.rs # Per-tool rate limiting
│ ├── error.rs # WASM-specific error types
│ └── storage.rs # Linear memory persistence
├── db/ # Dual-backend persistence (PostgreSQL + libSQL) — see src/db/CLAUDE.md ├── db/ # Dual-backend persistence (PostgreSQL + libSQL) — see src/db/CLAUDE.md
@@ -176,8 +144,6 @@ Dual-backend: PostgreSQL + libSQL/Turso. **All new persistence features must sup
When modifying a module with a spec, read the spec first. Code follows spec; spec is the tiebreaker. When modifying a module with a spec, read the spec first. Code follows spec; spec is the tiebreaker.
**Module-owned initialization:** Module-specific initialization logic (database connection, transport creation, channel setup) must live in the owning module as a public factory function — not in `main.rs` or `app.rs`. These entry-point files orchestrate calls to module factories. Feature-flag branching (`#[cfg(feature = ...)]`) must be confined to the module that owns the abstraction.
| Module | Spec | | Module | Spec |
|--------|------| |--------|------|
| `src/agent/` | `src/agent/CLAUDE.md` | | `src/agent/` | `src/agent/CLAUDE.md` |
-49
View File
@@ -1,34 +1,5 @@
# Contributing # Contributing
## Getting Started
```bash
git clone https://github.com/nearai/ironclaw.git
cd ironclaw
./scripts/dev-setup.sh
```
This installs the Rust toolchain, WASM targets, git hooks, and runs initial checks.
## Development Workflow
```bash
cargo fmt # format
cargo clippy --all --benches --tests --examples --all-features # lint (zero warnings)
cargo test # unit tests
cargo test --features integration # + PostgreSQL tests
```
## Code Style
- Zero clippy warnings policy
- No `.unwrap()` or `.expect()` in production code (tests are fine)
- Use `thiserror` for error types, map errors with context
- Prefer `crate::` for cross-module imports
- Comments for non-obvious logic only
See `CLAUDE.md` for full style guidelines.
## Feature Parity Requirement ## Feature Parity Requirement
When your change affects a tracked capability, update `FEATURE_PARITY.md` in the same branch. When your change affects a tracked capability, update `FEATURE_PARITY.md` in the same branch.
@@ -38,23 +9,3 @@ When your change affects a tracked capability, update `FEATURE_PARITY.md` in the
1. Review the relevant parity rows in `FEATURE_PARITY.md`. 1. Review the relevant parity rows in `FEATURE_PARITY.md`.
2. Update status/notes if behavior changed. 2. Update status/notes if behavior changed.
3. Include the `FEATURE_PARITY.md` diff in your commit when applicable. 3. Include the `FEATURE_PARITY.md` diff in your commit when applicable.
## Review Tracks
All PRs follow a risk-based review process:
| Track | Scope | Requirements |
|-------|-------|-------------|
| **A** | Docs, tests, chore, dependency bumps | 1 approval + CI green |
| **B** | Features, refactors, new tools/channels | 1 approval + CI green + test evidence |
| **C** | Security (`src/safety/`, `src/secrets/`), runtime (`src/agent/`, `src/worker/`), database schema, CI workflows | 2 approvals + rollback plan documented |
Select the appropriate track in the PR template based on what your changes touch.
## Database Changes
IronClaw uses dual-backend persistence (PostgreSQL + libSQL). All new persistence features must support both backends. See `src/db/CLAUDE.md`.
## Adding Dependencies
Run `cargo deny check` before adding new dependencies to verify license compatibility and check for known advisories.
+4 -4
View File
@@ -63,12 +63,12 @@ These files account for the vast majority of the coverage gap:
| `src/main.rs` | 740 | 522 | 29.4% | 485 | | `src/main.rs` | 740 | 522 | 29.4% | 485 |
| `src/channels/web/handlers/jobs.rs` | 513 | 456 | 11.1% | 430 | | `src/channels/web/handlers/jobs.rs` | 513 | 456 | 11.1% | 430 |
| `src/tools/builder/core.rs` | 524 | 456 | 13.0% | 429 | | `src/tools/builder/core.rs` | 524 | 456 | 13.0% | 429 |
| `src/worker/job.rs` | 1,078 | 467 | 56.7% | 413 | | `src/agent/worker.rs` | 1,078 | 467 | 56.7% | 413 |
| `src/channels/web/handlers/chat.rs` | 564 | 417 | 26.1% | 388 | | `src/channels/web/handlers/chat.rs` | 564 | 417 | 26.1% | 388 |
| `src/tools/wasm/wrapper.rs` | 1,005 | 436 | 56.6% | 385 | | `src/tools/wasm/wrapper.rs` | 1,005 | 436 | 56.6% | 385 |
| `src/channels/signal.rs` | 1,814 | 472 | 74.0% | 381 | | `src/channels/signal.rs` | 1,814 | 472 | 74.0% | 381 |
| `src/tools/mcp/auth.rs` | 472 | 378 | 19.9% | 354 | | `src/tools/mcp/auth.rs` | 472 | 378 | 19.9% | 354 |
| `src/worker/container.rs` | 350 | 330 | 5.7% | 312 | | `src/worker/runtime.rs` | 350 | 330 | 5.7% | 312 |
| `src/tools/builtin/job.rs` | 1,014 | 359 | 64.6% | 308 | | `src/tools/builtin/job.rs` | 1,014 | 359 | 64.6% | 308 |
| `src/cli/mcp.rs` | 322 | 319 | 0.9% | 302 | | `src/cli/mcp.rs` | 322 | 319 | 0.9% | 302 |
| `src/cli/oauth_defaults.rs` | 730 | 335 | 54.1% | 298 | | `src/cli/oauth_defaults.rs` | 730 | 335 | 54.1% | 298 |
@@ -346,7 +346,7 @@ Test slash commands through the agent loop.
### Trace: Worker Multi-Turn Execution ### Trace: Worker Multi-Turn Execution
**Covers:** `worker/job.rs` (+413 lines), `agent/agent_loop.rs` (+207 lines) **Covers:** `agent/worker.rs` (+413 lines), `agent/agent_loop.rs` (+207 lines)
Test multi-turn tool calling, error recovery, and completion flows. Test multi-turn tool calling, error recovery, and completion flows.
@@ -769,7 +769,7 @@ HTTP proxy for container network access.
- `test_proxy_connect_tunnel` -- HTTPS CONNECT method handling - `test_proxy_connect_tunnel` -- HTTPS CONNECT method handling
- `test_proxy_logging` -- request/response logging - `test_proxy_logging` -- request/response logging
### `src/worker/container.rs` -- 5.7% -> 95% (+312 lines) ### `src/worker/runtime.rs` -- 5.7% -> 95% (+312 lines)
Worker execution loop (runs inside containers). Worker execution loop (runs inside containers).
Generated
+1 -1
View File
@@ -3350,7 +3350,7 @@ dependencies = [
[[package]] [[package]]
name = "ironclaw" name = "ironclaw"
version = "0.18.0" version = "0.17.0"
dependencies = [ dependencies = [
"aes-gcm", "aes-gcm",
"aho-corasick", "aho-corasick",
+2 -7
View File
@@ -14,12 +14,11 @@ exclude = [
"tools-src/google-slides", "tools-src/google-slides",
"tools-src/slack", "tools-src/slack",
"tools-src/telegram", "tools-src/telegram",
"fuzz",
] ]
[package] [package]
name = "ironclaw" name = "ironclaw"
version = "0.18.0" version = "0.17.0"
edition = "2024" edition = "2024"
rust-version = "1.92" rust-version = "1.92"
description = "Secure personal AI assistant that protects your data and expands its capabilities on the fly" description = "Secure personal AI assistant that protects your data and expands its capabilities on the fly"
@@ -215,14 +214,10 @@ bedrock = ["dep:aws-config", "dep:aws-sdk-bedrockruntime", "dep:aws-smithy-types
name = "html_to_markdown" name = "html_to_markdown"
required-features = ["html-to-markdown"] required-features = ["html-to-markdown"]
[profile.release]
strip = true # Remove debug symbols from release binaries
# The profile that 'cargo dist' will build with # The profile that 'cargo dist' will build with
[profile.dist] [profile.dist]
inherits = "release" inherits = "release"
lto = "fat" # Full cross-crate LTO (slow build, better codegen) lto = "thin"
codegen-units = 1 # Single codegen unit for maximum optimization
# Config for 'dist' # Config for 'dist'
[workspace.metadata.dist] [workspace.metadata.dist]
+3 -3
View File
@@ -46,14 +46,14 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
| Bonjour/mDNS discovery | ✅ | ❌ | | | Bonjour/mDNS discovery | ✅ | ❌ | |
| Tailscale integration | ✅ | ❌ | | | Tailscale integration | ✅ | ❌ | |
| Health check endpoints | ✅ | ✅ | /api/health + /api/gateway/status + /healthz + /readyz, with channel-backed readiness probes | | Health check endpoints | ✅ | ✅ | /api/health + /api/gateway/status + /healthz + /readyz, with channel-backed readiness probes |
| `doctor` diagnostics | ✅ | 🚧 | 16 checks: settings, LLM, DB, embeddings, routines, gateway, MCP, skills, secrets, service, Docker daemon, tunnel binaries | | `doctor` diagnostics | ✅ | | |
| Agent event broadcast | ✅ | 🚧 | SSE broadcast manager exists (SseManager) but tool/job-state events not fully wired | | Agent event broadcast | ✅ | 🚧 | SSE broadcast manager exists (SseManager) but tool/job-state events not fully wired |
| Channel health monitor | ✅ | ❌ | Auto-restart with configurable interval | | Channel health monitor | ✅ | ❌ | Auto-restart with configurable interval |
| Presence system | ✅ | ❌ | Beacons on connect, system presence for agents | | Presence system | ✅ | ❌ | Beacons on connect, system presence for agents |
| Trusted-proxy auth mode | ✅ | ❌ | Header-based auth for reverse proxies | | Trusted-proxy auth mode | ✅ | ❌ | Header-based auth for reverse proxies |
| APNs push pipeline | ✅ | ❌ | Wake disconnected iOS nodes via push | | APNs push pipeline | ✅ | ❌ | Wake disconnected iOS nodes via push |
| Oversized payload guard | ✅ | 🚧 | HTTP webhook has 64KB body limit + Content-Length check; no chat.history cap | | Oversized payload guard | ✅ | 🚧 | HTTP webhook has 64KB body limit + Content-Length check; no chat.history cap |
| Pre-prompt context diagnostics | ✅ | 🚧 | Token breakdown logged before LLM call (conversational dispatcher path); other LLM entry points not yet covered | | Pre-prompt context diagnostics | ✅ | | Context size logging before prompt |
### Owner: _Unassigned_ ### Owner: _Unassigned_
@@ -175,7 +175,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
| `message send` | ✅ | ❌ | P2 | Send to channels | | `message send` | ✅ | ❌ | P2 | Send to channels |
| `browser` | ✅ | ❌ | P3 | Browser automation | | `browser` | ✅ | ❌ | P3 | Browser automation |
| `sandbox` | ✅ | ✅ | - | WASM sandbox | | `sandbox` | ✅ | ✅ | - | WASM sandbox |
| `doctor` | ✅ | 🚧 | P2 | 16 subsystem checks | | `doctor` | ✅ | | P2 | Diagnostics |
| `logs` | ✅ | ❌ | P3 | Query logs | | `logs` | ✅ | ❌ | P3 | Query logs |
| `update` | ✅ | ❌ | P3 | Self-update | | `update` | ✅ | ❌ | P3 | Self-update |
| `completion` | ✅ | ✅ | - | Shell completion | | `completion` | ✅ | ✅ | - | Shell completion |
-40
View File
@@ -1,40 +0,0 @@
[package]
name = "ironclaw-fuzz"
version = "0.0.0"
publish = false
edition = "2021"
[package.metadata]
cargo-fuzz = true
[dependencies]
libfuzzer-sys = "0.4"
serde_json = "1"
[dependencies.ironclaw]
path = ".."
[[bin]]
name = "fuzz_safety_sanitizer"
path = "fuzz_targets/fuzz_safety_sanitizer.rs"
doc = false
[[bin]]
name = "fuzz_safety_validator"
path = "fuzz_targets/fuzz_safety_validator.rs"
doc = false
[[bin]]
name = "fuzz_leak_detector"
path = "fuzz_targets/fuzz_leak_detector.rs"
doc = false
[[bin]]
name = "fuzz_tool_params"
path = "fuzz_targets/fuzz_tool_params.rs"
doc = false
[[bin]]
name = "fuzz_config_env"
path = "fuzz_targets/fuzz_config_env.rs"
doc = false
-43
View File
@@ -1,43 +0,0 @@
# IronClaw Fuzz Targets
Fuzz testing for security-critical input parsing paths using [cargo-fuzz](https://github.com/rust-fuzz/cargo-fuzz) (libFuzzer).
## Targets
| Target | What it exercises |
|--------|-------------------|
| `fuzz_safety_sanitizer` | Prompt injection pattern detection (Aho-Corasick + regex) |
| `fuzz_safety_validator` | Input validation (length, encoding, forbidden patterns) |
| `fuzz_leak_detector` | Secret leak detection (API keys, tokens, credentials) |
| `fuzz_tool_params` | Tool parameter and schema JSON validation |
| `fuzz_config_env` | SafetyLayer end-to-end (sanitize, validate, policy check) |
## Setup
```bash
cargo install cargo-fuzz
rustup install nightly
```
## Running
```bash
# Run a specific target (runs until stopped or crash found)
cargo +nightly fuzz run fuzz_safety_sanitizer
# Run with a time limit (5 minutes)
cargo +nightly fuzz run fuzz_leak_detector -- -max_total_time=300
# Run all targets for 60 seconds each
for target in fuzz_safety_sanitizer fuzz_safety_validator fuzz_leak_detector fuzz_tool_params fuzz_config_env; do
echo "==> $target"
cargo +nightly fuzz run "$target" -- -max_total_time=60
done
```
## Adding New Targets
1. Create `fuzz/fuzz_targets/fuzz_<name>.rs` following the existing pattern
2. Add a `[[bin]]` entry in `fuzz/Cargo.toml`
3. Create `fuzz/corpus/fuzz_<name>/` for seed inputs
4. Exercise real IronClaw code paths, not just generic serde
-55
View File
@@ -1,55 +0,0 @@
#![no_main]
use libfuzzer_sys::fuzz_target;
use ironclaw::safety::{LeakDetector, Sanitizer, Validator};
fuzz_target!(|data: &[u8]| {
if let Ok(input) = std::str::from_utf8(data) {
// Exercise Sanitizer: detect and neutralize prompt injection attempts.
let sanitizer = Sanitizer::new();
let sanitized = sanitizer.sanitize(input);
// The sanitized content must never be empty when input is non-empty,
// because sanitization wraps/escapes rather than deleting.
if !input.is_empty() {
assert!(
!sanitized.content.is_empty(),
"sanitize() produced empty content for non-empty input"
);
}
// If no modification occurred, content must equal input.
if !sanitized.was_modified {
assert_eq!(sanitized.content, input);
}
// Exercise Validator: input validation (length, encoding, patterns).
let validator = Validator::new();
let result = validator.validate(input);
// ValidationResult must always be well-formed: if valid, no errors.
if result.is_valid {
assert!(
result.errors.is_empty(),
"valid result should have no errors"
);
}
// Exercise LeakDetector: secret detection (API keys, tokens, etc.).
let detector = LeakDetector::new();
let scan = detector.scan(input);
// scan_and_clean must not panic and must return valid UTF-8.
let cleaned = detector.scan_and_clean(input);
if let Ok(ref clean_str) = cleaned {
// Cleaned output must never be longer than original + redaction markers.
// At minimum it should be valid UTF-8 (guaranteed by String type).
let _ = clean_str.len();
}
// If scan found no matches, scan_and_clean should return the input unchanged.
if scan.matches.is_empty() {
if let Ok(ref clean_str) = cleaned {
assert_eq!(
clean_str, input,
"scan_and_clean changed content despite no matches"
);
}
}
}
});
-23
View File
@@ -1,23 +0,0 @@
#![no_main]
use libfuzzer_sys::fuzz_target;
use ironclaw::safety::LeakDetector;
fuzz_target!(|data: &[u8]| {
if let Ok(s) = std::str::from_utf8(data) {
let detector = LeakDetector::new();
// Exercise scan path
let result = detector.scan(s);
// Invariant: if should_block, there must be matches
if result.should_block {
assert!(!result.matches.is_empty());
}
// Invariant: match locations must be valid
for m in &result.matches {
assert!(m.location.end <= s.len());
}
// Exercise scan_and_clean path
let _ = detector.scan_and_clean(s);
}
});
@@ -1,23 +0,0 @@
#![no_main]
use libfuzzer_sys::fuzz_target;
use ironclaw::safety::Sanitizer;
fuzz_target!(|data: &[u8]| {
if let Ok(s) = std::str::from_utf8(data) {
let sanitizer = Sanitizer::new();
// Exercise the main sanitization path
let result = sanitizer.sanitize(s);
// Verify invariant: warnings should have valid ranges
for w in &result.warnings {
assert!(w.location.end <= s.len());
}
// Verify invariant: critical severity triggers modification
let has_critical = result.warnings.iter().any(|w| {
w.severity == ironclaw::safety::Severity::Critical
});
if has_critical {
assert!(result.was_modified);
}
}
});
@@ -1,21 +0,0 @@
#![no_main]
use libfuzzer_sys::fuzz_target;
use ironclaw::safety::Validator;
fuzz_target!(|data: &[u8]| {
if let Ok(s) = std::str::from_utf8(data) {
let validator = Validator::new();
// Exercise input validation
let result = validator.validate(s);
// Invariant: empty input is always invalid
if s.is_empty() {
assert!(!result.is_valid);
}
// Exercise tool parameter validation with arbitrary JSON
if let Ok(value) = serde_json::from_str::<serde_json::Value>(s) {
let _ = validator.validate_tool_params(&value);
}
}
});
-22
View File
@@ -1,22 +0,0 @@
#![no_main]
use libfuzzer_sys::fuzz_target;
use ironclaw::safety::Validator;
use ironclaw::tools::validate_tool_schema;
fuzz_target!(|data: &[u8]| {
if let Ok(s) = std::str::from_utf8(data) {
// Try parsing as JSON and validating as tool parameters
if let Ok(value) = serde_json::from_str::<serde_json::Value>(s) {
// Exercise Validator::validate_tool_params with arbitrary JSON
let validator = Validator::new();
let result = validator.validate_tool_params(&value);
// Invariant: result should always be well-formed
if !result.is_valid {
assert!(!result.errors.is_empty());
}
// Exercise validate_tool_schema with arbitrary JSON as a schema
let _ = validate_tool_schema(&value, "fuzz");
}
}
});
-7
View File
@@ -1,7 +0,0 @@
-- Add token budget tracking columns to agent_jobs.
--
-- Tracks max_tokens (configured limit per job) and total_tokens_used (running total)
-- to enforce job-level token budgets and prevent budget bypass via user-supplied metadata.
ALTER TABLE agent_jobs ADD COLUMN max_tokens BIGINT NOT NULL DEFAULT 0;
ALTER TABLE agent_jobs ADD COLUMN total_tokens_used BIGINT NOT NULL DEFAULT 0;
+1 -1
View File
@@ -19,7 +19,7 @@
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/discord-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/discord-wasm32-wasip2.tar.gz",
"sha256": null "sha256": "85b424604482da3fb9badb56a0360ff4c93670bc7be0ad7f57ef9d85ff972b6f"
} }
}, },
"auth_summary": { "auth_summary": {
+1 -1
View File
@@ -19,7 +19,7 @@
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/slack-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/slack-wasm32-wasip2.tar.gz",
"sha256": null "sha256": "9190b8250bd20c22a8c97b1ea19a6590624a69d6c63a5f5c240a7840a4966286"
} }
}, },
"auth_summary": { "auth_summary": {
+1 -1
View File
@@ -19,7 +19,7 @@
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/telegram-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/telegram-wasm32-wasip2.tar.gz",
"sha256": null "sha256": "55f2a56e7afd129a48fd49b019f12f9638705defa53fa323ad3b8978d7c59664"
} }
}, },
"auth_summary": { "auth_summary": {
+1 -1
View File
@@ -19,7 +19,7 @@
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/whatsapp-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/whatsapp-wasm32-wasip2.tar.gz",
"sha256": null "sha256": "06bcf315df93af9f683134f4055eb810c602863d8c4a632e3733a10217cc5a89"
} }
}, },
"auth_summary": { "auth_summary": {
+1 -1
View File
@@ -20,7 +20,7 @@
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/github-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/github-wasm32-wasip2.tar.gz",
"sha256": null "sha256": "c443328a3f10b6a4cf4d3d62c9217aca204f6467ef753d986b58ca966ca53514"
} }
}, },
"auth_summary": { "auth_summary": {
+1 -1
View File
@@ -19,7 +19,7 @@
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/gmail-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/gmail-wasm32-wasip2.tar.gz",
"sha256": null "sha256": "e4f0095890d22e3de8e9d516f2e1e91964f8ff4acdaaa19f0a7094a1f2d7786b"
} }
}, },
"auth_summary": { "auth_summary": {
+1 -1
View File
@@ -19,7 +19,7 @@
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-calendar-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/google-calendar-wasm32-wasip2.tar.gz",
"sha256": null "sha256": "2d202bd838de94677c91ea6473c7155f021c0500cf91794d17639b1b27446b3d"
} }
}, },
"auth_summary": { "auth_summary": {
+1 -1
View File
@@ -19,7 +19,7 @@
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-docs-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/google-docs-wasm32-wasip2.tar.gz",
"sha256": null "sha256": "7a5e40fe58199e34f7625e11d22e5601cdfd2a94a10193a83f1925180bbb66df"
} }
}, },
"auth_summary": { "auth_summary": {
+1 -1
View File
@@ -19,7 +19,7 @@
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-drive-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/google-drive-wasm32-wasip2.tar.gz",
"sha256": null "sha256": "d19f856fde0ae0320fd3f636a34116af1df0b59698c3684b686e8412a60e887f"
} }
}, },
"auth_summary": { "auth_summary": {
+1 -1
View File
@@ -19,7 +19,7 @@
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-sheets-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/google-sheets-wasm32-wasip2.tar.gz",
"sha256": null "sha256": "e113c317f9fa21ea68d0ec8accbba4a62a8222ff3c4655ae85e1e58e01de3250"
} }
}, },
"auth_summary": { "auth_summary": {
+1 -1
View File
@@ -18,7 +18,7 @@
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-slides-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/google-slides-wasm32-wasip2.tar.gz",
"sha256": null "sha256": "7875a5ae1283e57937e0618bf14465f4bb4ee7f49110312382670202f4c567a5"
} }
}, },
"auth_summary": { "auth_summary": {
+1 -1
View File
@@ -18,7 +18,7 @@
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/slack-tool-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/slack-tool-wasm32-wasip2.tar.gz",
"sha256": null "sha256": "9190b8250bd20c22a8c97b1ea19a6590624a69d6c63a5f5c240a7840a4966286"
} }
}, },
"auth_summary": { "auth_summary": {
+1 -1
View File
@@ -19,7 +19,7 @@
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/telegram-mtproto-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/telegram-mtproto-wasm32-wasip2.tar.gz",
"sha256": null "sha256": "55f2a56e7afd129a48fd49b019f12f9638705defa53fa323ad3b8978d7c59664"
} }
}, },
"auth_summary": { "auth_summary": {
+1 -1
View File
@@ -19,7 +19,7 @@
"artifacts": { "artifacts": {
"wasm32-wasip2": { "wasm32-wasip2": {
"url": "https://github.com/nearai/ironclaw/releases/latest/download/web-search-wasm32-wasip2.tar.gz", "url": "https://github.com/nearai/ironclaw/releases/latest/download/web-search-wasm32-wasip2.tar.gz",
"sha256": null "sha256": "dd7e54956ee0b3037ca3506dbcbd20efcc4cd2749175ed511b3640b09f77506a"
} }
}, },
"auth_summary": { "auth_summary": {
@@ -1,80 +0,0 @@
---
name: ironclaw-workflow-orchestrator
description: "Install and operate a full GitHub issue-to-merge workflow in IronClaw using event-driven and cron routines. Use when setting up or tuning autonomous project orchestration: issue intake, planning, maintainer feedback handling, branch/PR execution, CI/comment follow-up, batched staging review every 8 hours, and memory updates from merge outcomes."
---
# IronClaw Workflow Orchestrator
## Overview
Use this skill to install and maintain a complete project workflow as routines, not core code changes. It maps GitHub webhook events plus scheduled checks into plan/update/implement/review/merge loops with explicit staging-batch analysis.
## Workflow
1. Gather workflow parameters.
2. Verify runtime prerequisites.
3. Install or update routine set from templates.
4. Run a dry test with `event_emit`.
5. Monitor outcomes and tune prompts/filters.
## Parameters
Collect these values before creating routines:
- `repository`: `owner/repo` (required)
- `maintainers`: GitHub handles allowed to trigger implement/replan actions
- `staging_branch`: default `staging`
- `main_branch`: default `main`
- `batch_interval_hours`: default `8`
- `implementation_label`: default `autonomous-impl`
## Prerequisites
Before installing routines, verify:
- Routines system enabled.
- GitHub tool authenticated (for issue/PR/comment/status operations).
- Events are emitted via `event_emit` tool calls (a future HTTP webhook ingestion endpoint is planned but not yet available).
## Install Procedure
1. Open [`workflow-routines.md`](references/workflow-routines.md).
2. For each template block:
- replace placeholders (`{{repository}}`, `{{maintainers}}`, branch names)
- call `routine_create`
3. If a routine already exists:
- use `routine_update` instead of creating duplicates
- keep names stable so long-lived metrics/history stay intact
4. Confirm install with `routine_list` and `routine_history`.
## Routine Set
Install these routines:
- `wf-issue-plan`: on `issue.opened` or `issue.reopened`, generate implementation plan comment/checklist.
- `wf-maintainer-comment-gate`: on maintainer comments, decide update-plan vs start implementation.
- `wf-pr-monitor-loop`: on PR open/sync/review-comment/review, address feedback and refresh branch.
- `wf-ci-fix-loop`: on CI status/check failures, apply fixes and push updates.
- `wf-staging-batch-review`: every 8h, review ready PRs, merge into staging, run deep batch correctness analysis, fix findings, then merge staging -> main.
- `wf-learning-memory`: on merged PRs, extract mistakes/lessons and write to shared memory.
## Event Filters
Prefer top-level filters for stability:
- `repository` (string)
- `sender` (string)
- `issue_number` / `pr_number`
- `ci_status`, `ci_conclusion`
- `review_state`, `comment_author`
Use narrow filters to avoid accidental triggers across repos.
## Operating Rules
- All implementation work must occur on non-main branches.
- PR loop must resolve both human and AI review comments.
- On conflicts with `origin/main`, refresh branch before continuing.
- Staging-batch routine is the only path for bulk correctness verification before mainline merge.
- Memory update routine runs only after successful merge.
## Validation
After install, run:
1. `event_emit` with a synthetic `issue.opened` payload for the target repo.
2. Confirm at least one routine fired.
3. Check corresponding `routine_history` entries.
4. Confirm no unrelated routines fired.
## When To Update Templates
Update this skill when:
- GitHub event names/payload fields change.
- Team review policy changes (e.g., staging cadence, maintainer gates).
- New CI policy requires different failure routing.
@@ -1,4 +0,0 @@
interface:
display_name: "IronClaw Workflow Orchestrator"
short_description: "Install and run event-driven GitHub workflow routines"
default_prompt: "Set up the full issue-to-merge workflow using routines and event triggers."
@@ -1,128 +0,0 @@
# Workflow Routine Templates
Replace `{{...}}` placeholders before use.
## 1) Issue -> Plan
```json
{
"name": "wf-issue-plan",
"description": "Create implementation plan when a new issue arrives",
"trigger_type": "system_event",
"event_source": "github",
"event_type": "issue.opened",
"event_filters": {
"repository": "{{repository}}"
},
"action_type": "full_job",
"prompt": "For issue #{{issue_number}} in {{repository}}, produce a concrete implementation plan with milestones, edge cases, and tests. Post/update an issue comment with the plan.",
"cooldown_secs": 30
}
```
## 2) Maintainer Comment Gate (Update Plan vs Implement)
Trigger per-maintainer by creating one routine per handle, or maintain a shared author convention.
```json
{
"name": "wf-maintainer-comment-gate-{{maintainer}}",
"description": "React to maintainer guidance comments on issues/PRs",
"trigger_type": "system_event",
"event_source": "github",
"event_type": "pr.comment.created",
"event_filters": {
"repository": "{{repository}}",
"comment_author": "{{maintainer}}"
},
"action_type": "full_job",
"prompt": "Read the maintainer comment and decide: update plan or start/continue implementation. If plan changes are requested, edit the plan artifact first. If implementation is requested, continue on the feature branch and update PR status/comment.",
"cooldown_secs": 20
}
```
## 3) PR Monitor Loop
```json
{
"name": "wf-pr-monitor-loop",
"description": "Keep PR healthy: address review comments and refresh branch",
"trigger_type": "system_event",
"event_source": "github",
"event_type": "pr.synchronize",
"event_filters": {
"repository": "{{repository}}"
},
"action_type": "full_job",
"prompt": "For PR #{{pr_number}}, collect open review comments and unresolved threads, apply fixes, push branch updates, and summarize remaining blockers. If conflict with {{main_branch}}, rebase/merge from origin/{{main_branch}} and resolve safely.",
"cooldown_secs": 20
}
```
## 4) CI Failure Fix Loop
```json
{
"name": "wf-ci-fix-loop",
"description": "Fix failing CI checks on active PRs",
"trigger_type": "system_event",
"event_source": "github",
"event_type": "ci.check_run.completed",
"event_filters": {
"repository": "{{repository}}",
"ci_conclusion": "failure"
},
"action_type": "full_job",
"prompt": "Find failing check details for PR #{{pr_number}}, implement minimal safe fixes, rerun or await CI, and post concise status updates. Prioritize deterministic and test-backed fixes.",
"cooldown_secs": 20
}
```
## 5) Staging Batch Review (Every 8h)
```json
{
"name": "wf-staging-batch-review",
"description": "Batch correctness review through staging, then merge to main",
"trigger_type": "cron",
"schedule": "0 0 */{{batch_interval_hours}} * * *",
"action_type": "full_job",
"prompt": "Every cycle: list ready PRs, merge ready ones into {{staging_branch}}, run deep correctness analysis in batch, fix discovered issues on affected branches, ensure CI green, then merge {{staging_branch}} into {{main_branch}} if clean.",
"cooldown_secs": 120
}
```
## 6) Post-Merge Learning -> Common Memory
```json
{
"name": "wf-learning-memory",
"description": "Capture merge learnings into shared memory",
"trigger_type": "system_event",
"event_source": "github",
"event_type": "pr.closed",
"event_filters": {
"repository": "{{repository}}",
"pr_merged": "true"
},
"action_type": "full_job",
"prompt": "From merged PR #{{pr_number}}, extract preventable mistakes, reviewer themes, CI failure causes, and successful patterns. Write/update a shared memory doc with actionable rules to reduce cycle time and regressions.",
"cooldown_secs": 30
}
```
## Optional: Synthetic Event Test
```json
{
"source": "github",
"event_type": "issue.opened",
"payload": {
"repository": "{{repository}}",
"issue_number": 99999,
"sender": "test-bot"
}
}
```
Use with `event_emit` after routine install.
+17 -20
View File
@@ -14,15 +14,14 @@ Core agent logic. This is the most complex subsystem — read this before workin
| `session_manager.rs` | Lifecycle: create/lookup sessions, map external thread IDs to internal UUIDs, prune stale sessions, manage undo managers. | | `session_manager.rs` | Lifecycle: create/lookup sessions, map external thread IDs to internal UUIDs, prune stale sessions, manage undo managers. |
| `router.rs` | Routes explicit `/commands` to `MessageIntent`. Natural language bypasses the router entirely. | | `router.rs` | Routes explicit `/commands` to `MessageIntent`. Natural language bypasses the router entirely. |
| `scheduler.rs` | Parallel job scheduling. Maintains `jobs` map (full LLM-driven) and `subtasks` map (tool-exec/background). | | `scheduler.rs` | Parallel job scheduling. Maintains `jobs` map (full LLM-driven) and `subtasks` map (tool-exec/background). |
| *(moved to `src/worker/job.rs`)* | Per-job execution now lives in `src/worker/job.rs` as `JobDelegate`, using the shared `run_agentic_loop()` engine. | | `worker.rs` | Per-job execution for background scheduler jobs: calls LLM, runs tools, handles the reasoning loop. Distinct from `dispatcher.rs`. |
| `agentic_loop.rs` | Shared agentic loop engine: `run_agentic_loop()`, `LoopDelegate` trait, `LoopOutcome`, `LoopSignal`, `TextAction`. All three execution paths (chat, job, container) delegate to this. |
| `compaction.rs` | Context window management: summarize old turns, write to workspace daily log, trim context. Three strategies. | | `compaction.rs` | Context window management: summarize old turns, write to workspace daily log, trim context. Three strategies. |
| `context_monitor.rs` | Detects memory pressure. Suggests `CompactionStrategy` based on usage level. | | `context_monitor.rs` | Detects memory pressure. Suggests `CompactionStrategy` based on usage level. |
| `self_repair.rs` | Detects stuck jobs and broken tools, attempts recovery. | | `self_repair.rs` | Detects stuck jobs and broken tools, attempts recovery. |
| `heartbeat.rs` | Proactive periodic execution. Reads `HEARTBEAT.md`, notifies via channel if findings. | | `heartbeat.rs` | Proactive periodic execution. Reads `HEARTBEAT.md`, notifies via channel if findings. |
| `submission.rs` | Parses all user submissions into typed variants before routing. | | `submission.rs` | Parses all user submissions into typed variants before routing. |
| `undo.rs` | Turn-based undo/redo with checkpoints. Checkpoints store message lists (max 20 by default). | | `undo.rs` | Turn-based undo/redo with checkpoints. Checkpoints store message lists (max 20 by default). |
| `routine.rs` | `Routine` types: `Trigger` (cron/event/system_event/manual) + `RoutineAction` (lightweight/full_job) + `RoutineGuardrails`. | | `routine.rs` | `Routine` types: `Trigger` (cron/event/webhook/manual) + `RoutineAction` (lightweight/full_job) + `RoutineGuardrails`. |
| `routine_engine.rs` | Cron ticker and event matcher. Fires routines when triggers match. Lightweight runs inline; full_job dispatches to `Scheduler`. | | `routine_engine.rs` | Cron ticker and event matcher. Fires routines when triggers match. Lightweight runs inline; full_job dispatches to `Scheduler`. |
| `task.rs` | Task types for the scheduler: `Job`, `ToolExec`, `Background`. Used by `spawn_subtask` and `spawn_batch`. | | `task.rs` | Task types for the scheduler: `Job`, `ToolExec`, `Background`. Used by `spawn_subtask` and `spawn_batch`. |
| `cost_guard.rs` | LLM spend and action-rate enforcement. Tracks daily budget (cents) and hourly call rate. Lives in `AgentDeps`. | | `cost_guard.rs` | LLM spend and action-rate enforcement. Tracks daily budget (cents) and hourly call rate. Lives in `AgentDeps`. |
@@ -50,28 +49,26 @@ Session (per user)
## Agentic Loop (dispatcher.rs) ## Agentic Loop (dispatcher.rs)
All three execution paths (chat, job, container) now use the shared `run_agentic_loop()` engine in `agentic_loop.rs`, each providing their own `LoopDelegate` implementation: The `dispatcher.rs` module handles **direct conversational turns** (user messages processed inline by the main agent). Background scheduler jobs use `worker.rs` instead — these are two separate execution paths.
- **`ChatDelegate`** (`dispatcher.rs`) — conversational turns, tool approval, skill context injection
- **`JobDelegate`** (`src/worker/job.rs`) — background scheduler jobs, planning support, completion detection
- **`ContainerDelegate`** (`src/worker/container.rs`) — Docker container worker, sequential tool exec, HTTP event streaming
``` ```
run_agentic_loop(delegate, reasoning, reason_ctx, config) run_agentic_loop() [dispatcher.rs — conversational turns]
1. Check signals (stop/cancel) via delegate.check_signals() 1. Load workspace system prompt (identity files: AGENTS.md, SOUL.md, etc.)
2. Pre-LLM hook via delegate.before_llm_call() 2. Detect group chat from metadata; exclude MEMORY.md if group chat
3. LLM call via delegate.call_llm() 3. Select active skills (keyword/pattern scoring against message content)
4. If text response → delegate.handle_text_response() → Continue or Return 4. Build skill context block (injected before user message)
5. If tool callsdelegate.execute_tool_calls() → Continue or Return 5. LLM call → text response OR tool calls
6. Post-iteration hook via delegate.after_iteration() 6. If tool calls:
7. Repeat until LoopOutcome returned or max_iterations reached a. Check tool approval (session auto-approvals, pending approval queue)
b. Execute tools (parallel via JoinSet)
c. Sanitize results through SafetyLayer
d. Feed results back → goto 5
7. Return AgenticLoopResult::Response or NeedApproval
``` ```
**Tool approval:** Tools flagged `requires_approval` pause the loop `ChatDelegate` returns `LoopOutcome::NeedApproval(pending)`. The web gateway stores the `PendingApproval` in session state and sends an `approval_needed` SSE event. The user's approval/deny resumes the loop. **Tool approval:** Tools flagged `requires_approval` pause the loop and return `NeedApproval`. The web gateway stores the `PendingApproval` in session state and sends an `approval_needed` SSE event. The user's approval/deny resumes the loop.
**Shared tool execution:** `tools/execute.rs` provides `execute_tool_with_safety()` (validate → timeout → execute → serialize) and `process_tool_result()` (sanitize → wrap → ChatMessage), used by all three delegates. **worker.rs vs dispatcher.rs:** `dispatcher.rs` runs the agentic loop for user-initiated conversational turns (holds session lock, tracks turns). `worker.rs` is spawned by the `Scheduler` for background jobs created via `CreateJob` / `/job` — it runs independently of the session and has its own LLM reasoning loop with planning support (`use_planning` flag).
**ChatDelegate vs JobDelegate:** `ChatDelegate` runs for user-initiated conversational turns (holds session lock, tracks turns). `JobDelegate` is spawned by the `Scheduler` for background jobs created via `CreateJob` / `/job` — it runs independently of the session and has planning support (`use_planning` flag).
## Command Routing (router.rs) ## Command Routing (router.rs)
+8 -34
View File
@@ -516,7 +516,7 @@ impl Agent {
*slot.write().await = Some(Arc::clone(&engine)); *slot.write().await = Some(Arc::clone(&engine));
} }
tracing::debug!( tracing::info!(
"Routines enabled: cron ticker every {}s, max {} concurrent", "Routines enabled: cron ticker every {}s, max {} concurrent",
rt_config.cron_check_interval_secs, rt_config.cron_check_interval_secs,
rt_config.max_concurrent_routines rt_config.max_concurrent_routines
@@ -538,20 +538,20 @@ impl Agent {
let routine_engine_for_loop = routine_handle.as_ref().map(|(_, e)| Arc::clone(e)); let routine_engine_for_loop = routine_handle.as_ref().map(|(_, e)| Arc::clone(e));
// Main message loop // Main message loop
tracing::debug!("Agent {} ready and listening", self.config.name); tracing::info!("Agent {} ready and listening", self.config.name);
loop { loop {
let message = tokio::select! { let message = tokio::select! {
biased; biased;
_ = tokio::signal::ctrl_c() => { _ = tokio::signal::ctrl_c() => {
tracing::debug!("Ctrl+C received, shutting down..."); tracing::info!("Ctrl+C received, shutting down...");
break; break;
} }
msg = message_stream.next() => { msg = message_stream.next() => {
match msg { match msg {
Some(m) => m, Some(m) => m,
None => { None => {
tracing::debug!("All channel streams ended, shutting down..."); tracing::info!("All channel streams ended, shutting down...");
break; break;
} }
} }
@@ -626,7 +626,7 @@ impl Agent {
} }
Ok(None) => { Ok(None) => {
// Shutdown signal received (/quit, /exit, /shutdown) // Shutdown signal received (/quit, /exit, /shutdown)
tracing::debug!("Shutdown command received, exiting..."); tracing::info!("Shutdown command received, exiting...");
break; break;
} }
Err(e) => { Err(e) => {
@@ -655,7 +655,7 @@ impl Agent {
} }
// Cleanup // Cleanup
tracing::debug!("Agent shutting down..."); tracing::info!("Agent shutting down...");
repair_handle.abort(); repair_handle.abort();
pruning_handle.abort(); pruning_handle.abort();
if let Some(handle) = heartbeat_handle { if let Some(handle) = heartbeat_handle {
@@ -738,18 +738,6 @@ impl Agent {
} }
async fn handle_message(&self, message: &IncomingMessage) -> Result<Option<String>, Error> { async fn handle_message(&self, message: &IncomingMessage) -> Result<Option<String>, Error> {
// Log at info level only for tracking without exposing PII (user_id can be a phone number)
tracing::info!(message_id = %message.id, "Processing message");
// Log sensitive details at debug level for troubleshooting
tracing::debug!(
message_id = %message.id,
user_id = %message.user_id,
channel = %message.channel,
thread_id = ?message.thread_id,
"Message details"
);
// Set message tool context for this turn (current channel and target) // Set message tool context for this turn (current channel and target)
// For Signal, use signal_target from metadata (group:ID or phone number), // For Signal, use signal_target from metadata (group:ID or phone number),
// otherwise fall back to user_id // otherwise fall back to user_id
@@ -765,7 +753,7 @@ impl Agent {
// Parse submission type first // Parse submission type first
let mut submission = SubmissionParser::parse(&message.content); let mut submission = SubmissionParser::parse(&message.content);
tracing::trace!( tracing::debug!(
"[agent_loop] Parsed submission: {:?}", "[agent_loop] Parsed submission: {:?}",
std::any::type_name_of_val(&submission) std::any::type_name_of_val(&submission)
); );
@@ -798,19 +786,10 @@ impl Agent {
// Hydrate thread from DB if it's a historical thread not in memory // Hydrate thread from DB if it's a historical thread not in memory
if let Some(ref external_thread_id) = message.thread_id { if let Some(ref external_thread_id) = message.thread_id {
tracing::trace!(
message_id = %message.id,
thread_id = %external_thread_id,
"Hydrating thread from DB"
);
self.maybe_hydrate_thread(message, external_thread_id).await; self.maybe_hydrate_thread(message, external_thread_id).await;
} }
// Resolve session and thread // Resolve session and thread
tracing::debug!(
message_id = %message.id,
"Resolving session and thread"
);
let (session, thread_id) = self let (session, thread_id) = self
.session_manager .session_manager
.resolve_thread( .resolve_thread(
@@ -819,11 +798,6 @@ impl Agent {
message.thread_id.as_deref(), message.thread_id.as_deref(),
) )
.await; .await;
tracing::debug!(
message_id = %message.id,
thread_id = %thread_id,
"Resolved session and thread"
);
// Auth mode interception: if the thread is awaiting a token, route // Auth mode interception: if the thread is awaiting a token, route
// the message directly to the credential store. Nothing touches // the message directly to the credential store. Nothing touches
@@ -853,7 +827,7 @@ impl Agent {
} }
} }
tracing::trace!( tracing::debug!(
"Received message from {} on {} ({} chars)", "Received message from {} on {} ({} chars)",
message.user_id, message.user_id,
message.channel, message.channel,
-587
View File
@@ -1,587 +0,0 @@
//! Unified agentic loop engine.
//!
//! Provides a single implementation of the core LLM call → tool execution →
//! result processing → context update → repeat cycle. Three consumers
//! (chat dispatcher, job worker, container runtime) customize behavior
//! via the `LoopDelegate` trait.
use async_trait::async_trait;
use crate::agent::session::PendingApproval;
use crate::error::Error;
use crate::llm::{ChatMessage, Reasoning, ReasoningContext, RespondResult};
/// Signal from the delegate indicating how the loop should proceed.
pub enum LoopSignal {
/// Continue normally.
Continue,
/// Stop the loop gracefully.
Stop,
/// Inject a user message into context and continue.
InjectMessage(String),
}
/// Outcome of a text response from the LLM.
pub enum TextAction {
/// Return this as the final loop result.
Return(LoopOutcome),
/// Continue the loop (text was handled but loop should proceed).
Continue,
}
/// Final outcome of the agentic loop.
pub enum LoopOutcome {
/// Completed with a text response.
Response(String),
/// Loop was stopped by a signal.
Stopped,
/// Max iterations exceeded.
MaxIterations,
/// A tool requires user approval before continuing (chat delegate only).
NeedApproval(Box<PendingApproval>),
}
/// Configuration for the agentic loop.
pub struct AgenticLoopConfig {
pub max_iterations: usize,
pub enable_tool_intent_nudge: bool,
pub max_tool_intent_nudges: u32,
}
impl Default for AgenticLoopConfig {
fn default() -> Self {
Self {
max_iterations: 50,
enable_tool_intent_nudge: true,
max_tool_intent_nudges: 2,
}
}
}
/// Strategy trait — each consumer implements this to customize I/O and lifecycle.
///
/// The shared loop calls these methods at well-defined points. Consumers
/// implement only the behavior that differs between chat, job, and container
/// contexts. The loop itself handles the common logic: tool intent nudge,
/// iteration counting, tool definition refresh, and the respond → execute → process cycle.
///
/// # `Send + Sync` requirement
///
/// This trait requires `Send + Sync` because the loop accepts `&dyn LoopDelegate`.
/// Delegates using borrowed references (e.g. `ChatDelegate<'a>`) must ensure all
/// borrowed fields are `Send + Sync`. This is a load-bearing constraint: if a
/// delegate needs to be spawned into a detached task, it must use `Arc`-based
/// ownership instead of borrows (as `JobDelegate` and `ContainerDelegate` do).
#[async_trait]
pub trait LoopDelegate: Send + Sync {
/// Called at the start of each iteration. Check for external signals
/// (cancellation, user messages, stop requests).
async fn check_signals(&self) -> LoopSignal;
/// Called before the LLM call. Allows the delegate to refresh tool
/// definitions, enforce cost guards, or inject messages.
/// Return `Some(outcome)` to break the loop early.
async fn before_llm_call(
&self,
reason_ctx: &mut ReasoningContext,
iteration: usize,
) -> Option<LoopOutcome>;
/// Call the LLM and return the result. Delegates own the LLM call
/// to handle consumer-specific concerns (rate limiting, auto-compaction,
/// cost tracking, force_text mode).
async fn call_llm(
&self,
reasoning: &Reasoning,
reason_ctx: &mut ReasoningContext,
iteration: usize,
) -> Result<crate::llm::RespondOutput, Error>;
/// Handle a text-only response from the LLM.
/// Return `TextAction::Return` to exit the loop, `TextAction::Continue` to proceed.
async fn handle_text_response(
&self,
text: &str,
reason_ctx: &mut ReasoningContext,
) -> TextAction;
/// Execute tool calls and add results to context.
/// Return `Some(outcome)` to break the loop (e.g. approval needed).
async fn execute_tool_calls(
&self,
tool_calls: Vec<crate::llm::ToolCall>,
content: Option<String>,
reason_ctx: &mut ReasoningContext,
) -> Result<Option<LoopOutcome>, Error>;
/// Called when the LLM expresses tool intent without actually calling a tool.
/// Delegates can use this to emit events or log the nudge for observability.
async fn on_tool_intent_nudge(&self, _text: &str, _reason_ctx: &mut ReasoningContext) {}
/// Called after each successful iteration (no error, no early return).
async fn after_iteration(&self, _iteration: usize) {}
}
/// Run the unified agentic loop.
///
/// This is the single implementation used by all three consumers (chat, job, container).
/// The `delegate` provides consumer-specific behavior via the `LoopDelegate` trait.
pub async fn run_agentic_loop(
delegate: &dyn LoopDelegate,
reasoning: &Reasoning,
reason_ctx: &mut ReasoningContext,
config: &AgenticLoopConfig,
) -> Result<LoopOutcome, Error> {
let mut consecutive_tool_intent_nudges: u32 = 0;
for iteration in 1..=config.max_iterations {
// Check for external signals (stop, cancellation, user messages)
match delegate.check_signals().await {
LoopSignal::Continue => {}
LoopSignal::Stop => return Ok(LoopOutcome::Stopped),
LoopSignal::InjectMessage(msg) => {
reason_ctx.messages.push(ChatMessage::user(&msg));
}
}
// Pre-LLM call hook (cost guard, tool refresh, iteration limit nudge)
if let Some(outcome) = delegate.before_llm_call(reason_ctx, iteration).await {
return Ok(outcome);
}
// Call LLM
let output = delegate.call_llm(reasoning, reason_ctx, iteration).await?;
match output.result {
RespondResult::Text(text) => {
// Tool intent nudge: if the LLM says "let me search..." without
// actually calling a tool, inject a nudge message.
if config.enable_tool_intent_nudge
&& !reason_ctx.available_tools.is_empty()
&& !reason_ctx.force_text
&& consecutive_tool_intent_nudges < config.max_tool_intent_nudges
&& crate::llm::llm_signals_tool_intent(&text)
{
consecutive_tool_intent_nudges += 1;
tracing::info!(
iteration,
"LLM expressed tool intent without calling a tool, nudging"
);
delegate.on_tool_intent_nudge(&text, reason_ctx).await;
reason_ctx.messages.push(ChatMessage::assistant(&text));
reason_ctx
.messages
.push(ChatMessage::user(crate::llm::TOOL_INTENT_NUDGE));
delegate.after_iteration(iteration).await;
continue;
}
// Reset nudge counter since we got a non-intent text response
if !crate::llm::llm_signals_tool_intent(&text) {
consecutive_tool_intent_nudges = 0;
}
match delegate.handle_text_response(&text, reason_ctx).await {
TextAction::Return(outcome) => return Ok(outcome),
TextAction::Continue => {}
}
}
RespondResult::ToolCalls {
tool_calls,
content,
} => {
consecutive_tool_intent_nudges = 0;
if let Some(outcome) = delegate
.execute_tool_calls(tool_calls, content, reason_ctx)
.await?
{
return Ok(outcome);
}
}
}
delegate.after_iteration(iteration).await;
}
Ok(LoopOutcome::MaxIterations)
}
/// Truncate a string for log/status previews.
///
/// `max` is a byte budget. The result is truncated at the last valid char
/// boundary at or before `max` bytes, so it is always valid UTF-8.
pub fn truncate_for_preview(s: &str, max: usize) -> String {
if s.len() <= max {
s.to_string()
} else {
let end = crate::util::floor_char_boundary(s, max);
format!("{}...", &s[..end])
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::llm::{RespondOutput, TokenUsage, ToolCall};
use crate::testing::StubLlm;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::sync::Mutex;
fn stub_reasoning() -> Reasoning {
Reasoning::new(Arc::new(StubLlm::default()))
}
fn zero_usage() -> TokenUsage {
TokenUsage {
input_tokens: 0,
output_tokens: 0,
cache_read_input_tokens: 0,
cache_creation_input_tokens: 0,
}
}
fn text_output(text: &str) -> RespondOutput {
RespondOutput {
result: RespondResult::Text(text.to_string()),
usage: zero_usage(),
}
}
fn tool_calls_output(calls: Vec<ToolCall>) -> RespondOutput {
RespondOutput {
result: RespondResult::ToolCalls {
tool_calls: calls,
content: None,
},
usage: zero_usage(),
}
}
/// Configurable mock delegate for testing run_agentic_loop.
struct MockDelegate {
signal: Mutex<LoopSignal>,
llm_responses: Mutex<Vec<RespondOutput>>,
tool_exec_count: AtomicUsize,
tool_exec_outcome: Mutex<Option<LoopOutcome>>,
iterations_seen: Mutex<Vec<usize>>,
early_exit: Mutex<Option<(usize, LoopOutcome)>>,
nudge_count: AtomicUsize,
}
impl MockDelegate {
fn new(responses: Vec<RespondOutput>) -> Self {
Self {
signal: Mutex::new(LoopSignal::Continue),
llm_responses: Mutex::new(responses),
tool_exec_count: AtomicUsize::new(0),
tool_exec_outcome: Mutex::new(None),
iterations_seen: Mutex::new(Vec::new()),
early_exit: Mutex::new(None),
nudge_count: AtomicUsize::new(0),
}
}
fn with_signal(mut self, signal: LoopSignal) -> Self {
self.signal = Mutex::new(signal);
self
}
fn with_early_exit(mut self, iteration: usize, outcome: LoopOutcome) -> Self {
self.early_exit = Mutex::new(Some((iteration, outcome)));
self
}
}
#[async_trait]
impl LoopDelegate for MockDelegate {
async fn check_signals(&self) -> LoopSignal {
let mut sig = self.signal.lock().await;
std::mem::replace(&mut *sig, LoopSignal::Continue)
}
async fn before_llm_call(
&self,
_reason_ctx: &mut ReasoningContext,
iteration: usize,
) -> Option<LoopOutcome> {
let mut guard = self.early_exit.lock().await;
let should_take = guard
.as_ref()
.is_some_and(|(target, _)| *target == iteration);
if should_take {
guard.take().map(|(_, o)| o)
} else {
None
}
}
async fn call_llm(
&self,
_reasoning: &Reasoning,
_reason_ctx: &mut ReasoningContext,
_iteration: usize,
) -> Result<crate::llm::RespondOutput, crate::error::Error> {
let mut responses = self.llm_responses.lock().await;
if responses.is_empty() {
panic!("MockDelegate: no more LLM responses queued");
}
Ok(responses.remove(0))
}
async fn handle_text_response(
&self,
text: &str,
_reason_ctx: &mut ReasoningContext,
) -> TextAction {
TextAction::Return(LoopOutcome::Response(text.to_string()))
}
async fn execute_tool_calls(
&self,
_tool_calls: Vec<ToolCall>,
_content: Option<String>,
reason_ctx: &mut ReasoningContext,
) -> Result<Option<LoopOutcome>, crate::error::Error> {
self.tool_exec_count.fetch_add(1, Ordering::SeqCst);
reason_ctx
.messages
.push(ChatMessage::user("tool result stub"));
let outcome = self.tool_exec_outcome.lock().await.take();
Ok(outcome)
}
async fn on_tool_intent_nudge(&self, _text: &str, _reason_ctx: &mut ReasoningContext) {
self.nudge_count.fetch_add(1, Ordering::SeqCst);
}
async fn after_iteration(&self, iteration: usize) {
self.iterations_seen.lock().await.push(iteration);
}
}
// --- Tests ---
#[tokio::test]
async fn test_text_response_returns_immediately() {
let delegate = MockDelegate::new(vec![text_output("Hello, world!")]);
let reasoning = stub_reasoning();
let mut ctx = ReasoningContext::new();
let config = AgenticLoopConfig::default();
let outcome = run_agentic_loop(&delegate, &reasoning, &mut ctx, &config)
.await
.unwrap();
match outcome {
LoopOutcome::Response(text) => assert_eq!(text, "Hello, world!"),
_ => panic!("Expected LoopOutcome::Response"),
}
// after_iteration is NOT called when handle_text_response returns Return
// (the loop exits before reaching after_iteration).
assert!(delegate.iterations_seen.lock().await.is_empty());
}
#[tokio::test]
async fn test_tool_call_then_text_response() {
let tool_call = ToolCall {
id: "call_1".to_string(),
name: "echo".to_string(),
arguments: serde_json::json!({}),
};
let delegate = MockDelegate::new(vec![
tool_calls_output(vec![tool_call]),
text_output("Done!"),
]);
let reasoning = stub_reasoning();
let mut ctx = ReasoningContext::new();
let config = AgenticLoopConfig::default();
let outcome = run_agentic_loop(&delegate, &reasoning, &mut ctx, &config)
.await
.unwrap();
match outcome {
LoopOutcome::Response(text) => assert_eq!(text, "Done!"),
_ => panic!("Expected LoopOutcome::Response"),
}
assert_eq!(delegate.tool_exec_count.load(Ordering::SeqCst), 1);
// after_iteration called for iteration 1 (tool call), but not 2
// (text response exits before after_iteration).
assert_eq!(*delegate.iterations_seen.lock().await, vec![1]);
}
#[tokio::test]
async fn test_stop_signal_exits_immediately() {
let delegate =
MockDelegate::new(vec![text_output("unreachable")]).with_signal(LoopSignal::Stop);
let reasoning = stub_reasoning();
let mut ctx = ReasoningContext::new();
let config = AgenticLoopConfig::default();
let outcome = run_agentic_loop(&delegate, &reasoning, &mut ctx, &config)
.await
.unwrap();
assert!(matches!(outcome, LoopOutcome::Stopped));
assert!(delegate.iterations_seen.lock().await.is_empty());
}
#[tokio::test]
async fn test_inject_message_adds_user_message() {
let delegate = MockDelegate::new(vec![text_output("Got it")])
.with_signal(LoopSignal::InjectMessage("injected prompt".to_string()));
let reasoning = stub_reasoning();
let mut ctx = ReasoningContext::new();
let config = AgenticLoopConfig::default();
let outcome = run_agentic_loop(&delegate, &reasoning, &mut ctx, &config)
.await
.unwrap();
assert!(matches!(outcome, LoopOutcome::Response(_)));
assert!(
ctx.messages
.iter()
.any(|m| m.role == crate::llm::Role::User && m.content.contains("injected prompt")),
"Injected message should appear in context"
);
}
#[tokio::test]
async fn test_max_iterations_reached() {
struct ContinueDelegate;
#[async_trait]
impl LoopDelegate for ContinueDelegate {
async fn check_signals(&self) -> LoopSignal {
LoopSignal::Continue
}
async fn before_llm_call(
&self,
_: &mut ReasoningContext,
_: usize,
) -> Option<LoopOutcome> {
None
}
async fn call_llm(
&self,
_: &Reasoning,
_: &mut ReasoningContext,
_: usize,
) -> Result<crate::llm::RespondOutput, crate::error::Error> {
Ok(text_output("still working"))
}
async fn handle_text_response(
&self,
_: &str,
ctx: &mut ReasoningContext,
) -> TextAction {
ctx.messages.push(ChatMessage::assistant("still working"));
TextAction::Continue
}
async fn execute_tool_calls(
&self,
_: Vec<ToolCall>,
_: Option<String>,
_: &mut ReasoningContext,
) -> Result<Option<LoopOutcome>, crate::error::Error> {
Ok(None)
}
}
let delegate = ContinueDelegate;
let reasoning = stub_reasoning();
let mut ctx = ReasoningContext::new();
let config = AgenticLoopConfig {
max_iterations: 3,
..Default::default()
};
let outcome = run_agentic_loop(&delegate, &reasoning, &mut ctx, &config)
.await
.unwrap();
assert!(matches!(outcome, LoopOutcome::MaxIterations));
let assistant_count = ctx
.messages
.iter()
.filter(|m| m.role == crate::llm::Role::Assistant)
.count();
assert_eq!(assistant_count, 3);
}
#[tokio::test]
async fn test_tool_intent_nudge_fires_and_caps() {
let delegate = MockDelegate::new(vec![
text_output("Let me search for that file"),
text_output("Let me search for that file"),
text_output("Let me search for that file"),
]);
let reasoning = stub_reasoning();
let mut ctx = ReasoningContext::new();
ctx.available_tools.push(crate::llm::ToolDefinition {
name: "search".to_string(),
description: "Search files".to_string(),
parameters: serde_json::json!({"type": "object"}),
});
let config = AgenticLoopConfig {
max_iterations: 10,
enable_tool_intent_nudge: true,
max_tool_intent_nudges: 2,
};
let outcome = run_agentic_loop(&delegate, &reasoning, &mut ctx, &config)
.await
.unwrap();
assert!(matches!(outcome, LoopOutcome::Response(_)));
assert_eq!(delegate.nudge_count.load(Ordering::SeqCst), 2);
let nudge_messages = ctx
.messages
.iter()
.filter(|m| {
m.role == crate::llm::Role::User
&& m.content.contains("you did not include any tool calls")
})
.count();
assert_eq!(
nudge_messages, 2,
"Should have exactly 2 nudge messages in context"
);
}
#[tokio::test]
async fn test_before_llm_call_early_exit() {
let delegate = MockDelegate::new(vec![text_output("unreachable")])
.with_early_exit(1, LoopOutcome::Stopped);
let reasoning = stub_reasoning();
let mut ctx = ReasoningContext::new();
let config = AgenticLoopConfig::default();
let outcome = run_agentic_loop(&delegate, &reasoning, &mut ctx, &config)
.await
.unwrap();
assert!(matches!(outcome, LoopOutcome::Stopped));
assert!(delegate.iterations_seen.lock().await.is_empty());
}
#[test]
fn test_truncate_short_string_unchanged() {
assert_eq!(truncate_for_preview("hello", 10), "hello");
}
#[test]
fn test_truncate_long_string_adds_ellipsis() {
let result = truncate_for_preview("hello world", 5);
assert_eq!(result, "hello...");
}
#[test]
fn test_truncate_multibyte_safe() {
let result = truncate_for_preview("café", 4);
assert_eq!(result, "caf...");
}
}
+2 -4
View File
@@ -405,8 +405,7 @@ impl Agent {
.with_max_tokens(512) .with_max_tokens(512)
.with_temperature(0.3); .with_temperature(0.3);
let reasoning = let reasoning = Reasoning::new(self.llm().clone());
Reasoning::new(self.llm().clone()).with_model_name(self.llm().active_model_name());
match reasoning.complete(request).await { match reasoning.complete(request).await {
Ok((text, _usage)) => Ok(SubmissionResult::response(format!( Ok((text, _usage)) => Ok(SubmissionResult::response(format!(
"Thread Summary:\n\n{}", "Thread Summary:\n\n{}",
@@ -454,8 +453,7 @@ impl Agent {
.with_max_tokens(512) .with_max_tokens(512)
.with_temperature(0.5); .with_temperature(0.5);
let reasoning = let reasoning = Reasoning::new(self.llm().clone());
Reasoning::new(self.llm().clone()).with_model_name(self.llm().active_model_name());
match reasoning.complete(request).await { match reasoning.complete(request).await {
Ok((text, _usage)) => Ok(SubmissionResult::response(format!( Ok((text, _usage)) => Ok(SubmissionResult::response(format!(
"Suggested Next Steps:\n\n{}", "Suggested Next Steps:\n\n{}",
+1 -2
View File
@@ -227,8 +227,7 @@ Be brief but capture all important details. Use bullet points."#,
.with_max_tokens(1024) .with_max_tokens(1024)
.with_temperature(0.3); .with_temperature(0.3);
let reasoning = let reasoning = Reasoning::new(self.llm.clone());
Reasoning::new(self.llm.clone()).with_model_name(self.llm.active_model_name());
let (text, _) = reasoning.complete(request).await?; let (text, _) = reasoning.complete(request).await?;
Ok(text) Ok(text)
} }
+311 -240
View File
@@ -14,12 +14,7 @@ use crate::agent::session::{PendingApproval, Session, ThreadState};
use crate::channels::{IncomingMessage, StatusUpdate}; use crate::channels::{IncomingMessage, StatusUpdate};
use crate::context::JobContext; use crate::context::JobContext;
use crate::error::Error; use crate::error::Error;
use async_trait::async_trait; use crate::llm::{ChatMessage, Reasoning, ReasoningContext, RespondResult};
use crate::agent::agentic_loop::{
AgenticLoopConfig, LoopDelegate, LoopOutcome, LoopSignal, TextAction,
};
use crate::llm::{ChatMessage, Reasoning, ReasoningContext};
use crate::tools::redact_params; use crate::tools::redact_params;
/// Result of the agentic loop execution. /// Result of the agentic loop execution.
@@ -90,7 +85,7 @@ impl Agent {
crate::skills::SkillTrust::Installed => "INSTALLED", crate::skills::SkillTrust::Installed => "INSTALLED",
}; };
tracing::debug!( tracing::info!(
skill_name = skill.name(), skill_name = skill.name(),
skill_version = skill.version(), skill_version = skill.version(),
trust = %skill.trust, trust = %skill.trust,
@@ -138,6 +133,9 @@ impl Agent {
reasoning = reasoning.with_skill_context(ctx); reasoning = reasoning.with_skill_context(ctx);
} }
// Build context with messages that we'll mutate during the loop
let mut context_messages = initial_messages;
// Create a JobContext for tool execution (chat doesn't have a real job) // Create a JobContext for tool execution (chat doesn't have a real job)
let mut job_ctx = let mut job_ctx =
JobContext::with_user(&message.user_id, "chat", "Interactive chat session"); JobContext::with_user(&message.user_id, "chat", "Interactive chat session");
@@ -156,118 +154,53 @@ impl Agent {
let cached_prompt_no_tools = reasoning.build_system_prompt_with_tools(&[]); let cached_prompt_no_tools = reasoning.build_system_prompt_with_tools(&[]);
let max_tool_iterations = self.config.max_tool_iterations; let max_tool_iterations = self.config.max_tool_iterations;
// Force a text-only response on the last iteration to guarantee termination
// instead of hard-erroring. The penultimate iteration also gets a nudge
// message so the LLM knows it should wrap up.
let force_text_at = max_tool_iterations; let force_text_at = max_tool_iterations;
let nudge_at = max_tool_iterations.saturating_sub(1); let nudge_at = max_tool_iterations.saturating_sub(1);
let mut iteration = 0;
let delegate = ChatDelegate { const MAX_TOOL_INTENT_NUDGES: u32 = 2;
agent: self, let mut consecutive_tool_intent_nudges: u32 = 0;
session: session.clone(), loop {
thread_id, iteration += 1;
message, // Hard ceiling one past the forced-text iteration (should never be reached
job_ctx, // since force_text_at guarantees a text response, but kept as a safety net).
active_skills, if iteration > max_tool_iterations + 1 {
cached_prompt, return Err(crate::error::LlmError::InvalidResponse {
cached_prompt_no_tools,
nudge_at,
force_text_at,
user_tz,
};
let mut reason_ctx = ReasoningContext::new()
.with_messages(initial_messages)
.with_tools(initial_tool_defs)
.with_system_prompt(delegate.cached_prompt.clone())
.with_metadata({
let mut m = std::collections::HashMap::new();
m.insert("thread_id".to_string(), thread_id.to_string());
m
});
let loop_config = AgenticLoopConfig {
// Hard ceiling: one past force_text_at (safety net).
max_iterations: max_tool_iterations + 1,
enable_tool_intent_nudge: true,
max_tool_intent_nudges: 2,
};
let outcome = crate::agent::agentic_loop::run_agentic_loop(
&delegate,
&reasoning,
&mut reason_ctx,
&loop_config,
)
.await?;
match outcome {
LoopOutcome::Response(text) => Ok(AgenticLoopResult::Response(text)),
LoopOutcome::Stopped => Err(crate::error::JobError::ContextError {
id: thread_id,
reason: "Interrupted".to_string(),
}
.into()),
LoopOutcome::MaxIterations => Err(crate::error::LlmError::InvalidResponse {
provider: "agent".to_string(), provider: "agent".to_string(),
reason: format!("Exceeded maximum tool iterations ({max_tool_iterations})"), reason: format!("Exceeded maximum tool iterations ({max_tool_iterations})"),
} }
.into()), .into());
LoopOutcome::NeedApproval(pending) => {
Ok(AgenticLoopResult::NeedApproval { pending: *pending })
}
}
} }
/// Execute a tool for chat (without full job context). // Check if interrupted
pub(super) async fn execute_chat_tool( {
&self, let sess = session.lock().await;
tool_name: &str, if let Some(thread) = sess.threads.get(&thread_id)
params: &serde_json::Value,
job_ctx: &JobContext,
) -> Result<String, Error> {
execute_chat_tool_standalone(self.tools(), self.safety(), tool_name, params, job_ctx).await
}
}
/// Delegate for the chat (dispatcher) context.
///
/// Implements `LoopDelegate` to customize the shared agentic loop for
/// interactive chat sessions with the full 3-phase tool execution
/// (preflight → parallel exec → post-flight), approval flow, hooks,
/// auth intercept, and cost tracking.
struct ChatDelegate<'a> {
agent: &'a Agent,
session: Arc<Mutex<Session>>,
thread_id: Uuid,
message: &'a IncomingMessage,
job_ctx: JobContext,
active_skills: Vec<crate::skills::LoadedSkill>,
cached_prompt: String,
cached_prompt_no_tools: String,
nudge_at: usize,
force_text_at: usize,
user_tz: chrono_tz::Tz,
}
#[async_trait]
impl<'a> LoopDelegate for ChatDelegate<'a> {
async fn check_signals(&self) -> LoopSignal {
let sess = self.session.lock().await;
if let Some(thread) = sess.threads.get(&self.thread_id)
&& thread.state == ThreadState::Interrupted && thread.state == ThreadState::Interrupted
{ {
return LoopSignal::Stop; return Err(crate::error::JobError::ContextError {
id: thread_id,
reason: "Interrupted".to_string(),
}
.into());
} }
LoopSignal::Continue
} }
async fn before_llm_call( // Enforce cost guardrails before the LLM call
&self, if let Err(limit) = self.cost_guard().check_allowed().await {
reason_ctx: &mut ReasoningContext, return Err(crate::error::LlmError::InvalidResponse {
iteration: usize, provider: "agent".to_string(),
) -> Option<LoopOutcome> { reason: limit.to_string(),
}
.into());
}
// Inject a nudge message when approaching the iteration limit so the // Inject a nudge message when approaching the iteration limit so the
// LLM is aware it should produce a final answer on the next turn. // LLM is aware it should produce a final answer on the next turn.
if iteration == self.nudge_at { if iteration == nudge_at {
reason_ctx.messages.push(ChatMessage::system( context_messages.push(ChatMessage::system(
"You are approaching the tool call limit. \ "You are approaching the tool call limit. \
Provide your best final answer on the next response \ Provide your best final answer on the next response \
using the information you have gathered so far. \ using the information you have gathered so far. \
@@ -275,15 +208,15 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
)); ));
} }
let force_text = iteration >= self.force_text_at; let force_text = iteration >= force_text_at;
// Refresh tool definitions each iteration so newly built tools become visible // Refresh tool definitions each iteration so newly built tools become visible
let tool_defs = self.agent.tools().tool_definitions().await; let tool_defs = self.tools().tool_definitions().await;
// Apply trust-based tool attenuation if skills are active. // Apply trust-based tool attenuation if skills are active.
let tool_defs = if !self.active_skills.is_empty() { let tool_defs = if !active_skills.is_empty() {
let result = crate::skills::attenuate_tools(&tool_defs, &self.active_skills); let result = crate::skills::attenuate_tools(&tool_defs, &active_skills);
tracing::debug!( tracing::info!(
min_trust = %result.min_trust, min_trust = %result.min_trust,
tools_available = result.tools.len(), tools_available = result.tools.len(),
tools_removed = result.removed_tools.len(), tools_removed = result.removed_tools.len(),
@@ -296,14 +229,23 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
tool_defs tool_defs
}; };
// Update context for this iteration // Call LLM with current context; force_text drops tools to guarantee a
reason_ctx.available_tools = tool_defs; // text response on the final iteration. The pre-built system prompt
reason_ctx.system_prompt = Some(if force_text { // avoids rebuilding the same ~1,500-token string each iteration.
self.cached_prompt_no_tools.clone() let mut context = ReasoningContext::new()
.with_messages(context_messages.clone())
.with_tools(tool_defs)
.with_system_prompt(if force_text {
cached_prompt_no_tools.clone()
} else { } else {
self.cached_prompt.clone() cached_prompt.clone()
})
.with_metadata({
let mut m = std::collections::HashMap::new();
m.insert("thread_id".to_string(), thread_id.to_string());
m
}); });
reason_ctx.force_text = force_text; context.force_text = force_text;
if force_text { if force_text {
tracing::info!( tracing::info!(
@@ -313,34 +255,15 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
} }
let _ = self let _ = self
.agent
.channels .channels
.send_status( .send_status(
&self.message.channel, &message.channel,
StatusUpdate::Thinking("Calling LLM...".into()), StatusUpdate::Thinking("Calling LLM...".into()),
&self.message.metadata, &message.metadata,
) )
.await; .await;
None let output = match reasoning.respond_with_tools(&context).await {
}
async fn call_llm(
&self,
reasoning: &Reasoning,
reason_ctx: &mut ReasoningContext,
iteration: usize,
) -> Result<crate::llm::RespondOutput, Error> {
// Enforce cost guardrails before the LLM call
if let Err(limit) = self.agent.cost_guard().check_allowed().await {
return Err(crate::error::LlmError::InvalidResponse {
provider: "agent".to_string(),
reason: limit.to_string(),
}
.into());
}
let output = match reasoning.respond_with_tools(reason_ctx).await {
Ok(output) => output, Ok(output) => output,
Err(crate::error::LlmError::ContextLengthExceeded { used, limit }) => { Err(crate::error::LlmError::ContextLengthExceeded { used, limit }) => {
tracing::warn!( tracing::warn!(
@@ -350,16 +273,23 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
"Context length exceeded, compacting messages and retrying" "Context length exceeded, compacting messages and retrying"
); );
// Compact messages in place and retry // Compact: keep system messages + last user message + current turn
reason_ctx.messages = compact_messages_for_retry(&reason_ctx.messages); context_messages = compact_messages_for_retry(&context_messages);
// When force_text, clear tools to further reduce token count // Rebuild context with compacted messages, reusing cached prompt
if reason_ctx.force_text { let mut retry_context = ReasoningContext::new()
reason_ctx.available_tools.clear(); .with_messages(context_messages.clone())
} .with_tools(if force_text {
Vec::new()
} else {
context.available_tools.clone()
})
.with_metadata(context.metadata.clone());
retry_context.force_text = force_text;
retry_context.system_prompt = context.system_prompt.clone();
reasoning reasoning
.respond_with_tools(reason_ctx) .respond_with_tools(&retry_context)
.await .await
.map_err(|retry_err| { .map_err(|retry_err| {
tracing::error!( tracing::error!(
@@ -368,6 +298,7 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
retry_error = %retry_err, retry_error = %retry_err,
"Retry after auto-compaction also failed" "Retry after auto-compaction also failed"
); );
// Propagate the actual retry error so callers see the real failure
crate::error::Error::from(retry_err) crate::error::Error::from(retry_err)
})? })?
} }
@@ -375,11 +306,10 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
}; };
// Record cost and track token usage // Record cost and track token usage
let model_name = self.agent.llm().active_model_name(); let model_name = self.llm().active_model_name();
let read_discount = self.agent.llm().cache_read_discount(); let read_discount = self.llm().cache_read_discount();
let write_multiplier = self.agent.llm().cache_write_multiplier(); let write_multiplier = self.llm().cache_write_multiplier();
let call_cost = self let call_cost = self
.agent
.cost_guard() .cost_guard()
.record_llm_call( .record_llm_call(
&model_name, &model_name,
@@ -389,7 +319,7 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
output.usage.cache_creation_input_tokens, output.usage.cache_creation_input_tokens,
read_discount, read_discount,
write_multiplier, write_multiplier,
Some(self.agent.llm().cost_per_token()), Some(self.llm().cost_per_token()),
) )
.await; .await;
tracing::debug!( tracing::debug!(
@@ -399,60 +329,72 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
call_cost, call_cost,
); );
Ok(output) match output.result {
RespondResult::Text(text) => {
// Nudge the LLM if it expressed tool intent without calling tools.
// This is common with non-Anthropic models (e.g. GLM-5 via NEAR AI)
// that output "Let me search…" but don't issue tool_calls.
if !force_text
&& !context.available_tools.is_empty()
&& consecutive_tool_intent_nudges < MAX_TOOL_INTENT_NUDGES
&& crate::llm::llm_signals_tool_intent(&text)
{
consecutive_tool_intent_nudges += 1;
tracing::info!(
iteration,
"LLM expressed tool intent without calling a tool, nudging"
);
context_messages.push(ChatMessage::assistant(&text));
context_messages.push(ChatMessage::user(crate::llm::TOOL_INTENT_NUDGE));
continue;
} }
async fn handle_text_response(
&self,
text: &str,
_reason_ctx: &mut ReasoningContext,
) -> TextAction {
// Strip internal "[Called tool ...]" text that can leak when // Strip internal "[Called tool ...]" text that can leak when
// provider flattening (e.g. NEAR AI) converts tool_calls to // provider flattening (e.g. NEAR AI) converts tool_calls to
// plain text and the LLM echoes it back. // plain text and the LLM echoes it back.
let sanitized = strip_internal_tool_call_text(text); let sanitized = strip_internal_tool_call_text(&text);
TextAction::Return(LoopOutcome::Response(sanitized)) return Ok(AgenticLoopResult::Response(sanitized));
} }
RespondResult::ToolCalls {
async fn execute_tool_calls( tool_calls,
&self, content,
tool_calls: Vec<crate::llm::ToolCall>, } => {
content: Option<String>, consecutive_tool_intent_nudges = 0;
reason_ctx: &mut ReasoningContext,
) -> Result<Option<LoopOutcome>, Error> {
// Add the assistant message with tool_calls to context. // Add the assistant message with tool_calls to context.
// OpenAI protocol requires this before tool-result messages. // OpenAI protocol requires this before tool-result messages.
reason_ctx context_messages.push(ChatMessage::assistant_with_tool_calls(
.messages
.push(ChatMessage::assistant_with_tool_calls(
content, content,
tool_calls.clone(), tool_calls.clone(),
)); ));
// Execute tools and add results to context // Execute tools and add results to context
let _ = self let _ = self
.agent
.channels .channels
.send_status( .send_status(
&self.message.channel, &message.channel,
StatusUpdate::Thinking(format!("Executing {} tool(s)...", tool_calls.len())), StatusUpdate::Thinking(format!(
&self.message.metadata, "Executing {} tool(s)...",
tool_calls.len()
)),
&message.metadata,
) )
.await; .await;
// Record tool calls in the thread with sensitive params redacted. // 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()); let mut redacted_args: Vec<serde_json::Value> =
Vec::with_capacity(tool_calls.len());
for tc in &tool_calls { for tc in &tool_calls {
let safe = if let Some(tool) = self.agent.tools().get(&tc.name).await { let safe = if let Some(tool) = self.tools().get(&tc.name).await {
redact_params(&tc.arguments, tool.sensitive_params()) redact_params(&tc.arguments, tool.sensitive_params())
} else { } else {
tc.arguments.clone() tc.arguments.clone()
}; };
redacted_args.push(safe); redacted_args.push(safe);
} }
let mut sess = self.session.lock().await; let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&self.thread_id) if let Some(thread) = sess.threads.get_mut(&thread_id)
&& let Some(turn) = thread.last_turn_mut() && let Some(turn) = thread.last_turn_mut()
{ {
for (tc, safe_args) in tool_calls.iter().zip(redacted_args) { for (tc, safe_args) in tool_calls.iter().zip(redacted_args) {
@@ -465,8 +407,13 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
// Walk tool_calls checking approval and hooks. Classify // Walk tool_calls checking approval and hooks. Classify
// each tool as Rejected (by hook) or Runnable. Stop at the // each tool as Rejected (by hook) or Runnable. Stop at the
// first tool that needs approval. // first tool that needs approval.
//
// Outcomes are indexed by original tool_calls position so
// Phase 3 can emit results in the correct order.
enum PreflightOutcome { enum PreflightOutcome {
/// Hook rejected/blocked this tool; contains the error message.
Rejected(String), Rejected(String),
/// Tool passed preflight and will be executed.
Runnable, Runnable,
} }
let mut preflight: Vec<(crate::llm::ToolCall, PreflightOutcome)> = Vec::new(); let mut preflight: Vec<(crate::llm::ToolCall, PreflightOutcome)> = Vec::new();
@@ -480,21 +427,26 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
for (idx, original_tc) in tool_calls.iter().enumerate() { for (idx, original_tc) in tool_calls.iter().enumerate() {
let mut tc = original_tc.clone(); let mut tc = original_tc.clone();
let tool_opt = self.agent.tools().get(&tc.name).await; // 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 let sensitive = tool_opt
.as_ref() .as_ref()
.map(|t| t.sensitive_params()) .map(|t| t.sensitive_params())
.unwrap_or(&[]); .unwrap_or(&[]);
// Hook: BeforeToolCall // Hook: BeforeToolCall (runs before approval so hooks can
// 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 hook_params = redact_params(&tc.arguments, sensitive);
let event = crate::hooks::HookEvent::ToolCall { let event = crate::hooks::HookEvent::ToolCall {
tool_name: tc.name.clone(), tool_name: tc.name.clone(),
parameters: hook_params, parameters: hook_params,
user_id: self.message.user_id.clone(), user_id: message.user_id.clone(),
context: "chat".to_string(), context: "chat".to_string(),
}; };
match self.agent.hooks().run(&event).await { match self.hooks().run(&event).await {
Err(crate::hooks::HookError::Rejected { reason }) => { Err(crate::hooks::HookError::Rejected { reason }) => {
preflight.push(( preflight.push((
tc, tc,
@@ -503,7 +455,7 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
reason reason
)), )),
)); ));
continue; continue; // skip to next tool (not infinite: using for loop)
} }
Err(err) => { Err(err) => {
preflight.push(( preflight.push((
@@ -519,9 +471,12 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
modified: Some(new_params), modified: Some(new_params),
}) => match serde_json::from_str::<serde_json::Value>(&new_params) { }) => match serde_json::from_str::<serde_json::Value>(&new_params) {
Ok(mut parsed) => { 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() { if let Some(obj) = parsed.as_object_mut() {
for key in sensitive { for key in sensitive {
if let Some(orig_val) = original_tc.arguments.get(*key) { if let Some(orig_val) = original_tc.arguments.get(*key)
{
obj.insert((*key).to_string(), orig_val.clone()); obj.insert((*key).to_string(), orig_val.clone());
} }
} }
@@ -539,15 +494,16 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
_ => {} _ => {}
} }
// Check if tool requires approval // Check if tool requires approval on the final (post-hook)
if !self.agent.config.auto_approve_tools // parameters. Skipped when auto_approve_tools is set.
if !self.config.auto_approve_tools
&& let Some(tool) = tool_opt && let Some(tool) = tool_opt
{ {
use crate::tools::ApprovalRequirement; use crate::tools::ApprovalRequirement;
let needs_approval = match tool.requires_approval(&tc.arguments) { let needs_approval = match tool.requires_approval(&tc.arguments) {
ApprovalRequirement::Never => false, ApprovalRequirement::Never => false,
ApprovalRequirement::UnlessAutoApproved => { ApprovalRequirement::UnlessAutoApproved => {
let sess = self.session.lock().await; let sess = session.lock().await;
!sess.is_tool_auto_approved(&tc.name) !sess.is_tool_auto_approved(&tc.name)
} }
ApprovalRequirement::Always => true, ApprovalRequirement::Always => true,
@@ -555,7 +511,7 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
if needs_approval { if needs_approval {
approval_needed = Some((idx, tc, tool)); approval_needed = Some((idx, tc, tool));
break; break; // remaining tools are deferred
} }
} }
@@ -565,58 +521,59 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
} }
// === Phase 2: Parallel execution === // === Phase 2: Parallel execution ===
// Execute runnable tools and slot results back by preflight
// index so Phase 3 can iterate in original order.
let mut exec_results: Vec<Option<Result<String, Error>>> = let mut exec_results: Vec<Option<Result<String, Error>>> =
(0..preflight.len()).map(|_| None).collect(); (0..preflight.len()).map(|_| None).collect();
if runnable.len() <= 1 { if runnable.len() <= 1 {
// Single tool (or none): execute inline
for (pf_idx, tc) in &runnable { for (pf_idx, tc) in &runnable {
let _ = self let _ = self
.agent
.channels .channels
.send_status( .send_status(
&self.message.channel, &message.channel,
StatusUpdate::ToolStarted { StatusUpdate::ToolStarted {
name: tc.name.clone(), name: tc.name.clone(),
}, },
&self.message.metadata, &message.metadata,
) )
.await; .await;
let result = self let result = self
.agent .execute_chat_tool(&tc.name, &tc.arguments, &job_ctx)
.execute_chat_tool(&tc.name, &tc.arguments, &self.job_ctx)
.await; .await;
let disp_tool = self.agent.tools().get(&tc.name).await; let disp_tool = self.tools().get(&tc.name).await;
let _ = self let _ = self
.agent
.channels .channels
.send_status( .send_status(
&self.message.channel, &message.channel,
StatusUpdate::tool_completed( StatusUpdate::tool_completed(
tc.name.clone(), tc.name.clone(),
&result, &result,
&tc.arguments, &tc.arguments,
disp_tool.as_deref(), disp_tool.as_deref(),
), ),
&self.message.metadata, &message.metadata,
) )
.await; .await;
exec_results[*pf_idx] = Some(result); exec_results[*pf_idx] = Some(result);
} }
} else { } else {
// Multiple tools: execute in parallel via JoinSet
let mut join_set = JoinSet::new(); let mut join_set = JoinSet::new();
for (pf_idx, tc) in &runnable { for (pf_idx, tc) in &runnable {
let pf_idx = *pf_idx; let pf_idx = *pf_idx;
let tools = self.agent.tools().clone(); let tools = self.tools().clone();
let safety = self.agent.safety().clone(); let safety = self.safety().clone();
let channels = self.agent.channels.clone(); let channels = self.channels.clone();
let job_ctx = self.job_ctx.clone(); let job_ctx = job_ctx.clone();
let tc = tc.clone(); let tc = tc.clone();
let channel = self.message.channel.clone(); let channel = message.channel.clone();
let metadata = self.message.metadata.clone(); let metadata = message.metadata.clone();
join_set.spawn(async move { join_set.spawn(async move {
let _ = channels let _ = channels
@@ -665,20 +622,25 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
if e.is_panic() { if e.is_panic() {
tracing::error!("Chat tool execution task panicked: {}", e); tracing::error!("Chat tool execution task panicked: {}", e);
} else { } else {
tracing::error!("Chat tool execution task cancelled: {}", e); tracing::error!(
"Chat tool execution task cancelled: {}",
e
);
} }
} }
} }
} }
// Fill panicked slots with error results // Fill panicked slots with error results
for (pf_idx, tc) in runnable.iter() { for (runnable_idx, (pf_idx, tc)) in runnable.iter().enumerate() {
if exec_results[*pf_idx].is_none() { if exec_results[*pf_idx].is_none() {
tracing::error!( tracing::error!(
tool = %tc.name, tool = %tc.name,
runnable_idx,
"Filling failed task slot with error" "Filling failed task slot with error"
); );
exec_results[*pf_idx] = Some(Err(crate::error::ToolError::ExecutionFailed { exec_results[*pf_idx] =
Some(Err(crate::error::ToolError::ExecutionFailed {
name: tc.name.clone(), name: tc.name.clone(),
reason: "Task failed during execution".to_string(), reason: "Task failed during execution".to_string(),
} }
@@ -688,25 +650,30 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
} }
// === Phase 3: Post-flight (sequential, in original order) === // === Phase 3: Post-flight (sequential, in original order) ===
// Process all results — both hook rejections and execution
// results — in the original tool_calls order. Auth intercept
// is deferred until after every result is recorded.
let mut deferred_auth: Option<String> = None; let mut deferred_auth: Option<String> = None;
for (pf_idx, (tc, outcome)) in preflight.into_iter().enumerate() { for (pf_idx, (tc, outcome)) in preflight.into_iter().enumerate() {
match outcome { match outcome {
PreflightOutcome::Rejected(error_msg) => { PreflightOutcome::Rejected(error_msg) => {
// Record hook rejection in thread
{ {
let mut sess = self.session.lock().await; let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&self.thread_id) if let Some(thread) = sess.threads.get_mut(&thread_id)
&& let Some(turn) = thread.last_turn_mut() && let Some(turn) = thread.last_turn_mut()
{ {
turn.record_tool_error(error_msg.clone()); turn.record_tool_error(error_msg.clone());
} }
} }
reason_ctx context_messages
.messages
.push(ChatMessage::tool_result(&tc.id, &tc.name, error_msg)); .push(ChatMessage::tool_result(&tc.id, &tc.name, error_msg));
} }
PreflightOutcome::Runnable => { PreflightOutcome::Runnable => {
let tool_result = exec_results[pf_idx].take().unwrap_or_else(|| { // Retrieve the execution result for this slot
let tool_result =
exec_results[pf_idx].take().unwrap_or_else(|| {
Err(crate::error::ToolError::ExecutionFailed { Err(crate::error::ToolError::ExecutionFailed {
name: tc.name.clone(), name: tc.name.clone(),
reason: "No result available".to_string(), reason: "No result available".to_string(),
@@ -714,11 +681,13 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
.into()) .into())
}); });
// Detect image generation sentinel // Detect image generation sentinel in tool output
// (only from image tools — avoids parsing all tool outputs)
let is_image_sentinel = if let Ok(ref output) = tool_result let is_image_sentinel = if let Ok(ref output) = tool_result
&& matches!(tc.name.as_str(), "image_generate" | "image_edit") && matches!(tc.name.as_str(), "image_generate" | "image_edit")
{ {
if let Ok(sentinel) = serde_json::from_str::<serde_json::Value>(output) if let Ok(sentinel) =
serde_json::from_str::<serde_json::Value>(output)
&& sentinel.get("type").and_then(|v| v.as_str()) && sentinel.get("type").and_then(|v| v.as_str())
== Some("image_generated") == Some("image_generated")
{ {
@@ -731,18 +700,19 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
.get("path") .get("path")
.and_then(|v| v.as_str()) .and_then(|v| v.as_str())
.map(String::from); .map(String::from);
// Skip broadcasting if data_url is empty to avoid
// sending a broken ImageGenerated SSE event.
if data_url.is_empty() { if data_url.is_empty() {
tracing::warn!( tracing::warn!(
"Image generation sentinel has empty data URL, skipping broadcast" "Image generation sentinel has empty data URL, skipping broadcast"
); );
} else { } else {
let _ = self let _ = self
.agent
.channels .channels
.send_status( .send_status(
&self.message.channel, &message.channel,
StatusUpdate::ImageGenerated { data_url, path }, StatusUpdate::ImageGenerated { data_url, path },
&self.message.metadata, &message.metadata,
) )
.await; .await;
} }
@@ -754,49 +724,49 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
false false
}; };
// Send ToolResult preview // Send ToolResult preview (skip for image sentinels to avoid
// broadcasting multi-MB base64 data as a preview)
if !is_image_sentinel if !is_image_sentinel
&& let Ok(ref output) = tool_result && let Ok(ref output) = tool_result
&& !output.is_empty() && !output.is_empty()
{ {
let _ = self let _ = self
.agent
.channels .channels
.send_status( .send_status(
&self.message.channel, &message.channel,
StatusUpdate::ToolResult { StatusUpdate::ToolResult {
name: tc.name.clone(), name: tc.name.clone(),
preview: output.clone(), preview: output.clone(),
}, },
&self.message.metadata, &message.metadata,
) )
.await; .await;
} }
// Check for auth awaiting // Check for auth awaiting — defer the return
// until all results are recorded.
if deferred_auth.is_none() if deferred_auth.is_none()
&& let Some((ext_name, instructions)) = && let Some((ext_name, instructions)) =
check_auth_required(&tc.name, &tool_result) check_auth_required(&tc.name, &tool_result)
{ {
let auth_data = parse_auth_result(&tool_result); let auth_data = parse_auth_result(&tool_result);
{ {
let mut sess = self.session.lock().await; let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&self.thread_id) { if let Some(thread) = sess.threads.get_mut(&thread_id) {
thread.enter_auth_mode(ext_name.clone()); thread.enter_auth_mode(ext_name.clone());
} }
} }
let _ = self let _ = self
.agent
.channels .channels
.send_status( .send_status(
&self.message.channel, &message.channel,
StatusUpdate::AuthRequired { StatusUpdate::AuthRequired {
extension_name: ext_name, extension_name: ext_name,
instructions: Some(instructions.clone()), instructions: Some(instructions.clone()),
auth_url: auth_data.auth_url, auth_url: auth_data.auth_url,
setup_url: auth_data.setup_url, setup_url: auth_data.setup_url,
}, },
&self.message.metadata, &message.metadata,
) )
.await; .await;
deferred_auth = Some(instructions); deferred_auth = Some(instructions);
@@ -804,7 +774,7 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
// Stash full output so subsequent tools can reference it // Stash full output so subsequent tools can reference it
if let Ok(ref output) = tool_result { if let Ok(ref output) = tool_result {
self.job_ctx job_ctx
.tool_output_stash .tool_output_stash
.write() .write()
.await .await
@@ -816,8 +786,8 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
let result_content = match tool_result { let result_content = match tool_result {
Ok(output) => { Ok(output) => {
let sanitized = let sanitized =
self.agent.safety().sanitize_tool_output(&tc.name, &output); self.safety().sanitize_tool_output(&tc.name, &output);
self.agent.safety().wrap_for_llm( self.safety().wrap_for_llm(
&tc.name, &tc.name,
&sanitized.content, &sanitized.content,
sanitized.was_modified, sanitized.was_modified,
@@ -826,21 +796,24 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
Err(e) => format!("Tool '{}' failed: {}", tc.name, e), Err(e) => format!("Tool '{}' failed: {}", tc.name, e),
}; };
// Record sanitized result in thread // Record sanitized result in thread so messages()
// and persist_tool_calls() use cleaned content.
{ {
let mut sess = self.session.lock().await; let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&self.thread_id) if let Some(thread) = sess.threads.get_mut(&thread_id)
&& let Some(turn) = thread.last_turn_mut() && let Some(turn) = thread.last_turn_mut()
{ {
if is_tool_error { if is_tool_error {
turn.record_tool_error(result_content.clone()); turn.record_tool_error(result_content.clone());
} else { } else {
turn.record_tool_result(serde_json::json!(result_content)); turn.record_tool_result(serde_json::json!(
result_content
));
} }
} }
} }
reason_ctx.messages.push(ChatMessage::tool_result( context_messages.push(ChatMessage::tool_result(
&tc.id, &tc.id,
&tc.name, &tc.name,
result_content, result_content,
@@ -851,11 +824,14 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
// Return auth response after all results are recorded // Return auth response after all results are recorded
if let Some(instructions) = deferred_auth { if let Some(instructions) = deferred_auth {
return Ok(Some(LoopOutcome::Response(instructions))); return Ok(AgenticLoopResult::Response(instructions));
} }
// Handle approval if a tool needed it // Handle approval if a tool needed it
if let Some((approval_idx, tc, tool)) = approval_needed { 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 display_params = redact_params(&tc.arguments, tool.sensitive_params());
let pending = PendingApproval { let pending = PendingApproval {
request_id: Uuid::new_v4(), request_id: Uuid::new_v4(),
@@ -864,23 +840,34 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
display_parameters: display_params, display_parameters: display_params,
description: tool.description().to_string(), description: tool.description().to_string(),
tool_call_id: tc.id.clone(), tool_call_id: tc.id.clone(),
context_messages: reason_ctx.messages.clone(), context_messages: context_messages.clone(),
deferred_tool_calls: tool_calls[approval_idx + 1..].to_vec(), deferred_tool_calls: tool_calls[approval_idx + 1..].to_vec(),
user_timezone: Some(self.user_tz.name().to_string()), user_timezone: Some(user_tz.name().to_string()),
}; };
return Ok(Some(LoopOutcome::NeedApproval(Box::new(pending)))); return Ok(AgenticLoopResult::NeedApproval { pending });
}
}
}
}
} }
Ok(None) /// Execute a tool for chat (without full job context).
pub(super) async fn execute_chat_tool(
&self,
tool_name: &str,
params: &serde_json::Value,
job_ctx: &JobContext,
) -> Result<String, Error> {
execute_chat_tool_standalone(self.tools(), self.safety(), tool_name, params, job_ctx).await
} }
} }
/// Execute a chat tool without requiring `&Agent`. /// Execute a chat tool without requiring `&Agent`.
/// ///
/// This standalone function enables parallel invocation from spawned JoinSet /// This standalone function enables parallel invocation from spawned JoinSet
/// tasks, which cannot borrow `&self`. Delegates to the shared /// tasks, which cannot borrow `&self`. It replicates the logic from
/// `execute_tool_with_safety` pipeline. /// `Agent::execute_chat_tool`.
pub(super) async fn execute_chat_tool_standalone( pub(super) async fn execute_chat_tool_standalone(
tools: &crate::tools::ToolRegistry, tools: &crate::tools::ToolRegistry,
safety: &crate::safety::SafetyLayer, safety: &crate::safety::SafetyLayer,
@@ -888,7 +875,91 @@ pub(super) async fn execute_chat_tool_standalone(
params: &serde_json::Value, params: &serde_json::Value,
job_ctx: &crate::context::JobContext, job_ctx: &crate::context::JobContext,
) -> Result<String, Error> { ) -> Result<String, Error> {
crate::tools::execute::execute_tool_with_safety(tools, safety, tool_name, params, job_ctx).await let tool = tools
.get(tool_name)
.await
.ok_or_else(|| crate::error::ToolError::NotFound {
name: tool_name.to_string(),
})?;
// Validate tool parameters
let validation = safety.validator().validate_tool_params(params);
if !validation.is_valid {
let details = validation
.errors
.iter()
.map(|e| format!("{}: {}", e.field, e.message))
.collect::<Vec<_>>()
.join("; ");
return Err(crate::error::ToolError::InvalidParameters {
name: tool_name.to_string(),
reason: format!("Invalid tool parameters: {}", details),
}
.into());
}
let safe_params = redact_params(params, tool.sensitive_params());
tracing::debug!(
tool = %tool_name,
params = %safe_params,
"Tool call started"
);
// Execute with per-tool timeout
let timeout = tool.execution_timeout();
let start = std::time::Instant::now();
let result = tokio::time::timeout(timeout, async {
tool.execute(params.clone(), job_ctx).await
})
.await;
let elapsed = start.elapsed();
match &result {
Ok(Ok(output)) => {
let result_str = serde_json::to_string(&output.result)
.unwrap_or_else(|_| "<serialize error>".to_string());
tracing::debug!(
tool = %tool_name,
elapsed_ms = elapsed.as_millis() as u64,
result = %result_str,
"Tool call succeeded"
);
}
Ok(Err(e)) => {
tracing::debug!(
tool = %tool_name,
elapsed_ms = elapsed.as_millis() as u64,
error = %e,
"Tool call failed"
);
}
Err(_) => {
tracing::debug!(
tool = %tool_name,
elapsed_ms = elapsed.as_millis() as u64,
timeout_secs = timeout.as_secs(),
"Tool call timed out"
);
}
}
let result = result
.map_err(|_| crate::error::ToolError::Timeout {
name: tool_name.to_string(),
timeout,
})?
.map_err(|e| crate::error::ToolError::ExecutionFailed {
name: tool_name.to_string(),
reason: e.to_string(),
})?;
serde_json::to_string_pretty(&result.result).map_err(|e| {
crate::error::ToolError::ExecutionFailed {
name: tool_name.to_string(),
reason: format!("Failed to serialize result: {}", e),
}
.into()
})
} }
/// Parsed auth result fields for emitting StatusUpdate::AuthRequired. /// Parsed auth result fields for emitting StatusUpdate::AuthRequired.
+4 -5
View File
@@ -189,7 +189,7 @@ impl HeartbeatRunner {
// Skip during quiet hours // Skip during quiet hours
if self.config.is_quiet_hours() { if self.config.is_quiet_hours() {
tracing::trace!("Heartbeat skipped: quiet hours"); tracing::debug!("Heartbeat skipped: quiet hours");
continue; continue;
} }
@@ -212,7 +212,7 @@ impl HeartbeatRunner {
match self.check_heartbeat().await { match self.check_heartbeat().await {
HeartbeatResult::Ok => { HeartbeatResult::Ok => {
tracing::trace!("Heartbeat OK"); tracing::debug!("Heartbeat OK");
self.consecutive_failures = 0; self.consecutive_failures = 0;
} }
HeartbeatResult::NeedsAttention(message) => { HeartbeatResult::NeedsAttention(message) => {
@@ -221,7 +221,7 @@ impl HeartbeatRunner {
self.send_notification(&message).await; self.send_notification(&message).await;
} }
HeartbeatResult::Skipped => { HeartbeatResult::Skipped => {
tracing::trace!("Heartbeat skipped"); tracing::debug!("Heartbeat skipped");
} }
HeartbeatResult::Failed(error) => { HeartbeatResult::Failed(error) => {
tracing::error!("Heartbeat failed: {}", error); tracing::error!("Heartbeat failed: {}", error);
@@ -303,8 +303,7 @@ impl HeartbeatRunner {
.with_max_tokens(max_tokens) .with_max_tokens(max_tokens)
.with_temperature(0.3); .with_temperature(0.3);
let reasoning = let reasoning = Reasoning::new(self.llm.clone());
Reasoning::new(self.llm.clone()).with_model_name(self.llm.active_model_name());
let (content, _usage) = match reasoning.complete(request).await { let (content, _usage) = match reasoning.complete(request).await {
Ok(r) => r, Ok(r) => r,
Err(e) => return HeartbeatResult::Failed(format!("LLM call failed: {}", e)), Err(e) => return HeartbeatResult::Failed(format!("LLM call failed: {}", e)),
+3 -3
View File
@@ -11,7 +11,6 @@
//! - Context compaction for long conversations //! - Context compaction for long conversations
mod agent_loop; mod agent_loop;
pub mod agentic_loop;
mod attachments; mod attachments;
mod commands; mod commands;
pub mod compaction; pub mod compaction;
@@ -23,7 +22,7 @@ pub mod job_monitor;
mod router; mod router;
pub mod routine; pub mod routine;
pub mod routine_engine; pub mod routine_engine;
pub(crate) mod scheduler; mod scheduler;
mod self_repair; mod self_repair;
pub mod session; pub mod session;
mod session_manager; mod session_manager;
@@ -31,8 +30,8 @@ pub mod submission;
pub mod task; pub mod task;
mod thread_ops; mod thread_ops;
pub mod undo; pub mod undo;
pub mod worker;
pub use crate::worker::{Worker, WorkerDeps};
pub(crate) use agent_loop::truncate_for_preview; pub(crate) use agent_loop::truncate_for_preview;
pub use agent_loop::{Agent, AgentDeps}; pub use agent_loop::{Agent, AgentDeps};
pub use compaction::{CompactionResult, ContextCompactor}; pub use compaction::{CompactionResult, ContextCompactor};
@@ -48,3 +47,4 @@ pub use session_manager::SessionManager;
pub use submission::{Submission, SubmissionParser, SubmissionResult}; pub use submission::{Submission, SubmissionParser, SubmissionResult};
pub use task::{Task, TaskContext, TaskHandler, TaskOutput}; pub use task::{Task, TaskContext, TaskHandler, TaskOutput};
pub use undo::{Checkpoint, UndoManager}; pub use undo::{Checkpoint, UndoManager};
pub use worker::{Worker, WorkerDeps};
+23 -86
View File
@@ -8,7 +8,7 @@
//! ┌──────────┐ ┌─────────┐ ┌──────────────────┐ //! ┌──────────┐ ┌─────────┐ ┌──────────────────┐
//! │ Trigger │────▶│ Engine │────▶│ Execution Mode │ //! │ Trigger │────▶│ Engine │────▶│ Execution Mode │
//! │ cron/event│ │guardrail│ │lightweight│full_job│ //! │ cron/event│ │guardrail│ │lightweight│full_job│
//! │ system │ │ check │ └──────────────────┘ //! │ webhook │ │ check │ └──────────────────┘
//! │ manual │ └─────────┘ │ //! │ manual │ └─────────┘ │
//! └──────────┘ ▼ //! └──────────┘ ▼
//! ┌──────────────┐ //! ┌──────────────┐
@@ -69,15 +69,12 @@ pub enum Trigger {
/// Regex pattern to match against message content. /// Regex pattern to match against message content.
pattern: String, pattern: String,
}, },
/// Fire when a structured system event is emitted. /// Fire on incoming webhook POST to /hooks/routine/{id}.
SystemEvent { Webhook {
/// Event source namespace (e.g. "github", "workflow", "tool"). /// Optional webhook path suffix (defaults to routine id).
source: String, path: Option<String>,
/// Event type within the source (e.g. "issue.opened"). /// Optional shared secret for HMAC validation.
event_type: String, secret: Option<String>,
/// Optional exact-match filters against payload top-level fields.
#[serde(default)]
filters: std::collections::HashMap<String, String>,
}, },
/// Only fires via tool call or CLI. /// Only fires via tool call or CLI.
Manual, Manual,
@@ -89,7 +86,7 @@ impl Trigger {
match self { match self {
Trigger::Cron { .. } => "cron", Trigger::Cron { .. } => "cron",
Trigger::Event { .. } => "event", Trigger::Event { .. } => "event",
Trigger::SystemEvent { .. } => "system_event", Trigger::Webhook { .. } => "webhook",
Trigger::Manual => "manual", Trigger::Manual => "manual",
} }
} }
@@ -137,39 +134,16 @@ impl Trigger {
.map(String::from); .map(String::from);
Ok(Trigger::Event { channel, pattern }) Ok(Trigger::Event { channel, pattern })
} }
"system_event" => { "webhook" => {
let source = config let path = config
.get("source") .get("path")
.and_then(|v| v.as_str()) .and_then(|v| v.as_str())
.ok_or_else(|| RoutineError::MissingField { .map(String::from);
context: "system_event trigger".into(), let secret = config
field: "source".into(), .get("secret")
})?
.to_string();
let event_type = config
.get("event_type")
.and_then(|v| v.as_str()) .and_then(|v| v.as_str())
.ok_or_else(|| RoutineError::MissingField { .map(String::from);
context: "system_event trigger".into(), Ok(Trigger::Webhook { path, secret })
field: "event_type".into(),
})?
.to_string();
let filters = config
.get("filters")
.and_then(|v| v.as_object())
.map(|m| {
m.iter()
.filter_map(|(k, v)| {
json_value_as_filter_string(v).map(|s| (k.clone(), s))
})
.collect()
})
.unwrap_or_default();
Ok(Trigger::SystemEvent {
source,
event_type,
filters,
})
} }
"manual" => Ok(Trigger::Manual), "manual" => Ok(Trigger::Manual),
other => Err(RoutineError::UnknownTriggerType { other => Err(RoutineError::UnknownTriggerType {
@@ -189,14 +163,9 @@ impl Trigger {
"pattern": pattern, "pattern": pattern,
"channel": channel, "channel": channel,
}), }),
Trigger::SystemEvent { Trigger::Webhook { path, secret } => serde_json::json!({
source, "path": path,
event_type, "secret": secret,
filters,
} => serde_json::json!({
"source": source,
"event_type": event_type,
"filters": filters,
}), }),
Trigger::Manual => serde_json::json!({}), Trigger::Manual => serde_json::json!({}),
} }
@@ -459,19 +428,6 @@ pub struct RoutineRun {
pub created_at: DateTime<Utc>, pub created_at: DateTime<Utc>,
} }
/// Convert a JSON value to a string for filter storage.
///
/// Handles strings, numbers, and booleans — consistent with the matching
/// logic in `routine_engine::json_value_as_string`.
pub fn json_value_as_filter_string(v: &serde_json::Value) -> Option<String> {
match v {
serde_json::Value::String(s) => Some(s.clone()),
serde_json::Value::Number(n) => Some(n.to_string()),
serde_json::Value::Bool(b) => Some(b.to_string()),
_ => None,
}
}
/// Compute a content hash for event dedup. /// Compute a content hash for event dedup.
pub fn content_hash(content: &str) -> u64 { pub fn content_hash(content: &str) -> u64 {
let mut hasher = DefaultHasher::new(); let mut hasher = DefaultHasher::new();
@@ -530,24 +486,6 @@ mod tests {
if channel == Some("telegram".to_string()) && pattern == r"deploy\s+\w+")); if channel == Some("telegram".to_string()) && pattern == r"deploy\s+\w+"));
} }
#[test]
fn test_system_event_trigger_roundtrip() {
let mut filters = std::collections::HashMap::new();
filters.insert("repo".to_string(), "nearai/ironclaw".to_string());
filters.insert("action".to_string(), "opened".to_string());
let trigger = Trigger::SystemEvent {
source: "github".to_string(),
event_type: "issue".to_string(),
filters: filters.clone(),
};
let json = trigger.to_config_json();
let parsed = Trigger::from_db("system_event", json).expect("parse system_event");
assert!(
matches!(parsed, Trigger::SystemEvent { source, event_type, filters: f }
if source == "github" && event_type == "issue" && f == filters)
);
}
#[test] #[test]
fn test_action_lightweight_roundtrip() { fn test_action_lightweight_roundtrip() {
let action = RoutineAction::Lightweight { let action = RoutineAction::Lightweight {
@@ -685,13 +623,12 @@ mod tests {
"event" "event"
); );
assert_eq!( assert_eq!(
Trigger::SystemEvent { Trigger::Webhook {
source: String::new(), path: None,
event_type: String::new(), secret: None
filters: std::collections::HashMap::new(),
} }
.type_tag(), .type_tag(),
"system_event" "webhook"
); );
assert_eq!(Trigger::Manual.type_tag(), "manual"); assert_eq!(Trigger::Manual.type_tag(), "manual");
} }
+20 -117
View File
@@ -32,14 +32,9 @@ use crate::llm::{
ChatMessage, CompletionRequest, FinishReason, LlmProvider, ToolCall, ToolCompletionRequest, ChatMessage, CompletionRequest, FinishReason, LlmProvider, ToolCall, ToolCompletionRequest,
}; };
use crate::safety::SafetyLayer; use crate::safety::SafetyLayer;
use crate::tools::{ApprovalContext, ApprovalRequirement, ToolError, ToolRegistry}; use crate::tools::{ApprovalContext, ApprovalRequirement, ToolError, ToolRegistry, redact_params};
use crate::workspace::Workspace; use crate::workspace::Workspace;
enum EventMatcher {
Message { routine: Routine, regex: Regex },
System { routine: Routine },
}
/// The routine execution engine. /// The routine execution engine.
pub struct RoutineEngine { pub struct RoutineEngine {
config: RoutineConfig, config: RoutineConfig,
@@ -50,8 +45,8 @@ pub struct RoutineEngine {
notify_tx: mpsc::Sender<OutgoingResponse>, notify_tx: mpsc::Sender<OutgoingResponse>,
/// Currently running routine count (across all routines). /// Currently running routine count (across all routines).
running_count: Arc<AtomicUsize>, running_count: Arc<AtomicUsize>,
/// Cached matchers for all event-driven routines. /// Compiled event regex cache: routine_id -> compiled regex.
event_cache: Arc<RwLock<Vec<EventMatcher>>>, event_cache: Arc<RwLock<Vec<(Uuid, Routine, Regex)>>>,
/// Scheduler for dispatching jobs (FullJob mode). /// Scheduler for dispatching jobs (FullJob mode).
scheduler: Option<Arc<Scheduler>>, scheduler: Option<Arc<Scheduler>>,
/// Tool registry for lightweight routine tool execution. /// Tool registry for lightweight routine tool execution.
@@ -92,12 +87,9 @@ impl RoutineEngine {
Ok(routines) => { Ok(routines) => {
let mut cache = Vec::new(); let mut cache = Vec::new();
for routine in routines { for routine in routines {
match &routine.trigger { if let Trigger::Event { ref pattern, .. } = routine.trigger {
Trigger::Event { pattern, .. } => match Regex::new(pattern) { match Regex::new(pattern) {
Ok(re) => cache.push(EventMatcher::Message { Ok(re) => cache.push((routine.id, routine.clone(), re)),
routine: routine.clone(),
regex: re,
}),
Err(e) => { Err(e) => {
tracing::warn!( tracing::warn!(
routine = %routine.name, routine = %routine.name,
@@ -105,18 +97,12 @@ impl RoutineEngine {
pattern, e pattern, e
); );
} }
},
Trigger::SystemEvent { .. } => {
cache.push(EventMatcher::System {
routine: routine.clone(),
});
} }
_ => {}
} }
} }
let count = cache.len(); let count = cache.len();
*self.event_cache.write().await = cache; *self.event_cache.write().await = cache;
tracing::trace!("Refreshed event cache: {} routines", count); tracing::debug!("Refreshed event cache: {} routines", count);
} }
Err(e) => { Err(e) => {
tracing::error!("Failed to refresh event cache: {}", e); tracing::error!("Failed to refresh event cache: {}", e);
@@ -132,11 +118,7 @@ impl RoutineEngine {
let cache = self.event_cache.read().await; let cache = self.event_cache.read().await;
let mut fired = 0; let mut fired = 0;
for matcher in cache.iter() { for (_, routine, re) in cache.iter() {
let (routine, re) = match matcher {
EventMatcher::Message { routine, regex } => (routine, regex),
EventMatcher::System { .. } => continue,
};
// Channel filter // Channel filter
if let Trigger::Event { if let Trigger::Event {
channel: Some(ch), .. channel: Some(ch), ..
@@ -153,13 +135,13 @@ impl RoutineEngine {
// Cooldown check // Cooldown check
if !self.check_cooldown(routine) { if !self.check_cooldown(routine) {
tracing::trace!(routine = %routine.name, "Skipped: cooldown active"); tracing::debug!(routine = %routine.name, "Skipped: cooldown active");
continue; continue;
} }
// Concurrent run check // Concurrent run check
if !self.check_concurrent(routine).await { if !self.check_concurrent(routine).await {
tracing::trace!(routine = %routine.name, "Skipped: max concurrent reached"); tracing::debug!(routine = %routine.name, "Skipped: max concurrent reached");
continue; continue;
} }
@@ -177,88 +159,6 @@ impl RoutineEngine {
fired fired
} }
/// Emit a structured event to system-event routines.
///
/// Returns the number of routines that were fired.
pub async fn emit_system_event(
&self,
source: &str,
event_type: &str,
payload: &serde_json::Value,
user_id: Option<&str>,
) -> usize {
let cache = self.event_cache.read().await;
let mut fired = 0;
for matcher in cache.iter() {
let routine = match matcher {
EventMatcher::System { routine } => routine,
EventMatcher::Message { .. } => continue,
};
let Trigger::SystemEvent {
source: expected_source,
event_type: expected_event,
filters,
} = &routine.trigger
else {
continue;
};
if !expected_source.eq_ignore_ascii_case(source)
|| !expected_event.eq_ignore_ascii_case(event_type)
{
continue;
}
if let Some(uid) = user_id
&& routine.user_id != uid
{
continue;
}
let mut matched = true;
for (key, expected) in filters {
let Some(actual) = payload
.get(key)
.and_then(crate::agent::routine::json_value_as_filter_string)
else {
tracing::debug!(routine = %routine.name, filter_key = %key, "Filter key not found in payload");
matched = false;
break;
};
if !actual.eq_ignore_ascii_case(expected) {
matched = false;
break;
}
}
if !matched {
continue;
}
if !self.check_cooldown(routine) {
tracing::debug!(routine = %routine.name, "Skipped: cooldown active");
continue;
}
if !self.check_concurrent(routine).await {
tracing::debug!(routine = %routine.name, "Skipped: max concurrent reached");
continue;
}
if self.running_count.load(Ordering::Relaxed) >= self.config.max_concurrent_routines {
tracing::warn!(routine = %routine.name, "Skipped: global max concurrent reached");
continue;
}
let detail = truncate(&format!("{source}:{event_type}"), 200);
self.spawn_fire(routine.clone(), "system_event", Some(detail));
fired += 1;
}
fired
}
/// Check all due cron routines and fire them. Called by the cron ticker. /// Check all due cron routines and fire them. Called by the cron ticker.
pub async fn check_cron_triggers(&self) { pub async fn check_cron_triggers(&self) {
let routines = match self.store.list_due_cron_routines().await { let routines = match self.store.list_due_cron_routines().await {
@@ -1013,6 +913,13 @@ async fn execute_routine_tool(
return Err(format!("Invalid tool parameters: {}", details).into()); return Err(format!("Invalid tool parameters: {}", details).into());
} }
let safe_params = redact_params(&tc.arguments, tool.sensitive_params());
tracing::debug!(
tool = %tc.name,
params = %safe_params,
"Lightweight routine tool call started"
);
// Execute with per-tool timeout // Execute with per-tool timeout
let timeout = tool.execution_timeout(); let timeout = tool.execution_timeout();
let start = std::time::Instant::now(); let start = std::time::Instant::now();
@@ -1022,14 +929,12 @@ async fn execute_routine_tool(
.await; .await;
let elapsed = start.elapsed(); let elapsed = start.elapsed();
// Log tool execution result (single consolidated log)
match &result { match &result {
Ok(Ok(_)) => { Ok(Ok(_)) => {
tracing::debug!( tracing::debug!(
tool = %tc.name, tool = %tc.name,
elapsed_ms = elapsed.as_millis() as u64, elapsed_ms = elapsed.as_millis() as u64,
status = "succeeded", "Lightweight routine tool call succeeded"
"Lightweight routine tool execution completed"
); );
} }
Ok(Err(e)) => { Ok(Err(e)) => {
@@ -1037,8 +942,7 @@ async fn execute_routine_tool(
tool = %tc.name, tool = %tc.name,
elapsed_ms = elapsed.as_millis() as u64, elapsed_ms = elapsed.as_millis() as u64,
error = %e, error = %e,
status = "failed", "Lightweight routine tool call failed"
"Lightweight routine tool execution completed"
); );
} }
Err(_) => { Err(_) => {
@@ -1046,8 +950,7 @@ async fn execute_routine_tool(
tool = %tc.name, tool = %tc.name,
elapsed_ms = elapsed.as_millis() as u64, elapsed_ms = elapsed.as_millis() as u64,
timeout_secs = timeout.as_secs(), timeout_secs = timeout.as_secs(),
status = "timeout", "Lightweight routine tool call timed out"
"Lightweight routine tool execution completed"
); );
} }
} }
+39 -169
View File
@@ -9,6 +9,7 @@ use tokio::task::JoinHandle;
use uuid::Uuid; use uuid::Uuid;
use crate::agent::task::{Task, TaskContext, TaskOutput}; use crate::agent::task::{Task, TaskContext, TaskOutput};
use crate::agent::worker::{Worker, WorkerDeps};
use crate::channels::web::types::SseEvent; use crate::channels::web::types::SseEvent;
use crate::config::AgentConfig; use crate::config::AgentConfig;
use crate::context::{ContextManager, JobContext, JobState}; use crate::context::{ContextManager, JobContext, JobState};
@@ -18,7 +19,6 @@ use crate::hooks::HookRegistry;
use crate::llm::LlmProvider; use crate::llm::LlmProvider;
use crate::safety::SafetyLayer; use crate::safety::SafetyLayer;
use crate::tools::{ApprovalContext, ToolRegistry}; use crate::tools::{ApprovalContext, ToolRegistry};
use crate::worker::job::{Worker, WorkerDeps};
/// Message to send to a worker. /// Message to send to a worker.
#[derive(Debug)] #[derive(Debug)]
@@ -160,36 +160,24 @@ impl Scheduler {
.create_job_for_user(user_id, title, description) .create_job_for_user(user_id, title, description)
.await?; .await?;
// Apply metadata and token budget in a single atomic update. // Apply token budget from config, allowing per-job metadata override.
// This prevents concurrent workers from observing partial state. let max_tokens = metadata
// Cap user-supplied max_tokens at the configured limit (Issue #815).
let user_max_tokens = metadata
.as_ref() .as_ref()
.and_then(|m| m.get("max_tokens")) .and_then(|m| m.get("max_tokens"))
.and_then(|v| v.as_u64()); .and_then(|v| v.as_u64())
let max_tokens = user_max_tokens
.map(|user_val| {
if self.config.max_tokens_per_job == 0 {
// Config is "unlimited": use the user-supplied value directly.
user_val
} else {
std::cmp::min(user_val, self.config.max_tokens_per_job)
}
})
.unwrap_or(self.config.max_tokens_per_job); .unwrap_or(self.config.max_tokens_per_job);
// Apply both metadata and token budget in one closure (Issue #813: atomic update) // Apply metadata if provided
if let Some(meta) = metadata { if let Some(meta) = metadata {
self.context_manager self.context_manager
.update_context(job_id, |ctx| { .update_context(job_id, |ctx| {
ctx.metadata = meta; ctx.metadata = meta;
if max_tokens > 0 {
ctx.max_tokens = max_tokens;
}
}) })
.await?; .await?;
} else if max_tokens > 0 { }
// Set token budget (separate update to avoid overwriting metadata)
if max_tokens > 0 {
self.context_manager self.context_manager
.update_context(job_id, |ctx| { .update_context(job_id, |ctx| {
ctx.max_tokens = max_tokens; ctx.max_tokens = max_tokens;
@@ -474,9 +462,6 @@ impl Scheduler {
} }
/// Execute a single tool as a subtask. /// Execute a single tool as a subtask.
///
/// Performs scheduler-specific checks (approval, cancellation) then
/// delegates to the shared `execute_tool_with_safety` pipeline.
async fn execute_tool_task( async fn execute_tool_task(
tools: Arc<ToolRegistry>, tools: Arc<ToolRegistry>,
context_manager: Arc<ContextManager>, context_manager: Arc<ContextManager>,
@@ -488,7 +473,7 @@ impl Scheduler {
) -> Result<TaskOutput, Error> { ) -> Result<TaskOutput, Error> {
let start = std::time::Instant::now(); let start = std::time::Instant::now();
// Get the tool for approval check // Get the tool
let tool = tools.get(tool_name).await.ok_or_else(|| { let tool = tools.get(tool_name).await.ok_or_else(|| {
Error::Tool(crate::error::ToolError::NotFound { Error::Tool(crate::error::ToolError::NotFound {
name: tool_name.to_string(), name: tool_name.to_string(),
@@ -505,7 +490,6 @@ impl Scheduler {
.into()); .into());
} }
// Scheduler-specific approval check
let requirement = tool.requires_approval(&params); let requirement = tool.requires_approval(&params);
let blocked = let blocked =
ApprovalContext::is_blocked_or_default(&approval_context, tool_name, requirement); ApprovalContext::is_blocked_or_default(&approval_context, tool_name, requirement);
@@ -516,23 +500,41 @@ impl Scheduler {
.into()); .into());
} }
// Delegate to shared tool execution pipeline // Validate tool parameters
let output_str = crate::tools::execute::execute_tool_with_safety( let validation = safety.validator().validate_tool_params(&params);
&tools, &safety, tool_name, &params, &job_ctx, if !validation.is_valid {
) let details = validation
.await?; .errors
.iter()
.map(|e| format!("{}: {}", e.field, e.message))
.collect::<Vec<_>>()
.join("; ");
return Err(crate::error::ToolError::InvalidParameters {
name: tool_name.to_string(),
reason: format!("Invalid tool parameters: {}", details),
}
.into());
}
// Parse back to Value for TaskOutput; this should be infallible given // Execute with per-tool timeout
// `execute_tool_with_safety` uses `serde_json::to_string_pretty`, but if it let tool_timeout = tool.execution_timeout();
// ever fails we surface a clear error instead of silently changing types. let result =
let result_value: serde_json::Value = serde_json::from_str(&output_str).map_err(|e| { tokio::time::timeout(tool_timeout, async { tool.execute(params, &job_ctx).await })
.await
.map_err(|_| {
Error::Tool(crate::error::ToolError::Timeout {
name: tool_name.to_string(),
timeout: tool_timeout,
})
})?
.map_err(|e| {
Error::Tool(crate::error::ToolError::ExecutionFailed { Error::Tool(crate::error::ToolError::ExecutionFailed {
name: tool_name.to_string(), name: tool_name.to_string(),
reason: format!("Failed to parse tool output as JSON: {}", e), reason: e.to_string(),
}) })
})?; })?;
Ok(TaskOutput::new(result_value, start.elapsed())) Ok(TaskOutput::new(result.result, start.elapsed()))
} }
/// Stop a running job. /// Stop a running job.
@@ -697,140 +699,8 @@ impl Scheduler {
mod tests { mod tests {
use super::*; use super::*;
use crate::config::SafetyConfig; use crate::config::SafetyConfig;
use crate::llm::{
CompletionRequest, CompletionResponse, LlmError, LlmProvider, ToolCompletionRequest,
ToolCompletionResponse,
};
use crate::safety::SafetyLayer; use crate::safety::SafetyLayer;
use crate::tools::{ApprovalRequirement, Tool, ToolError, ToolOutput}; use crate::tools::{ApprovalRequirement, Tool, ToolError, ToolOutput};
use rust_decimal_macros::dec;
/// Minimal LLM provider stub for scheduler tests that don't exercise LLM calls.
struct StubLlm;
#[async_trait::async_trait]
impl LlmProvider for StubLlm {
fn model_name(&self) -> &str {
"stub"
}
fn cost_per_token(&self) -> (rust_decimal::Decimal, rust_decimal::Decimal) {
(dec!(0), dec!(0))
}
async fn complete(&self, _req: CompletionRequest) -> Result<CompletionResponse, LlmError> {
Err(LlmError::RequestFailed {
provider: "stub".into(),
reason: "not implemented".into(),
})
}
async fn complete_with_tools(
&self,
_req: ToolCompletionRequest,
) -> Result<ToolCompletionResponse, LlmError> {
Err(LlmError::RequestFailed {
provider: "stub".into(),
reason: "not implemented".into(),
})
}
}
/// Create a Scheduler for token-budget tests. The LLM stub will fail if a
/// worker actually tries to call it, but `dispatch_job` sets the token
/// budget *before* spawning the worker so we can inspect the context
/// immediately after dispatch.
fn make_test_scheduler(max_tokens_per_job: u64) -> Scheduler {
let config = AgentConfig {
name: "test".to_string(),
max_parallel_jobs: 5,
job_timeout: std::time::Duration::from_secs(30),
stuck_threshold: std::time::Duration::from_secs(300),
repair_check_interval: std::time::Duration::from_secs(3600),
max_repair_attempts: 0,
use_planning: false,
session_idle_timeout: std::time::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,
default_timezone: "UTC".to_string(),
max_tokens_per_job,
};
let cm = Arc::new(ContextManager::new(5));
let llm: Arc<dyn LlmProvider> = Arc::new(StubLlm);
let safety = Arc::new(SafetyLayer::new(&SafetyConfig {
max_output_length: 100_000,
injection_check_enabled: false,
}));
let tools = Arc::new(ToolRegistry::new());
let hooks = Arc::new(HookRegistry::default());
Scheduler::new(config, cm, llm, safety, tools, None, hooks)
}
#[tokio::test]
async fn test_dispatch_job_caps_user_max_tokens() {
let sched = make_test_scheduler(1000);
let meta = serde_json::json!({ "max_tokens": 5000 });
let job_id = sched
.dispatch_job("user1", "test", "desc", Some(meta))
.await
.unwrap();
let ctx = sched.context_manager.get_context(job_id).await.unwrap();
assert_eq!(ctx.max_tokens, 1000, "should cap at configured limit");
}
#[tokio::test]
async fn test_dispatch_job_unlimited_config_preserves_user_tokens() {
let sched = make_test_scheduler(0); // 0 = unlimited
let meta = serde_json::json!({ "max_tokens": 5000 });
let job_id = sched
.dispatch_job("user1", "test", "desc", Some(meta))
.await
.unwrap();
let ctx = sched.context_manager.get_context(job_id).await.unwrap();
assert_eq!(
ctx.max_tokens, 5000,
"unlimited config should preserve user value"
);
}
#[tokio::test]
async fn test_dispatch_job_no_user_tokens_uses_config() {
let sched = make_test_scheduler(2000);
let job_id = sched
.dispatch_job("user1", "test", "desc", None)
.await
.unwrap();
let ctx = sched.context_manager.get_context(job_id).await.unwrap();
assert_eq!(
ctx.max_tokens, 2000,
"should use config default when no user value"
);
}
#[tokio::test]
async fn test_dispatch_job_atomic_metadata_and_tokens() {
let sched = make_test_scheduler(10_000);
let meta = serde_json::json!({
"max_tokens": 3000,
"custom_key": "custom_value"
});
let job_id = sched
.dispatch_job("user1", "test", "desc", Some(meta))
.await
.unwrap();
let ctx = sched.context_manager.get_context(job_id).await.unwrap();
assert_eq!(ctx.max_tokens, 3000, "should use user value within limit");
assert_eq!(
ctx.metadata.get("custom_key").and_then(|v| v.as_str()),
Some("custom_value"),
"metadata should be set atomically with token budget"
);
}
#[test] #[test]
fn test_scheduler_creation() { fn test_scheduler_creation() {
+9 -7
View File
@@ -334,21 +334,22 @@ impl RepairTask {
// Check for stuck jobs // Check for stuck jobs
let stuck_jobs = self.repair.detect_stuck_jobs().await; let stuck_jobs = self.repair.detect_stuck_jobs().await;
for job in stuck_jobs { for job in stuck_jobs {
tracing::info!("Attempting to repair stuck job {}", job.job_id);
match self.repair.repair_stuck_job(&job).await { match self.repair.repair_stuck_job(&job).await {
Ok(RepairResult::Success { message }) => { Ok(RepairResult::Success { message }) => {
tracing::info!(job = %job.job_id, status = "success", "Stuck job repair completed: {}", message); tracing::info!("Repair succeeded: {}", message);
} }
Ok(RepairResult::Retry { message }) => { Ok(RepairResult::Retry { message }) => {
tracing::debug!(job = %job.job_id, status = "retry", "Stuck job repair needs retry: {}", message); tracing::warn!("Repair needs retry: {}", message);
} }
Ok(RepairResult::Failed { message }) => { Ok(RepairResult::Failed { message }) => {
tracing::error!(job = %job.job_id, status = "failed", "Stuck job repair failed: {}", message); tracing::error!("Repair failed: {}", message);
} }
Ok(RepairResult::ManualRequired { message }) => { Ok(RepairResult::ManualRequired { message }) => {
tracing::warn!(job = %job.job_id, status = "manual", "Stuck job repair requires manual intervention: {}", message); tracing::warn!("Manual intervention needed: {}", message);
} }
Err(e) => { Err(e) => {
tracing::error!(job = %job.job_id, "Stuck job repair error: {}", e); tracing::error!("Repair error: {}", e);
} }
} }
} }
@@ -356,12 +357,13 @@ impl RepairTask {
// Check for broken tools // Check for broken tools
let broken_tools = self.repair.detect_broken_tools().await; let broken_tools = self.repair.detect_broken_tools().await;
for tool in broken_tools { for tool in broken_tools {
tracing::info!("Attempting to repair broken tool: {}", tool.name);
match self.repair.repair_broken_tool(&tool).await { match self.repair.repair_broken_tool(&tool).await {
Ok(result) => { Ok(result) => {
tracing::debug!(tool = %tool.name, status = "completed", "Tool repair completed: {:?}", result); tracing::info!("Tool repair result: {:?}", result);
} }
Err(e) => { Err(e) => {
tracing::error!(tool = %tool.name, "Tool repair error: {}", e); tracing::error!("Tool repair error: {}", e);
} }
} }
} }
+22 -50
View File
@@ -113,13 +113,6 @@ impl Agent {
thread_id: Uuid, thread_id: Uuid,
content: &str, content: &str,
) -> Result<SubmissionResult, Error> { ) -> Result<SubmissionResult, Error> {
tracing::debug!(
message_id = %message.id,
thread_id = %thread_id,
content_len = content.len(),
"Processing user input"
);
// First check thread state without holding lock during I/O // First check thread state without holding lock during I/O
let thread_state = { let thread_state = {
let sess = session.lock().await; let sess = session.lock().await;
@@ -130,41 +123,19 @@ impl Agent {
thread.state thread.state
}; };
tracing::debug!(
message_id = %message.id,
thread_id = %thread_id,
thread_state = ?thread_state,
"Checked thread state"
);
// Check thread state // Check thread state
match thread_state { match thread_state {
ThreadState::Processing => { ThreadState::Processing => {
tracing::warn!(
message_id = %message.id,
thread_id = %thread_id,
"Thread is processing, rejecting new input"
);
return Ok(SubmissionResult::error( return Ok(SubmissionResult::error(
"Turn in progress. Use /interrupt to cancel.", "Turn in progress. Use /interrupt to cancel.",
)); ));
} }
ThreadState::AwaitingApproval => { ThreadState::AwaitingApproval => {
tracing::warn!(
message_id = %message.id,
thread_id = %thread_id,
"Thread awaiting approval, rejecting new input"
);
return Ok(SubmissionResult::error( return Ok(SubmissionResult::error(
"Waiting for approval. Use /interrupt to cancel.", "Waiting for approval. Use /interrupt to cancel.",
)); ));
} }
ThreadState::Completed => { ThreadState::Completed => {
tracing::warn!(
message_id = %message.id,
thread_id = %thread_id,
"Thread completed, rejecting new input"
);
return Ok(SubmissionResult::error( return Ok(SubmissionResult::error(
"Thread completed. Use /thread new.", "Thread completed. Use /thread new.",
)); ));
@@ -298,20 +269,9 @@ impl Agent {
}; };
// Persist user message to DB immediately so it survives crashes // Persist user message to DB immediately so it survives crashes
tracing::debug!(
message_id = %message.id,
thread_id = %thread_id,
"Persisting user message to DB"
);
self.persist_user_message(thread_id, &message.user_id, effective_content) self.persist_user_message(thread_id, &message.user_id, effective_content)
.await; .await;
tracing::debug!(
message_id = %message.id,
thread_id = %thread_id,
"User message persisted, starting agentic loop"
);
// Send thinking status // Send thinking status
let _ = self let _ = self
.channels .channels
@@ -852,12 +812,19 @@ impl Agent {
// Sanitize tool result, then record the cleaned version in the // Sanitize tool result, then record the cleaned version in the
// thread. Must happen before auth intercept check which may return early. // thread. Must happen before auth intercept check which may return early.
let is_tool_error = tool_result.is_err(); let is_tool_error = tool_result.is_err();
let (result_content, _) = crate::tools::execute::process_tool_result( let result_content = match &tool_result {
self.safety(), Ok(output) => {
let sanitized = self
.safety()
.sanitize_tool_output(&pending.tool_name, output);
self.safety().wrap_for_llm(
&pending.tool_name, &pending.tool_name,
&pending.tool_call_id, &sanitized.content,
&tool_result, sanitized.was_modified,
); )
}
Err(e) => format!("Error: {}", e),
};
// Record sanitized result in thread // Record sanitized result in thread
{ {
@@ -1097,12 +1064,17 @@ impl Agent {
// Sanitize first, then record the cleaned version in thread. // Sanitize first, then record the cleaned version in thread.
// Must happen before auth detection which may set deferred_auth. // Must happen before auth detection which may set deferred_auth.
let is_deferred_error = deferred_result.is_err(); let is_deferred_error = deferred_result.is_err();
let (deferred_content, _) = crate::tools::execute::process_tool_result( let deferred_content = match &deferred_result {
self.safety(), Ok(output) => {
let sanitized = self.safety().sanitize_tool_output(&tc.name, output);
self.safety().wrap_for_llm(
&tc.name, &tc.name,
&tc.id, &sanitized.content,
&deferred_result, sanitized.was_modified,
); )
}
Err(e) => format!("Error: {}", e),
};
// Record sanitized result in thread // Record sanitized result in thread
{ {
File diff suppressed because it is too large Load Diff
+196 -47
View File
@@ -77,7 +77,10 @@ pub struct AppBuilder {
llm_override: Option<Arc<dyn LlmProvider>>, llm_override: Option<Arc<dyn LlmProvider>>,
// Backend-specific handles needed by secrets store // Backend-specific handles needed by secrets store
handles: Option<crate::db::DatabaseHandles>, #[cfg(feature = "postgres")]
pg_pool: Option<deadpool_postgres::Pool>,
#[cfg(feature = "libsql")]
libsql_db: Option<Arc<libsql::Database>>,
} }
impl AppBuilder { impl AppBuilder {
@@ -102,7 +105,10 @@ impl AppBuilder {
db: None, db: None,
secrets_store: None, secrets_store: None,
llm_override: None, llm_override: None,
handles: None, #[cfg(feature = "postgres")]
pg_pool: None,
#[cfg(feature = "libsql")]
libsql_db: None,
} }
} }
@@ -131,10 +137,71 @@ impl AppBuilder {
return Ok(()); return Ok(());
} }
let (db, handles) = crate::db::connect_with_handles(&self.config.database) let db: Arc<dyn Database> = match self.config.database.backend {
#[cfg(feature = "libsql")]
crate::config::DatabaseBackend::LibSql => {
use crate::db::Database as _;
use crate::db::libsql::LibSqlBackend;
use secrecy::ExposeSecret as _;
let default_path = crate::config::default_libsql_path();
let db_path = self
.config
.database
.libsql_path
.as_deref()
.unwrap_or(&default_path);
let backend = if let Some(ref url) = self.config.database.libsql_url {
let token =
self.config
.database
.libsql_auth_token
.as_ref()
.ok_or_else(|| {
anyhow::anyhow!(
"LIBSQL_AUTH_TOKEN is required when LIBSQL_URL is set"
)
})?;
LibSqlBackend::new_remote_replica(db_path, url, token.expose_secret()).await?
} else {
LibSqlBackend::new_local(db_path).await?
};
backend.run_migrations().await?;
tracing::info!("libSQL database connected and migrations applied");
#[cfg(feature = "libsql")]
{
self.libsql_db = Some(backend.shared_db());
}
Arc::new(backend) as Arc<dyn Database>
}
#[cfg(feature = "postgres")]
_ => {
use crate::db::Database as _;
let pg = crate::db::postgres::PgBackend::new(&self.config.database)
.await .await
.map_err(|e| anyhow::anyhow!("{}", e))?; .map_err(|e| anyhow::anyhow!("{}", e))?;
self.handles = Some(handles); pg.run_migrations()
.await
.map_err(|e| anyhow::anyhow!("{}", e))?;
tracing::info!("PostgreSQL database connected and migrations applied");
#[cfg(feature = "postgres")]
{
self.pg_pool = Some(pg.pool());
}
Arc::new(pg) as Arc<dyn Database>
}
#[cfg(not(feature = "postgres"))]
_ => {
anyhow::bail!(
"No database backend available. Enable 'postgres' or 'libsql' feature."
);
}
};
// Post-init: migrate disk config, reload config from DB, attach session, cleanup // Post-init: migrate disk config, reload config from DB, attach session, cleanup
if let Err(e) = crate::bootstrap::migrate_disk_to_db(db.as_ref(), "default").await { if let Err(e) = crate::bootstrap::migrate_disk_to_db(db.as_ref(), "default").await {
@@ -145,7 +212,7 @@ impl AppBuilder {
match Config::from_db_with_toml(db.as_ref(), "default", toml_path).await { match Config::from_db_with_toml(db.as_ref(), "default", toml_path).await {
Ok(db_config) => { Ok(db_config) => {
self.config = db_config; self.config = db_config;
tracing::debug!("Configuration reloaded from database"); tracing::info!("Configuration reloaded from database");
} }
Err(e) => { Err(e) => {
tracing::warn!( tracing::warn!(
@@ -184,7 +251,10 @@ impl AppBuilder {
crate::config::inject_os_credentials(); crate::config::inject_os_credentials();
// Consume unused handles // Consume unused handles
self.handles.take(); #[cfg(feature = "libsql")]
{
self.libsql_db.take();
}
// Re-resolve only the LLM config with OS credentials. // Re-resolve only the LLM config with OS credentials.
let store: Option<&(dyn crate::db::SettingsStore + Sync)> = let store: Option<&(dyn crate::db::SettingsStore + Sync)> =
@@ -208,16 +278,35 @@ impl AppBuilder {
Ok(c) => Arc::new(c), Ok(c) => Arc::new(c),
Err(e) => { Err(e) => {
tracing::warn!("Failed to initialize secrets crypto: {}", e); tracing::warn!("Failed to initialize secrets crypto: {}", e);
self.handles.take(); #[cfg(feature = "libsql")]
{
self.libsql_db.take();
}
return Ok(()); return Ok(());
} }
}; };
// Fallback covers the no-database path where `init_database` returned let store: Option<Arc<dyn SecretsStore + Send + Sync>> = None;
// early before populating `self.handles`.
let empty_handles = crate::db::DatabaseHandles::default(); #[cfg(feature = "libsql")]
let handles = self.handles.as_ref().unwrap_or(&empty_handles); let store = store.or_else(|| {
let store = crate::secrets::create_secrets_store(crypto, handles); self.libsql_db.take().map(|db| {
Arc::new(crate::secrets::LibSqlSecretsStore::new(
db,
Arc::clone(&crypto),
)) as Arc<dyn SecretsStore + Send + Sync>
})
});
#[cfg(feature = "postgres")]
let store = store.or_else(|| {
self.pg_pool.as_ref().map(|pool| {
Arc::new(crate::secrets::PostgresSecretsStore::new(
pool.clone(),
Arc::clone(&crypto),
)) as Arc<dyn SecretsStore + Send + Sync>
})
});
if let Some(ref secrets) = store { if let Some(ref secrets) = store {
// Inject LLM API keys from encrypted storage // Inject LLM API keys from encrypted storage
@@ -274,7 +363,7 @@ impl AppBuilder {
anyhow::Error, anyhow::Error,
> { > {
let safety = Arc::new(SafetyLayer::new(&self.config.safety)); let safety = Arc::new(SafetyLayer::new(&self.config.safety));
tracing::debug!("Safety layer initialized"); tracing::info!("Safety layer initialized");
// Initialize tool registry with credential injection support // Initialize tool registry with credential injection support
let credential_registry = Arc::new(SharedCredentialRegistry::new()); let credential_registry = Arc::new(SharedCredentialRegistry::new());
@@ -361,7 +450,7 @@ impl AppBuilder {
tools tools
.register_builder_tool(llm.clone(), Some(self.config.builder.to_builder_config())) .register_builder_tool(llm.clone(), Some(self.config.builder.to_builder_config()))
.await; .await;
tracing::debug!("Builder mode enabled"); tracing::info!("Builder mode enabled");
} }
Ok((safety, tools, embeddings, workspace)) Ok((safety, tools, embeddings, workspace))
@@ -383,7 +472,9 @@ impl AppBuilder {
), ),
anyhow::Error, anyhow::Error,
> { > {
use crate::tools::mcp::config::load_mcp_servers_from_db; use crate::tools::mcp::{
McpClient, McpTransport, config::load_mcp_servers_from_db, is_authenticated,
};
use crate::tools::wasm::{WasmToolLoader, load_dev_tools}; use crate::tools::wasm::{WasmToolLoader, load_dev_tools};
let mcp_session_manager = Arc::new(McpSessionManager::new()); let mcp_session_manager = Arc::new(McpSessionManager::new());
@@ -419,7 +510,7 @@ impl AppBuilder {
match loader.load_from_dir(&wasm_config.tools_dir).await { match loader.load_from_dir(&wasm_config.tools_dir).await {
Ok(results) => { Ok(results) => {
if !results.loaded.is_empty() { if !results.loaded.is_empty() {
tracing::debug!( tracing::info!(
"Loaded {} WASM tools from {}", "Loaded {} WASM tools from {}",
results.loaded.len(), results.loaded.len(),
wasm_config.tools_dir.display() wasm_config.tools_dir.display()
@@ -442,7 +533,7 @@ impl AppBuilder {
Ok(results) => { Ok(results) => {
dev_loaded_tool_names.extend(results.loaded.iter().cloned()); dev_loaded_tool_names.extend(results.loaded.iter().cloned());
if !dev_loaded_tool_names.is_empty() { if !dev_loaded_tool_names.is_empty() {
tracing::debug!( tracing::info!(
"Loaded {} dev WASM tools from build artifacts", "Loaded {} dev WASM tools from build artifacts",
dev_loaded_tool_names.len() dev_loaded_tool_names.len()
); );
@@ -474,10 +565,7 @@ impl AppBuilder {
Ok(servers) => { Ok(servers) => {
let enabled: Vec<_> = servers.enabled_servers().cloned().collect(); let enabled: Vec<_> = servers.enabled_servers().cloned().collect();
if !enabled.is_empty() { if !enabled.is_empty() {
tracing::debug!( tracing::info!("Loading {} configured MCP server(s)...", enabled.len());
"Loading {} configured MCP server(s)...",
enabled.len()
);
} }
let mut join_set = tokio::task::JoinSet::new(); let mut join_set = tokio::task::JoinSet::new();
@@ -490,24 +578,95 @@ impl AppBuilder {
join_set.spawn(async move { join_set.spawn(async move {
let server_name = server.name.clone(); let server_name = server.name.clone();
let client = match crate::tools::mcp::create_client_from_config( let client: McpClient = match server.effective_transport() {
server, crate::tools::mcp::config::EffectiveTransport::Stdio {
&mcp_sm, command,
&pm, args,
secrets, env,
"default", } => {
match pm
.spawn_stdio(
&server_name,
command,
args.to_vec(),
env.clone(),
) )
.await .await
{ {
Ok(c) => c, Ok(transport) => McpClient::new_with_transport(
&server_name,
transport as Arc<dyn McpTransport>,
None,
secrets,
"default",
Some(server),
),
Err(e) => { Err(e) => {
tracing::warn!( tracing::warn!(
"Failed to create MCP client for '{}': {}", "Failed to spawn stdio MCP server '{}': {}",
server_name, server_name,
e e
); );
return; return;
} }
}
}
#[cfg(unix)]
crate::tools::mcp::config::EffectiveTransport::Unix {
socket_path,
} => {
match crate::tools::mcp::unix_transport::UnixMcpTransport::connect(
&server_name,
socket_path,
)
.await
{
Ok(transport) => McpClient::new_with_transport(
&server_name,
Arc::new(transport) as Arc<dyn McpTransport>,
None,
secrets,
"default",
Some(server),
),
Err(e) => {
tracing::warn!(
"Failed to connect to Unix MCP server '{}': {}",
server_name,
e
);
return;
}
}
}
#[cfg(not(unix))]
crate::tools::mcp::config::EffectiveTransport::Unix { .. } => {
tracing::warn!(
"Unix socket transport is not supported on this platform (server '{}')",
server_name
);
return;
}
crate::tools::mcp::config::EffectiveTransport::Http => {
if let Some(ref secrets) = secrets {
let has_tokens =
is_authenticated(&server, secrets, "default")
.await;
if has_tokens || server.requires_auth() {
McpClient::new_authenticated(
server,
Arc::clone(&mcp_sm),
Arc::clone(secrets),
"default",
)
} else {
McpClient::new_with_config(server)
}
} else {
McpClient::new_with_config(server)
}
}
}; };
match client.list_tools().await { match client.list_tools().await {
@@ -518,7 +677,7 @@ impl AppBuilder {
for tool in tool_impls { for tool in tool_impls {
tools.register(tool).await; tools.register(tool).await;
} }
tracing::debug!( tracing::info!(
"Loaded {} tools from MCP server '{}'", "Loaded {} tools from MCP server '{}'",
tool_count, tool_count,
server_name server_name
@@ -572,14 +731,14 @@ impl AppBuilder {
let (dev_loaded_tool_names, _) = tokio::join!(wasm_tools_future, mcp_servers_future); let (dev_loaded_tool_names, _) = tokio::join!(wasm_tools_future, mcp_servers_future);
// Load registry catalog entries for extension discovery // Load registry catalog entries for extension discovery
let mut catalog_entries = match crate::registry::RegistryCatalog::load_or_embedded() { let catalog_entries = match crate::registry::RegistryCatalog::load_or_embedded() {
Ok(catalog) => { Ok(catalog) => {
let entries: Vec<_> = catalog let entries: Vec<_> = catalog
.all() .all()
.iter() .iter()
.map(|m| m.to_registry_entry()) .map(|m| m.to_registry_entry())
.collect(); .collect();
tracing::debug!( tracing::info!(
count = entries.len(), count = entries.len(),
"Loaded registry catalog entries for extension discovery" "Loaded registry catalog entries for extension discovery"
); );
@@ -591,15 +750,6 @@ impl AppBuilder {
} }
}; };
// Append builtin entries (e.g. channel-relay integrations) so they appear
// in the web UI's available extensions list.
let builtin = crate::extensions::registry::builtin_entries();
for entry in builtin {
if !catalog_entries.iter().any(|e| e.name == entry.name) {
catalog_entries.push(entry);
}
}
// Create extension manager. Use ephemeral in-memory secrets if no // Create extension manager. Use ephemeral in-memory secrets if no
// persistent store is configured (listing/install/activate still work). // persistent store is configured (listing/install/activate still work).
let ext_secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync> = if let Some(ref s) = let ext_secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync> = if let Some(ref s) =
@@ -617,7 +767,6 @@ impl AppBuilder {
let extension_manager = { let extension_manager = {
let manager = Arc::new(ExtensionManager::new( let manager = Arc::new(ExtensionManager::new(
Arc::clone(&mcp_session_manager), Arc::clone(&mcp_session_manager),
Arc::clone(&mcp_process_manager),
ext_secrets, ext_secrets,
Arc::clone(tools), Arc::clone(tools),
Some(Arc::clone(hooks)), Some(Arc::clone(hooks)),
@@ -630,7 +779,7 @@ impl AppBuilder {
catalog_entries.clone(), catalog_entries.clone(),
)); ));
tools.register_extension_tools(Arc::clone(&manager)); tools.register_extension_tools(Arc::clone(&manager));
tracing::debug!("Extension manager initialized with in-chat discovery tools"); tracing::info!("Extension manager initialized with in-chat discovery tools");
Some(manager) Some(manager)
}; };
@@ -701,7 +850,7 @@ impl AppBuilder {
let import_path = std::path::Path::new(&import_dir); let import_path = std::path::Path::new(&import_dir);
match ws.import_from_directory(import_path).await { match ws.import_from_directory(import_path).await {
Ok(count) if count > 0 => { Ok(count) if count > 0 => {
tracing::debug!("Imported {} workspace file(s) from {}", count, import_dir); tracing::info!("Imported {} workspace file(s) from {}", count, import_dir);
} }
Ok(_) => {} Ok(_) => {}
Err(e) => { Err(e) => {
@@ -726,7 +875,7 @@ impl AppBuilder {
tokio::spawn(async move { tokio::spawn(async move {
match ws_bg.backfill_embeddings().await { match ws_bg.backfill_embeddings().await {
Ok(count) if count > 0 => { Ok(count) if count > 0 => {
tracing::debug!("Backfilled embeddings for {} chunks", count); tracing::info!("Backfilled embeddings for {} chunks", count);
} }
Ok(_) => {} Ok(_) => {}
Err(e) => { Err(e) => {
@@ -743,7 +892,7 @@ impl AppBuilder {
.with_installed_dir(self.config.skills.installed_dir.clone()); .with_installed_dir(self.config.skills.installed_dir.clone());
let loaded = registry.discover_all().await; let loaded = registry.discover_all().await;
if !loaded.is_empty() { if !loaded.is_empty() {
tracing::debug!("Loaded {} skill(s): {}", loaded.len(), loaded.join(", ")); tracing::info!("Loaded {} skill(s): {}", loaded.len(), loaded.join(", "));
} }
let registry = Arc::new(std::sync::RwLock::new(registry)); let registry = Arc::new(std::sync::RwLock::new(registry));
let catalog = crate::skills::catalog::shared_catalog(); let catalog = crate::skills::catalog::shared_catalog();
@@ -761,7 +910,7 @@ impl AppBuilder {
}, },
)); ));
tracing::debug!( tracing::info!(
"Tool registry initialized with {} total tools", "Tool registry initialized with {} total tools",
tools.count() tools.count()
); );
-156
View File
@@ -198,58 +198,6 @@ pub fn save_bootstrap_env_to(path: &std::path::Path, vars: &[(&str, &str)]) -> s
Ok(()) Ok(())
} }
/// Update or add multiple variables in `~/.ironclaw/.env`, preserving existing content.
///
/// Like `upsert_bootstrap_var` but batched — replaces lines for any key in `vars`
/// and preserves all other existing lines. Use this instead of `save_bootstrap_env`
/// when you want to update specific keys without destroying user-added variables.
pub fn upsert_bootstrap_vars(vars: &[(&str, &str)]) -> std::io::Result<()> {
upsert_bootstrap_vars_to(&ironclaw_env_path(), vars)
}
/// Update or add multiple variables at an arbitrary path (testable variant).
pub fn upsert_bootstrap_vars_to(
path: &std::path::Path,
vars: &[(&str, &str)],
) -> std::io::Result<()> {
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
let keys_being_written: std::collections::HashSet<&str> =
vars.iter().map(|(k, _)| *k).collect();
let existing = match std::fs::read_to_string(path) {
Ok(contents) => contents,
Err(e) if e.kind() == std::io::ErrorKind::NotFound => String::new(),
Err(e) => return Err(e),
};
let mut result = String::new();
for line in existing.lines() {
// Extract key from lines matching `KEY=...`
let is_overwritten = line
.split_once('=')
.map(|(k, _)| keys_being_written.contains(k.trim()))
.unwrap_or(false);
if !is_overwritten {
result.push_str(line);
result.push('\n');
}
}
// Append all new key=value pairs
for (key, value) in vars {
let escaped = value.replace('\\', "\\\\").replace('"', "\\\"");
result.push_str(&format!("{}=\"{}\"\n", key, escaped));
}
std::fs::write(path, &result)?;
restrict_file_permissions(path)?;
Ok(())
}
/// Update or add a single variable in `~/.ironclaw/.env`, preserving existing content. /// Update or add a single variable in `~/.ironclaw/.env`, preserving existing content.
/// ///
/// Unlike `save_bootstrap_env` (which overwrites the entire file), this /// Unlike `save_bootstrap_env` (which overwrites the entire file), this
@@ -1289,108 +1237,4 @@ INJECTED="pwned"#;
let lock = PidLock::acquire_at(pid_path).unwrap(); let lock = PidLock::acquire_at(pid_path).unwrap();
drop(lock); drop(lock);
} }
#[test]
fn upsert_bootstrap_vars_preserves_unknown_keys() {
let dir = tempdir().unwrap();
let env_path = dir.path().join(".env");
// Simulate a user-edited .env with custom vars
let initial =
"HTTP_HOST=\"0.0.0.0\"\nDATABASE_BACKEND=\"postgres\"\nCUSTOM_VAR=\"keep_me\"\n";
std::fs::write(&env_path, initial).unwrap();
// Upsert wizard vars — should preserve HTTP_HOST and CUSTOM_VAR
let vars = [("DATABASE_BACKEND", "libsql"), ("LLM_BACKEND", "openai")];
upsert_bootstrap_vars_to(&env_path, &vars).unwrap();
let parsed: Vec<(String, String)> = dotenvy::from_path_iter(&env_path)
.unwrap()
.filter_map(|r| r.ok())
.collect();
assert_eq!(
parsed.len(),
4,
"should have 4 vars (2 preserved + 2 upserted)"
);
// User-added vars must be preserved
assert!(
parsed
.iter()
.any(|(k, v)| k == "HTTP_HOST" && v == "0.0.0.0"),
"HTTP_HOST must be preserved"
);
assert!(
parsed
.iter()
.any(|(k, v)| k == "CUSTOM_VAR" && v == "keep_me"),
"CUSTOM_VAR must be preserved"
);
// Wizard vars must be updated/added
assert!(
parsed
.iter()
.any(|(k, v)| k == "DATABASE_BACKEND" && v == "libsql"),
"DATABASE_BACKEND must be updated to libsql"
);
assert!(
parsed
.iter()
.any(|(k, v)| k == "LLM_BACKEND" && v == "openai"),
"LLM_BACKEND must be added"
);
// Now update LLM_BACKEND and verify HTTP_HOST still preserved
let vars2 = [("LLM_BACKEND", "anthropic")];
upsert_bootstrap_vars_to(&env_path, &vars2).unwrap();
let parsed2: Vec<(String, String)> = dotenvy::from_path_iter(&env_path)
.unwrap()
.filter_map(|r| r.ok())
.collect();
assert_eq!(
parsed2.len(),
4,
"should still have 4 vars after second upsert"
);
assert!(
parsed2
.iter()
.any(|(k, v)| k == "HTTP_HOST" && v == "0.0.0.0"),
"HTTP_HOST must still be preserved after second upsert"
);
assert!(
parsed2
.iter()
.any(|(k, v)| k == "LLM_BACKEND" && v == "anthropic"),
"LLM_BACKEND must be updated to anthropic"
);
}
#[test]
fn upsert_bootstrap_vars_creates_file_if_missing() {
let dir = tempdir().unwrap();
let env_path = dir.path().join("subdir").join(".env");
// File doesn't exist yet
assert!(!env_path.exists());
let vars = [("DATABASE_BACKEND", "libsql")];
upsert_bootstrap_vars_to(&env_path, &vars).unwrap();
assert!(env_path.exists());
let parsed: Vec<(String, String)> = dotenvy::from_path_iter(&env_path)
.unwrap()
.filter_map(|r| r.ok())
.collect();
assert_eq!(parsed.len(), 1);
assert_eq!(
parsed[0],
("DATABASE_BACKEND".to_string(), "libsql".to_string())
);
}
} }
+2 -21
View File
@@ -344,28 +344,9 @@ pub trait Channel: Send + Sync {
} }
} }
/// Trait for channels that support hot-secret-swapping during SIGHUP reload.
///
/// This allows channels to update authentication credentials without restarting,
/// enabling zero-downtime configuration reloads. Channels that don't support
/// secret updates can simply not implement this trait.
#[async_trait]
pub trait ChannelSecretUpdater: Send + Sync {
/// Update the secret for this channel.
///
/// Called during SIGHUP configuration reload. Implementation should:
/// - Apply the new secret atomically
/// - Not fail the entire reload if secret update fails
/// - Log appropriate errors/info messages
///
/// The secret is optional (may be None if secret is no longer configured).
async fn update_secret(&self, new_secret: Option<secrecy::SecretString>);
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use crate::testing::credentials::TEST_REDACT_SECRET_123;
/// Stub tool that marks `"value"` as sensitive. /// Stub tool that marks `"value"` as sensitive.
struct SecretTool; struct SecretTool;
@@ -395,7 +376,7 @@ mod tests {
#[test] #[test]
fn tool_completed_redacts_sensitive_params_on_failure() { fn tool_completed_redacts_sensitive_params_on_failure() {
let params = serde_json::json!({"name": "api_key", "value": TEST_REDACT_SECRET_123}); let params = serde_json::json!({"name": "api_key", "value": "sk-secret-123"});
let err: Result<String, crate::error::Error> = let err: Result<String, crate::error::Error> =
Err(crate::error::ToolError::ExecutionFailed { Err(crate::error::ToolError::ExecutionFailed {
name: "secret_save".into(), name: "secret_save".into(),
@@ -430,7 +411,7 @@ mod tests {
param_str param_str
); );
assert!( assert!(
!param_str.contains(TEST_REDACT_SECRET_123), !param_str.contains("sk-secret-123"),
"raw secret should not appear: {}", "raw secret should not appear: {}",
param_str param_str
); );
+13 -202
View File
@@ -10,7 +10,7 @@ use axum::{
response::IntoResponse, response::IntoResponse,
routing::{get, post}, routing::{get, post},
}; };
use secrecy::{ExposeSecret, SecretString}; use secrecy::ExposeSecret;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use subtle::ConstantTimeEq; use subtle::ConstantTimeEq;
use tokio::sync::{RwLock, mpsc, oneshot}; use tokio::sync::{RwLock, mpsc, oneshot};
@@ -18,8 +18,7 @@ use tokio_stream::wrappers::ReceiverStream;
use uuid::Uuid; use uuid::Uuid;
use crate::channels::{ use crate::channels::{
AttachmentKind, Channel, ChannelSecretUpdater, IncomingAttachment, IncomingMessage, AttachmentKind, Channel, IncomingAttachment, IncomingMessage, MessageStream, OutgoingResponse,
MessageStream, OutgoingResponse,
}; };
use crate::config::HttpConfig; use crate::config::HttpConfig;
use crate::error::ChannelError; use crate::error::ChannelError;
@@ -30,16 +29,13 @@ pub struct HttpChannel {
state: Arc<HttpChannelState>, state: Arc<HttpChannelState>,
} }
pub struct HttpChannelState { struct HttpChannelState {
/// Sender for incoming messages. /// Sender for incoming messages.
tx: RwLock<Option<mpsc::Sender<IncomingMessage>>>, tx: RwLock<Option<mpsc::Sender<IncomingMessage>>>,
/// Pending responses keyed by message ID. /// Pending responses keyed by message ID.
pending_responses: RwLock<std::collections::HashMap<Uuid, oneshot::Sender<String>>>, pending_responses: RwLock<std::collections::HashMap<Uuid, oneshot::Sender<String>>>,
/// Expected webhook secret for authentication (if configured). /// Expected webhook secret for authentication (if configured).
/// Stored in a separate Arc<RwLock<>> to avoid contending with other state operations. webhook_secret: Option<String>,
/// Rarely changes (only on SIGHUP), so isolated from hot-path state accesses.
/// Uses SecretString to prevent accidental logging and memory dump exposure.
webhook_secret: Arc<RwLock<Option<SecretString>>>,
/// Fixed user ID for this HTTP channel. /// Fixed user ID for this HTTP channel.
user_id: String, user_id: String,
/// Rate limiting state. /// Rate limiting state.
@@ -52,14 +48,6 @@ struct RateLimitState {
request_count: u32, request_count: u32,
} }
impl HttpChannelState {
/// Update the webhook secret in-place without restarting the listener.
/// Called during SIGHUP to hot-swap credentials.
pub async fn update_secret(&self, new_secret: Option<SecretString>) {
*self.webhook_secret.write().await = new_secret;
}
}
/// Maximum JSON body size for webhook requests (15 MB, to support base64 image attachments /// Maximum JSON body size for webhook requests (15 MB, to support base64 image attachments
/// with ~33% overhead from base64 encoding). /// with ~33% overhead from base64 encoding).
const MAX_BODY_BYTES: usize = 15 * 1024 * 1024; const MAX_BODY_BYTES: usize = 15 * 1024 * 1024;
@@ -79,7 +67,7 @@ impl HttpChannel {
let webhook_secret = config let webhook_secret = config
.webhook_secret .webhook_secret
.as_ref() .as_ref()
.map(|s| SecretString::from(s.expose_secret().to_string())); .map(|s| s.expose_secret().to_string());
let user_id = config.user_id.clone(); let user_id = config.user_id.clone();
Self { Self {
@@ -87,7 +75,7 @@ impl HttpChannel {
state: Arc::new(HttpChannelState { state: Arc::new(HttpChannelState {
tx: RwLock::new(None), tx: RwLock::new(None),
pending_responses: RwLock::new(std::collections::HashMap::new()), pending_responses: RwLock::new(std::collections::HashMap::new()),
webhook_secret: Arc::new(RwLock::new(webhook_secret)), webhook_secret,
user_id, user_id,
rate_limit: tokio::sync::Mutex::new(RateLimitState { rate_limit: tokio::sync::Mutex::new(RateLimitState {
window_start: std::time::Instant::now(), window_start: std::time::Instant::now(),
@@ -114,16 +102,6 @@ impl HttpChannel {
pub fn addr(&self) -> (&str, u16) { pub fn addr(&self) -> (&str, u16) {
(&self.config.host, self.config.port) (&self.config.host, self.config.port)
} }
/// Return a shared handle to the channel state for out-of-band updates.
pub fn shared_state(&self) -> Arc<HttpChannelState> {
Arc::clone(&self.state)
}
/// Update the webhook secret in-place without restarting the listener.
pub async fn update_secret(&self, new_secret: Option<SecretString>) {
self.state.update_secret(new_secret).await;
}
} }
#[derive(Debug, Deserialize)] #[derive(Debug, Deserialize)]
@@ -223,10 +201,9 @@ async fn webhook_handler(
}); });
// Validate secret if configured // Validate secret if configured
if let Some(ref expected_secret) = *state.webhook_secret.read().await { if let Some(ref expected_secret) = state.webhook_secret {
let expected_bytes = expected_secret.expose_secret().as_bytes();
match &req.secret { match &req.secret {
Some(provided) if bool::from(provided.as_bytes().ct_eq(expected_bytes)) => { Some(provided) if bool::from(provided.as_bytes().ct_eq(expected_secret.as_bytes())) => {
// Secret matches, continue // Secret matches, continue
} }
Some(_) => { Some(_) => {
@@ -395,14 +372,9 @@ async fn process_message(
None None
}; };
// Clone sender while holding read lock, then release lock before async send. // Send message to the channel
// This prevents blocking other webhook handlers during the async I/O. let tx_guard = state.tx.read().await;
let tx = { if let Some(tx) = tx_guard.as_ref() {
let guard = state.tx.read().await;
guard.as_ref().cloned()
};
if let Some(tx) = tx {
if tx.send(msg).await.is_err() { if tx.send(msg).await.is_err() {
return ( return (
StatusCode::INTERNAL_SERVER_ERROR, StatusCode::INTERNAL_SERVER_ERROR,
@@ -423,6 +395,7 @@ async fn process_message(
}), }),
); );
} }
drop(tx_guard);
// Wait for response if requested // Wait for response if requested
let response = if let Some(rx) = response_rx { let response = if let Some(rx) = response_rx {
@@ -455,7 +428,7 @@ impl Channel for HttpChannel {
} }
async fn start(&self) -> Result<MessageStream, ChannelError> { async fn start(&self) -> Result<MessageStream, ChannelError> {
if self.state.webhook_secret.read().await.is_none() { if self.state.webhook_secret.is_none() {
return Err(ChannelError::StartupFailed { return Err(ChannelError::StartupFailed {
name: "http".to_string(), name: "http".to_string(),
reason: "HTTP webhook secret is required (set HTTP_WEBHOOK_SECRET)".to_string(), reason: "HTTP webhook secret is required (set HTTP_WEBHOOK_SECRET)".to_string(),
@@ -502,16 +475,6 @@ impl Channel for HttpChannel {
} }
} }
/// Implement secret update for HTTP channel state.
/// This allows SIGHUP handler to update secrets generically via the trait.
#[async_trait]
impl ChannelSecretUpdater for HttpChannelState {
async fn update_secret(&self, new_secret: Option<SecretString>) {
*self.webhook_secret.write().await = new_secret;
tracing::info!("HTTP webhook secret updated");
}
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use axum::body::Body; use axum::body::Body;
@@ -599,156 +562,4 @@ mod tests {
let resp = app.oneshot(req).await.unwrap(); let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
} }
#[tokio::test]
async fn test_update_secret_hot_swap() {
let channel = test_channel(Some("old-secret"));
let _stream = channel.start().await.unwrap();
let app1 = channel.routes();
// Request with old-secret should succeed
let body_old = serde_json::json!({
"content": "hello",
"secret": "old-secret"
});
let req1 = Request::builder()
.method("POST")
.uri("/webhook")
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&body_old).unwrap()))
.unwrap();
let resp1 = app1.oneshot(req1).await.unwrap();
assert_eq!(
resp1.status(),
StatusCode::OK,
"old secret should work initially"
);
// Update secret to new-secret
channel
.update_secret(Some(SecretString::from("new-secret".to_string())))
.await;
let app2 = channel.routes();
// Request with old-secret should fail
let req2 = Request::builder()
.method("POST")
.uri("/webhook")
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&body_old).unwrap()))
.unwrap();
let resp2 = app2.oneshot(req2).await.unwrap();
assert_eq!(
resp2.status(),
StatusCode::UNAUTHORIZED,
"old secret should fail after update"
);
let app3 = channel.routes();
// Request with new-secret should succeed
let body_new = serde_json::json!({
"content": "hello",
"secret": "new-secret"
});
let req3 = Request::builder()
.method("POST")
.uri("/webhook")
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&body_new).unwrap()))
.unwrap();
let resp3 = app3.oneshot(req3).await.unwrap();
assert_eq!(
resp3.status(),
StatusCode::OK,
"new secret should work after update"
);
}
#[tokio::test]
async fn test_concurrent_requests_during_secret_update() {
use std::sync::Arc as StdArc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
let channel = test_channel(Some("initial-secret"));
let _stream = channel.start().await.unwrap();
let app = channel.routes();
// Counters for request outcomes
let success_count = StdArc::new(AtomicUsize::new(0));
let mut handles = vec![];
// Spawn 5 concurrent tasks that keep making requests with the initial secret
for i in 0..5 {
let app = app.clone();
let success = StdArc::clone(&success_count);
let handle = tokio::spawn(async move {
let body = serde_json::json!({
"content": format!("test-{}", i),
"secret": "initial-secret"
});
let req = Request::builder()
.method("POST")
.uri("/webhook")
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&body).unwrap()))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
if resp.status() == StatusCode::OK {
success.fetch_add(1, Ordering::SeqCst);
}
});
handles.push(handle);
}
// Update secret mid-flight (tests that RwLock allows readers while writer holds lock)
tokio::time::sleep(Duration::from_millis(5)).await;
channel
.update_secret(Some(SecretString::from("updated-secret".to_string())))
.await;
// Spawn 5 more tasks that use the new secret
for i in 5..10 {
let app = app.clone();
let success = StdArc::clone(&success_count);
let handle = tokio::spawn(async move {
let body = serde_json::json!({
"content": format!("test-{}", i),
"secret": "updated-secret"
});
let req = Request::builder()
.method("POST")
.uri("/webhook")
.header("content-type", "application/json")
.body(Body::from(serde_json::to_vec(&body).unwrap()))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
if resp.status() == StatusCode::OK {
success.fetch_add(1, Ordering::SeqCst);
}
});
handles.push(handle);
}
// Wait for all tasks to complete
for handle in handles {
let _ = handle.await;
}
// Verify all requests succeeded with their respective secrets
assert_eq!(
success_count.load(Ordering::SeqCst),
10,
"All concurrent requests should succeed with correct secrets after update"
);
}
} }
+2 -39
View File
@@ -56,17 +56,6 @@ impl ChannelManager {
/// the agent loop. /// the agent loop.
pub async fn hot_add(&self, channel: Box<dyn Channel>) -> Result<(), ChannelError> { pub async fn hot_add(&self, channel: Box<dyn Channel>) -> Result<(), ChannelError> {
let name = channel.name().to_string(); let name = channel.name().to_string();
// Shut down any existing channel with the same name to avoid parallel consumers.
// The old forwarding task will stop when the channel's stream ends after shutdown.
{
let channels = self.channels.read().await;
if let Some(existing) = channels.get(&name) {
tracing::debug!(channel = %name, "Shutting down existing channel before hot-add replacement");
let _ = existing.shutdown().await;
}
}
let stream = channel.start().await?; let stream = channel.start().await?;
// Register for respond/broadcast/send_status // Register for respond/broadcast/send_status
@@ -86,7 +75,7 @@ impl ChannelManager {
break; break;
} }
} }
tracing::debug!(channel = %name, "Hot-added channel stream ended"); tracing::info!(channel = %name, "Hot-added channel stream ended");
}); });
Ok(()) Ok(())
@@ -103,7 +92,7 @@ impl ChannelManager {
for (name, channel) in channels.iter() { for (name, channel) in channels.iter() {
match channel.start().await { match channel.start().await {
Ok(stream) => { Ok(stream) => {
tracing::debug!("Started channel: {}", name); tracing::info!("Started channel: {}", name);
streams.push(stream); streams.push(stream);
} }
Err(e) => { Err(e) => {
@@ -348,30 +337,4 @@ mod tests {
let msg = stream.next().await.expect("stream ended"); let msg = stream.next().await.expect("stream ended");
assert_eq!(msg.content, "background alert"); assert_eq!(msg.content, "background alert");
} }
#[tokio::test]
async fn test_hot_add_replaces_existing_channel() {
// Regression: hot_add must shut down the existing channel before replacing it,
// to prevent duplicate SSE consumers from running in parallel.
let manager = ChannelManager::new();
let (stub1, _tx1) = StubChannel::new("relay");
manager.add(Box::new(stub1)).await;
let mut stream = manager.start_all().await.expect("start_all");
// Hot-add a replacement channel with the same name
let (stub2, tx2) = StubChannel::new("relay");
manager.hot_add(Box::new(stub2)).await.expect("hot_add");
// Send through the new channel — should arrive in the merged stream
tx2.send(IncomingMessage::new("relay", "u1", "from new"))
.await
.expect("send");
let msg = stream.next().await.expect("stream");
assert_eq!(msg.content, "from new");
// Verify only one channel entry exists
let channels = manager.channels.read().await;
assert_eq!(channels.len(), 1);
assert!(channels.contains_key("relay"));
}
} }
+3 -4
View File
@@ -30,7 +30,6 @@
mod channel; mod channel;
mod http; mod http;
mod manager; mod manager;
pub mod relay;
mod repl; mod repl;
mod signal; mod signal;
pub mod wasm; pub mod wasm;
@@ -38,10 +37,10 @@ pub mod web;
mod webhook_server; mod webhook_server;
pub use channel::{ pub use channel::{
AttachmentKind, Channel, ChannelSecretUpdater, IncomingAttachment, IncomingMessage, AttachmentKind, Channel, IncomingAttachment, IncomingMessage, MessageStream, OutgoingResponse,
MessageStream, OutgoingResponse, StatusUpdate, StatusUpdate,
}; };
pub use http::{HttpChannel, HttpChannelState}; pub use http::HttpChannel;
pub use manager::ChannelManager; pub use manager::ChannelManager;
pub use repl::ReplChannel; pub use repl::ReplChannel;
pub use signal::SignalChannel; pub use signal::SignalChannel;
-642
View File
@@ -1,642 +0,0 @@
//! Channel trait implementation for channel-relay SSE streams.
//!
//! `RelayChannel` connects to a channel-relay service via SSE, converts
//! incoming events to `IncomingMessage`s, and sends responses via the
//! relay's provider-specific proxy API (Slack).
use std::collections::HashMap;
use std::sync::Arc;
use async_trait::async_trait;
use tokio::sync::{RwLock, mpsc};
use crate::channels::relay::client::{RelayClient, RelayError};
use crate::channels::{Channel, IncomingMessage, MessageStream, OutgoingResponse, StatusUpdate};
use crate::error::ChannelError;
/// Default channel name for the Slack relay integration.
pub const DEFAULT_RELAY_NAME: &str = "slack-relay";
/// The messaging provider backing a relay channel.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RelayProvider {
Slack,
}
impl RelayProvider {
/// Provider string used in proxy API routes and metadata.
pub fn as_str(&self) -> &'static str {
match self {
Self::Slack => "slack",
}
}
/// The default channel name for this provider.
pub fn channel_name(&self) -> &'static str {
match self {
Self::Slack => DEFAULT_RELAY_NAME,
}
}
}
/// Channel implementation that connects to a channel-relay SSE stream.
pub struct RelayChannel {
client: RelayClient,
provider: RelayProvider,
stream_token: Arc<RwLock<String>>,
team_id: String,
instance_id: String,
user_id: String,
/// SSE stream long-poll timeout in seconds.
stream_timeout_secs: u64,
/// Initial exponential backoff in milliseconds.
backoff_initial_ms: u64,
/// Maximum exponential backoff in milliseconds.
backoff_max_ms: u64,
/// Handle to the reconnect task for clean shutdown.
reconnect_handle: RwLock<Option<tokio::task::JoinHandle<()>>>,
/// Handle to the SSE parser task for clean shutdown.
parser_handle: Arc<RwLock<Option<tokio::task::JoinHandle<()>>>>,
/// Maximum consecutive reconnect failures before giving up.
max_consecutive_failures: u64,
}
impl RelayChannel {
/// Create a new relay channel for Slack (default provider).
pub fn new(
client: RelayClient,
stream_token: String,
team_id: String,
instance_id: String,
user_id: String,
) -> Self {
Self::new_with_provider(
client,
RelayProvider::Slack,
stream_token,
team_id,
instance_id,
user_id,
)
}
/// Create a new relay channel with a specific provider.
pub fn new_with_provider(
client: RelayClient,
provider: RelayProvider,
stream_token: String,
team_id: String,
instance_id: String,
user_id: String,
) -> Self {
Self {
client,
provider,
stream_token: Arc::new(RwLock::new(stream_token)),
team_id,
instance_id,
user_id,
stream_timeout_secs: 86400,
backoff_initial_ms: 1000,
backoff_max_ms: 60000,
reconnect_handle: RwLock::new(None),
parser_handle: Arc::new(RwLock::new(None)),
max_consecutive_failures: 50,
}
}
/// Set backoff/timeout parameters from relay config values.
pub fn with_timeouts(
mut self,
stream_timeout_secs: u64,
backoff_initial_ms: u64,
backoff_max_ms: u64,
) -> Self {
self.stream_timeout_secs = stream_timeout_secs;
self.backoff_initial_ms = backoff_initial_ms;
self.backoff_max_ms = backoff_max_ms;
self
}
/// Set the maximum number of consecutive reconnect failures before giving up.
pub fn with_max_failures(mut self, max: u64) -> Self {
self.max_consecutive_failures = max;
self
}
/// Build a provider-appropriate proxy body for sending a message.
fn build_send_body(
&self,
channel_id: &str,
text: &str,
thread_id: Option<&str>,
) -> (String, serde_json::Value) {
match self.provider {
RelayProvider::Slack => {
let mut body = serde_json::json!({
"channel": channel_id,
"text": text,
});
if let Some(tid) = thread_id {
body["thread_ts"] = serde_json::Value::String(tid.to_string());
}
("chat.postMessage".to_string(), body)
}
}
}
/// Send a message via the provider proxy.
async fn proxy_send(
&self,
team_id: &str,
method: &str,
body: serde_json::Value,
) -> Result<serde_json::Value, RelayError> {
self.client
.proxy_provider(
self.provider.as_str(),
team_id,
method,
body,
Some(&self.instance_id),
)
.await
}
}
#[async_trait]
impl Channel for RelayChannel {
fn name(&self) -> &str {
self.provider.channel_name()
}
async fn start(&self) -> Result<MessageStream, ChannelError> {
let channel_name = self.name().to_string();
let token = self.stream_token.read().await.clone();
let (stream, initial_parser_handle) = self
.client
.connect_stream(&token, self.stream_timeout_secs)
.await
.map_err(|e| ChannelError::StartupFailed {
name: channel_name.clone(),
reason: e.to_string(),
})?;
*self.parser_handle.write().await = Some(initial_parser_handle);
let (tx, rx) = mpsc::channel(64);
// Spawn the stream reader + reconnect task
let client = self.client.clone();
let stream_token = Arc::clone(&self.stream_token);
let instance_id = self.instance_id.clone();
let user_id = self.user_id.clone();
let team_id = self.team_id.clone();
let stream_timeout_secs = self.stream_timeout_secs;
let backoff_initial_ms = self.backoff_initial_ms;
let backoff_max_ms = self.backoff_max_ms;
let max_consecutive_failures = self.max_consecutive_failures;
let parser_handle = Arc::clone(&self.parser_handle);
let provider_str = self.provider.as_str().to_string();
let relay_name = channel_name.clone();
let handle = tokio::spawn(async move {
use futures::StreamExt;
let mut current_stream = stream;
let mut backoff_ms = backoff_initial_ms;
let mut consecutive_failures: u64 = 0;
loop {
// Read events from the current stream
while let Some(event) = current_stream.next().await {
// Reset backoff and failure count on successful event
backoff_ms = backoff_initial_ms;
consecutive_failures = 0;
// Validate required fields
if event.sender_id.is_empty()
|| event.channel_id.is_empty()
|| event.provider_scope.is_empty()
{
tracing::debug!(
event_type = %event.event_type,
sender_id = %event.sender_id,
channel_id = %event.channel_id,
"Relay: skipping event with missing required fields"
);
continue;
}
// Skip non-message events
if !event.is_message() {
tracing::debug!(
event_type = %event.event_type,
"Relay: skipping non-message event"
);
continue;
}
tracing::info!(
event_type = %event.event_type,
sender = %event.sender_id,
channel = %event.channel_id,
provider = %provider_str,
"Relay: received message from {}", provider_str
);
let msg = IncomingMessage::new(&relay_name, &event.sender_id, event.text())
.with_user_name(event.display_name())
.with_metadata(serde_json::json!({
"team_id": event.team_id(),
"channel_id": event.channel_id,
"sender_id": event.sender_id,
"sender_name": event.display_name(),
"event_type": event.event_type,
"thread_id": event.thread_id,
"provider": event.provider,
}));
let msg = if let Some(ref thread_id) = event.thread_id {
msg.with_thread(thread_id)
} else {
msg.with_thread(&event.channel_id)
};
if tx.send(msg).await.is_err() {
tracing::info!("Relay channel receiver dropped, stopping");
return;
}
}
// Stream ended, attempt reconnect with backoff
consecutive_failures += 1;
if consecutive_failures >= max_consecutive_failures {
tracing::error!(
channel = %relay_name,
failures = consecutive_failures,
"Relay channel giving up after {} consecutive failures",
consecutive_failures
);
break;
}
tracing::warn!(
backoff_ms = backoff_ms,
failures = consecutive_failures,
"Relay SSE stream ended, reconnecting..."
);
tokio::time::sleep(std::time::Duration::from_millis(backoff_ms)).await;
backoff_ms = (backoff_ms * 2).min(backoff_max_ms);
// Try to reconnect
let token = stream_token.read().await.clone();
match client.connect_stream(&token, stream_timeout_secs).await {
Ok((new_stream, new_parser)) => {
tracing::info!("Relay SSE stream reconnected");
current_stream = new_stream;
// Abort old parser before replacing
if let Some(old) = parser_handle.write().await.take() {
old.abort();
}
*parser_handle.write().await = Some(new_parser);
}
Err(RelayError::TokenExpired) => {
// Attempt token renewal
tracing::info!("Relay stream token expired, renewing...");
match client.renew_token(&instance_id, &user_id).await {
Ok(new_token) => {
*stream_token.write().await = new_token.clone();
match client.connect_stream(&new_token, stream_timeout_secs).await {
Ok((new_stream, new_parser)) => {
tracing::info!(
"Relay SSE stream reconnected with new token"
);
current_stream = new_stream;
if let Some(old) = parser_handle.write().await.take() {
old.abort();
}
*parser_handle.write().await = Some(new_parser);
}
Err(e) => {
tracing::error!(
error = %e,
"Failed to reconnect after token renewal"
);
}
}
}
Err(e) => {
tracing::error!(
error = %e,
"Failed to renew relay stream token"
);
}
}
}
Err(e) => {
tracing::error!(error = %e, "Failed to reconnect relay SSE stream");
}
}
// Check if the team is still valid (skip when team_id is unknown,
// e.g. when no DB store was available at activation time)
if !team_id.is_empty() {
match client.list_connections(&instance_id).await {
Ok(conns) => {
let has_team =
conns.iter().any(|c| c.team_id == team_id && c.connected);
if !has_team {
tracing::warn!(
team_id = %team_id,
"Team no longer connected, stopping relay channel"
);
return;
}
}
Err(e) => {
tracing::warn!(
error = %e,
"Could not verify team connection, will retry next iteration"
);
}
}
}
}
});
*self.reconnect_handle.write().await = Some(handle);
let stream = tokio_stream::wrappers::ReceiverStream::new(rx);
Ok(Box::pin(stream))
}
async fn respond(
&self,
msg: &IncomingMessage,
response: OutgoingResponse,
) -> Result<(), ChannelError> {
let channel_name = self.name().to_string();
let metadata = &msg.metadata;
let team_id = metadata
.get("team_id")
.and_then(|v| v.as_str())
.unwrap_or(&self.team_id);
let channel_id = metadata
.get("channel_id")
.and_then(|v| v.as_str())
.ok_or_else(|| ChannelError::SendFailed {
name: channel_name.clone(),
reason: "Missing channel_id in message metadata".to_string(),
})?;
// Determine thread_id from response or metadata
let thread_id = response
.thread_id
.as_deref()
.or_else(|| metadata.get("thread_id").and_then(|v| v.as_str()));
let (method, body) = self.build_send_body(channel_id, &response.content, thread_id);
self.proxy_send(team_id, &method, body)
.await
.map_err(|e| ChannelError::SendFailed {
name: channel_name,
reason: e.to_string(),
})?;
Ok(())
}
/// Status updates are not forwarded to messaging providers to avoid noise.
async fn send_status(
&self,
_status: StatusUpdate,
_metadata: &serde_json::Value,
) -> Result<(), ChannelError> {
Ok(())
}
async fn broadcast(
&self,
target: &str,
response: OutgoingResponse,
) -> Result<(), ChannelError> {
let channel_name = self.name().to_string();
// Determine thread_id from response or metadata
let thread_id = response
.thread_id
.as_deref()
.or_else(|| response.metadata.get("thread_ts").and_then(|v| v.as_str()));
let (method, body) = self.build_send_body(target, &response.content, thread_id);
self.proxy_send(&self.team_id, &method, body)
.await
.map_err(|e| ChannelError::SendFailed {
name: channel_name,
reason: e.to_string(),
})?;
Ok(())
}
async fn health_check(&self) -> Result<(), ChannelError> {
self.client
.list_connections(&self.instance_id)
.await
.map_err(|_| ChannelError::HealthCheckFailed {
name: self.name().to_string(),
})?;
Ok(())
}
fn conversation_context(&self, metadata: &serde_json::Value) -> HashMap<String, String> {
let mut ctx = HashMap::new();
if let Some(sender) = metadata.get("sender_name").and_then(|v| v.as_str()) {
ctx.insert("sender".to_string(), sender.to_string());
}
if let Some(sender_id) = metadata.get("sender_id").and_then(|v| v.as_str()) {
ctx.insert("sender_uuid".to_string(), sender_id.to_string());
}
if let Some(channel_id) = metadata.get("channel_id").and_then(|v| v.as_str()) {
ctx.insert("group".to_string(), channel_id.to_string());
}
ctx.insert("platform".to_string(), self.provider.as_str().to_string());
ctx
}
async fn shutdown(&self) -> Result<(), ChannelError> {
if let Some(handle) = self.reconnect_handle.write().await.take() {
handle.abort();
}
if let Some(handle) = self.parser_handle.write().await.take() {
handle.abort();
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
fn test_client() -> RelayClient {
RelayClient::new(
"http://localhost:3001".into(),
secrecy::SecretString::from("key".to_string()),
30,
)
.expect("client")
}
#[test]
fn relay_channel_name() {
let channel = RelayChannel::new(
test_client(),
"token".into(),
"T123".into(),
"inst1".into(),
"user1".into(),
);
assert_eq!(channel.name(), DEFAULT_RELAY_NAME);
}
#[test]
fn conversation_context_extracts_metadata() {
let channel = RelayChannel::new(
test_client(),
"token".into(),
"T123".into(),
"inst1".into(),
"user1".into(),
);
let metadata = serde_json::json!({
"sender_name": "bob",
"sender_id": "U123",
"channel_id": "C456",
});
let ctx = channel.conversation_context(&metadata);
assert_eq!(ctx.get("sender"), Some(&"bob".to_string()));
assert_eq!(ctx.get("sender_uuid"), Some(&"U123".to_string()));
assert_eq!(ctx.get("platform"), Some(&"slack".to_string()));
}
#[test]
fn metadata_shape_includes_event_type_and_sender_name() {
// Regression: metadata JSON must include event_type and sender_name
// for downstream routing (DM vs channel) and conversation_context().
let metadata = serde_json::json!({
"team_id": "T123",
"channel_id": "C456",
"sender_id": "U789",
"sender_name": "alice",
"event_type": "direct_message",
"thread_id": null,
"provider": "slack",
});
// event_type must be present for DM-vs-channel routing
assert_eq!(
metadata.get("event_type").and_then(|v| v.as_str()),
Some("direct_message")
);
// sender_name must be present for conversation_context
assert_eq!(
metadata.get("sender_name").and_then(|v| v.as_str()),
Some("alice")
);
}
#[test]
fn with_timeouts_sets_values() {
let channel = RelayChannel::new(
test_client(),
"token".into(),
"T123".into(),
"inst1".into(),
"user1".into(),
)
.with_timeouts(43200, 2000, 120000);
assert_eq!(channel.stream_timeout_secs, 43200);
assert_eq!(channel.backoff_initial_ms, 2000);
assert_eq!(channel.backoff_max_ms, 120000);
}
#[test]
fn build_send_body_slack() {
let channel = RelayChannel::new(
test_client(),
"token".into(),
"T123".into(),
"inst1".into(),
"user1".into(),
);
let (method, body) = channel.build_send_body("C456", "hello", Some("1234567.890"));
assert_eq!(method, "chat.postMessage");
assert_eq!(body["channel"], "C456");
assert_eq!(body["text"], "hello");
assert_eq!(body["thread_ts"], "1234567.890");
}
#[test]
fn parser_handle_is_shared_arc() {
let channel = RelayChannel::new(
test_client(),
"token".into(),
"T123".into(),
"inst1".into(),
"user1".into(),
);
// parser_handle should be an Arc — cloning should give a second reference
let handle_clone = Arc::clone(&channel.parser_handle);
// Both point to the same allocation
assert!(Arc::ptr_eq(&channel.parser_handle, &handle_clone));
}
#[test]
fn with_max_failures_sets_value() {
let channel = RelayChannel::new(
test_client(),
"token".into(),
"T123".into(),
"inst1".into(),
"user1".into(),
)
.with_max_failures(10);
assert_eq!(channel.max_consecutive_failures, 10);
}
#[test]
fn default_max_failures_is_50() {
let channel = RelayChannel::new(
test_client(),
"token".into(),
"T123".into(),
"inst1".into(),
"user1".into(),
);
assert_eq!(channel.max_consecutive_failures, 50);
}
#[test]
fn empty_team_id_accepted_at_construction() {
// Regression: empty team_id (when no DB store is available) must not
// prevent channel construction or cause immediate shutdown.
let channel = RelayChannel::new(
test_client(),
"token".into(),
String::new(), // empty team_id
"inst1".into(),
"user1".into(),
);
assert_eq!(channel.team_id, "");
// The reconnect loop now skips team validation when team_id is empty,
// so the channel remains alive.
}
}
-549
View File
@@ -1,549 +0,0 @@
//! HTTP client for the channel-relay service.
//!
//! Wraps reqwest for all channel-relay API calls: OAuth initiation,
//! SSE streaming, token renewal, and Slack API proxy.
use std::pin::Pin;
use std::task::{Context, Poll};
use futures::Stream;
use secrecy::{ExposeSecret, SecretString};
use serde::{Deserialize, Serialize};
use tokio::sync::mpsc;
/// Known relay event types.
pub mod event_types {
pub const MESSAGE: &str = "message";
pub const DIRECT_MESSAGE: &str = "direct_message";
pub const MENTION: &str = "mention";
}
/// A parsed SSE event from the channel-relay stream.
///
/// Field names match the channel-relay `ChannelEvent` struct exactly.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ChannelEvent {
/// Unique event ID.
#[serde(default)]
pub id: String,
/// Event type enum from channel-relay (e.g., "direct_message", "message", "mention").
pub event_type: String,
/// Provider (e.g., "slack").
#[serde(default)]
pub provider: String,
/// Team/workspace ID (called `provider_scope` in channel-relay).
#[serde(alias = "team_id", default)]
pub provider_scope: String,
/// Channel or DM conversation ID.
#[serde(default)]
pub channel_id: String,
/// Sender user ID.
#[serde(default)]
pub sender_id: String,
/// Sender display name.
#[serde(default)]
pub sender_name: Option<String>,
/// Message text content (called `content` in channel-relay).
#[serde(alias = "text", default)]
pub content: Option<String>,
/// Thread ID (for threaded replies, called `thread_id` in channel-relay).
#[serde(alias = "thread_ts", default)]
pub thread_id: Option<String>,
/// Full raw event data.
#[serde(default)]
pub raw: serde_json::Value,
/// Event timestamp (ISO 8601 from channel-relay).
#[serde(default)]
pub timestamp: Option<String>,
}
impl ChannelEvent {
/// Get the team_id (provider_scope).
pub fn team_id(&self) -> &str {
&self.provider_scope
}
/// Get the message text content.
pub fn text(&self) -> &str {
self.content.as_deref().unwrap_or("")
}
/// Get the sender name or fallback to sender_id.
pub fn display_name(&self) -> &str {
self.sender_name.as_deref().unwrap_or(&self.sender_id)
}
/// Check if this is a message-like event that should be forwarded to the agent.
pub fn is_message(&self) -> bool {
matches!(
self.event_type.as_str(),
event_types::MESSAGE | event_types::DIRECT_MESSAGE | event_types::MENTION
)
}
}
/// Connection info returned by list_connections.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Connection {
pub provider: String,
pub team_id: String,
pub team_name: Option<String>,
pub connected: bool,
}
/// HTTP client for the channel-relay service.
#[derive(Clone)]
pub struct RelayClient {
http: reqwest::Client,
base_url: String,
api_key: SecretString,
}
impl RelayClient {
/// Create a new relay client.
pub fn new(
base_url: String,
api_key: SecretString,
request_timeout_secs: u64,
) -> Result<Self, RelayError> {
let http = reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(request_timeout_secs))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| RelayError::Network(format!("Failed to build HTTP client: {e}")))?;
Ok(Self {
http,
base_url: base_url.trim_end_matches('/').to_string(),
api_key,
})
}
/// Initiate Slack OAuth flow via channel-relay.
///
/// Calls `GET /oauth/slack/auth` with `redirect(Policy::none())` and
/// returns the `Location` header (Slack OAuth URL) without following it.
pub async fn initiate_oauth(
&self,
instance_id: &str,
user_id: &str,
callback_url: &str,
) -> Result<String, RelayError> {
let resp = self
.http
.get(format!("{}/oauth/slack/auth", self.base_url))
.header("X-API-Key", self.api_key.expose_secret())
.query(&[
("instance_id", instance_id),
("user_id", user_id),
("callback", callback_url),
])
.send()
.await
.map_err(|e| RelayError::Network(e.to_string()))?;
let status = resp.status();
if status.is_redirection() {
let location = resp
.headers()
.get(reqwest::header::LOCATION)
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string())
.ok_or_else(|| {
RelayError::Protocol("Redirect response missing Location header".to_string())
})?;
Ok(location)
} else if status.is_success() {
// Some relay implementations return the URL in JSON body instead
let body: serde_json::Value = resp
.json()
.await
.map_err(|e| RelayError::Protocol(e.to_string()))?;
body.get("auth_url")
.or_else(|| body.get("url"))
.and_then(|v| v.as_str())
.map(|s| s.to_string())
.ok_or_else(|| RelayError::Protocol("Response missing auth_url field".to_string()))
} else {
let body = resp.text().await.unwrap_or_default();
Err(RelayError::Api {
status: status.as_u16(),
message: body,
})
}
}
/// Connect to the SSE event stream.
///
/// Returns a stream of parsed `ChannelEvent`s and the `JoinHandle` of the
/// background SSE parser task. The caller is responsible for reconnection
/// logic on stream end/error and for aborting the handle on shutdown.
pub async fn connect_stream(
&self,
stream_token: &str,
stream_timeout_secs: u64,
) -> Result<(ChannelEventStream, tokio::task::JoinHandle<()>), RelayError> {
let resp = self
.http
.get(format!("{}/stream", self.base_url))
.query(&[("token", stream_token)])
.timeout(std::time::Duration::from_secs(stream_timeout_secs))
.send()
.await
.map_err(|e| RelayError::Network(e.to_string()))?;
let status = resp.status();
if status == reqwest::StatusCode::UNAUTHORIZED {
return Err(RelayError::TokenExpired);
}
if !status.is_success() {
let body = resp.text().await.unwrap_or_default();
return Err(RelayError::Api {
status: status.as_u16(),
message: body,
});
}
// Spawn a background task that reads the SSE stream and sends parsed events
let (tx, rx) = mpsc::channel(64);
let byte_stream = resp.bytes_stream();
let handle = tokio::spawn(parse_sse_stream(byte_stream, tx));
Ok((ChannelEventStream { rx }, handle))
}
/// Renew an expired stream token.
///
/// Calls `POST /stream/renew` with API key auth, returns a new stream token.
pub async fn renew_token(
&self,
instance_id: &str,
user_id: &str,
) -> Result<String, RelayError> {
let resp = self
.http
.post(format!("{}/stream/renew", self.base_url))
.header("X-API-Key", self.api_key.expose_secret())
.json(&serde_json::json!({
"instance_id": instance_id,
"user_id": user_id,
}))
.send()
.await
.map_err(|e| RelayError::Network(e.to_string()))?;
let status = resp.status();
if !status.is_success() {
let body = resp.text().await.unwrap_or_default();
return Err(RelayError::Api {
status: status.as_u16(),
message: body,
});
}
let body: serde_json::Value = resp
.json()
.await
.map_err(|e| RelayError::Protocol(e.to_string()))?;
body.get("stream_token")
.or_else(|| body.get("token"))
.and_then(|v| v.as_str())
.map(|s| s.to_string())
.ok_or_else(|| RelayError::Protocol("Response missing stream_token field".to_string()))
}
/// Proxy an API call through channel-relay for any provider.
///
/// Calls `POST /proxy/{provider}/{method}?team_id=X&instance_id=Y` with the given JSON body.
pub async fn proxy_provider(
&self,
provider: &str,
team_id: &str,
method: &str,
body: serde_json::Value,
instance_id: Option<&str>,
) -> Result<serde_json::Value, RelayError> {
let mut query: Vec<(&str, &str)> = vec![("team_id", team_id)];
if let Some(iid) = instance_id {
query.push(("instance_id", iid));
}
let resp = self
.http
.post(format!("{}/proxy/{}/{}", self.base_url, provider, method))
.header("X-API-Key", self.api_key.expose_secret())
.query(&query)
.json(&body)
.send()
.await
.map_err(|e| RelayError::Network(e.to_string()))?;
if !resp.status().is_success() {
let status = resp.status().as_u16();
let body = resp.text().await.unwrap_or_default();
return Err(RelayError::Api {
status,
message: body,
});
}
resp.json()
.await
.map_err(|e| RelayError::Protocol(e.to_string()))
}
/// List active connections for an instance.
pub async fn list_connections(&self, instance_id: &str) -> Result<Vec<Connection>, RelayError> {
let resp = self
.http
.get(format!("{}/connections", self.base_url))
.header("X-API-Key", self.api_key.expose_secret())
.query(&[("instance_id", instance_id)])
.send()
.await
.map_err(|e| RelayError::Network(e.to_string()))?;
if !resp.status().is_success() {
let status = resp.status().as_u16();
let body = resp.text().await.unwrap_or_default();
return Err(RelayError::Api {
status,
message: body,
});
}
resp.json()
.await
.map_err(|e| RelayError::Protocol(e.to_string()))
}
}
/// Async stream of parsed channel events from SSE.
pub struct ChannelEventStream {
rx: mpsc::Receiver<ChannelEvent>,
}
impl Stream for ChannelEventStream {
type Item = ChannelEvent;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
self.rx.poll_recv(cx)
}
}
/// Parse SSE format from a reqwest bytes stream.
///
/// SSE format:
/// ```text
/// event: message
/// data: {"key": "value"}
///
/// ```
/// Blank line terminates an event.
async fn parse_sse_stream(
byte_stream: impl futures::Stream<Item = Result<bytes::Bytes, reqwest::Error>> + Send + 'static,
tx: mpsc::Sender<ChannelEvent>,
) {
use futures::StreamExt;
let mut buffer = Vec::<u8>::new();
let mut event_type = String::new();
let mut data_lines = Vec::new();
let mut byte_stream = std::pin::pin!(byte_stream);
while let Some(chunk_result) = byte_stream.next().await {
let chunk = match chunk_result {
Ok(c) => c,
Err(e) => {
tracing::debug!(error = %e, "SSE stream chunk error");
break;
}
};
buffer.extend_from_slice(&chunk);
// Process complete lines (decode UTF-8 only on full lines to avoid
// corruption when multi-byte characters span chunk boundaries)
while let Some(newline_pos) = buffer.iter().position(|&b| b == b'\n') {
let line = String::from_utf8_lossy(&buffer[..newline_pos])
.trim_end_matches('\r')
.to_string();
buffer.drain(..=newline_pos);
if line.is_empty() {
// Blank line = end of event
if !data_lines.is_empty() {
let data = data_lines.join("\n");
if let Ok(mut event) = serde_json::from_str::<ChannelEvent>(&data) {
if event.event_type.is_empty() && !event_type.is_empty() {
event.event_type = event_type.clone();
}
if tx.send(event).await.is_err() {
return; // receiver dropped
}
} else {
tracing::debug!(
event_type = %event_type,
data_len = data.len(),
"Failed to parse SSE event data as ChannelEvent"
);
}
}
event_type.clear();
data_lines.clear();
} else if let Some(value) = line.strip_prefix("event:") {
event_type = value.trim().to_string();
} else if let Some(value) = line.strip_prefix("data:") {
data_lines.push(value.trim().to_string());
}
// Ignore other fields (id:, retry:, comments)
}
}
tracing::debug!("SSE stream ended");
}
/// Errors from relay client operations.
#[derive(Debug, thiserror::Error)]
pub enum RelayError {
#[error("Network error: {0}")]
Network(String),
#[error("API error (HTTP {status}): {message}")]
Api { status: u16, message: String },
#[error("Protocol error: {0}")]
Protocol(String),
#[error("Stream token expired")]
TokenExpired,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn channel_event_deserialize_minimal() {
let json = r#"{"event_type": "message", "content": "hello"}"#;
let event: ChannelEvent = serde_json::from_str(json).expect("parse failed");
assert_eq!(event.event_type, "message");
assert_eq!(event.text(), "hello");
assert!(event.provider_scope.is_empty());
}
#[test]
fn channel_event_deserialize_relay_format() {
// Matches the actual channel-relay ChannelEvent serialization format.
let json = r#"{
"id": "evt_123",
"event_type": "direct_message",
"provider": "slack",
"provider_scope": "T123",
"channel_id": "D456",
"sender_id": "U789",
"sender_name": "bob",
"content": "hi there",
"thread_id": "1234567890.123456",
"raw": {},
"timestamp": "2026-03-09T21:00:00Z"
}"#;
let event: ChannelEvent = serde_json::from_str(json).expect("parse failed");
assert_eq!(event.provider, "slack");
assert_eq!(event.team_id(), "T123");
assert_eq!(event.display_name(), "bob");
assert_eq!(event.thread_id, Some("1234567890.123456".to_string()));
assert!(event.is_message());
}
#[test]
fn channel_event_is_message() {
let make = |et: &str| ChannelEvent {
id: String::new(),
event_type: et.to_string(),
provider: String::new(),
provider_scope: String::new(),
channel_id: String::new(),
sender_id: String::new(),
sender_name: None,
content: None,
thread_id: None,
raw: serde_json::Value::Null,
timestamp: None,
};
assert!(make("message").is_message());
assert!(make("direct_message").is_message());
assert!(make("mention").is_message());
assert!(!make("reaction").is_message());
}
#[test]
fn connection_deserialize() {
let json = r#"{"provider": "slack", "team_id": "T123", "team_name": "My Team", "connected": true}"#;
let conn: Connection = serde_json::from_str(json).expect("parse failed");
assert_eq!(conn.provider, "slack");
assert!(conn.connected);
}
#[test]
fn relay_error_display() {
let err = RelayError::Network("timeout".into());
assert_eq!(err.to_string(), "Network error: timeout");
let err = RelayError::Api {
status: 401,
message: "unauthorized".into(),
};
assert_eq!(err.to_string(), "API error (HTTP 401): unauthorized");
let err = RelayError::TokenExpired;
assert_eq!(err.to_string(), "Stream token expired");
}
#[test]
fn event_type_constants_match_is_message() {
let make = |et: &str| ChannelEvent {
id: String::new(),
event_type: et.to_string(),
provider: String::new(),
provider_scope: String::new(),
channel_id: String::new(),
sender_id: String::new(),
sender_name: None,
content: None,
thread_id: None,
raw: serde_json::Value::Null,
timestamp: None,
};
assert!(make(event_types::MESSAGE).is_message());
assert!(make(event_types::DIRECT_MESSAGE).is_message());
assert!(make(event_types::MENTION).is_message());
}
#[tokio::test]
async fn parse_sse_handles_multibyte_utf8_across_chunks() {
// The crab emoji (🦀) is 4 bytes: [0xF0, 0x9F, 0xA6, 0x80].
// Split it across two chunks to verify no U+FFFD corruption.
let event_json = r#"{"event_type":"message","content":"hello 🦀 world","provider_scope":"T1","channel_id":"C1","sender_id":"U1"}"#;
let full = format!("event: message\ndata: {}\n\n", event_json);
let bytes = full.as_bytes();
// Find the crab emoji and split mid-character
let crab_pos = bytes
.windows(4)
.position(|w| w == [0xF0, 0x9F, 0xA6, 0x80])
.expect("crab emoji not found");
let split_at = crab_pos + 2; // split in the middle of the 4-byte emoji
let chunk1 = bytes::Bytes::copy_from_slice(&bytes[..split_at]);
let chunk2 = bytes::Bytes::copy_from_slice(&bytes[split_at..]);
let chunks: Vec<Result<bytes::Bytes, reqwest::Error>> = vec![Ok(chunk1), Ok(chunk2)];
let stream = futures::stream::iter(chunks);
let (tx, mut rx) = mpsc::channel(8);
parse_sse_stream(stream, tx).await;
let event = rx.recv().await.expect("should receive event");
assert_eq!(event.text(), "hello 🦀 world");
}
}
-12
View File
@@ -1,12 +0,0 @@
//! Channel-relay integration for connecting to external messaging platforms
//! (Slack) via the channel-relay service.
//!
//! The relay service handles OAuth, credential storage, webhook ingestion,
//! and SSE event streaming. IronClaw consumes the SSE stream and sends
//! messages via the relay's proxy API.
pub mod channel;
pub mod client;
pub use channel::{DEFAULT_RELAY_NAME, RelayChannel};
pub use client::RelayClient;
+2 -33
View File
@@ -184,32 +184,18 @@ impl WasmChannelLoader {
/// └── telegram.capabilities.json /// └── telegram.capabilities.json
/// ``` /// ```
pub async fn load_from_dir(&self, dir: &Path) -> Result<LoadResults, WasmChannelError> { pub async fn load_from_dir(&self, dir: &Path) -> Result<LoadResults, WasmChannelError> {
match fs::metadata(dir).await { if !dir.is_dir() {
Ok(meta) if meta.is_dir() => {}
Ok(_) => {
return Err(WasmChannelError::Io(std::io::Error::new( return Err(WasmChannelError::Io(std::io::Error::new(
std::io::ErrorKind::NotADirectory, std::io::ErrorKind::NotADirectory,
format!("{} is not a directory", dir.display()), format!("{} is not a directory", dir.display()),
))); )));
} }
Err(e) if e.kind() == std::io::ErrorKind::NotFound => {
return Ok(LoadResults::default());
}
Err(e) => return Err(WasmChannelError::Io(e)),
}
let mut results = LoadResults::default(); let mut results = LoadResults::default();
// Collect all .wasm entries first, then load in parallel // Collect all .wasm entries first, then load in parallel
let mut channel_entries = Vec::new(); let mut channel_entries = Vec::new();
// Handle TOCTOU: if read_dir fails with NotFound, treat as empty let mut entries = fs::read_dir(dir).await?;
let mut entries = match fs::read_dir(dir).await {
Ok(entries) => entries,
Err(e) if e.kind() == std::io::ErrorKind::NotFound => {
return Ok(LoadResults::default());
}
Err(e) => return Err(WasmChannelError::Io(e)),
};
while let Some(entry) = entries.next_entry().await? { while let Some(entry) = entries.next_entry().await? {
let path = entry.path(); let path = entry.path();
@@ -500,21 +486,4 @@ mod tests {
let result = loader.load_from_files("", &wasm_path, None).await; let result = loader.load_from_files("", &wasm_path, None).await;
assert!(result.is_err()); assert!(result.is_err());
} }
#[tokio::test]
async fn load_from_dir_returns_empty_when_dir_missing() {
let config = WasmChannelRuntimeConfig::for_testing();
let runtime = Arc::new(WasmChannelRuntime::new(config).unwrap());
let loader = WasmChannelLoader::new(runtime, Arc::new(PairingStore::new()), None);
let dir = TempDir::new().unwrap();
let missing = dir.path().join("nonexistent_channels_dir");
let results = loader.load_from_dir(&missing).await;
// Must succeed with empty results, not error
let results = results.expect("missing dir should return Ok, not Err");
assert!(results.loaded.is_empty());
assert!(results.errors.is_empty());
}
} }
-2
View File
@@ -86,7 +86,6 @@ mod loader;
mod router; mod router;
mod runtime; mod runtime;
mod schema; mod schema;
pub mod setup;
pub(crate) mod signature; pub(crate) mod signature;
#[allow(dead_code)] #[allow(dead_code)]
pub(crate) mod storage; pub(crate) mod storage;
@@ -106,5 +105,4 @@ pub use runtime::{PreparedChannelModule, WasmChannelRuntime, WasmChannelRuntimeC
pub use schema::{ pub use schema::{
ChannelCapabilitiesFile, ChannelConfig, SecretSetupSchema, SetupSchema, WebhookSchema, ChannelCapabilitiesFile, ChannelConfig, SecretSetupSchema, SetupSchema, WebhookSchema,
}; };
pub use setup::{WasmChannelSetup, inject_channel_credentials, setup_wasm_channels};
pub use wrapper::{HttpResponse, SharedWasmChannel, WasmChannel}; pub use wrapper::{HttpResponse, SharedWasmChannel, WasmChannel};
-350
View File
@@ -1,350 +0,0 @@
//! WASM channel setup and credential injection.
//!
//! Encapsulates the logic for loading WASM channels, registering their
//! webhook routes, and injecting credentials from the secrets store.
use std::collections::HashSet;
use std::sync::Arc;
use crate::channels::wasm::{
LoadedChannel, RegisteredEndpoint, SharedWasmChannel, WasmChannel, WasmChannelLoader,
WasmChannelRouter, WasmChannelRuntime, WasmChannelRuntimeConfig, create_wasm_channel_router,
};
use crate::config::Config;
use crate::db::Database;
use crate::extensions::ExtensionManager;
use crate::pairing::PairingStore;
use crate::secrets::SecretsStore;
/// Result of WASM channel setup.
pub struct WasmChannelSetup {
pub channels: Vec<(String, Box<dyn crate::channels::Channel>)>,
pub channel_names: Vec<String>,
pub webhook_routes: Option<axum::Router>,
/// Runtime objects needed for hot-activation via ExtensionManager.
pub wasm_channel_runtime: Arc<WasmChannelRuntime>,
pub pairing_store: Arc<PairingStore>,
pub wasm_channel_router: Arc<WasmChannelRouter>,
}
/// Load WASM channels and register their webhook routes.
pub async fn setup_wasm_channels(
config: &Config,
secrets_store: &Option<Arc<dyn SecretsStore + Send + Sync>>,
extension_manager: Option<&Arc<ExtensionManager>>,
database: Option<&Arc<dyn Database>>,
) -> Option<WasmChannelSetup> {
let runtime = match WasmChannelRuntime::new(WasmChannelRuntimeConfig::default()) {
Ok(r) => Arc::new(r),
Err(e) => {
tracing::warn!("Failed to initialize WASM channel runtime: {}", e);
return None;
}
};
let pairing_store = Arc::new(PairingStore::new());
let settings_store: Option<Arc<dyn crate::db::SettingsStore>> =
database.map(|db| Arc::clone(db) as Arc<dyn crate::db::SettingsStore>);
let mut loader = WasmChannelLoader::new(
Arc::clone(&runtime),
Arc::clone(&pairing_store),
settings_store,
);
if let Some(secrets) = secrets_store {
loader = loader.with_secrets_store(Arc::clone(secrets));
}
let results = match loader
.load_from_dir(&config.channels.wasm_channels_dir)
.await
{
Ok(r) => r,
Err(e) => {
tracing::warn!("Failed to scan WASM channels directory: {}", e);
return None;
}
};
let wasm_router = Arc::new(WasmChannelRouter::new());
let mut channels: Vec<(String, Box<dyn crate::channels::Channel>)> = Vec::new();
let mut channel_names: Vec<String> = Vec::new();
for loaded in results.loaded {
let (name, channel) = register_channel(loaded, config, secrets_store, &wasm_router).await;
channel_names.push(name.clone());
channels.push((name, channel));
}
for (path, err) in &results.errors {
tracing::warn!("Failed to load WASM channel {}: {}", path.display(), err);
}
// Always create webhook routes (even with no channels loaded) so that
// channels hot-added at runtime can receive webhooks without a restart.
let webhook_routes = {
Some(create_wasm_channel_router(
Arc::clone(&wasm_router),
extension_manager.map(Arc::clone),
))
};
Some(WasmChannelSetup {
channels,
channel_names,
webhook_routes,
wasm_channel_runtime: runtime,
pairing_store,
wasm_channel_router: wasm_router,
})
}
/// Process a single loaded WASM channel: retrieve secrets, inject config,
/// register with the router, and set up signing keys and credentials.
async fn register_channel(
loaded: LoadedChannel,
config: &Config,
secrets_store: &Option<Arc<dyn SecretsStore + Send + Sync>>,
wasm_router: &Arc<WasmChannelRouter>,
) -> (String, Box<dyn crate::channels::Channel>) {
let channel_name = loaded.name().to_string();
tracing::info!("Loaded WASM channel: {}", channel_name);
let secret_name = loaded.webhook_secret_name();
let sig_key_secret_name = loaded.signature_key_secret_name();
let hmac_secret_name = loaded.hmac_secret_name();
let webhook_secret = if let Some(secrets) = secrets_store {
secrets
.get_decrypted("default", &secret_name)
.await
.ok()
.map(|s| s.expose().to_string())
} else {
None
};
let secret_header = loaded.webhook_secret_header().map(|s| s.to_string());
let webhook_path = format!("/webhook/{}", channel_name);
let endpoints = vec![RegisteredEndpoint {
channel_name: channel_name.clone(),
path: webhook_path,
methods: vec!["POST".to_string()],
require_secret: webhook_secret.is_some(),
}];
let channel_arc = Arc::new(loaded.channel);
// Inject runtime config (tunnel URL, webhook secret, owner_id).
{
let mut config_updates = std::collections::HashMap::new();
if let Some(ref tunnel_url) = config.tunnel.public_url {
config_updates.insert(
"tunnel_url".to_string(),
serde_json::Value::String(tunnel_url.clone()),
);
}
if let Some(ref secret) = webhook_secret {
config_updates.insert(
"webhook_secret".to_string(),
serde_json::Value::String(secret.clone()),
);
}
if let Some(&owner_id) = config
.channels
.wasm_channel_owner_ids
.get(channel_name.as_str())
{
config_updates.insert("owner_id".to_string(), serde_json::json!(owner_id));
}
if !config_updates.is_empty() {
channel_arc.update_config(config_updates).await;
tracing::info!(
channel = %channel_name,
has_tunnel = config.tunnel.public_url.is_some(),
has_webhook_secret = webhook_secret.is_some(),
"Injected runtime config into channel"
);
}
}
tracing::info!(
channel = %channel_name,
has_webhook_secret = webhook_secret.is_some(),
secret_header = ?secret_header,
"Registering channel with router"
);
wasm_router
.register(
Arc::clone(&channel_arc),
endpoints,
webhook_secret.clone(),
secret_header,
)
.await;
// Register Ed25519 signature key if declared in capabilities.
if let Some(ref sig_key_name) = sig_key_secret_name
&& let Some(secrets) = secrets_store
&& let Ok(key_secret) = secrets.get_decrypted("default", sig_key_name).await
{
match wasm_router
.register_signature_key(&channel_name, key_secret.expose())
.await
{
Ok(()) => {
tracing::info!(channel = %channel_name, "Registered Ed25519 signature key")
}
Err(e) => {
tracing::error!(channel = %channel_name, error = %e, "Invalid signature key in secrets store")
}
}
}
// Register HMAC signing secret if declared in capabilities.
if let Some(ref hmac_secret_name) = hmac_secret_name
&& let Some(secrets) = secrets_store
&& let Ok(secret) = secrets.get_decrypted("default", hmac_secret_name).await
{
wasm_router
.register_hmac_secret(&channel_name, secret.expose())
.await;
tracing::info!(channel = %channel_name, "Registered HMAC signing secret");
}
// Inject credentials from secrets store / environment.
match inject_channel_credentials(
&channel_arc,
secrets_store
.as_ref()
.map(|s| s.as_ref() as &dyn SecretsStore),
&channel_name,
)
.await
{
Ok(count) => {
if count > 0 {
tracing::info!(
channel = %channel_name,
credentials_injected = count,
"Channel credentials injected"
);
}
}
Err(e) => {
tracing::error!(
channel = %channel_name,
error = %e,
"Failed to inject channel credentials"
);
}
}
(channel_name, Box::new(SharedWasmChannel::new(channel_arc)))
}
/// Inject credentials for a channel based on naming convention.
///
/// Looks for secrets matching the pattern `{channel_name}_*` and injects them
/// as credential placeholders (e.g., `telegram_bot_token` -> `{TELEGRAM_BOT_TOKEN}`).
///
/// Falls back to environment variables starting with the uppercase channel name
/// prefix (e.g., `TELEGRAM_` for channel `telegram`) for missing credentials.
///
/// Returns the number of credentials injected.
pub async fn inject_channel_credentials(
channel: &Arc<WasmChannel>,
secrets: Option<&dyn SecretsStore>,
channel_name: &str,
) -> anyhow::Result<usize> {
if channel_name.trim().is_empty() {
return Ok(0);
}
let mut count = 0;
let mut injected_placeholders = HashSet::new();
// 1. Try injecting from persistent secrets store if available
if let Some(secrets) = secrets {
let all_secrets = secrets
.list("default")
.await
.map_err(|e| anyhow::anyhow!("Failed to list secrets: {}", e))?;
let prefix = format!("{}_", channel_name.to_ascii_lowercase());
for secret_meta in all_secrets {
if !secret_meta.name.to_ascii_lowercase().starts_with(&prefix) {
continue;
}
let decrypted = match secrets.get_decrypted("default", &secret_meta.name).await {
Ok(d) => d,
Err(e) => {
tracing::warn!(
secret = %secret_meta.name,
error = %e,
"Failed to decrypt secret for channel credential injection"
);
continue;
}
};
let placeholder = secret_meta.name.to_uppercase();
tracing::debug!(
channel = %channel_name,
secret = %secret_meta.name,
placeholder = %placeholder,
"Injecting credential"
);
channel
.set_credential(&placeholder, decrypted.expose().to_string())
.await;
injected_placeholders.insert(placeholder);
count += 1;
}
}
// 2. Fall back to environment variables for credentials not in the secrets store.
// Only env vars starting with the channel's uppercase prefix are allowed
// (e.g., TELEGRAM_ for channel "telegram") to prevent reading unrelated host
// credentials like AWS_SECRET_ACCESS_KEY.
let prefix = format!("{}_", channel_name.to_ascii_uppercase());
let caps = channel.capabilities();
if let Some(ref http_cap) = caps.tool_capabilities.http {
for cred_mapping in http_cap.credentials.values() {
let placeholder = cred_mapping.secret_name.to_uppercase();
if injected_placeholders.contains(&placeholder) {
continue;
}
if !placeholder.starts_with(&prefix) {
tracing::warn!(
channel = %channel_name,
placeholder = %placeholder,
"Ignoring non-prefixed credential placeholder in environment fallback"
);
continue;
}
if let Ok(env_value) = std::env::var(&placeholder)
&& !env_value.is_empty()
{
tracing::debug!(
channel = %channel_name,
placeholder = %placeholder,
"Injecting credential from environment variable"
);
channel.set_credential(&placeholder, env_value).await;
count += 1;
}
}
}
Ok(count)
}
+5 -8
View File
@@ -3059,7 +3059,6 @@ mod tests {
}; };
use crate::channels::wasm::wrapper::{HttpResponse, WasmChannel}; use crate::channels::wasm::wrapper::{HttpResponse, WasmChannel};
use crate::pairing::PairingStore; use crate::pairing::PairingStore;
use crate::testing::credentials::TEST_TELEGRAM_BOT_TOKEN;
use crate::tools::wasm::ResourceLimits; use crate::tools::wasm::ResourceLimits;
fn create_test_channel() -> WasmChannel { fn create_test_channel() -> WasmChannel {
@@ -4010,7 +4009,7 @@ mod tests {
let mut creds = std::collections::HashMap::new(); let mut creds = std::collections::HashMap::new();
creds.insert( creds.insert(
"TELEGRAM_BOT_TOKEN".to_string(), "TELEGRAM_BOT_TOKEN".to_string(),
TEST_TELEGRAM_BOT_TOKEN.to_string(), "8218490433:AAEZeUxwqZ5OO3mOCXv7fKvpdhDgsmBBNis".to_string(),
); );
creds.insert("OTHER_SECRET".to_string(), "s3cret".to_string()); creds.insert("OTHER_SECRET".to_string(), "s3cret".to_string());
@@ -4023,15 +4022,13 @@ mod tests {
Arc::new(PairingStore::new()), Arc::new(PairingStore::new()),
); );
let error = format!( let error = "HTTP request failed: error sending request for url \
"HTTP request failed: error sending request for url \ (https://api.telegram.org/bot8218490433:AAEZeUxwqZ5OO3mOCXv7fKvpdhDgsmBBNis/getUpdates)";
(https://api.telegram.org/bot{TEST_TELEGRAM_BOT_TOKEN}/getUpdates)"
);
let redacted = store.redact_credentials(&error); let redacted = store.redact_credentials(error);
assert!( assert!(
!redacted.contains(TEST_TELEGRAM_BOT_TOKEN), !redacted.contains("8218490433:AAEZeUxwqZ5OO3mOCXv7fKvpdhDgsmBBNis"),
"credential value should be redacted" "credential value should be redacted"
); );
assert!( assert!(
+26 -27
View File
@@ -83,15 +83,14 @@ pub async fn auth_middleware(
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use crate::testing::credentials::{TEST_AUTH_SECRET_TOKEN, TEST_BEARER_TOKEN};
#[test] #[test]
fn test_auth_state_clone() { fn test_auth_state_clone() {
let state = AuthState { let state = AuthState {
token: TEST_BEARER_TOKEN.to_string(), token: "test-token".to_string(),
}; };
let cloned = state.clone(); let cloned = state.clone();
assert_eq!(cloned.token, TEST_BEARER_TOKEN); assert_eq!(cloned.token, "test-token");
} }
use axum::Router; use axum::Router;
@@ -121,10 +120,10 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_valid_bearer_token_passes() { async fn test_valid_bearer_token_passes() {
let app = test_app(TEST_AUTH_SECRET_TOKEN); let app = test_app("secret-token");
let req = Request::builder() let req = Request::builder()
.uri("/api/chat/events") .uri("/api/chat/events")
.header("Authorization", format!("Bearer {TEST_AUTH_SECRET_TOKEN}")) .header("Authorization", "Bearer secret-token")
.body(Body::empty()) .body(Body::empty())
.unwrap(); .unwrap();
let resp = app.oneshot(req).await.unwrap(); let resp = app.oneshot(req).await.unwrap();
@@ -133,7 +132,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_invalid_bearer_token_rejected() { async fn test_invalid_bearer_token_rejected() {
let app = test_app(TEST_AUTH_SECRET_TOKEN); let app = test_app("secret-token");
let req = Request::builder() let req = Request::builder()
.uri("/api/chat/events") .uri("/api/chat/events")
.header("Authorization", "Bearer wrong-token") .header("Authorization", "Bearer wrong-token")
@@ -145,9 +144,9 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_query_token_allowed_for_chat_events() { async fn test_query_token_allowed_for_chat_events() {
let app = test_app(TEST_AUTH_SECRET_TOKEN); let app = test_app("secret-token");
let req = Request::builder() let req = Request::builder()
.uri(format!("/api/chat/events?token={TEST_AUTH_SECRET_TOKEN}")) .uri("/api/chat/events?token=secret-token")
.body(Body::empty()) .body(Body::empty())
.unwrap(); .unwrap();
let resp = app.oneshot(req).await.unwrap(); let resp = app.oneshot(req).await.unwrap();
@@ -156,9 +155,9 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_query_token_allowed_for_logs_events() { async fn test_query_token_allowed_for_logs_events() {
let app = test_app(TEST_AUTH_SECRET_TOKEN); let app = test_app("secret-token");
let req = Request::builder() let req = Request::builder()
.uri(format!("/api/logs/events?token={TEST_AUTH_SECRET_TOKEN}")) .uri("/api/logs/events?token=secret-token")
.body(Body::empty()) .body(Body::empty())
.unwrap(); .unwrap();
let resp = app.oneshot(req).await.unwrap(); let resp = app.oneshot(req).await.unwrap();
@@ -167,9 +166,9 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_query_token_allowed_for_ws_upgrade() { async fn test_query_token_allowed_for_ws_upgrade() {
let app = test_app(TEST_AUTH_SECRET_TOKEN); let app = test_app("secret-token");
let req = Request::builder() let req = Request::builder()
.uri(format!("/api/chat/ws?token={TEST_AUTH_SECRET_TOKEN}")) .uri("/api/chat/ws?token=secret-token")
.body(Body::empty()) .body(Body::empty())
.unwrap(); .unwrap();
let resp = app.oneshot(req).await.unwrap(); let resp = app.oneshot(req).await.unwrap();
@@ -203,9 +202,9 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_query_token_rejected_for_non_sse_get() { async fn test_query_token_rejected_for_non_sse_get() {
let app = test_app(TEST_AUTH_SECRET_TOKEN); let app = test_app("secret-token");
let req = Request::builder() let req = Request::builder()
.uri(format!("/api/chat/history?token={TEST_AUTH_SECRET_TOKEN}")) .uri("/api/chat/history?token=secret-token")
.body(Body::empty()) .body(Body::empty())
.unwrap(); .unwrap();
let resp = app.oneshot(req).await.unwrap(); let resp = app.oneshot(req).await.unwrap();
@@ -214,10 +213,10 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_query_token_rejected_for_post() { async fn test_query_token_rejected_for_post() {
let app = test_app(TEST_AUTH_SECRET_TOKEN); let app = test_app("secret-token");
let req = Request::builder() let req = Request::builder()
.method(Method::POST) .method(Method::POST)
.uri(format!("/api/chat/send?token={TEST_AUTH_SECRET_TOKEN}")) .uri("/api/chat/send?token=secret-token")
.body(Body::empty()) .body(Body::empty())
.unwrap(); .unwrap();
let resp = app.oneshot(req).await.unwrap(); let resp = app.oneshot(req).await.unwrap();
@@ -226,7 +225,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_query_token_invalid_rejected() { async fn test_query_token_invalid_rejected() {
let app = test_app(TEST_AUTH_SECRET_TOKEN); let app = test_app("secret-token");
let req = Request::builder() let req = Request::builder()
.uri("/api/chat/events?token=wrong-token") .uri("/api/chat/events?token=wrong-token")
.body(Body::empty()) .body(Body::empty())
@@ -237,7 +236,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_no_auth_at_all_rejected() { async fn test_no_auth_at_all_rejected() {
let app = test_app(TEST_AUTH_SECRET_TOKEN); let app = test_app("secret-token");
let req = Request::builder() let req = Request::builder()
.uri("/api/chat/events") .uri("/api/chat/events")
.body(Body::empty()) .body(Body::empty())
@@ -248,11 +247,11 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_bearer_header_works_for_post() { async fn test_bearer_header_works_for_post() {
let app = test_app(TEST_AUTH_SECRET_TOKEN); let app = test_app("secret-token");
let req = Request::builder() let req = Request::builder()
.method(Method::POST) .method(Method::POST)
.uri("/api/chat/send") .uri("/api/chat/send")
.header("Authorization", format!("Bearer {TEST_AUTH_SECRET_TOKEN}")) .header("Authorization", "Bearer secret-token")
.body(Body::empty()) .body(Body::empty())
.unwrap(); .unwrap();
let resp = app.oneshot(req).await.unwrap(); let resp = app.oneshot(req).await.unwrap();
@@ -261,10 +260,10 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_bearer_prefix_case_insensitive() { async fn test_bearer_prefix_case_insensitive() {
let app = test_app(TEST_AUTH_SECRET_TOKEN); let app = test_app("secret-token");
let req = Request::builder() let req = Request::builder()
.uri("/api/chat/events") .uri("/api/chat/events")
.header("Authorization", format!("bearer {TEST_AUTH_SECRET_TOKEN}")) .header("Authorization", "bearer secret-token")
.body(Body::empty()) .body(Body::empty())
.unwrap(); .unwrap();
let resp = app.oneshot(req).await.unwrap(); let resp = app.oneshot(req).await.unwrap();
@@ -273,10 +272,10 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_bearer_prefix_mixed_case() { async fn test_bearer_prefix_mixed_case() {
let app = test_app(TEST_AUTH_SECRET_TOKEN); let app = test_app("secret-token");
let req = Request::builder() let req = Request::builder()
.uri("/api/chat/events") .uri("/api/chat/events")
.header("Authorization", format!("BEARER {TEST_AUTH_SECRET_TOKEN}")) .header("Authorization", "BEARER secret-token")
.body(Body::empty()) .body(Body::empty())
.unwrap(); .unwrap();
let resp = app.oneshot(req).await.unwrap(); let resp = app.oneshot(req).await.unwrap();
@@ -285,7 +284,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_empty_bearer_token_rejected() { async fn test_empty_bearer_token_rejected() {
let app = test_app(TEST_AUTH_SECRET_TOKEN); let app = test_app("secret-token");
let req = Request::builder() let req = Request::builder()
.uri("/api/chat/events") .uri("/api/chat/events")
.header("Authorization", "Bearer ") .header("Authorization", "Bearer ")
@@ -297,10 +296,10 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_token_with_whitespace_rejected() { async fn test_token_with_whitespace_rejected() {
let app = test_app(TEST_AUTH_SECRET_TOKEN); let app = test_app("secret-token");
let req = Request::builder() let req = Request::builder()
.uri("/api/chat/events") .uri("/api/chat/events")
.header("Authorization", format!("Bearer {TEST_AUTH_SECRET_TOKEN}")) .header("Authorization", "Bearer secret-token")
.body(Body::empty()) .body(Body::empty())
.unwrap(); .unwrap();
let resp = app.oneshot(req).await.unwrap(); let resp = app.oneshot(req).await.unwrap();
+6 -30
View File
@@ -35,7 +35,6 @@ pub async fn chat_send_handler(
} }
let msg_id = msg.id; let msg_id = msg.id;
let thread_id = msg.thread_id.clone();
let tx_guard = state.msg_tx.read().await; let tx_guard = state.msg_tx.read().await;
let tx = tx_guard.as_ref().ok_or(( let tx = tx_guard.as_ref().ok_or((
@@ -50,13 +49,6 @@ pub async fn chat_send_handler(
) )
})?; })?;
tracing::debug!(
message_id = %msg_id,
thread_id = ?thread_id,
content_len = req.content.len(),
"Message queued to agent loop"
);
Ok(( Ok((
StatusCode::ACCEPTED, StatusCode::ACCEPTED,
Json(SendMessageResponse { Json(SendMessageResponse {
@@ -271,6 +263,7 @@ pub async fn chat_history_handler(
))?; ))?;
let session = session_manager.get_or_create_session(&state.user_id).await; let session = session_manager.get_or_create_session(&state.user_id).await;
let sess = session.lock().await;
let limit = query.limit.unwrap_or(50); let limit = query.limit.unwrap_or(50);
let before_cursor = query let before_cursor = query
@@ -288,12 +281,11 @@ pub async fn chat_history_handler(
}) })
.transpose()?; .transpose()?;
// Find the thread (lock only briefly to get active_thread if needed) // Find the thread
let thread_id = if let Some(ref tid) = query.thread_id { let thread_id = if let Some(ref tid) = query.thread_id {
Uuid::parse_str(tid) Uuid::parse_str(tid)
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid thread_id".to_string()))? .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid thread_id".to_string()))?
} else { } else {
let sess = session.lock().await;
sess.active_thread sess.active_thread
.ok_or((StatusCode::NOT_FOUND, "No active thread".to_string()))? .ok_or((StatusCode::NOT_FOUND, "No active thread".to_string()))?
}; };
@@ -306,13 +298,10 @@ pub async fn chat_history_handler(
.conversation_belongs_to_user(thread_id, &state.user_id) .conversation_belongs_to_user(thread_id, &state.user_id)
.await .await
.unwrap_or(false); .unwrap_or(false);
if !owned { if !owned && !sess.threads.contains_key(&thread_id) {
let sess = session.lock().await;
if !sess.threads.contains_key(&thread_id) {
return Err((StatusCode::NOT_FOUND, "Thread not found".to_string())); return Err((StatusCode::NOT_FOUND, "Thread not found".to_string()));
} }
} }
}
// For paginated requests (before cursor set), always go to DB // For paginated requests (before cursor set), always go to DB
if before_cursor.is_some() if before_cursor.is_some()
@@ -335,9 +324,6 @@ pub async fn chat_history_handler(
} }
// Try in-memory first (freshest data for active threads) // Try in-memory first (freshest data for active threads)
// Lock only when checking in-memory state
{
let sess = session.lock().await;
if let Some(thread) = sess.threads.get(&thread_id) if let Some(thread) = sess.threads.get(&thread_id)
&& (!thread.turns.is_empty() || thread.pending_approval.is_some()) && (!thread.turns.is_empty() || thread.pending_approval.is_some())
{ {
@@ -389,7 +375,6 @@ pub async fn chat_history_handler(
pending_approval, pending_approval,
})); }));
} }
}
// Fall back to DB for historical threads not in memory (paginated) // Fall back to DB for historical threads not in memory (paginated)
if let Some(ref store) = state.store { if let Some(ref store) = state.store {
@@ -430,6 +415,7 @@ pub async fn chat_threads_handler(
))?; ))?;
let session = session_manager.get_or_create_session(&state.user_id).await; let session = session_manager.get_or_create_session(&state.user_id).await;
let sess = session.lock().await;
// Try DB first for persistent thread list // Try DB first for persistent thread list
if let Some(ref store) = state.store { if let Some(ref store) = state.store {
@@ -479,22 +465,15 @@ pub async fn chat_threads_handler(
}); });
} }
// Read active thread while holding minimal lock (just before return)
let active_thread = {
let sess = session.lock().await;
sess.active_thread
};
return Ok(Json(ThreadListResponse { return Ok(Json(ThreadListResponse {
assistant_thread, assistant_thread,
threads, threads,
active_thread, active_thread: sess.active_thread,
})); }));
} }
} }
// Fallback: in-memory only (no assistant thread without DB) // Fallback: in-memory only (no assistant thread without DB)
let sess = session.lock().await;
let mut sorted_threads: Vec<_> = sess.threads.values().collect(); let mut sorted_threads: Vec<_> = sess.threads.values().collect();
sorted_threads.sort_by(|a, b| b.updated_at.cmp(&a.updated_at)); sorted_threads.sort_by(|a, b| b.updated_at.cmp(&a.updated_at));
let threads: Vec<ThreadInfo> = sorted_threads let threads: Vec<ThreadInfo> = sorted_threads
@@ -511,13 +490,10 @@ pub async fn chat_threads_handler(
}) })
.collect(); .collect();
let active_thread = sess.active_thread;
drop(sess); // Explicit drop to release lock
Ok(Json(ThreadListResponse { Ok(Json(ThreadListResponse {
assistant_thread: None, assistant_thread: None,
threads, threads,
active_thread, active_thread: sess.active_thread,
})) }))
} }
+56 -9
View File
@@ -46,14 +46,6 @@ pub async fn extensions_list_handler(
} else { } else {
"configured".to_string() "configured".to_string()
}) })
} else if ext.kind == crate::extensions::ExtensionKind::ChannelRelay {
Some(if ext.active {
"active".to_string()
} else if ext.authenticated {
"configured".to_string()
} else {
"installed".to_string()
})
} else { } else {
None None
}; };
@@ -111,7 +103,6 @@ pub async fn extensions_install_handler(
"mcp_server" => Some(crate::extensions::ExtensionKind::McpServer), "mcp_server" => Some(crate::extensions::ExtensionKind::McpServer),
"wasm_tool" => Some(crate::extensions::ExtensionKind::WasmTool), "wasm_tool" => Some(crate::extensions::ExtensionKind::WasmTool),
"wasm_channel" => Some(crate::extensions::ExtensionKind::WasmChannel), "wasm_channel" => Some(crate::extensions::ExtensionKind::WasmChannel),
"channel_relay" => Some(crate::extensions::ExtensionKind::ChannelRelay),
_ => None, _ => None,
}); });
@@ -124,6 +115,62 @@ pub async fn extensions_install_handler(
} }
} }
pub async fn extensions_activate_handler(
State(state): State<Arc<GatewayState>>,
Path(name): Path<String>,
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
let ext_mgr = state.extension_manager.as_ref().ok_or((
StatusCode::NOT_IMPLEMENTED,
"Extension manager not available (secrets store required)".to_string(),
))?;
match ext_mgr.activate(&name).await {
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")
|| err_str.contains("401")
|| err_str.contains("Unauthorized");
if !needs_auth {
return Ok(Json(ActionResponse::fail(err_str)));
}
// Activation failed due to auth; try authenticating first.
match ext_mgr.auth(&name, None).await {
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))),
Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))),
}
}
Ok(auth_result) => {
// Auth in progress (OAuth URL or awaiting manual token).
let mut resp = ActionResponse::fail(
auth_result
.instructions()
.map(String::from)
.unwrap_or_else(|| format!("'{}' requires authentication.", name)),
);
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!(
"Authentication failed: {}",
auth_err
)))),
}
}
}
}
pub async fn extensions_remove_handler( pub async fn extensions_remove_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
Path(name): Path<String>, Path(name): Path<String>,
+49 -1
View File
@@ -27,7 +27,7 @@ pub async fn routines_list_handler(
.await .await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
let items: Vec<RoutineInfo> = routines.iter().map(RoutineInfo::from_routine).collect(); let items: Vec<RoutineInfo> = routines.iter().map(routine_to_info).collect();
Ok(Json(RoutineListResponse { routines: items })) Ok(Json(RoutineListResponse { routines: items }))
} }
@@ -263,6 +263,54 @@ pub async fn routines_runs_handler(
}))) })))
} }
/// Convert a Routine to the trimmed RoutineInfo for list display.
fn routine_to_info(r: &crate::agent::routine::Routine) -> RoutineInfo {
let (trigger_type, trigger_summary) = match &r.trigger {
crate::agent::routine::Trigger::Cron { schedule, .. } => {
("cron".to_string(), format!("cron: {}", schedule))
}
crate::agent::routine::Trigger::Event {
pattern, channel, ..
} => {
let ch = channel.as_deref().unwrap_or("any");
("event".to_string(), format!("on {} /{}/", ch, pattern))
}
crate::agent::routine::Trigger::Webhook { path, .. } => {
let p = path.as_deref().unwrap_or("/");
("webhook".to_string(), format!("webhook: {}", p))
}
crate::agent::routine::Trigger::Manual => ("manual".to_string(), "manual only".to_string()),
};
let action_type = match &r.action {
crate::agent::routine::RoutineAction::Lightweight { .. } => "lightweight",
crate::agent::routine::RoutineAction::FullJob { .. } => "full_job",
};
let status = if !r.enabled {
"disabled"
} else if r.consecutive_failures > 0 {
"failing"
} else {
"active"
};
RoutineInfo {
id: r.id,
name: r.name.clone(),
description: r.description.clone(),
enabled: r.enabled,
trigger_type,
trigger_summary,
action_type: action_type.to_string(),
last_run_at: r.last_run_at.map(|dt| dt.to_rfc3339()),
next_fire_at: r.next_fire_at.map(|dt| dt.to_rfc3339()),
run_count: r.run_count,
consecutive_failures: r.consecutive_failures,
status: status.to_string(),
}
}
/// Map `RoutineError` variants to appropriate HTTP status codes. /// Map `RoutineError` variants to appropriate HTTP status codes.
fn routine_error_status(err: &RoutineError) -> StatusCode { fn routine_error_status(err: &RoutineError) -> StatusCode {
match err { match err {
-2
View File
@@ -97,7 +97,6 @@ impl GatewayChannel {
skill_registry: None, skill_registry: None,
skill_catalog: None, skill_catalog: None,
chat_rate_limiter: server::RateLimiter::new(30, 60), chat_rate_limiter: server::RateLimiter::new(30, 60),
oauth_rate_limiter: server::RateLimiter::new(10, 60),
registry_entries: Vec::new(), registry_entries: Vec::new(),
cost_guard: None, cost_guard: None,
routine_engine: Arc::new(tokio::sync::RwLock::new(None)), routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
@@ -134,7 +133,6 @@ impl GatewayChannel {
skill_registry: self.state.skill_registry.clone(), skill_registry: self.state.skill_registry.clone(),
skill_catalog: self.state.skill_catalog.clone(), skill_catalog: self.state.skill_catalog.clone(),
chat_rate_limiter: server::RateLimiter::new(30, 60), chat_rate_limiter: server::RateLimiter::new(30, 60),
oauth_rate_limiter: server::RateLimiter::new(10, 60),
registry_entries: self.state.registry_entries.clone(), registry_entries: self.state.registry_entries.clone(),
cost_guard: self.state.cost_guard.clone(), cost_guard: self.state.cost_guard.clone(),
routine_engine: Arc::clone(&self.state.routine_engine), routine_engine: Arc::clone(&self.state.routine_engine),
+65 -404
View File
@@ -28,7 +28,6 @@ use uuid::Uuid;
use crate::agent::SessionManager; use crate::agent::SessionManager;
use crate::bootstrap::ironclaw_base_dir; use crate::bootstrap::ironclaw_base_dir;
use crate::channels::IncomingMessage; use crate::channels::IncomingMessage;
use crate::channels::relay::DEFAULT_RELAY_NAME;
use crate::channels::web::auth::{AuthState, auth_middleware}; use crate::channels::web::auth::{AuthState, auth_middleware};
use crate::channels::web::handlers::jobs::{ use crate::channels::web::handlers::jobs::{
job_files_list_handler, job_files_read_handler, jobs_cancel_handler, jobs_detail_handler, job_files_list_handler, job_files_read_handler, jobs_cancel_handler, jobs_detail_handler,
@@ -165,8 +164,6 @@ pub struct GatewayState {
pub scheduler: Option<crate::tools::builtin::SchedulerSlot>, pub scheduler: Option<crate::tools::builtin::SchedulerSlot>,
/// Rate limiter for chat endpoints (30 messages per 60 seconds). /// Rate limiter for chat endpoints (30 messages per 60 seconds).
pub chat_rate_limiter: RateLimiter, pub chat_rate_limiter: RateLimiter,
/// Rate limiter for OAuth callback endpoints (10 requests per 60 seconds).
pub oauth_rate_limiter: RateLimiter,
/// Registry catalog entries for the available extensions API. /// Registry catalog entries for the available extensions API.
/// Populated at startup from `registry/` manifests, independent of extension manager. /// Populated at startup from `registry/` manifests, independent of extension manager.
pub registry_entries: Vec<crate::extensions::RegistryEntry>, pub registry_entries: Vec<crate::extensions::RegistryEntry>,
@@ -203,11 +200,7 @@ pub async fn start_server(
// Public routes (no auth) // Public routes (no auth)
let public = Router::new() let public = Router::new()
.route("/api/health", get(health_handler)) .route("/api/health", get(health_handler))
.route("/oauth/callback", get(oauth_callback_handler)) .route("/oauth/callback", get(oauth_callback_handler));
.route(
"/oauth/slack/callback",
get(slack_relay_oauth_callback_handler),
);
// Protected routes (require auth) // Protected routes (require auth)
let auth_state = AuthState { token: auth_token }; let auth_state = AuthState { token: auth_token };
@@ -377,7 +370,7 @@ pub async fn start_server(
if let Err(e) = axum::serve(listener, app) if let Err(e) = axum::serve(listener, app)
.with_graceful_shutdown(async { .with_graceful_shutdown(async {
let _ = shutdown_rx.await; let _ = shutdown_rx.await;
tracing::debug!("Web gateway shutting down"); tracing::info!("Web gateway shutting down");
}) })
.await .await
{ {
@@ -613,208 +606,6 @@ async fn oauth_callback_handler(
axum::response::Html(html).into_response() axum::response::Html(html).into_response()
} }
/// OAuth callback for Slack via channel-relay.
///
/// This is a PUBLIC route (no Bearer token required) because channel-relay
/// redirects the user's browser here after Slack OAuth completes.
/// Query params: `stream_token`, `provider`, `team_id`.
async fn slack_relay_oauth_callback_handler(
State(state): State<Arc<GatewayState>>,
Query(params): Query<std::collections::HashMap<String, String>>,
) -> impl IntoResponse {
// Rate limit
if !state.oauth_rate_limiter.check() {
return axum::response::Html(
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
<h2>Too Many Requests</h2>\
<p>Please try again later.</p>\
</body></html>"
.to_string(),
)
.into_response();
}
// Validate stream_token: required, non-empty, max 2048 bytes
let stream_token = match params.get("stream_token") {
Some(t) if !t.is_empty() && t.len() <= 2048 => t.clone(),
Some(t) if t.len() > 2048 => {
return axum::response::Html(
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
<h2>Error</h2><p>Invalid callback parameters.</p></body></html>"
.to_string(),
)
.into_response();
}
_ => {
return axum::response::Html(
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
<h2>Error</h2><p>Invalid callback parameters.</p></body></html>"
.to_string(),
)
.into_response();
}
};
// Validate team_id format: empty or T followed by alphanumeric (max 20 chars)
let team_id = params.get("team_id").cloned().unwrap_or_default();
if !team_id.is_empty() {
let valid_team_id = team_id.len() <= 21
&& team_id.starts_with('T')
&& team_id[1..].chars().all(|c| c.is_ascii_alphanumeric());
if !valid_team_id {
return axum::response::Html(
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
<h2>Error</h2><p>Invalid callback parameters.</p></body></html>"
.to_string(),
)
.into_response();
}
}
// Validate provider: must be "slack" (only supported provider)
let provider = params
.get("provider")
.cloned()
.unwrap_or_else(|| "slack".into());
if provider != "slack" {
return axum::response::Html(
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
<h2>Error</h2><p>Invalid callback parameters.</p></body></html>"
.to_string(),
)
.into_response();
}
let ext_mgr = match state.extension_manager.as_ref() {
Some(mgr) => mgr,
None => {
return axum::response::Html(
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
<h2>Error</h2><p>Extension manager not available.</p></body></html>"
.to_string(),
)
.into_response();
}
};
// Validate CSRF state parameter
let state_param = match params.get("state") {
Some(s) if !s.is_empty() && s.len() <= 128 => s.clone(),
_ => {
return axum::response::Html(
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
<h2>Error</h2><p>Invalid or expired authorization.</p></body></html>"
.to_string(),
)
.into_response();
}
};
let state_key = format!("relay:{}:oauth_state", DEFAULT_RELAY_NAME);
let stored_state = match ext_mgr
.secrets()
.get_decrypted(&state.user_id, &state_key)
.await
{
Ok(secret) => secret.expose().to_string(),
Err(_) => {
return axum::response::Html(
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
<h2>Error</h2><p>Invalid or expired authorization.</p></body></html>"
.to_string(),
)
.into_response();
}
};
if state_param != stored_state {
return axum::response::Html(
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
<h2>Error</h2><p>Invalid or expired authorization.</p></body></html>"
.to_string(),
)
.into_response();
}
// Delete the nonce (one-time use)
let _ = ext_mgr.secrets().delete(&state.user_id, &state_key).await;
let result: Result<(), String> = async {
// Store the stream token as a secret
let token_key = format!("relay:{}:stream_token", DEFAULT_RELAY_NAME);
let _ = ext_mgr.secrets().delete(&state.user_id, &token_key).await;
ext_mgr
.secrets()
.create(
&state.user_id,
crate::secrets::CreateSecretParams {
name: token_key,
value: secrecy::SecretString::from(stream_token),
provider: Some(provider.clone()),
expires_at: None,
},
)
.await
.map_err(|e| format!("Failed to store stream token: {}", e))?;
// Store team_id in settings
if let Some(ref store) = state.store {
let team_id_key = format!("relay:{}:team_id", DEFAULT_RELAY_NAME);
let _ = store
.set_setting(&state.user_id, &team_id_key, &serde_json::json!(team_id))
.await;
}
// Activate the relay channel
ext_mgr
.activate_stored_relay(DEFAULT_RELAY_NAME)
.await
.map_err(|e| format!("Failed to activate relay channel: {}", e))?;
Ok(())
}
.await;
let (success, message) = match &result {
Ok(()) => (true, "Slack connected successfully!".to_string()),
Err(e) => {
tracing::error!(error = %e, "Slack relay OAuth callback failed");
(
false,
"Connection failed. Check server logs for details.".to_string(),
)
}
};
// Broadcast SSE event to notify the web UI
state.sse.broadcast(SseEvent::AuthCompleted {
extension_name: DEFAULT_RELAY_NAME.to_string(),
success,
message: message.clone(),
});
if success {
axum::response::Html(
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
<h2>Slack Connected!</h2>\
<p>You can close this tab and return to IronClaw.</p>\
<script>window.close()</script>\
</body></html>"
.to_string(),
)
.into_response()
} else {
axum::response::Html(format!(
"<html><body style='font-family: system-ui; text-align: center; padding: 60px;'>\
<h2>Connection Failed</h2>\
<p>{}</p>\
</body></html>",
message
))
.into_response()
}
}
// --- Chat handlers --- // --- Chat handlers ---
/// Convert web gateway `ImageData` to `IncomingAttachment` objects. /// Convert web gateway `ImageData` to `IncomingAttachment` objects.
@@ -872,9 +663,9 @@ async fn chat_send_handler(
headers: axum::http::HeaderMap, headers: axum::http::HeaderMap,
Json(req): Json<SendMessageRequest>, Json(req): Json<SendMessageRequest>,
) -> Result<(StatusCode, Json<SendMessageResponse>), (StatusCode, String)> { ) -> Result<(StatusCode, Json<SendMessageResponse>), (StatusCode, String)> {
tracing::trace!( tracing::debug!(
"[chat_send_handler] Received message: content_len={}, thread_id={:?}", "[chat_send_handler] Received message: content={:?}, thread_id={:?}",
req.content.len(), req.content,
req.thread_id req.thread_id
); );
@@ -907,10 +698,10 @@ async fn chat_send_handler(
} }
let msg_id = msg.id; let msg_id = msg.id;
tracing::trace!( tracing::debug!(
"[chat_send_handler] Created message id={}, content_len={}, images={}", "[chat_send_handler] Created message id={}, content={:?}, images={}",
msg_id, msg_id,
req.content.len(), req.content,
req.images.len() req.images.len()
); );
@@ -1848,13 +1639,13 @@ async fn extensions_activate_handler(
Ok(Json(resp)) Ok(Json(resp))
} }
Err(activate_err) => { Err(activate_err) => {
let needs_auth = matches!( let err_str = activate_err.to_string();
&activate_err, let needs_auth = err_str.contains("authentication")
crate::extensions::ExtensionError::AuthRequired || err_str.contains("401")
); || err_str.contains("Unauthorized");
if !needs_auth { if !needs_auth {
return Ok(Json(ActionResponse::fail(activate_err.to_string()))); return Ok(Json(ActionResponse::fail(err_str)));
} }
// Activation failed due to auth; try authenticating first. // Activation failed due to auth; try authenticating first.
@@ -2145,7 +1936,7 @@ async fn routines_list_handler(
.await .await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
let items: Vec<RoutineInfo> = routines.iter().map(RoutineInfo::from_routine).collect(); let items: Vec<RoutineInfo> = routines.iter().map(routine_to_info).collect();
Ok(Json(RoutineListResponse { routines: items })) Ok(Json(RoutineListResponse { routines: items }))
} }
@@ -2389,6 +2180,54 @@ async fn routines_runs_handler(
}))) })))
} }
/// Convert a Routine to the trimmed RoutineInfo for list display.
fn routine_to_info(r: &crate::agent::routine::Routine) -> RoutineInfo {
let (trigger_type, trigger_summary) = match &r.trigger {
crate::agent::routine::Trigger::Cron { schedule, .. } => {
("cron".to_string(), format!("cron: {}", schedule))
}
crate::agent::routine::Trigger::Event {
pattern, channel, ..
} => {
let ch = channel.as_deref().unwrap_or("any");
("event".to_string(), format!("on {} /{}/", ch, pattern))
}
crate::agent::routine::Trigger::Webhook { path, .. } => {
let p = path.as_deref().unwrap_or("/");
("webhook".to_string(), format!("webhook: {}", p))
}
crate::agent::routine::Trigger::Manual => ("manual".to_string(), "manual only".to_string()),
};
let action_type = match &r.action {
crate::agent::routine::RoutineAction::Lightweight { .. } => "lightweight",
crate::agent::routine::RoutineAction::FullJob { .. } => "full_job",
};
let status = if !r.enabled {
"disabled"
} else if r.consecutive_failures > 0 {
"failing"
} else {
"active"
};
RoutineInfo {
id: r.id,
name: r.name.clone(),
description: r.description.clone(),
enabled: r.enabled,
trigger_type,
trigger_summary,
action_type: action_type.to_string(),
last_run_at: r.last_run_at.map(|dt| dt.to_rfc3339()),
next_fire_at: r.next_fire_at.map(|dt| dt.to_rfc3339()),
run_count: r.run_count,
consecutive_failures: r.consecutive_failures,
status: status.to_string(),
}
}
// --- Settings handlers --- // --- Settings handlers ---
async fn settings_list_handler( async fn settings_list_handler(
@@ -2588,7 +2427,6 @@ struct GatewayStatusResponse {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use crate::testing::credentials::TEST_GATEWAY_CRYPTO_KEY;
#[test] #[test]
fn test_build_turns_from_db_messages_complete() { fn test_build_turns_from_db_messages_complete() {
@@ -2690,7 +2528,6 @@ mod tests {
skill_catalog: None, skill_catalog: None,
scheduler: None, scheduler: None,
chat_rate_limiter: RateLimiter::new(30, 60), chat_rate_limiter: RateLimiter::new(30, 60),
oauth_rate_limiter: RateLimiter::new(10, 60),
registry_entries: vec![], registry_entries: vec![],
cost_guard: None, cost_guard: None,
routine_engine: Arc::new(tokio::sync::RwLock::new(None)), routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
@@ -2763,7 +2600,7 @@ mod tests {
// Build an ExtensionManager so the handler can look up flows // Build an ExtensionManager so the handler can look up flows
let secrets = Arc::new(crate::secrets::InMemorySecretsStore::new(Arc::new( let secrets = Arc::new(crate::secrets::InMemorySecretsStore::new(Arc::new(
crate::secrets::SecretsCrypto::new(secrecy::SecretString::from( crate::secrets::SecretsCrypto::new(secrecy::SecretString::from(
TEST_GATEWAY_CRYPTO_KEY.to_string(), "test-key-at-least-32-chars-long!!".to_string(),
)) ))
.expect("crypto"), .expect("crypto"),
))); )));
@@ -2772,7 +2609,6 @@ mod tests {
let ext_mgr = Arc::new(ExtensionManager::new( let ext_mgr = Arc::new(ExtensionManager::new(
mcp_sm, mcp_sm,
Arc::new(crate::tools::mcp::process::McpProcessManager::new()),
secrets, secrets,
tool_registry, tool_registry,
None, None,
@@ -2813,7 +2649,7 @@ mod tests {
let secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync> = let secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync> =
Arc::new(crate::secrets::InMemorySecretsStore::new(Arc::new( Arc::new(crate::secrets::InMemorySecretsStore::new(Arc::new(
crate::secrets::SecretsCrypto::new(secrecy::SecretString::from( crate::secrets::SecretsCrypto::new(secrecy::SecretString::from(
TEST_GATEWAY_CRYPTO_KEY.to_string(), "test-key-at-least-32-chars-long!!".to_string(),
)) ))
.expect("crypto"), .expect("crypto"),
))); )));
@@ -2822,7 +2658,6 @@ mod tests {
let ext_mgr = Arc::new(ExtensionManager::new( let ext_mgr = Arc::new(ExtensionManager::new(
mcp_sm, mcp_sm,
Arc::new(crate::tools::mcp::process::McpProcessManager::new()),
secrets.clone(), secrets.clone(),
tool_registry, tool_registry,
None, None,
@@ -2919,7 +2754,7 @@ mod tests {
let secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync> = let secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync> =
Arc::new(crate::secrets::InMemorySecretsStore::new(Arc::new( Arc::new(crate::secrets::InMemorySecretsStore::new(Arc::new(
crate::secrets::SecretsCrypto::new(secrecy::SecretString::from( crate::secrets::SecretsCrypto::new(secrecy::SecretString::from(
TEST_GATEWAY_CRYPTO_KEY.to_string(), "test-key-at-least-32-chars-long!!".to_string(),
)) ))
.expect("crypto"), .expect("crypto"),
))); )));
@@ -2928,7 +2763,6 @@ mod tests {
let ext_mgr = Arc::new(ExtensionManager::new( let ext_mgr = Arc::new(ExtensionManager::new(
mcp_sm, mcp_sm,
Arc::new(crate::tools::mcp::process::McpProcessManager::new()),
secrets.clone(), secrets.clone(),
tool_registry, tool_registry,
None, None,
@@ -3013,177 +2847,4 @@ mod tests {
.is_none() .is_none()
); );
} }
// --- Slack relay OAuth CSRF tests ---
fn test_relay_oauth_router(state: Arc<GatewayState>) -> Router {
Router::new()
.route(
"/oauth/slack/callback",
get(slack_relay_oauth_callback_handler),
)
.with_state(state)
}
fn test_secrets_store() -> 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"),
)))
}
fn test_ext_mgr(
secrets: Arc<dyn crate::secrets::SecretsStore + Send + Sync>,
) -> Arc<ExtensionManager> {
let tool_registry = Arc::new(ToolRegistry::new());
let mcp_sm = Arc::new(crate::tools::mcp::session::McpSessionManager::new());
let mcp_pm = Arc::new(crate::tools::mcp::process::McpProcessManager::new());
Arc::new(ExtensionManager::new(
mcp_sm,
mcp_pm,
secrets,
tool_registry,
None,
None,
std::path::PathBuf::from("/tmp/wasm_tools"),
std::path::PathBuf::from("/tmp/wasm_channels"),
None,
"test".to_string(),
None,
vec![],
))
}
#[tokio::test]
async fn test_relay_oauth_callback_missing_state_param() {
use axum::body::Body;
use tower::ServiceExt;
let secrets = test_secrets_store();
let ext_mgr = test_ext_mgr(secrets);
let state = test_gateway_state(Some(ext_mgr));
let app = test_relay_oauth_router(state);
// Callback without state param should be rejected
let req = axum::http::Request::builder()
.uri("/oauth/slack/callback?stream_token=tok123&team_id=T123&provider=slack")
.body(Body::empty())
.expect("request");
let resp = ServiceExt::<axum::http::Request<Body>>::oneshot(app, req)
.await
.expect("response");
let body = axum::body::to_bytes(resp.into_body(), 1024 * 64)
.await
.expect("body");
let html = String::from_utf8_lossy(&body);
assert!(
html.contains("Invalid or expired authorization"),
"Expected CSRF error, got: {}",
&html[..html.len().min(300)]
);
}
#[tokio::test]
async fn test_relay_oauth_callback_wrong_state_param() {
use axum::body::Body;
use tower::ServiceExt;
let secrets = test_secrets_store();
// Store a valid nonce
secrets
.create(
"test",
crate::secrets::CreateSecretParams::new(
format!("relay:{}:oauth_state", DEFAULT_RELAY_NAME),
"correct-nonce-value",
),
)
.await
.expect("store nonce");
let ext_mgr = test_ext_mgr(secrets);
let state = test_gateway_state(Some(ext_mgr));
let app = test_relay_oauth_router(state);
// Callback with wrong state param
let req = axum::http::Request::builder()
.uri("/oauth/slack/callback?stream_token=tok123&team_id=T123&provider=slack&state=wrong-nonce")
.body(Body::empty())
.expect("request");
let resp = ServiceExt::<axum::http::Request<Body>>::oneshot(app, req)
.await
.expect("response");
let body = axum::body::to_bytes(resp.into_body(), 1024 * 64)
.await
.expect("body");
let html = String::from_utf8_lossy(&body);
assert!(
html.contains("Invalid or expired authorization"),
"Expected CSRF error for wrong nonce, got: {}",
&html[..html.len().min(300)]
);
}
#[tokio::test]
async fn test_relay_oauth_callback_correct_state_proceeds() {
use axum::body::Body;
use tower::ServiceExt;
let secrets = test_secrets_store();
let nonce = "valid-test-nonce-12345";
// Store the correct nonce
secrets
.create(
"test",
crate::secrets::CreateSecretParams::new(
format!("relay:{}:oauth_state", DEFAULT_RELAY_NAME),
nonce,
),
)
.await
.expect("store nonce");
let ext_mgr = test_ext_mgr(secrets.clone());
let state = test_gateway_state(Some(ext_mgr));
let app = test_relay_oauth_router(state);
// Callback with correct state param — will pass CSRF check
// but may fail downstream (no real relay service) — that's OK,
// we just verify it doesn't return a CSRF error.
let req = axum::http::Request::builder()
.uri(format!(
"/oauth/slack/callback?stream_token=tok123&team_id=T123&provider=slack&state={}",
nonce
))
.body(Body::empty())
.expect("request");
let resp = ServiceExt::<axum::http::Request<Body>>::oneshot(app, req)
.await
.expect("response");
let body = axum::body::to_bytes(resp.into_body(), 1024 * 64)
.await
.expect("body");
let html = String::from_utf8_lossy(&body);
// Should NOT contain the CSRF error message
assert!(
!html.contains("Invalid or expired authorization"),
"Should have passed CSRF check, got: {}",
&html[..html.len().min(300)]
);
// Verify the nonce was consumed (deleted)
let state_key = format!("relay:{}:oauth_state", DEFAULT_RELAY_NAME);
let exists = secrets.exists("test", &state_key).await.unwrap_or(true);
assert!(!exists, "CSRF nonce should be deleted after use");
}
} }
+2 -2
View File
@@ -2350,8 +2350,8 @@ function renderExtensionCard(ext) {
activeLabel.textContent = ext.active ? 'Active' : 'Installed'; activeLabel.textContent = ext.active ? 'Active' : 'Installed';
actions.appendChild(activeLabel); actions.appendChild(activeLabel);
// MCP servers and channel-relay extensions may be installed but inactive — show Activate button // MCP servers may be installed but inactive — show Activate button
if ((ext.kind === 'mcp_server' || ext.kind === 'channel_relay') && !ext.active) { if (ext.kind === 'mcp_server' && !ext.active) {
const activateBtn = document.createElement('button'); const activateBtn = document.createElement('button');
activateBtn.className = 'btn-ext activate'; activateBtn.className = 'btn-ext activate';
activateBtn.textContent = 'Activate'; activateBtn.textContent = 'Activate';
-17
View File
@@ -1277,8 +1277,6 @@ body {
gap: 8px; gap: 8px;
background: var(--bg-secondary); background: var(--bg-secondary);
border-top: 1px solid var(--border); border-top: 1px solid var(--border);
flex-shrink: 0;
min-height: 56px;
} }
.chat-input textarea { .chat-input textarea {
@@ -3722,21 +3720,6 @@ mark {
.ext-install-form input { .ext-install-form input {
width: 100%; width: 100%;
} }
/* Chat input: ensure visibility on mobile */
.chat-input {
min-height: 52px;
}
.chat-input textarea {
min-height: 36px;
max-height: 100px;
}
.chat-input button {
padding: 6px 16px;
font-size: 14px;
}
} }
/* Slash command autocomplete dropdown */ /* Slash command autocomplete dropdown */
-1
View File
@@ -82,7 +82,6 @@ impl TestGatewayBuilder {
skill_catalog: None, skill_catalog: None,
scheduler: None, scheduler: None,
chat_rate_limiter: RateLimiter::new(30, 60), chat_rate_limiter: RateLimiter::new(30, 60),
oauth_rate_limiter: RateLimiter::new(10, 60),
registry_entries: Vec::new(), registry_entries: Vec::new(),
cost_guard: None, cost_guard: None,
routine_engine: Arc::new(tokio::sync::RwLock::new(None)), routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
-54
View File
@@ -735,60 +735,6 @@ pub struct RoutineInfo {
pub status: String, pub status: String,
} }
impl RoutineInfo {
/// Convert a `Routine` to the trimmed `RoutineInfo` for list display.
pub fn from_routine(r: &crate::agent::routine::Routine) -> Self {
let (trigger_type, trigger_summary) = match &r.trigger {
crate::agent::routine::Trigger::Cron { schedule, .. } => {
("cron".to_string(), format!("cron: {}", schedule))
}
crate::agent::routine::Trigger::Event {
pattern, channel, ..
} => {
let ch = channel.as_deref().unwrap_or("any");
("event".to_string(), format!("on {} /{}/", ch, pattern))
}
crate::agent::routine::Trigger::SystemEvent {
source, event_type, ..
} => (
"system_event".to_string(),
format!("event: {}.{}", source, event_type),
),
crate::agent::routine::Trigger::Manual => {
("manual".to_string(), "manual only".to_string())
}
};
let action_type = match &r.action {
crate::agent::routine::RoutineAction::Lightweight { .. } => "lightweight",
crate::agent::routine::RoutineAction::FullJob { .. } => "full_job",
};
let status = if !r.enabled {
"disabled"
} else if r.consecutive_failures > 0 {
"failing"
} else {
"active"
};
RoutineInfo {
id: r.id,
name: r.name.clone(),
description: r.description.clone(),
enabled: r.enabled,
trigger_type,
trigger_summary,
action_type: action_type.to_string(),
last_run_at: r.last_run_at.map(|dt| dt.to_rfc3339()),
next_fire_at: r.next_fire_at.map(|dt| dt.to_rfc3339()),
run_count: r.run_count,
consecutive_failures: r.consecutive_failures,
status: status.to_string(),
}
}
}
#[derive(Debug, Serialize)] #[derive(Debug, Serialize)]
pub struct RoutineListResponse { pub struct RoutineListResponse {
pub routines: Vec<RoutineInfo>, pub routines: Vec<RoutineInfo>,
-1
View File
@@ -509,7 +509,6 @@ mod tests {
skill_registry: None, skill_registry: None,
skill_catalog: None, skill_catalog: None,
chat_rate_limiter: crate::channels::web::server::RateLimiter::new(30, 60), chat_rate_limiter: crate::channels::web::server::RateLimiter::new(30, 60),
oauth_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60),
registry_entries: Vec::new(), registry_entries: Vec::new(),
cost_guard: None, cost_guard: None,
routine_engine: Arc::new(tokio::sync::RwLock::new(None)), routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
+1 -237
View File
@@ -24,8 +24,6 @@ pub struct WebhookServerConfig {
pub struct WebhookServer { pub struct WebhookServer {
config: WebhookServerConfig, config: WebhookServerConfig,
routes: Vec<Router>, routes: Vec<Router>,
/// Merged router saved after start() for restart_with_addr().
merged_router: Option<Router>,
shutdown_tx: Option<oneshot::Sender<()>>, shutdown_tx: Option<oneshot::Sender<()>>,
handle: Option<JoinHandle<()>>, handle: Option<JoinHandle<()>>,
} }
@@ -36,7 +34,6 @@ impl WebhookServer {
Self { Self {
config, config,
routes: Vec::new(), routes: Vec::new(),
merged_router: None,
shutdown_tx: None, shutdown_tx: None,
handle: None, handle: None,
} }
@@ -54,13 +51,7 @@ impl WebhookServer {
for fragment in self.routes.drain(..) { for fragment in self.routes.drain(..) {
app = app.merge(fragment); app = app.merge(fragment);
} }
self.merged_router = Some(app.clone());
self.bind_and_spawn(app).await
}
/// Bind a listener to the configured address and spawn the server task.
/// Private helper used by both start() and restart_with_addr().
async fn bind_and_spawn(&mut self, app: Router) -> Result<(), ChannelError> {
let listener = tokio::net::TcpListener::bind(self.config.addr) let listener = tokio::net::TcpListener::bind(self.config.addr)
.await .await
.map_err(|e| ChannelError::StartupFailed { .map_err(|e| ChannelError::StartupFailed {
@@ -77,7 +68,7 @@ impl WebhookServer {
if let Err(e) = axum::serve(listener, app) if let Err(e) = axum::serve(listener, app)
.with_graceful_shutdown(async { .with_graceful_shutdown(async {
let _ = shutdown_rx.await; let _ = shutdown_rx.await;
tracing::debug!("Webhook server shutting down"); tracing::info!("Webhook server shutting down");
}) })
.await .await
{ {
@@ -89,54 +80,6 @@ impl WebhookServer {
Ok(()) Ok(())
} }
/// Gracefully shut down the current listener and rebind to a new address.
/// The merged router from the original `start()` call is reused.
///
/// If binding to the new address fails, the old listener remains active and
/// state is restored. This prevents a denial-of-service if the new address
/// is invalid or already in use.
pub async fn restart_with_addr(&mut self, new_addr: SocketAddr) -> Result<(), ChannelError> {
let app = self
.merged_router
.clone()
.ok_or_else(|| ChannelError::StartupFailed {
name: "webhook_server".to_string(),
reason: "restart_with_addr called before start()".to_string(),
})?;
// Save old state for rollback if new bind fails
let old_addr = self.config.addr;
let old_shutdown_tx = self.shutdown_tx.take();
let old_handle = self.handle.take();
// Update config to new address and try to bind
self.config.addr = new_addr;
match self.bind_and_spawn(app).await {
Ok(()) => {
// New listener is running, gracefully shut down the old one
if let Some(tx) = old_shutdown_tx {
let _ = tx.send(());
}
if let Some(handle) = old_handle {
let _ = handle.await;
}
Ok(())
}
Err(e) => {
// Restore old state; old listener remains active
self.config.addr = old_addr;
self.shutdown_tx = old_shutdown_tx;
self.handle = old_handle;
Err(e)
}
}
}
/// Return the current bind address.
pub fn current_addr(&self) -> SocketAddr {
self.config.addr
}
/// Signal graceful shutdown and wait for the server task to finish. /// Signal graceful shutdown and wait for the server task to finish.
pub async fn shutdown(&mut self) { pub async fn shutdown(&mut self) {
if let Some(tx) = self.shutdown_tx.take() { if let Some(tx) = self.shutdown_tx.take() {
@@ -147,182 +90,3 @@ impl WebhookServer {
} }
} }
} }
#[cfg(test)]
mod tests {
use super::*;
use axum::Json;
use serde_json::json;
#[tokio::test]
async fn test_restart_with_addr_rebinds_listener() {
use std::net::TcpListener as StdTcpListener;
// Find two available ports by binding and immediately closing
let port1 = {
let listener =
StdTcpListener::bind("127.0.0.1:0").expect("Failed to find available port 1");
listener
.local_addr()
.expect("Failed to get local addr")
.port()
};
let port2 = {
let listener =
StdTcpListener::bind("127.0.0.1:0").expect("Failed to find available port 2");
listener
.local_addr()
.expect("Failed to get local addr")
.port()
};
assert_ne!(port1, port2, "Should have different ports");
assert_ne!(port1, 0, "Port 1 should be non-zero");
assert_ne!(port2, 0, "Port 2 should be non-zero");
// Start server on first port
let addr1 = format!("127.0.0.1:{}", port1).parse().unwrap();
let mut server = WebhookServer::new(WebhookServerConfig { addr: addr1 });
// Create a test router that responds to health checks
let test_router = axum::Router::new().route(
"/health",
axum::routing::get(|| async { Json(json!({"status": "ok"})) }),
);
server.add_routes(test_router);
// Start the server on first port
server.start().await.expect("Failed to start server");
assert_eq!(
server.current_addr(),
addr1,
"Server should be bound to initial address"
);
// Verify the first server is actually listening
let client = reqwest::Client::new();
let response = client
.get(format!("http://{}/health", addr1))
.send()
.await
.expect("Failed to send request to first server");
assert_eq!(
response.status(),
200,
"First server should respond to health check"
);
// Restart on second port
let addr2 = format!("127.0.0.1:{}", port2).parse().unwrap();
server
.restart_with_addr(addr2)
.await
.expect("Failed to restart with new addr");
// Assert the address changed
assert_eq!(
server.current_addr(),
addr2,
"Server address should be updated after restart"
);
assert_ne!(
addr1, addr2,
"Address should change after restart_with_addr"
);
// Verify the new server is actually listening on the new address
let response = client
.get(format!("http://{}/health", addr2))
.send()
.await
.expect("Failed to send request to restarted server");
assert_eq!(
response.status(),
200,
"Restarted server should respond to health check on new address"
);
// Verify the old address is no longer responding
let old_result = tokio::time::timeout(
std::time::Duration::from_millis(200),
client.get(format!("http://{}/health", addr1)).send(),
)
.await;
assert!(
old_result.is_err() || old_result.as_ref().unwrap().is_err(),
"Old address should not respond after server restarts"
);
// Clean up
server.shutdown().await;
}
#[tokio::test]
async fn test_restart_with_addr_rollback_on_bind_failure() {
use std::net::TcpListener as StdTcpListener;
// Find an available port
let port1 = {
let listener =
StdTcpListener::bind("127.0.0.1:0").expect("Failed to find available port");
listener
.local_addr()
.expect("Failed to get local addr")
.port()
};
// Start server on first port
let addr1 = format!("127.0.0.1:{}", port1).parse().unwrap();
let mut server = WebhookServer::new(WebhookServerConfig { addr: addr1 });
// Create a test router
let test_router = axum::Router::new().route(
"/health",
axum::routing::get(|| async { Json(json!({"status": "ok"})) }),
);
server.add_routes(test_router);
// Start the server on first port
server.start().await.expect("Failed to start server");
// Verify the server is listening
let client = reqwest::Client::new();
let response = client
.get(format!("http://{}/health", addr1))
.send()
.await
.expect("Failed to send request");
assert_eq!(response.status(), 200, "Server should be listening");
// Try to restart on an invalid address (port 0 is reserved, won't bind)
// Use port 1 which typically requires elevated privileges
let invalid_addr: SocketAddr = "127.0.0.1:1".parse().unwrap();
// Attempt restart (should fail)
let result = server.restart_with_addr(invalid_addr).await;
assert!(result.is_err(), "Restart with invalid address should fail");
// Verify the old address is still responding (rollback succeeded)
let response = client
.get(format!("http://{}/health", addr1))
.send()
.await
.expect("Failed to send request to old address");
assert_eq!(
response.status(),
200,
"Old listener should still be running after failed restart"
);
// Verify the server address is unchanged
assert_eq!(
server.current_addr(),
addr1,
"Server address should be restored after failed restart"
);
// Clean up
server.shutdown().await;
}
}
+5 -541
View File
@@ -7,7 +7,6 @@
use std::path::PathBuf; use std::path::PathBuf;
use crate::bootstrap::ironclaw_base_dir; use crate::bootstrap::ironclaw_base_dir;
use crate::settings::Settings;
/// Run all diagnostic checks and print results. /// Run all diagnostic checks and print results.
pub async fn run_doctor_command() -> anyhow::Result<()> { pub async fn run_doctor_command() -> anyhow::Result<()> {
@@ -16,35 +15,14 @@ pub async fn run_doctor_command() -> anyhow::Result<()> {
let mut passed = 0u32; let mut passed = 0u32;
let mut failed = 0u32; let mut failed = 0u32;
let mut skipped = 0u32;
// Load settings once for checks that need them. // ── Configuration checks ──────────────────────────────────
let settings = Settings::load();
// ── Settings & core config ─────────────────────────────────
check(
"Settings file",
check_settings_file(),
&mut passed,
&mut failed,
&mut skipped,
);
check( check(
"NEAR AI session", "NEAR AI session",
check_nearai_session().await, check_nearai_session().await,
&mut passed, &mut passed,
&mut failed, &mut failed,
&mut skipped,
);
check(
"LLM configuration",
check_llm_config(&settings),
&mut passed,
&mut failed,
&mut skipped,
); );
check( check(
@@ -52,7 +30,6 @@ pub async fn run_doctor_command() -> anyhow::Result<()> {
check_database().await, check_database().await,
&mut passed, &mut passed,
&mut failed, &mut failed,
&mut skipped,
); );
check( check(
@@ -60,75 +37,15 @@ pub async fn run_doctor_command() -> anyhow::Result<()> {
check_workspace_dir(), check_workspace_dir(),
&mut passed, &mut passed,
&mut failed, &mut failed,
&mut skipped,
);
// ── Subsystem configuration checks ─────────────────────────
check(
"Embeddings",
check_embeddings(&settings),
&mut passed,
&mut failed,
&mut skipped,
);
check(
"Routines config",
check_routines_config(),
&mut passed,
&mut failed,
&mut skipped,
);
check(
"Gateway config",
check_gateway_config(&settings),
&mut passed,
&mut failed,
&mut skipped,
);
check(
"MCP servers",
check_mcp_config().await,
&mut passed,
&mut failed,
&mut skipped,
);
check(
"Skills",
check_skills().await,
&mut passed,
&mut failed,
&mut skipped,
);
check(
"Secrets",
check_secrets(&settings),
&mut passed,
&mut failed,
&mut skipped,
);
check(
"Service",
check_service_installed(),
&mut passed,
&mut failed,
&mut skipped,
); );
// ── External binary checks ──────────────────────────────── // ── External binary checks ────────────────────────────────
check( check(
"Docker daemon", "Docker",
check_docker_daemon().await, check_binary("docker", &["--version"]),
&mut passed, &mut passed,
&mut failed, &mut failed,
&mut skipped,
); );
check( check(
@@ -136,7 +53,6 @@ pub async fn run_doctor_command() -> anyhow::Result<()> {
check_binary("cloudflared", &["--version"]), check_binary("cloudflared", &["--version"]),
&mut passed, &mut passed,
&mut failed, &mut failed,
&mut skipped,
); );
check( check(
@@ -144,7 +60,6 @@ pub async fn run_doctor_command() -> anyhow::Result<()> {
check_binary("ngrok", &["version"]), check_binary("ngrok", &["version"]),
&mut passed, &mut passed,
&mut failed, &mut failed,
&mut skipped,
); );
check( check(
@@ -152,13 +67,12 @@ pub async fn run_doctor_command() -> anyhow::Result<()> {
check_binary("tailscale", &["version"]), check_binary("tailscale", &["version"]),
&mut passed, &mut passed,
&mut failed, &mut failed,
&mut skipped,
); );
// ── Summary ─────────────────────────────────────────────── // ── Summary ───────────────────────────────────────────────
println!(); println!();
println!(" {passed} passed, {failed} failed, {skipped} skipped"); println!(" {passed} passed, {failed} failed");
if failed > 0 { if failed > 0 {
println!("\n Some checks failed. This is normal if you don't use those features."); println!("\n Some checks failed. This is normal if you don't use those features.");
@@ -169,7 +83,7 @@ pub async fn run_doctor_command() -> anyhow::Result<()> {
// ── Individual checks ─────────────────────────────────────── // ── Individual checks ───────────────────────────────────────
fn check(name: &str, result: CheckResult, passed: &mut u32, failed: &mut u32, skipped: &mut u32) { fn check(name: &str, result: CheckResult, passed: &mut u32, failed: &mut u32) {
match result { match result {
CheckResult::Pass(detail) => { CheckResult::Pass(detail) => {
*passed += 1; *passed += 1;
@@ -180,7 +94,6 @@ fn check(name: &str, result: CheckResult, passed: &mut u32, failed: &mut u32, sk
println!(" [FAIL] {name}: {detail}"); println!(" [FAIL] {name}: {detail}");
} }
CheckResult::Skip(reason) => { CheckResult::Skip(reason) => {
*skipped += 1;
println!(" [skip] {name}: {reason}"); println!(" [skip] {name}: {reason}");
} }
} }
@@ -192,29 +105,6 @@ enum CheckResult {
Skip(String), Skip(String),
} }
// ── Settings file ───────────────────────────────────────────
fn check_settings_file() -> CheckResult {
let path = Settings::default_path();
if !path.exists() {
return CheckResult::Pass("no settings file (defaults will be used)".into());
}
match std::fs::read_to_string(&path) {
Ok(data) => match serde_json::from_str::<serde_json::Value>(&data) {
Ok(_) => CheckResult::Pass(format!("valid ({})", path.display())),
Err(e) => CheckResult::Fail(format!(
"settings.json is malformed: {}. Fix or delete {}",
e,
path.display()
)),
},
Err(e) => CheckResult::Fail(format!("cannot read {}: {}", path.display(), e)),
}
}
// ── NEAR AI session ─────────────────────────────────────────
async fn check_nearai_session() -> CheckResult { async fn check_nearai_session() -> CheckResult {
// Check if session file exists // Check if session file exists
let session_path = crate::config::llm::default_session_path(); let session_path = crate::config::llm::default_session_path();
@@ -239,27 +129,6 @@ async fn check_nearai_session() -> CheckResult {
} }
} }
// ── LLM configuration ──────────────────────────────────────
fn check_llm_config(settings: &Settings) -> CheckResult {
match crate::llm::LlmConfig::resolve(settings) {
Ok(config) => {
// Show the model for the active backend, not always nearai.model.
let model = if let Some(ref bedrock) = config.bedrock {
&bedrock.model
} else if let Some(ref provider) = config.provider {
&provider.model
} else {
&config.nearai.model
};
CheckResult::Pass(format!("backend={}, model={}", config.backend, model))
}
Err(e) => CheckResult::Fail(format!("LLM config error: {e}")),
}
}
// ── Database ────────────────────────────────────────────────
async fn check_database() -> CheckResult { async fn check_database() -> CheckResult {
let backend = std::env::var("DATABASE_BACKEND") let backend = std::env::var("DATABASE_BACKEND")
.ok() .ok()
@@ -323,8 +192,6 @@ async fn try_pg_connect() -> Result<(), String> {
Err("postgres feature not compiled in".into()) Err("postgres feature not compiled in".into())
} }
// ── Workspace directory ─────────────────────────────────────
fn check_workspace_dir() -> CheckResult { fn check_workspace_dir() -> CheckResult {
let dir = ironclaw_base_dir(); let dir = ironclaw_base_dir();
@@ -339,222 +206,6 @@ fn check_workspace_dir() -> CheckResult {
} }
} }
// ── Embeddings ──────────────────────────────────────────────
fn check_embeddings(settings: &Settings) -> CheckResult {
match crate::config::EmbeddingsConfig::resolve(settings) {
Ok(config) => {
if !config.enabled {
return CheckResult::Skip("disabled (set EMBEDDING_ENABLED=true)".into());
}
let has_creds = match config.provider.as_str() {
"openai" => config.openai_api_key().is_some(),
"nearai" => {
// NearAiEmbeddings uses SessionManager::get_token() which
// only returns session tokens, NOT NEARAI_API_KEY
// (src/workspace/embeddings.rs:309, src/llm/session.rs:132).
let session_path = crate::config::llm::default_session_path();
session_path.exists()
&& std::fs::read_to_string(&session_path)
.map(|s| !s.trim().is_empty())
.unwrap_or(false)
}
"ollama" => true, // local, no creds needed
_ => config.openai_api_key().is_some(),
};
if has_creds {
CheckResult::Pass(format!(
"provider={}, model={}",
config.provider, config.model
))
} else {
let hint = match config.provider.as_str() {
"nearai" => "run `ironclaw onboard` to create a session",
_ => "set OPENAI_API_KEY",
};
CheckResult::Fail(format!(
"provider={} but credentials missing ({})",
config.provider, hint
))
}
}
Err(e) => CheckResult::Fail(format!("config error: {e}")),
}
}
// ── Routines config ─────────────────────────────────────────
fn check_routines_config() -> CheckResult {
match crate::config::RoutineConfig::resolve() {
Ok(config) => {
if config.enabled {
CheckResult::Pass(format!(
"enabled (interval={}s, max_concurrent={})",
config.cron_check_interval_secs, config.max_concurrent_routines
))
} else {
CheckResult::Skip("disabled".into())
}
}
Err(e) => CheckResult::Fail(format!("config error: {e}")),
}
}
// ── Gateway config ──────────────────────────────────────────
fn check_gateway_config(settings: &Settings) -> CheckResult {
// Use the same resolve() path as runtime so invalid env values
// (e.g. GATEWAY_PORT=abc) are caught here too.
match crate::config::ChannelsConfig::resolve(settings) {
Ok(channels) => match channels.gateway {
Some(gw) => {
if gw.auth_token.is_some() {
CheckResult::Pass(format!(
"enabled at {}:{} (auth token set)",
gw.host, gw.port
))
} else {
CheckResult::Pass(format!(
"enabled at {}:{} (no auth token — random token will be generated)",
gw.host, gw.port
))
}
}
None => CheckResult::Skip("disabled (GATEWAY_ENABLED=false)".into()),
},
Err(e) => CheckResult::Fail(format!("config error: {e}")),
}
}
// ── MCP servers ─────────────────────────────────────────────
async fn check_mcp_config() -> CheckResult {
match crate::tools::mcp::config::load_mcp_servers().await {
Ok(file) => {
let servers: Vec<_> = file.enabled_servers().collect();
if servers.is_empty() {
return CheckResult::Skip("no MCP servers configured".into());
}
let mut invalid = Vec::new();
for server in &servers {
if let Err(e) = server.validate() {
invalid.push(format!("{}: {}", server.name, e));
}
}
if invalid.is_empty() {
CheckResult::Pass(format!("{} server(s) configured, all valid", servers.len()))
} else {
CheckResult::Fail(format!(
"{} server(s), {} invalid: {}",
servers.len(),
invalid.len(),
invalid.join("; ")
))
}
}
Err(e) => {
// Distinguish no config from corrupted config
let msg = e.to_string();
if msg.contains("not found") || msg.contains("No such file") {
CheckResult::Skip("no MCP config file".into())
} else {
CheckResult::Fail(format!("config error: {e}"))
}
}
}
}
// ── Skills ──────────────────────────────────────────────────
async fn check_skills() -> CheckResult {
let user_dir = ironclaw_base_dir().join("skills");
let installed_dir = ironclaw_base_dir().join("installed_skills");
let mut registry = crate::skills::SkillRegistry::new(user_dir.clone());
registry = registry.with_installed_dir(installed_dir);
// discover_all() returns loaded skill names (not warnings).
let _loaded_names = registry.discover_all().await;
let count = registry.count();
if count == 0 {
return CheckResult::Skip("no skills discovered".into());
}
CheckResult::Pass(format!("{count} skill(s) loaded"))
}
// ── Secrets ─────────────────────────────────────────────────
fn check_secrets(settings: &Settings) -> CheckResult {
match settings.secrets_master_key_source {
crate::settings::KeySource::Keychain => {
CheckResult::Pass("master key source: OS keychain".into())
}
crate::settings::KeySource::Env => {
if std::env::var("SECRETS_MASTER_KEY").is_ok() {
CheckResult::Pass("master key source: env var (set)".into())
} else {
CheckResult::Fail(
"master key source: env var but SECRETS_MASTER_KEY not set".into(),
)
}
}
crate::settings::KeySource::None => {
CheckResult::Skip("secrets not configured (run `ironclaw onboard`)".into())
}
}
}
// ── Service ─────────────────────────────────────────────────
fn check_service_installed() -> CheckResult {
if cfg!(target_os = "macos") {
let plist =
dirs::home_dir().map(|h| h.join("Library/LaunchAgents/com.ironclaw.daemon.plist"));
match plist {
Some(path) if path.exists() => {
CheckResult::Pass(format!("launchd plist installed ({})", path.display()))
}
Some(_) => CheckResult::Skip("not installed (run `ironclaw service install`)".into()),
None => CheckResult::Skip("cannot determine home directory".into()),
}
} else if cfg!(target_os = "linux") {
let unit = dirs::home_dir().map(|h| h.join(".config/systemd/user/ironclaw.service"));
match unit {
Some(path) if path.exists() => {
CheckResult::Pass(format!("systemd unit installed ({})", path.display()))
}
Some(_) => CheckResult::Skip("not installed (run `ironclaw service install`)".into()),
None => CheckResult::Skip("cannot determine home directory".into()),
}
} else {
CheckResult::Skip("service management not supported on this platform".into())
}
}
// ── Docker daemon ───────────────────────────────────────────
async fn check_docker_daemon() -> CheckResult {
let detection = crate::sandbox::check_docker().await;
match detection.status {
crate::sandbox::DockerStatus::Available => CheckResult::Pass("running".into()),
crate::sandbox::DockerStatus::NotInstalled => CheckResult::Skip(format!(
"not installed. {}",
detection.platform.install_hint()
)),
crate::sandbox::DockerStatus::NotRunning => CheckResult::Fail(format!(
"installed but not running. {}",
detection.platform.start_hint()
)),
crate::sandbox::DockerStatus::Disabled => CheckResult::Skip("sandbox disabled".into()),
}
}
// ── External binary ─────────────────────────────────────────
fn check_binary(name: &str, args: &[&str]) -> CheckResult { fn check_binary(name: &str, args: &[&str]) -> CheckResult {
match std::process::Command::new(name) match std::process::Command::new(name)
.args(args) .args(args)
@@ -622,193 +273,6 @@ mod tests {
} }
} }
#[test]
fn check_settings_file_handles_missing() {
// Settings::default_path() might or might not exist, but must not panic
let result = check_settings_file();
match result {
CheckResult::Pass(_) | CheckResult::Fail(_) | CheckResult::Skip(_) => {}
}
}
#[test]
fn check_llm_config_does_not_panic() {
let settings = Settings::default();
let result = check_llm_config(&settings);
match result {
CheckResult::Pass(_) | CheckResult::Fail(_) | CheckResult::Skip(_) => {}
}
}
#[test]
fn check_routines_config_does_not_panic() {
let result = check_routines_config();
match result {
CheckResult::Pass(_) | CheckResult::Fail(_) | CheckResult::Skip(_) => {}
}
}
#[test]
fn check_gateway_config_does_not_panic() {
let settings = Settings::default();
let result = check_gateway_config(&settings);
match result {
CheckResult::Pass(_) | CheckResult::Fail(_) | CheckResult::Skip(_) => {}
}
}
#[test]
fn check_embeddings_does_not_panic() {
let settings = Settings::default();
let result = check_embeddings(&settings);
match result {
CheckResult::Pass(_) | CheckResult::Fail(_) | CheckResult::Skip(_) => {}
}
}
#[test]
fn check_secrets_none_returns_skip() {
let settings = Settings::default();
match check_secrets(&settings) {
CheckResult::Skip(msg) => {
assert!(
msg.contains("not configured"),
"expected 'not configured' in skip message, got: {msg}"
);
}
other => panic!(
"expected Skip for default settings, got: {}",
format_result(&other)
),
}
}
#[test]
fn check_service_installed_does_not_panic() {
let result = check_service_installed();
match result {
CheckResult::Pass(_) | CheckResult::Fail(_) | CheckResult::Skip(_) => {}
}
}
#[tokio::test]
async fn check_docker_daemon_does_not_panic() {
let result = check_docker_daemon().await;
match result {
CheckResult::Pass(_) | CheckResult::Fail(_) | CheckResult::Skip(_) => {}
}
}
#[tokio::test]
async fn check_mcp_config_does_not_panic() {
let result = check_mcp_config().await;
match result {
CheckResult::Pass(_) | CheckResult::Fail(_) | CheckResult::Skip(_) => {}
}
}
#[tokio::test]
async fn check_skills_does_not_panic() {
let result = check_skills().await;
match result {
CheckResult::Pass(_) | CheckResult::Fail(_) | CheckResult::Skip(_) => {}
}
}
#[test]
fn check_llm_config_shows_nearai_model_for_nearai_backend() {
let _guard = crate::config::helpers::ENV_MUTEX.lock().expect("env mutex");
// SAFETY: Under ENV_MUTEX, no concurrent env access.
unsafe {
std::env::remove_var("LLM_BACKEND");
}
let settings = Settings::default();
match check_llm_config(&settings) {
CheckResult::Pass(msg) => {
assert!(
msg.contains("backend=nearai"),
"expected nearai backend, got: {msg}"
);
// Must NOT show a bedrock or registry model when backend is nearai
assert!(
!msg.contains("anthropic.claude"),
"should not show bedrock model for nearai backend: {msg}"
);
}
other => panic!(
"expected Pass for default LLM config, got: {}",
format_result(&other)
),
}
}
#[test]
fn check_embeddings_disabled_by_default_returns_skip() {
let _guard = crate::config::helpers::ENV_MUTEX.lock().expect("env mutex");
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::remove_var("EMBEDDING_ENABLED");
}
let settings = Settings::default();
match check_embeddings(&settings) {
CheckResult::Skip(msg) => {
assert!(
msg.contains("disabled"),
"expected 'disabled' in skip message, got: {msg}"
);
}
other => panic!(
"expected Skip for disabled embeddings, got: {}",
format_result(&other)
),
}
}
#[test]
fn check_routines_enabled_by_default() {
let _guard = crate::config::helpers::ENV_MUTEX.lock().expect("env mutex");
// SAFETY: Under ENV_MUTEX.
unsafe {
std::env::remove_var("ROUTINES_ENABLED");
}
match check_routines_config() {
CheckResult::Pass(msg) => {
assert!(
msg.contains("enabled"),
"routines should be enabled by default, got: {msg}"
);
}
other => panic!(
"expected Pass for default routines, got: {}",
format_result(&other)
),
}
}
#[test]
fn check_secrets_env_without_var_returns_fail() {
let settings = Settings {
secrets_master_key_source: crate::settings::KeySource::Env,
..Default::default()
};
match check_secrets(&settings) {
CheckResult::Fail(msg) => {
assert!(
msg.contains("SECRETS_MASTER_KEY not set"),
"expected mention of missing env var, got: {msg}"
);
}
CheckResult::Pass(_) => {
// If SECRETS_MASTER_KEY happens to be set in the environment,
// Pass is correct — don't fail the test.
}
other => panic!(
"expected Fail or Pass for env key source, got: {}",
format_result(&other)
),
}
}
fn format_result(r: &CheckResult) -> String { fn format_result(r: &CheckResult) -> String {
match r { match r {
CheckResult::Pass(s) => format!("Pass({s})"), CheckResult::Pass(s) => format!("Pass({s})"),
+12 -2
View File
@@ -10,7 +10,7 @@ use clap::{Args, Subcommand};
use crate::config::Config; use crate::config::Config;
use crate::db::Database; use crate::db::Database;
use crate::secrets::SecretsStore; use crate::secrets::{SecretsCrypto, SecretsStore};
use crate::tools::mcp::{ use crate::tools::mcp::{
McpClient, McpServerConfig, McpSessionManager, OAuthConfig, McpClient, McpServerConfig, McpSessionManager, OAuthConfig,
auth::{authorize_mcp_server, is_authenticated}, auth::{authorize_mcp_server, is_authenticated},
@@ -628,7 +628,17 @@ async fn save_servers(
/// Initialize and return the secrets store. /// Initialize and return the secrets store.
async fn get_secrets_store() -> anyhow::Result<Arc<dyn SecretsStore + Send + Sync>> { async fn get_secrets_store() -> anyhow::Result<Arc<dyn SecretsStore + Send + Sync>> {
crate::cli::init_secrets_store().await let config = Config::from_env().await?;
let master_key = config.secrets.master_key().ok_or_else(|| {
anyhow::anyhow!(
"SECRETS_MASTER_KEY not set. Run 'ironclaw onboard' first or set it in .env"
)
})?;
let crypto = Arc::new(SecretsCrypto::new(master_key.clone())?);
Ok(crate::db::create_secrets_store(&config.database, crypto).await?)
} }
#[cfg(test)] #[cfg(test)]
+4 -45
View File
@@ -28,6 +28,8 @@ pub use config::{ConfigCommand, run_config_command};
pub use doctor::run_doctor_command; pub use doctor::run_doctor_command;
pub use mcp::{McpCommand, run_mcp_command}; pub use mcp::{McpCommand, run_mcp_command};
pub use memory::MemoryCommand; pub use memory::MemoryCommand;
#[cfg(feature = "postgres")]
pub use memory::run_memory_command;
pub use memory::run_memory_command_with_db; pub use memory::run_memory_command_with_db;
pub use pairing::{PairingCommand, run_pairing_command, run_pairing_command_with_store}; pub use pairing::{PairingCommand, run_pairing_command, run_pairing_command_with_store};
pub use registry::{RegistryCommand, run_registry_command}; pub use registry::{RegistryCommand, run_registry_command};
@@ -35,8 +37,6 @@ pub use service::{ServiceCommand, run_service_command};
pub use status::run_status_command; pub use status::run_status_command;
pub use tool::{ToolCommand, run_tool_command}; pub use tool::{ToolCommand, run_tool_command};
use std::sync::Arc;
use clap::{ColorChoice, Parser, Subcommand}; use clap::{ColorChoice, Parser, Subcommand};
#[derive(Parser, Debug)] #[derive(Parser, Debug)]
@@ -94,16 +94,12 @@ pub enum Command {
skip_auth: bool, skip_auth: bool,
/// Reconfigure channels only /// Reconfigure channels only
#[arg(long, conflicts_with_all = ["provider_only", "quick"])] #[arg(long, conflicts_with = "provider_only")]
channels_only: bool, channels_only: bool,
/// Reconfigure LLM provider and model only /// Reconfigure LLM provider and model only
#[arg(long, conflicts_with_all = ["channels_only", "quick"])] #[arg(long, conflicts_with = "channels_only")]
provider_only: bool, provider_only: bool,
/// Quick setup: auto-defaults everything except LLM provider and model
#[arg(long, conflicts_with_all = ["channels_only", "provider_only"])]
quick: bool,
}, },
/// Manage configuration settings /// Manage configuration settings
@@ -229,43 +225,6 @@ impl Cli {
} }
} }
/// Initialize a secrets store from environment config.
///
/// Shared helper for CLI subcommands (`mcp auth`, `tool auth`, etc.) that need
/// access to encrypted secrets without spinning up the full AppBuilder.
pub async fn init_secrets_store()
-> anyhow::Result<Arc<dyn crate::secrets::SecretsStore + Send + Sync>> {
let config = crate::config::Config::from_env().await?;
let master_key = config.secrets.master_key().ok_or_else(|| {
anyhow::anyhow!(
"SECRETS_MASTER_KEY not set. Run 'ironclaw onboard' first or set it in .env"
)
})?;
let crypto = Arc::new(crate::secrets::SecretsCrypto::new(master_key.clone())?);
Ok(crate::db::create_secrets_store(&config.database, crypto).await?)
}
/// Run the Memory CLI subcommand.
pub async fn run_memory_command(mem_cmd: &MemoryCommand) -> anyhow::Result<()> {
let config = crate::config::Config::from_env()
.await
.map_err(|e| anyhow::anyhow!("{}", e))?;
let session = crate::llm::create_session_manager(config.llm.session.clone()).await;
let embeddings = config
.embeddings
.create_provider(&config.llm.nearai.base_url, session);
let db: Arc<dyn crate::db::Database> = crate::db::connect_from_config(&config.database)
.await
.map_err(|e| anyhow::anyhow!("{}", e))?;
run_memory_command_with_db(mem_cmd.clone(), db, embeddings).await
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
+12 -2
View File
@@ -10,7 +10,8 @@ use clap::Subcommand;
use tokio::fs; use tokio::fs;
use crate::bootstrap::ironclaw_base_dir; use crate::bootstrap::ironclaw_base_dir;
use crate::secrets::{CreateSecretParams, SecretsStore}; use crate::config::Config;
use crate::secrets::{CreateSecretParams, SecretsCrypto, SecretsStore};
use crate::tools::wasm::{CapabilitiesFile, compute_binary_hash}; use crate::tools::wasm::{CapabilitiesFile, compute_binary_hash};
/// Default tools directory. /// Default tools directory.
@@ -551,7 +552,16 @@ fn validate_tool_name(name: &str) -> anyhow::Result<()> {
/// Initialize the secrets store from environment config. /// Initialize the secrets store from environment config.
async fn init_secrets_store() -> anyhow::Result<Arc<dyn SecretsStore + Send + Sync>> { async fn init_secrets_store() -> anyhow::Result<Arc<dyn SecretsStore + Send + Sync>> {
crate::cli::init_secrets_store().await let config = Config::from_env().await?;
let master_key = config.secrets.master_key().ok_or_else(|| {
anyhow::anyhow!(
"SECRETS_MASTER_KEY not set. Run 'ironclaw onboard' first or set it in .env"
)
})?;
let crypto = Arc::new(SecretsCrypto::new(master_key.clone())?);
Ok(crate::db::create_secrets_store(&config.database, crypto).await?)
} }
/// Configure authentication for a tool. /// Configure authentication for a tool.
+5 -6
View File
@@ -100,13 +100,13 @@ impl EmbeddingsConfig {
session: Arc<SessionManager>, session: Arc<SessionManager>,
) -> Option<Arc<dyn EmbeddingProvider>> { ) -> Option<Arc<dyn EmbeddingProvider>> {
if !self.enabled { if !self.enabled {
tracing::debug!("Embeddings disabled (set EMBEDDING_ENABLED=true to enable)"); tracing::info!("Embeddings disabled (set EMBEDDING_ENABLED=true to enable)");
return None; return None;
} }
match self.provider.as_str() { match self.provider.as_str() {
"nearai" => { "nearai" => {
tracing::debug!( tracing::info!(
"Embeddings enabled via NEAR AI (model: {}, dim: {})", "Embeddings enabled via NEAR AI (model: {}, dim: {})",
self.model, self.model,
self.dimension, self.dimension,
@@ -117,7 +117,7 @@ impl EmbeddingsConfig {
)) ))
} }
"ollama" => { "ollama" => {
tracing::debug!( tracing::info!(
"Embeddings enabled via Ollama (model: {}, url: {}, dim: {})", "Embeddings enabled via Ollama (model: {}, url: {}, dim: {})",
self.model, self.model,
self.ollama_base_url, self.ollama_base_url,
@@ -130,7 +130,7 @@ impl EmbeddingsConfig {
} }
_ => { _ => {
if let Some(api_key) = self.openai_api_key() { if let Some(api_key) = self.openai_api_key() {
tracing::debug!( tracing::info!(
"Embeddings enabled via OpenAI (model: {}, dim: {})", "Embeddings enabled via OpenAI (model: {}, dim: {})",
self.model, self.model,
self.dimension, self.dimension,
@@ -154,7 +154,6 @@ mod tests {
use super::*; use super::*;
use crate::config::helpers::ENV_MUTEX; use crate::config::helpers::ENV_MUTEX;
use crate::settings::{EmbeddingsSettings, Settings}; use crate::settings::{EmbeddingsSettings, Settings};
use crate::testing::credentials::*;
/// Clear all embedding-related env vars. /// Clear all embedding-related env vars.
fn clear_embedding_env() { fn clear_embedding_env() {
@@ -174,7 +173,7 @@ mod tests {
clear_embedding_env(); clear_embedding_env();
// SAFETY: Under ENV_MUTEX, no concurrent env access. // SAFETY: Under ENV_MUTEX, no concurrent env access.
unsafe { unsafe {
std::env::set_var("OPENAI_API_KEY", TEST_OPENAI_API_KEY_ISSUE_129); std::env::set_var("OPENAI_API_KEY", "sk-test-key-for-issue-129");
} }
let settings = Settings { let settings = Settings {
+7 -8
View File
@@ -389,7 +389,6 @@ mod tests {
use super::*; use super::*;
use crate::config::helpers::ENV_MUTEX; use crate::config::helpers::ENV_MUTEX;
use crate::settings::Settings; use crate::settings::Settings;
use crate::testing::credentials::*;
/// Clear all openai-compatible-related env vars. /// Clear all openai-compatible-related env vars.
fn clear_openai_compatible_env() { fn clear_openai_compatible_env() {
@@ -658,7 +657,7 @@ mod tests {
// SAFETY: Under ENV_MUTEX. // SAFETY: Under ENV_MUTEX.
unsafe { unsafe {
std::env::set_var("LLM_BACKEND", "open_ai"); std::env::set_var("LLM_BACKEND", "open_ai");
std::env::set_var("OPENAI_API_KEY", TEST_API_KEY); std::env::set_var("OPENAI_API_KEY", "test-key");
} }
let settings = Settings::default(); let settings = Settings::default();
@@ -792,7 +791,7 @@ mod tests {
clear_anthropic_env(); clear_anthropic_env();
// SAFETY: Under ENV_MUTEX. // SAFETY: Under ENV_MUTEX.
unsafe { unsafe {
std::env::set_var("ANTHROPIC_OAUTH_TOKEN", TEST_ANTHROPIC_OAUTH_TOKEN); std::env::set_var("ANTHROPIC_OAUTH_TOKEN", "sk-ant-oat01-test-token");
} }
let settings = Settings { let settings = Settings {
@@ -816,7 +815,7 @@ mod tests {
); );
assert_eq!( assert_eq!(
provider.oauth_token.as_ref().unwrap().expose_secret(), provider.oauth_token.as_ref().unwrap().expose_secret(),
TEST_ANTHROPIC_OAUTH_TOKEN "sk-ant-oat01-test-token"
); );
clear_anthropic_env(); clear_anthropic_env();
@@ -830,8 +829,8 @@ mod tests {
clear_anthropic_env(); clear_anthropic_env();
// SAFETY: Under ENV_MUTEX. // SAFETY: Under ENV_MUTEX.
unsafe { unsafe {
std::env::set_var("ANTHROPIC_API_KEY", TEST_ANTHROPIC_API_KEY); std::env::set_var("ANTHROPIC_API_KEY", "sk-ant-real-key");
std::env::set_var("ANTHROPIC_OAUTH_TOKEN", TEST_ANTHROPIC_OAUTH_TOKEN); std::env::set_var("ANTHROPIC_OAUTH_TOKEN", "sk-ant-oat01-test-token");
} }
let settings = Settings { let settings = Settings {
@@ -846,7 +845,7 @@ mod tests {
.api_key .api_key
.as_ref() .as_ref()
.map(|k| k.expose_secret().to_string()), .map(|k| k.expose_secret().to_string()),
Some(TEST_ANTHROPIC_API_KEY.to_string()), Some("sk-ant-real-key".to_string()),
"real API key should take priority over OAuth placeholder" "real API key should take priority over OAuth placeholder"
); );
assert!( assert!(
@@ -863,7 +862,7 @@ mod tests {
clear_anthropic_env(); clear_anthropic_env();
// SAFETY: Under ENV_MUTEX. // SAFETY: Under ENV_MUTEX.
unsafe { unsafe {
std::env::set_var("ANTHROPIC_OAUTH_TOKEN", TEST_ANTHROPIC_OAUTH_TOKEN); std::env::set_var("ANTHROPIC_OAUTH_TOKEN", "sk-ant-oat01-test-token");
} }
let settings = Settings { let settings = Settings {
-7
View File
@@ -14,7 +14,6 @@ mod heartbeat;
pub(crate) mod helpers; pub(crate) mod helpers;
mod hygiene; mod hygiene;
pub(crate) mod llm; pub(crate) mod llm;
pub mod relay;
mod routines; mod routines;
mod safety; mod safety;
mod sandbox; mod sandbox;
@@ -39,7 +38,6 @@ pub use self::embeddings::EmbeddingsConfig;
pub use self::heartbeat::HeartbeatConfig; pub use self::heartbeat::HeartbeatConfig;
pub use self::hygiene::HygieneConfig; pub use self::hygiene::HygieneConfig;
pub use self::llm::default_session_path; pub use self::llm::default_session_path;
pub use self::relay::RelayConfig;
pub use self::routines::RoutineConfig; pub use self::routines::RoutineConfig;
pub use self::safety::SafetyConfig; pub use self::safety::SafetyConfig;
pub use self::sandbox::{ClaudeCodeConfig, SandboxModeConfig}; pub use self::sandbox::{ClaudeCodeConfig, SandboxModeConfig};
@@ -87,9 +85,6 @@ pub struct Config {
pub skills: SkillsConfig, pub skills: SkillsConfig,
pub transcription: TranscriptionConfig, pub transcription: TranscriptionConfig,
pub observability: crate::observability::ObservabilityConfig, pub observability: crate::observability::ObservabilityConfig,
/// Channel-relay integration (Slack via external relay service).
/// Present only when both `CHANNEL_RELAY_URL` and `CHANNEL_RELAY_API_KEY` are set.
pub relay: Option<RelayConfig>,
} }
impl Config { impl Config {
@@ -162,7 +157,6 @@ impl Config {
}, },
transcription: TranscriptionConfig::default(), transcription: TranscriptionConfig::default(),
observability: crate::observability::ObservabilityConfig::default(), observability: crate::observability::ObservabilityConfig::default(),
relay: None,
} }
} }
@@ -316,7 +310,6 @@ impl Config {
observability: crate::observability::ObservabilityConfig { observability: crate::observability::ObservabilityConfig {
backend: std::env::var("OBSERVABILITY_BACKEND").unwrap_or_else(|_| "none".into()), backend: std::env::var("OBSERVABILITY_BACKEND").unwrap_or_else(|_| "none".into()),
}, },
relay: RelayConfig::from_env(),
}) })
} }
} }
-157
View File
@@ -1,157 +0,0 @@
//! Channel-relay service configuration.
use secrecy::SecretString;
/// Configuration for connecting to a channel-relay service.
#[derive(Clone)]
pub struct RelayConfig {
/// Base URL of the channel-relay service (e.g., `http://localhost:3001`).
pub url: String,
/// API key for authenticated channel-relay endpoints.
pub api_key: SecretString,
/// Override for the OAuth callback URL (e.g., a tunnel URL).
pub callback_url: Option<String>,
/// Override for the instance identifier.
pub instance_id: Option<String>,
/// HTTP request timeout in seconds (default: 30).
pub request_timeout_secs: u64,
/// SSE stream long-poll timeout in seconds (default: 86400 = 24 h).
pub stream_timeout_secs: u64,
/// Initial exponential backoff in milliseconds (default: 1000).
pub backoff_initial_ms: u64,
/// Maximum exponential backoff in milliseconds (default: 60000).
pub backoff_max_ms: u64,
}
impl std::fmt::Debug for RelayConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RelayConfig")
.field("url", &self.url)
.field("api_key", &"[REDACTED]")
.field("callback_url", &self.callback_url)
.field("instance_id", &self.instance_id)
.field("request_timeout_secs", &self.request_timeout_secs)
.field("stream_timeout_secs", &self.stream_timeout_secs)
.field("backoff_initial_ms", &self.backoff_initial_ms)
.field("backoff_max_ms", &self.backoff_max_ms)
.finish()
}
}
impl RelayConfig {
/// Load relay config from environment variables.
///
/// Returns `None` if either `CHANNEL_RELAY_URL` or `CHANNEL_RELAY_API_KEY`
/// is not set, making the relay integration opt-in.
pub fn from_env() -> Option<Self> {
Self::from_env_reader(|key| std::env::var(key).ok())
}
/// Build a config for tests without touching the process environment.
pub fn from_values(url: impl Into<String>, api_key: impl Into<String>) -> Self {
Self {
url: url.into(),
api_key: SecretString::from(api_key.into()),
callback_url: None,
instance_id: None,
request_timeout_secs: 30,
stream_timeout_secs: 86400,
backoff_initial_ms: 1000,
backoff_max_ms: 60000,
}
}
/// Internal constructor that reads values through a closure, enabling safe testing.
fn from_env_reader(env: impl Fn(&str) -> Option<String>) -> Option<Self> {
let url = env("CHANNEL_RELAY_URL")?;
let api_key = SecretString::from(env("CHANNEL_RELAY_API_KEY")?);
Some(Self {
url,
api_key,
callback_url: env("IRONCLAW_OAUTH_CALLBACK_URL"),
instance_id: env("IRONCLAW_INSTANCE_ID"),
request_timeout_secs: env("RELAY_REQUEST_TIMEOUT_SECS")
.and_then(|v| v.parse().ok())
.unwrap_or(30),
stream_timeout_secs: env("RELAY_STREAM_TIMEOUT_SECS")
.and_then(|v| v.parse().ok())
.unwrap_or(86400),
backoff_initial_ms: env("RELAY_BACKOFF_INITIAL_MS")
.and_then(|v| v.parse().ok())
.unwrap_or(1000),
backoff_max_ms: env("RELAY_BACKOFF_MAX_MS")
.and_then(|v| v.parse().ok())
.unwrap_or(60000),
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn from_env_reader_returns_none_when_unset() {
let config = RelayConfig::from_env_reader(|_| None);
assert!(config.is_none());
}
#[test]
fn from_env_reader_loads_defaults() {
let config = RelayConfig::from_env_reader(|key| match key {
"CHANNEL_RELAY_URL" => Some("http://localhost:3001".into()),
"CHANNEL_RELAY_API_KEY" => Some("test-key".into()),
_ => None,
})
.expect("config should be Some");
assert_eq!(config.url, "http://localhost:3001");
assert_eq!(config.request_timeout_secs, 30);
assert_eq!(config.stream_timeout_secs, 86400);
assert_eq!(config.backoff_initial_ms, 1000);
assert_eq!(config.backoff_max_ms, 60000);
assert!(config.callback_url.is_none());
assert!(config.instance_id.is_none());
}
#[test]
fn from_env_reader_loads_overrides() {
let config = RelayConfig::from_env_reader(|key| match key {
"CHANNEL_RELAY_URL" => Some("http://relay:3001".into()),
"CHANNEL_RELAY_API_KEY" => Some("secret".into()),
"IRONCLAW_OAUTH_CALLBACK_URL" => Some("https://tunnel.example.com".into()),
"IRONCLAW_INSTANCE_ID" => Some("my-instance".into()),
"RELAY_REQUEST_TIMEOUT_SECS" => Some("60".into()),
"RELAY_STREAM_TIMEOUT_SECS" => Some("43200".into()),
"RELAY_BACKOFF_INITIAL_MS" => Some("2000".into()),
"RELAY_BACKOFF_MAX_MS" => Some("120000".into()),
_ => None,
})
.expect("config should be Some");
assert_eq!(
config.callback_url.as_deref(),
Some("https://tunnel.example.com")
);
assert_eq!(config.instance_id.as_deref(), Some("my-instance"));
assert_eq!(config.request_timeout_secs, 60);
assert_eq!(config.stream_timeout_secs, 43200);
assert_eq!(config.backoff_initial_ms, 2000);
assert_eq!(config.backoff_max_ms, 120000);
}
#[test]
fn from_values_builds_with_defaults() {
let config = RelayConfig::from_values("http://localhost:3001", "key");
assert_eq!(config.url, "http://localhost:3001");
assert_eq!(config.request_timeout_secs, 30);
}
#[test]
fn debug_redacts_api_key() {
let config = RelayConfig::from_values("http://localhost:3001", "super-secret");
let debug = format!("{:?}", config);
assert!(debug.contains("[REDACTED]"));
assert!(!debug.contains("super-secret"));
}
}
+10 -17
View File
@@ -272,7 +272,6 @@ fn parse_oauth_access_token(json: &str) -> Option<String> {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use crate::config::sandbox::*; use crate::config::sandbox::*;
use crate::testing::credentials::*;
// ── SandboxModeConfig defaults ────────────────────────────────── // ── SandboxModeConfig defaults ──────────────────────────────────
@@ -406,12 +405,9 @@ mod tests {
#[test] #[test]
fn parse_oauth_token_valid() { fn parse_oauth_token_valid() {
let json = format!( let json = r#"{"claudeAiOauth": {"accessToken": "sk-ant-oat01-fake"}}"#;
r#"{{"claudeAiOauth": {{"accessToken": "{}"}}}}"#, let token = parse_oauth_access_token(json);
TEST_ANTHROPIC_OAUTH_BASIC assert_eq!(token, Some("sk-ant-oat01-fake".to_string()));
);
let token = parse_oauth_access_token(&json);
assert_eq!(token, Some(TEST_ANTHROPIC_OAUTH_BASIC.to_string()));
} }
#[test] #[test]
@@ -438,19 +434,16 @@ mod tests {
#[test] #[test]
fn parse_oauth_token_nested_extra_fields() { fn parse_oauth_token_nested_extra_fields() {
let json = format!( let json = r#"{
r#"{{ "claudeAiOauth": {
"claudeAiOauth": {{ "accessToken": "sk-ant-oat01-real-token",
"accessToken": "{}",
"refreshToken": "rt-abc", "refreshToken": "rt-abc",
"expiresAt": 1700000000 "expiresAt": 1700000000
}} }
}}"#, }"#;
TEST_ANTHROPIC_OAUTH_NESTED
);
assert_eq!( assert_eq!(
parse_oauth_access_token(&json), parse_oauth_access_token(json),
Some(TEST_ANTHROPIC_OAUTH_NESTED.to_string()) Some("sk-ant-oat01-real-token".to_string())
); );
} }
+8 -14
View File
@@ -30,9 +30,8 @@ impl JobStore for LibSqlBackend {
id, conversation_id, title, description, category, status, source, id, conversation_id, title, description, category, status, source,
user_id, user_id,
budget_amount, budget_token, bid_amount, estimated_cost, estimated_time_secs, budget_amount, budget_token, bid_amount, estimated_cost, estimated_time_secs,
actual_cost, repair_attempts, max_tokens, total_tokens_used, actual_cost, repair_attempts, created_at, started_at, completed_at
created_at, started_at, completed_at ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17, ?18)
) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17, ?18, ?19, ?20)
ON CONFLICT (id) DO UPDATE SET ON CONFLICT (id) DO UPDATE SET
title = excluded.title, title = excluded.title,
description = excluded.description, description = excluded.description,
@@ -43,8 +42,6 @@ impl JobStore for LibSqlBackend {
estimated_time_secs = excluded.estimated_time_secs, estimated_time_secs = excluded.estimated_time_secs,
actual_cost = excluded.actual_cost, actual_cost = excluded.actual_cost,
repair_attempts = excluded.repair_attempts, repair_attempts = excluded.repair_attempts,
max_tokens = excluded.max_tokens,
total_tokens_used = excluded.total_tokens_used,
started_at = excluded.started_at, started_at = excluded.started_at,
completed_at = excluded.completed_at completed_at = excluded.completed_at
"#, "#,
@@ -64,8 +61,6 @@ impl JobStore for LibSqlBackend {
estimated_time_secs, estimated_time_secs,
ctx.actual_cost.to_string(), ctx.actual_cost.to_string(),
ctx.repair_attempts as i64, ctx.repair_attempts as i64,
ctx.max_tokens as i64,
ctx.total_tokens_used as i64,
fmt_ts(&ctx.created_at), fmt_ts(&ctx.created_at),
fmt_opt_ts(&ctx.started_at), fmt_opt_ts(&ctx.started_at),
fmt_opt_ts(&ctx.completed_at), fmt_opt_ts(&ctx.completed_at),
@@ -83,8 +78,7 @@ impl JobStore for LibSqlBackend {
r#" r#"
SELECT id, conversation_id, title, description, category, status, user_id, SELECT id, conversation_id, title, description, category, status, user_id,
budget_amount, budget_token, bid_amount, estimated_cost, estimated_time_secs, budget_amount, budget_token, bid_amount, estimated_cost, estimated_time_secs,
actual_cost, repair_attempts, max_tokens, total_tokens_used, actual_cost, repair_attempts, created_at, started_at, completed_at
created_at, started_at, completed_at
FROM agent_jobs WHERE id = ?1 FROM agent_jobs WHERE id = ?1
"#, "#,
params![id.to_string()], params![id.to_string()],
@@ -117,12 +111,12 @@ impl JobStore for LibSqlBackend {
estimated_duration: estimated_time_secs estimated_duration: estimated_time_secs
.map(|s| std::time::Duration::from_secs(s as u64)), .map(|s| std::time::Duration::from_secs(s as u64)),
actual_cost: get_decimal(&row, 12), actual_cost: get_decimal(&row, 12),
max_tokens: get_i64(&row, 14) as u64, total_tokens_used: 0,
total_tokens_used: get_i64(&row, 15) as u64, max_tokens: 0,
repair_attempts: get_i64(&row, 13) as u32, repair_attempts: get_i64(&row, 13) as u32,
created_at: get_ts(&row, 16), created_at: get_ts(&row, 14),
started_at: get_opt_ts(&row, 17), started_at: get_opt_ts(&row, 15),
completed_at: get_opt_ts(&row, 18), completed_at: get_opt_ts(&row, 16),
transitions: Vec::new(), transitions: Vec::new(),
metadata: serde_json::Value::Null, metadata: serde_json::Value::Null,
extra_env: std::sync::Arc::new(std::collections::HashMap::new()), extra_env: std::sync::Arc::new(std::collections::HashMap::new()),
+1 -1
View File
@@ -167,7 +167,7 @@ impl RoutineStore for LibSqlBackend {
let mut rows = conn let mut rows = conn
.query( .query(
&format!( &format!(
"SELECT {} FROM routines WHERE enabled = 1 AND trigger_type IN ('event', 'system_event')", "SELECT {} FROM routines WHERE enabled = 1 AND trigger_type = 'event'",
ROUTINE_COLUMNS ROUTINE_COLUMNS
), ),
(), (),
+5 -21
View File
@@ -583,8 +583,7 @@ INSERT OR IGNORE INTO leak_detection_patterns (id, name, pattern, severity, acti
/// ///
/// Each entry is `(version, name, sql)`. Migrations are idempotent: the /// Each entry is `(version, name, sql)`. Migrations are idempotent: the
/// `_migrations` table tracks which versions have been applied. /// `_migrations` table tracks which versions have been applied.
pub const INCREMENTAL_MIGRATIONS: &[(i64, &str, &str)] = &[ pub const INCREMENTAL_MIGRATIONS: &[(i64, &str, &str)] = &[(
(
9, 9,
"flexible_embedding_dimension", "flexible_embedding_dimension",
// Rebuild memory_chunks to remove the fixed F32_BLOB(1536) type // Rebuild memory_chunks to remove the fixed F32_BLOB(1536) type
@@ -645,18 +644,7 @@ CREATE TRIGGER IF NOT EXISTS memory_chunks_fts_update AFTER UPDATE ON memory_chu
INSERT INTO memory_chunks_fts(rowid, content) VALUES (new._rowid, new.content); INSERT INTO memory_chunks_fts(rowid, content) VALUES (new._rowid, new.content);
END; END;
"#, "#,
), )];
(
12,
"job_token_budget",
// Add token budget tracking columns to agent_jobs.
// SQLite supports ALTER TABLE ADD COLUMN, so no table rebuild needed.
r#"
ALTER TABLE agent_jobs ADD COLUMN max_tokens INTEGER NOT NULL DEFAULT 0;
ALTER TABLE agent_jobs ADD COLUMN total_tokens_used INTEGER NOT NULL DEFAULT 0;
"#,
),
];
/// Run incremental migrations that haven't been applied yet. /// Run incremental migrations that haven't been applied yet.
/// ///
@@ -665,7 +653,6 @@ ALTER TABLE agent_jobs ADD COLUMN total_tokens_used INTEGER NOT NULL DEFAULT 0;
pub async fn run_incremental(conn: &libsql::Connection) -> Result<(), crate::error::DatabaseError> { pub async fn run_incremental(conn: &libsql::Connection) -> Result<(), crate::error::DatabaseError> {
use crate::error::DatabaseError; use crate::error::DatabaseError;
let mut applied_count = 0;
for &(version, name, sql) in INCREMENTAL_MIGRATIONS { for &(version, name, sql) in INCREMENTAL_MIGRATIONS {
// Check if already applied // Check if already applied
let mut rows = conn let mut rows = conn
@@ -682,6 +669,8 @@ pub async fn run_incremental(conn: &libsql::Connection) -> Result<(), crate::err
continue; // Already applied continue; // Already applied
} }
tracing::info!(version, name, "libSQL: applying incremental migration");
// Wrap migration + recording in a transaction for atomicity. // Wrap migration + recording in a transaction for atomicity.
// If the process crashes mid-migration, the transaction rolls back // If the process crashes mid-migration, the transaction rolls back
// and the migration will be retried on next startup. // and the migration will be retried on next startup.
@@ -713,12 +702,7 @@ pub async fn run_incremental(conn: &libsql::Connection) -> Result<(), crate::err
)) ))
})?; })?;
applied_count += 1; tracing::info!(version, name, "libSQL: migration applied successfully");
tracing::debug!(version, name, "libSQL: migration applied");
}
if applied_count > 0 {
tracing::info!("libSQL: applied {} incremental migrations", applied_count);
} }
Ok(()) Ok(())
+2 -33
View File
@@ -51,29 +51,6 @@ use crate::workspace::{SearchConfig, SearchResult};
pub async fn connect_from_config( pub async fn connect_from_config(
config: &crate::config::DatabaseConfig, config: &crate::config::DatabaseConfig,
) -> Result<Arc<dyn Database>, DatabaseError> { ) -> Result<Arc<dyn Database>, DatabaseError> {
let (db, _handles) = connect_with_handles(config).await?;
Ok(db)
}
/// Backend-specific handles retained after database connection.
///
/// These are needed by satellite stores (e.g., `SecretsStore`) that require
/// a backend-specific handle rather than the generic `Arc<dyn Database>`.
#[derive(Default)]
pub struct DatabaseHandles {
#[cfg(feature = "postgres")]
pub pg_pool: Option<deadpool_postgres::Pool>,
#[cfg(feature = "libsql")]
pub libsql_db: Option<Arc<::libsql::Database>>,
}
/// Connect to the database, run migrations, and return both the generic
/// `Database` trait object and the backend-specific handles.
pub async fn connect_with_handles(
config: &crate::config::DatabaseConfig,
) -> Result<(Arc<dyn Database>, DatabaseHandles), DatabaseError> {
let mut handles = DatabaseHandles::default();
match config.backend { match config.backend {
#[cfg(feature = "libsql")] #[cfg(feature = "libsql")]
crate::config::DatabaseBackend::LibSql => { crate::config::DatabaseBackend::LibSql => {
@@ -97,11 +74,7 @@ pub async fn connect_with_handles(
.map_err(|e| DatabaseError::Pool(e.to_string()))? .map_err(|e| DatabaseError::Pool(e.to_string()))?
}; };
backend.run_migrations().await?; backend.run_migrations().await?;
tracing::info!("libSQL database connected and migrations applied"); Ok(Arc::new(backend))
handles.libsql_db = Some(backend.shared_db());
Ok((Arc::new(backend) as Arc<dyn Database>, handles))
} }
#[cfg(feature = "postgres")] #[cfg(feature = "postgres")]
_ => { _ => {
@@ -109,11 +82,7 @@ pub async fn connect_with_handles(
.await .await
.map_err(|e| DatabaseError::Pool(e.to_string()))?; .map_err(|e| DatabaseError::Pool(e.to_string()))?;
pg.run_migrations().await?; pg.run_migrations().await?;
tracing::info!("PostgreSQL database connected and migrations applied"); Ok(Arc::new(pg))
handles.pg_pool = Some(pg.pool());
Ok((Arc::new(pg) as Arc<dyn Database>, handles))
} }
#[cfg(not(feature = "postgres"))] #[cfg(not(feature = "postgres"))]
_ => Err(DatabaseError::Pool( _ => Err(DatabaseError::Pool(
-1
View File
@@ -250,7 +250,6 @@ fn extract_source(source: &ExtensionSource) -> String {
ExtensionSource::Discovered { url } => url.clone(), ExtensionSource::Discovered { url } => url.clone(),
ExtensionSource::WasmDownload { wasm_url, .. } => wasm_url.clone(), ExtensionSource::WasmDownload { wasm_url, .. } => wasm_url.clone(),
ExtensionSource::WasmBuildable { source_dir, .. } => source_dir.clone(), ExtensionSource::WasmBuildable { source_dir, .. } => source_dir.clone(),
ExtensionSource::ChannelRelay { relay_url } => relay_url.clone(),
} }
} }
+27 -784
View File
File diff suppressed because it is too large Load Diff
-11
View File
@@ -37,8 +37,6 @@ pub enum ExtensionKind {
WasmTool, WasmTool,
/// WASM channel module with hot-activation support. /// WASM channel module with hot-activation support.
WasmChannel, WasmChannel,
/// External channel via channel-relay service (Slack, etc.).
ChannelRelay,
} }
impl std::fmt::Display for ExtensionKind { impl std::fmt::Display for ExtensionKind {
@@ -47,7 +45,6 @@ impl std::fmt::Display for ExtensionKind {
ExtensionKind::McpServer => write!(f, "mcp_server"), ExtensionKind::McpServer => write!(f, "mcp_server"),
ExtensionKind::WasmTool => write!(f, "wasm_tool"), ExtensionKind::WasmTool => write!(f, "wasm_tool"),
ExtensionKind::WasmChannel => write!(f, "wasm_channel"), ExtensionKind::WasmChannel => write!(f, "wasm_channel"),
ExtensionKind::ChannelRelay => write!(f, "channel_relay"),
} }
} }
} }
@@ -102,8 +99,6 @@ pub enum ExtensionSource {
}, },
/// Discovered online (not yet validated for a specific source type). /// Discovered online (not yet validated for a specific source type).
Discovered { url: String }, Discovered { url: String },
/// External channel via channel-relay service.
ChannelRelay { relay_url: String },
} }
/// Hint about what authentication method is needed. /// Hint about what authentication method is needed.
@@ -121,8 +116,6 @@ pub enum AuthHint {
CapabilitiesAuth, CapabilitiesAuth,
/// No authentication needed. /// No authentication needed.
None, None,
/// OAuth via channel-relay service.
ChannelRelayOAuth,
} }
/// Where a search result came from. /// Where a search result came from.
@@ -506,9 +499,6 @@ pub enum ExtensionError {
#[error("Activation failed: {0}")] #[error("Activation failed: {0}")]
ActivationFailed(String), ActivationFailed(String),
#[error("Authentication required")]
AuthRequired,
#[error("Installation failed: {0}")] #[error("Installation failed: {0}")]
InstallFailed(String), InstallFailed(String),
@@ -986,7 +976,6 @@ mod tests {
ExtensionError::Config("missing key".into()), ExtensionError::Config("missing key".into()),
"Config error: missing key", "Config error: missing key",
), ),
(ExtensionError::AuthRequired, "Authentication required"),
( (
ExtensionError::Other("something broke".into()), ExtensionError::Other("something broke".into()),
"something broke", "something broke",
+3 -59
View File
@@ -224,16 +224,8 @@ fn score_entry(entry: &RegistryEntry, tokens: &[String]) -> u32 {
} }
/// Well-known extensions that ship with ironclaw. /// Well-known extensions that ship with ironclaw.
/// fn builtin_entries() -> Vec<RegistryEntry> {
/// If `relay_url` is provided, a channel-relay Slack entry is included in the list. vec![
/// Pass `None` when the relay is not configured.
pub fn builtin_entries() -> Vec<RegistryEntry> {
builtin_entries_with_relay(std::env::var("CHANNEL_RELAY_URL").ok())
}
/// Well-known extensions, with an optional relay URL for the channel-relay entry.
pub fn builtin_entries_with_relay(relay_url: Option<String>) -> Vec<RegistryEntry> {
let mut entries = vec![
// -- MCP Servers -- // -- MCP Servers --
RegistryEntry { RegistryEntry {
name: "notion".to_string(), name: "notion".to_string(),
@@ -423,29 +415,7 @@ pub fn builtin_entries_with_relay(relay_url: Option<String>) -> Vec<RegistryEntr
// WASM channels (telegram, slack, discord, whatsapp) come from the embedded // WASM channels (telegram, slack, discord, whatsapp) come from the embedded
// registry catalog (registry/channels/*.json) with WasmDownload URLs pointing // registry catalog (registry/channels/*.json) with WasmDownload URLs pointing
// to GitHub release artifacts. See new_with_catalog() for merging. // to GitHub release artifacts. See new_with_catalog() for merging.
]; ]
// Conditionally add channel-relay entries when relay URL is configured
if let Some(relay_url) = relay_url {
entries.push(RegistryEntry {
name: crate::channels::relay::DEFAULT_RELAY_NAME.to_string(),
display_name: "Slack".to_string(),
kind: ExtensionKind::ChannelRelay,
description: "Connect Slack workspace via channel relay".to_string(),
keywords: vec![
"slack".into(),
"chat".into(),
"messaging".into(),
"relay".into(),
],
source: ExtensionSource::ChannelRelay { relay_url },
fallback_source: None,
auth_hint: AuthHint::ChannelRelayOAuth,
version: None,
});
}
entries
} }
#[cfg(test)] #[cfg(test)]
@@ -965,30 +935,4 @@ mod tests {
// The first catalog entry added is the channel. // The first catalog entry added is the channel.
assert_eq!(entry.unwrap().kind, ExtensionKind::WasmChannel); assert_eq!(entry.unwrap().kind, ExtensionKind::WasmChannel);
} }
#[test]
fn test_builtin_entries_with_relay_none_excludes_relay() {
let entries = super::builtin_entries_with_relay(None);
assert!(
!entries
.iter()
.any(|e| e.kind == ExtensionKind::ChannelRelay),
"No ChannelRelay entry when relay URL is None"
);
}
#[test]
fn test_builtin_entries_with_relay_some_includes_relay() {
let entries =
super::builtin_entries_with_relay(Some("http://relay.example.com".to_string()));
let relay = entries
.iter()
.find(|e| e.kind == ExtensionKind::ChannelRelay);
assert!(relay.is_some(), "ChannelRelay entry should be present");
if let ExtensionSource::ChannelRelay { relay_url } = &relay.unwrap().source {
assert_eq!(relay_url, "http://relay.example.com");
} else {
panic!("Expected ChannelRelay source");
}
}
} }
+6 -13
View File
@@ -151,9 +151,8 @@ impl Store {
id, conversation_id, title, description, category, status, source, id, conversation_id, title, description, category, status, source,
user_id, user_id,
budget_amount, budget_token, bid_amount, estimated_cost, estimated_time_secs, budget_amount, budget_token, bid_amount, estimated_cost, estimated_time_secs,
actual_cost, repair_attempts, max_tokens, total_tokens_used, actual_cost, repair_attempts, created_at, started_at, completed_at
created_at, started_at, completed_at ) VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15, $16, $17, $18)
) VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15, $16, $17, $18, $19, $20)
ON CONFLICT (id) DO UPDATE SET ON CONFLICT (id) DO UPDATE SET
title = EXCLUDED.title, title = EXCLUDED.title,
description = EXCLUDED.description, description = EXCLUDED.description,
@@ -164,8 +163,6 @@ impl Store {
estimated_time_secs = EXCLUDED.estimated_time_secs, estimated_time_secs = EXCLUDED.estimated_time_secs,
actual_cost = EXCLUDED.actual_cost, actual_cost = EXCLUDED.actual_cost,
repair_attempts = EXCLUDED.repair_attempts, repair_attempts = EXCLUDED.repair_attempts,
max_tokens = EXCLUDED.max_tokens,
total_tokens_used = EXCLUDED.total_tokens_used,
started_at = EXCLUDED.started_at, started_at = EXCLUDED.started_at,
completed_at = EXCLUDED.completed_at completed_at = EXCLUDED.completed_at
"#, "#,
@@ -185,8 +182,6 @@ impl Store {
&estimated_time_secs, &estimated_time_secs,
&ctx.actual_cost, &ctx.actual_cost,
&(ctx.repair_attempts as i32), &(ctx.repair_attempts as i32),
&(ctx.max_tokens as i64),
&(ctx.total_tokens_used as i64),
&ctx.created_at, &ctx.created_at,
&ctx.started_at, &ctx.started_at,
&ctx.completed_at, &ctx.completed_at,
@@ -206,8 +201,7 @@ impl Store {
r#" r#"
SELECT id, conversation_id, title, description, category, status, user_id, SELECT id, conversation_id, title, description, category, status, user_id,
budget_amount, budget_token, bid_amount, estimated_cost, estimated_time_secs, budget_amount, budget_token, bid_amount, estimated_cost, estimated_time_secs,
actual_cost, repair_attempts, max_tokens, total_tokens_used, actual_cost, repair_attempts, created_at, started_at, completed_at
created_at, started_at, completed_at
FROM agent_jobs WHERE id = $1 FROM agent_jobs WHERE id = $1
"#, "#,
&[&id], &[&id],
@@ -243,9 +237,8 @@ impl Store {
completed_at: row.get("completed_at"), completed_at: row.get("completed_at"),
transitions: Vec::new(), // Not loaded from DB for now transitions: Vec::new(), // Not loaded from DB for now
metadata: serde_json::Value::Null, metadata: serde_json::Value::Null,
max_tokens: row.get::<_, Option<i64>>("max_tokens").unwrap_or(0) as u64, total_tokens_used: 0,
total_tokens_used: row.get::<_, Option<i64>>("total_tokens_used").unwrap_or(0) max_tokens: 0,
as u64,
extra_env: std::sync::Arc::new(std::collections::HashMap::new()), extra_env: std::sync::Arc::new(std::collections::HashMap::new()),
http_interceptor: None, http_interceptor: None,
tool_output_stash: std::sync::Arc::new(tokio::sync::RwLock::new( tool_output_stash: std::sync::Arc::new(tokio::sync::RwLock::new(
@@ -1094,7 +1087,7 @@ impl Store {
let conn = self.conn().await?; let conn = self.conn().await?;
let rows = conn let rows = conn
.query( .query(
"SELECT * FROM routines WHERE enabled AND trigger_type IN ('event', 'system_event')", "SELECT * FROM routines WHERE enabled AND trigger_type = 'event'",
&[], &[],
) )
.await?; .await?;
+19 -4
View File
@@ -19,8 +19,7 @@ use crate::llm::costs;
use crate::llm::error::LlmError; use crate::llm::error::LlmError;
use crate::llm::provider::{ use crate::llm::provider::{
ChatMessage, CompletionRequest, CompletionResponse, FinishReason, LlmProvider, Role, ToolCall, ChatMessage, CompletionRequest, CompletionResponse, FinishReason, LlmProvider, Role, ToolCall,
ToolCompletionRequest, ToolCompletionResponse, strip_unsupported_completion_params, ToolCompletionRequest, ToolCompletionResponse,
strip_unsupported_tool_params,
}; };
const ANTHROPIC_API_URL: &str = "https://api.anthropic.com/v1/messages"; const ANTHROPIC_API_URL: &str = "https://api.anthropic.com/v1/messages";
@@ -81,12 +80,28 @@ impl AnthropicOAuthProvider {
/// Strip unsupported fields from a `CompletionRequest` in place. /// Strip unsupported fields from a `CompletionRequest` in place.
fn strip_unsupported_completion_params(&self, req: &mut CompletionRequest) { fn strip_unsupported_completion_params(&self, req: &mut CompletionRequest) {
strip_unsupported_completion_params(&self.unsupported_params, req); if self.unsupported_params.is_empty() {
return;
}
if self.unsupported_params.contains("temperature") {
req.temperature = None;
}
if self.unsupported_params.contains("max_tokens") {
req.max_tokens = None;
}
} }
/// Strip unsupported fields from a `ToolCompletionRequest` in place. /// Strip unsupported fields from a `ToolCompletionRequest` in place.
fn strip_unsupported_tool_params(&self, req: &mut ToolCompletionRequest) { fn strip_unsupported_tool_params(&self, req: &mut ToolCompletionRequest) {
strip_unsupported_tool_params(&self.unsupported_params, req); if self.unsupported_params.is_empty() {
return;
}
if self.unsupported_params.contains("temperature") {
req.temperature = None;
}
if self.unsupported_params.contains("max_tokens") {
req.max_tokens = None;
}
} }
fn api_url(&self) -> String { fn api_url(&self) -> String {

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