mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-25 14:53:34 +00:00
Compare commits
67
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
1d4f6f0fdc | ||
|
|
8c581e6240 | ||
|
|
b11b0331b4 | ||
|
|
3a2989d009 | ||
|
|
94d101924e | ||
|
|
a868b14221 | ||
|
|
a95f5ebb05 | ||
|
|
83950d11a4 | ||
|
|
764be8547f | ||
|
|
7de639e782 | ||
|
|
a5f88b32fd | ||
|
|
7d8576a464 | ||
|
|
f4b7309523 | ||
|
|
577e26eff4 | ||
|
|
bcbdc273a5 | ||
|
|
c541220ea4 | ||
|
|
14aadd3063 | ||
|
|
45923ef360 | ||
|
|
fcb152e408 | ||
|
|
e86b372fa6 | ||
|
|
63f140d391 | ||
|
|
ab0a2e05de | ||
|
|
290d925c7f | ||
|
|
d73e35cfb0 | ||
|
|
30d81fcdee | ||
|
|
d8dcc34319 | ||
|
|
652f30a826 | ||
|
|
98e9a40762 | ||
|
|
553c306c52 | ||
|
|
7fb2f47999 | ||
|
|
02f85a8ad5 | ||
|
|
9401ab0d58 | ||
|
|
7d1461fc74 | ||
|
|
605a4ba46e | ||
|
|
fe91ba2ab4 | ||
|
|
da2569bb77 | ||
|
|
732b3ecfeb | ||
|
|
461d7712e8 | ||
|
|
1c5117eded | ||
|
|
33b02eabb7 | ||
|
|
068ad2d4b7 | ||
|
|
56b7218897 | ||
|
|
200aed16cd | ||
|
|
4c0275bcdc | ||
|
|
272d31797e | ||
|
|
edff54b0b1 | ||
|
|
4d61d3eedf | ||
|
|
df3635d6be | ||
|
|
a20e19ab16 | ||
|
|
3b57d5bec9 | ||
|
|
11c5e25422 | ||
|
|
12ba79ffc3 | ||
|
|
d3cf637d4a | ||
|
|
b6cf2a6b73 | ||
|
|
9851f2a6ae | ||
|
|
8dc4ca5a98 | ||
|
|
9f71bd0d44 | ||
|
|
d144484b06 | ||
|
|
30790439ee | ||
|
|
424a0366a9 | ||
|
|
633b234e44 | ||
|
|
45ec691f4c | ||
|
|
cf96a3253c | ||
|
|
8fbb782090 | ||
|
|
3f22f4321d | ||
|
|
4ac78a5b1f | ||
|
|
ae89a52ac2 |
@@ -0,0 +1,303 @@
|
||||
---
|
||||
description: Full PR lifecycle — review, fix findings, address comments, quality gate, push, CI fix loop, merge
|
||||
disable-model-invocation: true
|
||||
allowed-tools: Bash(gh pr view:*), Bash(gh pr diff:*), Bash(gh pr comment:*), Bash(gh pr merge:*), Bash(gh pr checks:*), Bash(gh pr edit:*), Bash(gh pr list:*), Bash(gh pr checkout:*), Bash(gh api:*), Bash(gh repo view:*), Bash(gh run view:*), Bash(gh run watch:*), Bash(git diff:*), Bash(git log:*), Bash(git fetch:*), Bash(git checkout:*), Bash(git status:*), Bash(git branch:*), Bash(git add:*), Bash(git commit:*), Bash(git push:*), Bash(git merge:*), Bash(git rebase:*), Bash(cargo fmt:*), Bash(cargo clippy:*), Bash(cargo test:*), Bash(cargo check:*), Read, Edit, Write, Grep, Glob, Agent
|
||||
argument-hint: "<pr-number or url> [--fix] [--merge] [--review-only]"
|
||||
---
|
||||
|
||||
# PR Shepherd
|
||||
|
||||
Full PR lifecycle: review → fix → quality gate → push → CI → merge.
|
||||
|
||||
Parse `$ARGUMENTS`:
|
||||
- Extract PR number from bare number or `https://github.com/owner/repo/pull/123` URL.
|
||||
- Flags: `--fix` (auto-fix without asking), `--merge` (merge when CI green), `--review-only` (stop after review, don't fix).
|
||||
- If no PR number, detect from current branch: `gh pr list --head $(git branch --show-current) --json number --jq '.[0].number'`
|
||||
- If still nothing, stop and ask the user.
|
||||
|
||||
---
|
||||
|
||||
## Phase 1: Situational Awareness
|
||||
|
||||
Gather everything in parallel:
|
||||
|
||||
**PR metadata:**
|
||||
```
|
||||
gh pr view {number} --json number,title,body,author,baseRefName,headRefName,headRefOid,state,isDraft,mergeable,mergeStateStatus,files,additions,deletions,labels,reviewRequests
|
||||
```
|
||||
|
||||
**Diff:**
|
||||
```
|
||||
gh pr diff {number}
|
||||
gh pr diff {number} --name-only
|
||||
```
|
||||
|
||||
**CI status:**
|
||||
```
|
||||
gh pr checks {number} --json name,status,conclusion,detailsUrl
|
||||
```
|
||||
|
||||
**Review comments (human + bot):**
|
||||
```
|
||||
gh api --paginate repos/{owner}/{repo}/pulls/{number}/comments
|
||||
gh api --paginate repos/{owner}/{repo}/pulls/{number}/reviews
|
||||
```
|
||||
|
||||
Resolve `{owner}/{repo}`:
|
||||
```
|
||||
gh repo view --json owner,name --jq '"\(.owner.login)/\(.name)"'
|
||||
```
|
||||
|
||||
Save `headRefOid` — needed for posting line comments later.
|
||||
|
||||
**Assess the situation and print a status card:**
|
||||
|
||||
```
|
||||
PR #{number}: {title}
|
||||
Author: {author} Base: {base} ← {head}
|
||||
Size: +{additions} -{deletions} across {file_count} files
|
||||
CI: {PASS|FAIL|PENDING|NONE} Mergeable: {yes|no|conflict}
|
||||
Reviews: {N approved, N changes_requested, N comments-only, N bot-only}
|
||||
Unresolved comments: {N}
|
||||
Draft: {yes|no}
|
||||
```
|
||||
|
||||
**Decide the mode** based on situation:
|
||||
- **Has unresolved review comments** → Phase 2a (address comments first, then review remaining)
|
||||
- **No reviews yet / bot-only reviews** → Phase 2b (full deep review)
|
||||
- **CI failing, no review issues** → Phase 4 (jump to CI fix)
|
||||
- **Everything green + approved** → Phase 6 (ready to merge)
|
||||
|
||||
---
|
||||
|
||||
## Phase 2a: Address Existing Review Comments
|
||||
|
||||
For each unresolved review comment or review with CHANGES_REQUESTED:
|
||||
|
||||
1. **Read the referenced code** at the file and line mentioned. Never assess without reading.
|
||||
2. **Classify each comment:**
|
||||
- ✅ **Valid & unresolved** — needs a code fix
|
||||
- ✅ **Already fixed** — a later commit addressed it
|
||||
- ❌ **False positive** — explain why the code is correct
|
||||
- 🔧 **Nit** — optional improvement, not blocking
|
||||
|
||||
3. **Deduplicate** — bots (Copilot, Gemini) often post the same finding. Group by actual issue.
|
||||
|
||||
Present a table:
|
||||
|
||||
| # | Source | File:Line | Issue | Status | Planned Fix |
|
||||
|---|--------|-----------|-------|--------|-------------|
|
||||
|
||||
Wait for user confirmation (unless `--fix` flag set), then proceed to Phase 3.
|
||||
|
||||
---
|
||||
|
||||
## Phase 2b: Deep Review (6 Lenses)
|
||||
|
||||
Read EVERY changed file in full (not just diff hunks). For PRs touching >20 files, prioritize: service logic > handlers > types > tests > docs. Batch reads in parallel via Agent tool.
|
||||
|
||||
### IronClaw-specific checks (always)
|
||||
- No `.unwrap()` or `.expect()` in production code
|
||||
- Prefer `crate::` for cross-module imports (`super::` OK in tests/intra-module)
|
||||
- Error types use `thiserror`
|
||||
- If persistence touched, both backends updated (postgres.rs AND libsql/)
|
||||
- New tools implement `Tool` trait correctly and registered
|
||||
- External tool output passes through safety layer
|
||||
- Tool parameters redacted before logging/SSE
|
||||
- No byte-index slicing on external strings
|
||||
- Case-insensitive comparisons where needed
|
||||
|
||||
### Correctness
|
||||
Off-by-one, wrong operators, inverted conditions, unreachable code, type confusion, error propagation, broken invariants, TOCTOU races.
|
||||
|
||||
### Edge cases & failure handling
|
||||
Empty/None/zero-length input, external service failures, integer boundaries, malformed/adversarial input, partial failure handling.
|
||||
|
||||
### Security (assume adversarial actors)
|
||||
Auth/authz bypass, IDOR, injection (SQL/command/log/header), data leakage in logs/errors/API responses, resource exhaustion, replay/race conditions.
|
||||
|
||||
### Test coverage
|
||||
New public functions tested? Error paths tested? Edge cases covered? Existing tests still valid?
|
||||
|
||||
### Architecture
|
||||
Follows existing patterns? Unnecessary abstractions? Duplicated logic? Clean module dependencies?
|
||||
|
||||
**Present findings as a table:**
|
||||
|
||||
| # | Severity | Category | File:Line | Finding | Suggested Fix |
|
||||
|---|----------|----------|-----------|---------|---------------|
|
||||
|
||||
Severity: Critical > High > Medium > Low > Nit
|
||||
|
||||
If `--review-only` flag is set, post findings as GitHub comments (see Phase 2c) and STOP.
|
||||
|
||||
Otherwise, ask which findings to fix (default: all Critical + High + Medium). Then proceed to Phase 3.
|
||||
|
||||
---
|
||||
|
||||
## Phase 2c: Post Review Comments on GitHub
|
||||
|
||||
For each finding the user approved (or all Critical/High/Medium if `--fix`):
|
||||
|
||||
**Line-specific findings** — post as PR review comments:
|
||||
```
|
||||
gh api repos/{owner}/{repo}/pulls/{number}/comments \
|
||||
-f body="**{Severity}**: {finding}\n\n{explanation}\n\n**Suggested fix:** {suggestion}" \
|
||||
-f path="{file}" \
|
||||
-f commit_id="{headRefOid}" \
|
||||
-F line={line} \
|
||||
-f side="RIGHT"
|
||||
```
|
||||
|
||||
**Cross-cutting/architectural findings** — post as regular PR comment:
|
||||
```
|
||||
gh pr comment {number} --body "..."
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Phase 3: Fix
|
||||
|
||||
Checkout the PR branch if not already on it (handles fork PRs automatically):
|
||||
```
|
||||
gh pr checkout {number}
|
||||
```
|
||||
|
||||
**Implement fixes** for:
|
||||
1. All approved review comment fixes (from Phase 2a)
|
||||
2. All approved review findings (from Phase 2b)
|
||||
|
||||
Follow IronClaw conventions:
|
||||
- `thiserror` for errors
|
||||
- `crate::` imports
|
||||
- No `.unwrap()` in production
|
||||
- Both DB backends if persistence touched
|
||||
- Regression test for every bug fix (enforced by commit-msg hook; bypass only with `[skip-regression-check]` if genuinely not feasible)
|
||||
|
||||
After all fixes implemented, proceed to Phase 4.
|
||||
|
||||
---
|
||||
|
||||
## Phase 4: Quality Gate
|
||||
|
||||
Run the full IronClaw shipping checklist:
|
||||
|
||||
```bash
|
||||
cargo fmt
|
||||
```
|
||||
|
||||
```bash
|
||||
cargo clippy --all --benches --tests --examples --all-features
|
||||
```
|
||||
|
||||
```bash
|
||||
cargo test --lib
|
||||
```
|
||||
|
||||
If persistence changes are present, also verify feature isolation:
|
||||
```bash
|
||||
cargo check --no-default-features --features libsql
|
||||
cargo check --all-features
|
||||
```
|
||||
|
||||
**If any step fails:** fix the issue and re-run. Do NOT proceed past a failing step. Loop up to 3 times per step. If still failing after 3 attempts, report the failure and stop.
|
||||
|
||||
---
|
||||
|
||||
## Phase 5: Commit & Push
|
||||
|
||||
Stage changed files by name (never `git add -A` — it can include unintended files):
|
||||
```bash
|
||||
git add path/to/changed/file1 path/to/changed/file2
|
||||
git commit -m "{message}"
|
||||
```
|
||||
|
||||
Commit message format:
|
||||
- For review fixes: `fix: address review findings on PR #{number}`
|
||||
- For comment responses: `fix: address review comments on PR #{number}`
|
||||
- For CI fixes: `fix: resolve CI failures on PR #{number}`
|
||||
- Include specifics in the body (which findings/comments were addressed)
|
||||
|
||||
Push:
|
||||
```bash
|
||||
git push origin {headRefName}
|
||||
```
|
||||
|
||||
**Reply to addressed review comments on GitHub.** For each comment that was fixed, reply with the commit SHA and a brief description of what was done. For false positives, reply explaining why no change was needed.
|
||||
|
||||
---
|
||||
|
||||
## Phase 6: CI Monitor & Fix Loop
|
||||
|
||||
Wait briefly for CI to start, then poll (do NOT use `--watch` as it can hang indefinitely):
|
||||
```
|
||||
gh pr checks {number} --json name,status,conclusion
|
||||
```
|
||||
|
||||
Re-check every 30 seconds, up to 10 minutes. If still pending after 10 minutes, report status and ask the user whether to keep waiting.
|
||||
|
||||
**If CI passes** → proceed to Phase 7.
|
||||
|
||||
**If CI fails** (up to 3 fix attempts):
|
||||
|
||||
1. Identify the failing check:
|
||||
```
|
||||
gh run view {run_id} --log-failed
|
||||
```
|
||||
If `--log-failed` shows nothing useful:
|
||||
```
|
||||
gh run view {run_id} --log | tail -100
|
||||
```
|
||||
|
||||
2. Diagnose and fix the failure.
|
||||
3. Re-run Phase 4 (quality gate).
|
||||
4. Commit and push (Phase 5).
|
||||
5. Go back to top of Phase 6.
|
||||
|
||||
**After 3 failed CI fix attempts:** Report what's failing and why, then stop. Don't keep looping.
|
||||
|
||||
---
|
||||
|
||||
## Phase 7: Merge Decision
|
||||
|
||||
Print final status:
|
||||
```
|
||||
PR #{number}: {title}
|
||||
CI: ✅ PASS
|
||||
Reviews: {summary}
|
||||
Findings fixed: {N}
|
||||
Comments addressed: {N}
|
||||
Commits added: {N}
|
||||
```
|
||||
|
||||
**Auto-merge conditions** (if `--merge` flag or user confirms):
|
||||
- CI is passing
|
||||
- No unresolved CHANGES_REQUESTED reviews
|
||||
- PR is not draft
|
||||
- PR is mergeable (no conflicts)
|
||||
|
||||
If all conditions met, ask the user for merge strategy:
|
||||
|
||||
"CI is green. Merge this PR? [squash/rebase/merge/no]"
|
||||
|
||||
Then execute:
|
||||
```
|
||||
gh pr merge {number} --{strategy} --delete-branch
|
||||
```
|
||||
|
||||
If any condition NOT met, report what's blocking and let the user decide.
|
||||
|
||||
---
|
||||
|
||||
## Rules
|
||||
|
||||
- **Read before judging.** Never comment on code you haven't read in full. Verify line numbers.
|
||||
- **Be specific.** "Line 42 returns 404 but should return 400 because X" not "this might have issues."
|
||||
- **Fix the pattern, not just the instance.** When fixing a bug, grep for the same pattern across `src/`.
|
||||
- **Respect the commit-msg hook.** Bug fixes need regression tests. Use `[skip-regression-check]` only if genuinely not feasible.
|
||||
- **Don't over-fix.** Only change what was flagged. Don't refactor surrounding code or add improvements beyond the review scope.
|
||||
- **Credit original authors.** If taking over someone else's PR, credit them in commits and comments.
|
||||
- **No secrets in comments.** Never include customer data, credentials, or PII in GitHub comments.
|
||||
- **Distinguish certainty.** "This IS a bug" vs "This COULD be a bug if X." Be honest.
|
||||
- **Round up severity when uncertain.** Cheaper to dismiss a false alarm than miss a real bug.
|
||||
- **Parallel where possible.** Use Agent tool for parallel file reads on large PRs. Batch `gh api` calls.
|
||||
@@ -0,0 +1,63 @@
|
||||
---
|
||||
paths:
|
||||
- "src/db/**"
|
||||
- "src/history/**"
|
||||
- "migrations/**"
|
||||
---
|
||||
# Database Rules
|
||||
|
||||
Dual-backend persistence: PostgreSQL + libSQL/Turso. **All new persistence features must support both backends.**
|
||||
|
||||
See `src/db/CLAUDE.md` for full schema, dialect differences, and libSQL limitations.
|
||||
|
||||
## Adding a New Operation
|
||||
|
||||
1. Decide which sub-trait it belongs to (`ConversationStore`, `JobStore`, `SandboxStore`, `RoutineStore`, `ToolFailureStore`, `SettingsStore`, `WorkspaceStore`) or create a new one
|
||||
2. Add the async method signature to that sub-trait in `src/db/mod.rs`
|
||||
3. Implement in `src/db/postgres.rs` (delegate to `Store`/`Repository`)
|
||||
4. Implement in `src/db/libsql/<module>.rs` (use `self.connect().await?` per operation)
|
||||
5. Add migration if needed:
|
||||
- PostgreSQL: new `migrations/VN__description.sql`
|
||||
- libSQL: add `CREATE TABLE IF NOT EXISTS` to `libsql_migrations.rs`
|
||||
6. Test feature isolation:
|
||||
```bash
|
||||
cargo check # postgres (default)
|
||||
cargo check --no-default-features --features libsql # libsql only
|
||||
cargo check --all-features # both
|
||||
```
|
||||
|
||||
## SQL Dialect Translation Checklist
|
||||
|
||||
When writing SQL for both backends, translate these types:
|
||||
|
||||
| PostgreSQL | libSQL |
|
||||
|-----------|--------|
|
||||
| `UUID` | `TEXT` |
|
||||
| `TIMESTAMPTZ` | `TEXT` (ISO-8601, write with `fmt_ts()`, read with `get_ts()`) |
|
||||
| `JSONB` | `TEXT` (JSON string) |
|
||||
| `BOOLEAN` | `INTEGER` (0/1 -- use `get_i64(row, idx) != 0` to read) |
|
||||
| `NUMERIC` | `TEXT` (preserves `rust_decimal` precision) |
|
||||
| `TEXT[]` | `TEXT` (JSON-encoded array) |
|
||||
| `VECTOR` | `BLOB` (flexible dimensions; vector index dropped, brute-force search fallback) |
|
||||
| `jsonb_set(col, '{key}', val)` | `json_patch(col, '{"key": val}')` -- replaces top-level keys entirely, cannot do partial nested updates |
|
||||
| `DEFAULT NOW()` | `DEFAULT (datetime('now'))` |
|
||||
| `tsvector` + `ts_rank_cd` | FTS5 virtual table + sync triggers |
|
||||
|
||||
## Schema Translation Beyond DDL
|
||||
|
||||
Don't just translate `CREATE TABLE`. Also check:
|
||||
- **Indexes** -- diff `CREATE INDEX` statements between backends
|
||||
- **Seed data** -- check for `INSERT INTO` in migrations (e.g., `leak_detection_patterns`)
|
||||
- **Triggers** -- PostgreSQL functions vs SQLite triggers (no stored procs in SQLite)
|
||||
|
||||
## Transaction Safety
|
||||
|
||||
Multi-step operations (INSERT+INSERT, UPDATE+DELETE, read-modify-write) MUST be wrapped in a transaction. Ask: "If this crashes between step N and N+1, is the database consistent?" If not, wrap in a transaction. Applies to both backends.
|
||||
|
||||
## libSQL Connection Model
|
||||
|
||||
`LibSqlBackend::connect()` creates a fresh connection per operation with `PRAGMA busy_timeout = 5000`. This is intentional -- no pool exists. Never hold connections open across `await` points. Satellite stores (`LibSqlSecretsStore`, `LibSqlWasmToolStore`) receive `Arc<LibSqlDatabase>` via `shared_db()` and call `.connect()` themselves -- never pass a live `Connection`.
|
||||
|
||||
## Fix the Pattern, Not the Instance
|
||||
|
||||
When fixing a bug in one backend's SQL, always grep for the same pattern in the other. A fix to `postgres.rs` that doesn't also fix `libsql/jobs.rs` is half a fix. Same applies to satellite stores.
|
||||
@@ -0,0 +1,48 @@
|
||||
---
|
||||
paths:
|
||||
- "src/**/*.rs"
|
||||
---
|
||||
# Review & Fix Discipline
|
||||
|
||||
Hard-won lessons from code review -- follow these when fixing bugs or addressing review feedback.
|
||||
|
||||
**Fix the pattern, not just the instance:** When a reviewer flags a bug (e.g., TOCTOU race in INSERT + SELECT-back), search the entire codebase for all instances of that same pattern. A fix in `SecretsStore::create()` that doesn't also fix `WasmToolStore::store()` is half a fix.
|
||||
|
||||
**Propagate architectural fixes to satellite types:** If a core type changes its concurrency model (e.g., `LibSqlBackend` switches to connection-per-operation), every type that was handed a resource from the old model must also be updated. Grep for the old type across the codebase.
|
||||
|
||||
**Schema translation is more than DDL:** When translating a database schema between backends (PostgreSQL to libSQL, etc.), check for:
|
||||
- **Indexes** -- diff `CREATE INDEX` statements between the two schemas
|
||||
- **Seed data** -- check for `INSERT INTO` in migrations (e.g., `leak_detection_patterns`)
|
||||
- **Semantic differences** -- document where SQL functions behave differently (e.g., `json_patch` vs `jsonb_set`)
|
||||
|
||||
**Feature flag testing:** When adding feature-gated code, test compilation with each feature in isolation:
|
||||
```bash
|
||||
cargo check # default features
|
||||
cargo check --no-default-features --features libsql # libsql only
|
||||
cargo check --all-features # all features
|
||||
```
|
||||
|
||||
**Regression test with every fix:** Every bug fix must include a test that would have caught the bug. Add a `#[test]` or `#[tokio::test]` that reproduces the original failure. Exempt: changes limited to `src/channels/web/static/` or `.md` files. Use `[skip-regression-check]` in commit message or PR label if genuinely not feasible. The `commit-msg` hook and CI workflow enforce this automatically.
|
||||
|
||||
**Zero clippy warnings policy:** Fix ALL clippy warnings before committing, including pre-existing ones in files you didn't change. Never leave warnings behind.
|
||||
|
||||
**Transaction safety:** Multi-step database operations (INSERT+INSERT, UPDATE+DELETE, read-then-write) MUST be wrapped in a transaction. Never assume sequential calls are atomic. This applies to both postgres and libsql backends.
|
||||
|
||||
**UTF-8 string safety:** Never use byte-index slicing (`&s[..n]`) on user-supplied or external strings -- it panics on multi-byte characters. Use `is_char_boundary()` or `char_indices()`. Grep for `[..` in changed files.
|
||||
|
||||
**Case-insensitive comparisons:** When comparing user-supplied strings (file paths, media types, extension names), normalize to lowercase with `.to_ascii_lowercase()`. Path comparisons must be case-insensitive on macOS/Windows.
|
||||
|
||||
**Decorator/wrapper trait delegation:** When adding a new method to `LlmProvider` (or any trait with decorator wrappers), update ALL wrapper types to delegate. Grep for `impl LlmProvider for` to find all implementations. Test through the full provider chain.
|
||||
|
||||
**Sensitive data in logs & events:** Tool parameters and outputs MUST be redacted before logging or broadcasting via SSE/WebSocket. Use `redact_params()` before any `tracing::info!`, `JobEvent`, or SSE emission that includes tool call data.
|
||||
|
||||
**Test temporary files:** Use the `tempfile` crate. Never hardcode `/tmp/...` paths.
|
||||
|
||||
**Trust boundaries in multi-process architecture:** Data from worker containers is untrusted. The orchestrator MUST validate: tool domain, nesting depth (server-side tracking), and parameter sensitivity.
|
||||
|
||||
**Mechanical verification before committing:**
|
||||
- `cargo clippy --all --benches --tests --examples --all-features` -- zero warnings
|
||||
- `grep -rnE '\.unwrap\(|\.expect\(' <files>` -- no panics in production
|
||||
- `grep -rn 'super::' <files>` -- prefer `crate::` for cross-module imports (`super::` OK in tests/intra-module)
|
||||
- If you fixed a pattern bug, `grep` for other instances across `src/`
|
||||
- Run `scripts/pre-commit-safety.sh` to catch UTF-8, case-sensitivity, hardcoded /tmp, and logging issues
|
||||
@@ -0,0 +1,34 @@
|
||||
---
|
||||
paths:
|
||||
- "src/safety/**"
|
||||
- "src/sandbox/**"
|
||||
- "src/secrets/**"
|
||||
- "src/tools/wasm/**"
|
||||
---
|
||||
# Safety Layer & Sandbox Rules
|
||||
|
||||
## Safety Layer
|
||||
|
||||
All external tool output passes through `SafetyLayer`:
|
||||
1. **Sanitizer** - Detects injection patterns, escapes dangerous content
|
||||
2. **Validator** - Checks length, encoding, forbidden patterns
|
||||
3. **Policy** - Rules with severity (Critical/High/Medium/Low) and actions (Block/Warn/Review/Sanitize)
|
||||
4. **Leak Detector** - Scans for 15+ secret patterns at two points: tool output before LLM, and LLM responses before user
|
||||
|
||||
Tool outputs are wrapped in `<tool_output>` XML before reaching the LLM.
|
||||
|
||||
## Shell Environment Scrubbing
|
||||
|
||||
The shell tool scrubs sensitive env vars before executing commands. The sanitizer detects command injection patterns (chained commands, subshells, path traversal).
|
||||
|
||||
## Sandbox Policies
|
||||
|
||||
| Policy | Filesystem | Network |
|
||||
|--------|-----------|---------|
|
||||
| ReadOnly | Read-only workspace | Allowlisted domains |
|
||||
| WorkspaceWrite | Read-write workspace | Allowlisted domains |
|
||||
| FullAccess | Full filesystem | Unrestricted |
|
||||
|
||||
## Zero-Exposure Credential Model
|
||||
|
||||
Secrets are stored encrypted on the host and injected into HTTP requests by the proxy at transit time. Container processes never see raw credential values.
|
||||
@@ -0,0 +1,56 @@
|
||||
---
|
||||
paths:
|
||||
- "src/skills/**"
|
||||
- "skills/**"
|
||||
---
|
||||
# Skills System
|
||||
|
||||
SKILL.md files extend the agent's prompt with domain-specific instructions. Each skill is a YAML frontmatter block (metadata, activation criteria, required tools) followed by a markdown body injected into the LLM context.
|
||||
|
||||
## Trust Model
|
||||
|
||||
| Trust Level | Source | Tool Access |
|
||||
|-------------|--------|-------------|
|
||||
| **Trusted** | User-placed in `~/.ironclaw/skills/` or workspace `skills/` | All tools available to the agent |
|
||||
| **Installed** | Downloaded from ClawHub registry (`~/.ironclaw/installed_skills/`) | Read-only tools only (no shell, file write, HTTP) |
|
||||
|
||||
## SKILL.md Format
|
||||
|
||||
```yaml
|
||||
---
|
||||
name: my-skill
|
||||
version: 0.1.0
|
||||
description: Does something useful
|
||||
activation:
|
||||
patterns:
|
||||
- "deploy to.*production"
|
||||
keywords:
|
||||
- "deployment"
|
||||
exclude_keywords:
|
||||
- "rollback"
|
||||
tags:
|
||||
- "devops"
|
||||
max_context_tokens: 2000
|
||||
metadata:
|
||||
openclaw:
|
||||
requires:
|
||||
bins: [docker, kubectl]
|
||||
env: [KUBECONFIG]
|
||||
---
|
||||
|
||||
# Skill instructions here...
|
||||
```
|
||||
|
||||
## Selection Pipeline
|
||||
|
||||
1. **Gating** -- Check binary/env/config requirements; skip skills whose prerequisites are missing
|
||||
2. **Scoring** -- Deterministic scoring: keywords (10/5 pts, cap 30) + patterns (20 pts, cap 40) + tags (3 pts, cap 15). `exclude_keywords` veto (score = 0 if any present)
|
||||
3. **Budget** -- Select top-scoring skills within `SKILLS_MAX_TOKENS` prompt budget
|
||||
4. **Attenuation** -- Minimum trust across active skills determines tool ceiling; installed skills lose dangerous tools
|
||||
|
||||
## Skill Tools
|
||||
|
||||
- `skill_list` -- List all discovered skills with trust level and status
|
||||
- `skill_search` -- Search ClawHub registry for available skills
|
||||
- `skill_install` -- Download and install a skill from ClawHub
|
||||
- `skill_remove` -- Remove an installed skill
|
||||
@@ -0,0 +1,25 @@
|
||||
---
|
||||
paths:
|
||||
- "src/**/*.rs"
|
||||
- "tests/**"
|
||||
---
|
||||
# Testing Rules
|
||||
|
||||
## Test Tiers
|
||||
|
||||
| Tier | Command | External deps |
|
||||
|------|---------|---------------|
|
||||
| Unit | `cargo test` | None |
|
||||
| Integration | `cargo test --features integration` | Running PostgreSQL |
|
||||
| Live | `cargo test --features integration -- --ignored` | PostgreSQL + LLM API keys |
|
||||
|
||||
Run `bash scripts/check-boundaries.sh` to verify test tier gating.
|
||||
|
||||
## Key Patterns
|
||||
|
||||
- Unit tests in `mod tests {}` at the bottom of each file
|
||||
- Async tests with `#[tokio::test]`
|
||||
- No mocks, prefer real implementations or stubs
|
||||
- Use `tempfile` crate for test directories, never hardcode `/tmp/`
|
||||
- Regression test with every bug fix (enforced by commit-msg hook)
|
||||
- Integration tests (`--test workspace_integration`) require PostgreSQL; skipped if DB is unreachable
|
||||
@@ -0,0 +1,39 @@
|
||||
---
|
||||
paths:
|
||||
- "src/tools/**"
|
||||
- "tools-src/**"
|
||||
---
|
||||
# Tool Architecture
|
||||
|
||||
**Keep tool-specific logic out of the main agent codebase.** The main agent provides generic infrastructure; tools are self-contained units that declare requirements through `<name>.capabilities.json` sidecar files (in dev mode: `tools-src/<name>/<name>-tool.capabilities.json`).
|
||||
|
||||
Tools can be WASM (sandboxed, credential-injected, single binary) or MCP servers (ecosystem, any language, no sandbox). Both are first-class via `ironclaw tool install`.
|
||||
|
||||
See `src/tools/README.md` for full architecture, adding new tools, auth JSON examples, and WASM vs MCP decision guide.
|
||||
|
||||
## Tool Implementation Pattern
|
||||
|
||||
```rust
|
||||
#[async_trait]
|
||||
impl Tool for MyTool {
|
||||
fn name(&self) -> &str { "my_tool" }
|
||||
fn description(&self) -> &str { "Does something useful" }
|
||||
fn parameters_schema(&self) -> serde_json::Value {
|
||||
serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"param": { "type": "string", "description": "A parameter" }
|
||||
},
|
||||
"required": ["param"]
|
||||
})
|
||||
}
|
||||
async fn execute(&self, params: serde_json::Value, ctx: &JobContext)
|
||||
-> Result<ToolOutput, ToolError>
|
||||
{
|
||||
let start = std::time::Instant::now();
|
||||
// ... do work ...
|
||||
Ok(ToolOutput::text("result", start.elapsed()))
|
||||
}
|
||||
fn requires_sanitization(&self) -> bool { true } // External data
|
||||
}
|
||||
```
|
||||
@@ -5,6 +5,19 @@ DATABASE_POOL_SIZE=10
|
||||
# LLM Provider
|
||||
# LLM_BACKEND=nearai # default
|
||||
# Possible values: nearai, ollama, openai_compatible, openai, anthropic, tinfoil
|
||||
# LLM_REQUEST_TIMEOUT_SECS=120 # Increase for local LLMs (Ollama, vLLM, LM Studio)
|
||||
|
||||
# === Anthropic Direct ===
|
||||
# Two auth modes:
|
||||
# 1. API key: Set ANTHROPIC_API_KEY (from console.anthropic.com/settings/keys)
|
||||
# 2. OAuth token: Set ANTHROPIC_OAUTH_TOKEN (from `claude login`)
|
||||
# OAuth tokens use Authorization: Bearer instead of x-api-key header.
|
||||
# ANTHROPIC_API_KEY=sk-ant-...
|
||||
# ANTHROPIC_OAUTH_TOKEN=sk-ant-oat01-... # from `claude login` credentials
|
||||
# ANTHROPIC_MODEL=claude-sonnet-4-20250514
|
||||
|
||||
# === OpenAI Direct ===
|
||||
# OPENAI_API_KEY=sk-...
|
||||
|
||||
# === NEAR AI (Chat Completions API) ===
|
||||
# Two auth modes:
|
||||
@@ -57,6 +70,17 @@ NEARAI_AUTH_URL=https://private.near.ai
|
||||
# LLM_BASE_URL=https://api.fireworks.ai/inference/v1
|
||||
# LLM_API_KEY=fw_...
|
||||
|
||||
# === Anthropic Direct ===
|
||||
# LLM_BACKEND=anthropic
|
||||
# ANTHROPIC_MODEL=claude-sonnet-4-6
|
||||
# ANTHROPIC_API_KEY=sk-ant-...
|
||||
# ANTHROPIC_BASE_URL=https://api.anthropic.com # default
|
||||
# Prompt cache retention — controls Anthropic server-side prompt caching:
|
||||
# none = disabled (no cache_control injected)
|
||||
# short = 5-minute TTL, 1.25× (125%) write surcharge (default)
|
||||
# long = 1-hour TTL, 2.0× (200%) write surcharge
|
||||
# ANTHROPIC_CACHE_RETENTION=short
|
||||
|
||||
# For full provider setup guide see docs/LLM_PROVIDERS.md
|
||||
|
||||
# Channel Configuration
|
||||
@@ -91,6 +115,8 @@ AGENT_NAME=ironclaw
|
||||
AGENT_MAX_PARALLEL_JOBS=5
|
||||
AGENT_JOB_TIMEOUT_SECS=3600
|
||||
AGENT_STUCK_THRESHOLD_SECS=300
|
||||
# Maximum tokens per job (0 = unlimited, also settable via settings.json agent.max_tokens_per_job)
|
||||
# AGENT_MAX_TOKENS_PER_JOB=0
|
||||
# Enable planning phase before tool execution (default: true)
|
||||
AGENT_USE_PLANNING=true
|
||||
|
||||
|
||||
Symlink
+1
@@ -0,0 +1 @@
|
||||
../scripts/commit-msg-regression.sh
|
||||
Executable
+24
@@ -0,0 +1,24 @@
|
||||
#!/usr/bin/env bash
|
||||
set -euo pipefail
|
||||
|
||||
# Pre-commit hook: run version bump checks when WIT or extension sources change.
|
||||
# Install: git config core.hooksPath .githooks
|
||||
|
||||
# Only run the check if relevant files are staged
|
||||
STAGED=$(git diff --cached --name-only)
|
||||
|
||||
NEEDS_CHECK=false
|
||||
if echo "$STAGED" | grep -qE '^wit/|^channels-src/|^tools-src/'; then
|
||||
NEEDS_CHECK=true
|
||||
fi
|
||||
|
||||
if $NEEDS_CHECK; then
|
||||
echo "pre-commit: checking version bumps..."
|
||||
if ! ./scripts/check-version-bumps.sh; then
|
||||
echo ""
|
||||
echo "Commit blocked: version bump check failed."
|
||||
echo "Bump versions in the relevant registry JSON and/or WIT package declaration."
|
||||
echo "To bypass: git commit --no-verify"
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
@@ -0,0 +1,50 @@
|
||||
## 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) -->
|
||||
@@ -0,0 +1,100 @@
|
||||
name: Claude Code Review
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
types: [labeled]
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
pull-requests: write
|
||||
issues: write
|
||||
id-token: write
|
||||
|
||||
concurrency:
|
||||
group: claude-review-${{ github.event.pull_request.number || github.run_id }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
review:
|
||||
name: Claude Code Review
|
||||
if: contains(github.event.pull_request.labels.*.name, 'staging-promotion')
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
fetch-depth: 0
|
||||
|
||||
- name: Run Claude Code review
|
||||
uses: anthropics/claude-code-action@v1
|
||||
with:
|
||||
anthropic_api_key: ${{ secrets.ANTHROPIC_API_KEY }}
|
||||
allowed_bots: "ironclaw-ci[bot]"
|
||||
claude_args: "--max-turns 50 --model claude-haiku-4-5-20251001 --allowedTools 'Bash(gh pr comment:*),Bash(gh pr diff:*),Bash(gh pr view:*),Bash(gh pr list:*),Bash(gh issue view:*),Bash(gh issue list:*),Bash(gh search:*),Bash(git blame:*),Bash(git log:*),Bash(git diff:*)'"
|
||||
prompt: |
|
||||
Code review this pull request. Follow these steps precisely:
|
||||
|
||||
1. Use a Haiku agent to find relevant CLAUDE.md files: the root CLAUDE.md
|
||||
and any CLAUDE.md files in directories whose files this PR modifies.
|
||||
|
||||
2. Use a Haiku agent to summarize the PR change (use `gh pr diff`).
|
||||
|
||||
3. Launch 4 parallel agents to review the change independently. Each agent should
|
||||
read the PR diff with `gh pr diff` and the full source files for changed
|
||||
code, then return a list of issues found:
|
||||
|
||||
Agent 1 — Security & Safety
|
||||
Check for: command injection, path traversal, SSRF, XSS, auth bypass,
|
||||
secrets in logs, .unwrap()/.expect() in production code (not tests),
|
||||
race conditions, TOCTOU, unsafe blocks, panics in async, unbounded allocations.
|
||||
|
||||
Agent 2 — Architecture & Patterns
|
||||
Check for: extensible design (traits/enums over nested conditionals),
|
||||
clean abstractions, proper error types (thiserror), CLAUDE.md compliance,
|
||||
type-driven design over stringly-typed code, DRY violations.
|
||||
|
||||
Agent 3 — Bug Scan
|
||||
Shallow diff-only scan for obvious bugs: logic errors, off-by-one,
|
||||
missing error handling, division by zero, incorrect return values.
|
||||
Ignore nitpicks and likely false positives. Do NOT read extra context
|
||||
beyond the diff — focus only on the changes.
|
||||
|
||||
Agent 4 — Performance & Production
|
||||
Check for: blocking in async, N+1 queries, unbounded loops, missing
|
||||
timeouts, resource leaks (file handles, connections), large allocations
|
||||
in hot paths.
|
||||
|
||||
4. For each issue found, launch a parallel Haiku agent to:
|
||||
a. Assign a severity:
|
||||
- CRITICAL: security vulns, panics in prod (.unwrap/.expect), data exfiltration, race conditions
|
||||
- HIGH: logic bugs, missing error handling, breaking API/schema changes
|
||||
- MEDIUM: missing tests, unnecessary complexity, performance issues
|
||||
- LOW: documentation gaps, naming suggestions
|
||||
b. Score confidence 0-100 (give this rubric verbatim):
|
||||
0: False positive, doesn't stand up to scrutiny, or pre-existing issue.
|
||||
25: Might be real, but may be false positive. Stylistic issues not in CLAUDE.md.
|
||||
50: Real issue but nitpick or rare in practice. Not very important.
|
||||
75: Verified real issue, will be hit in practice. Directly impacts functionality
|
||||
or explicitly mentioned in CLAUDE.md.
|
||||
100: Certain, confirmed, will happen frequently. Evidence directly confirms.
|
||||
|
||||
5. Post a single comment on the PR using `gh pr comment` with this format.
|
||||
If no issues were found, post "No issues found." instead:
|
||||
|
||||
### Code review
|
||||
|
||||
Found N issues:
|
||||
|
||||
1. [SEVERITY:CONFIDENCE] <brief description>
|
||||
|
||||
<permalink to file:line using full SHA, eg https://github.com/owner/repo/blob/abc123def/src/file.rs#L10-L15>
|
||||
|
||||
Example: [CRITICAL:92] `.unwrap()` can panic in production when config is missing
|
||||
|
||||
You MUST use the full git SHA in links (not HEAD or branch name).
|
||||
Provide 1 line of context before and after each linked range.
|
||||
|
||||
Notes:
|
||||
- Use `gh` for all GitHub interactions, not web fetch
|
||||
- Do NOT check build signal or attempt to build/test the code
|
||||
- Ignore pre-existing issues not introduced by this PR
|
||||
- Ignore issues a linter/compiler would catch (formatting, imports, types)
|
||||
@@ -12,7 +12,6 @@ jobs:
|
||||
- name: Install Rust
|
||||
uses: dtolnay/rust-toolchain@stable
|
||||
with:
|
||||
profile: minimal
|
||||
components: rustfmt
|
||||
- name: Check formatting
|
||||
run: cargo fmt --all -- --check
|
||||
@@ -36,7 +35,6 @@ jobs:
|
||||
- name: Install Rust
|
||||
uses: dtolnay/rust-toolchain@stable
|
||||
with:
|
||||
profile: minimal
|
||||
components: clippy
|
||||
- uses: Swatinem/rust-cache@v2
|
||||
with:
|
||||
@@ -46,6 +44,7 @@ jobs:
|
||||
|
||||
clippy-windows:
|
||||
name: Clippy Windows (${{ matrix.name }})
|
||||
if: github.base_ref == 'main'
|
||||
runs-on: windows-latest
|
||||
strategy:
|
||||
fail-fast: false
|
||||
@@ -63,7 +62,6 @@ jobs:
|
||||
- name: Install Rust
|
||||
uses: dtolnay/rust-toolchain@stable
|
||||
with:
|
||||
profile: minimal
|
||||
components: clippy
|
||||
- uses: Swatinem/rust-cache@v2
|
||||
with:
|
||||
@@ -79,7 +77,12 @@ jobs:
|
||||
needs: [format, clippy, clippy-windows]
|
||||
steps:
|
||||
- run: |
|
||||
if [[ "${{ needs.format.result }}" != "success" || "${{ needs.clippy.result }}" != "success" || "${{ needs.clippy-windows.result }}" != "success" ]]; then
|
||||
if [[ "${{ needs.format.result }}" != "success" || "${{ needs.clippy.result }}" != "success" ]]; then
|
||||
echo "One or more jobs failed"
|
||||
exit 1
|
||||
fi
|
||||
# clippy-windows only runs on main PRs, so skip/success are both acceptable
|
||||
if [[ "${{ needs.clippy-windows.result }}" == "failure" ]]; then
|
||||
echo "Windows clippy failed"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
@@ -1,3 +1,31 @@
|
||||
# Code Coverage Workflow
|
||||
#
|
||||
# This workflow runs test coverage analysis and uploads reports to Codecov.
|
||||
# Coverage reports help identify untested code paths and maintain code quality.
|
||||
#
|
||||
# What it does:
|
||||
# - Runs unit and integration tests with coverage instrumentation
|
||||
# - Runs E2E tests with coverage instrumentation
|
||||
# - Uploads coverage reports to Codecov (https://codecov.io/gh/nearai/ironclaw)
|
||||
#
|
||||
# Viewing coverage reports:
|
||||
# - PRs automatically get coverage comments showing changes in coverage
|
||||
# - Visit https://codecov.io/gh/nearai/ironclaw for detailed coverage reports
|
||||
# - Coverage reports are generated for three configurations:
|
||||
# 1. all-features: Full feature set
|
||||
# 2. default: Default features
|
||||
# 3. libsql-only: Minimal libSQL-only configuration
|
||||
# - E2E coverage tracks end-to-end test coverage separately
|
||||
#
|
||||
# Coverage files:
|
||||
# - Unit/integration: lcov.info (uploaded to Codecov with "unit" flag)
|
||||
# - E2E: e2e-coverage.info (uploaded to Codecov with "e2e" flag)
|
||||
#
|
||||
# Requirements:
|
||||
# - Uses cargo-llvm-cov for coverage instrumentation
|
||||
# - Requires PostgreSQL for integration tests (pgvector/pgvector:pg16)
|
||||
# - E2E tests require Python 3.12 and Playwright
|
||||
|
||||
name: Code Coverage
|
||||
on:
|
||||
push:
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
name: E2E Tests
|
||||
on:
|
||||
workflow_call:
|
||||
schedule:
|
||||
- cron: "0 6 * * 1" # Weekly Monday 6 AM UTC
|
||||
workflow_dispatch:
|
||||
|
||||
@@ -0,0 +1,471 @@
|
||||
name: Staging CI (Batched)
|
||||
|
||||
on:
|
||||
schedule:
|
||||
- cron: "0 * * * *" # Every 60 minutes
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
force:
|
||||
description: "Force run even if no new commits"
|
||||
type: boolean
|
||||
default: false
|
||||
skip_claude_gate:
|
||||
description: "Skip Claude review gate (bypass blocking findings)"
|
||||
type: boolean
|
||||
default: false
|
||||
|
||||
permissions:
|
||||
contents: write
|
||||
issues: write
|
||||
pull-requests: write
|
||||
checks: read
|
||||
|
||||
concurrency:
|
||||
group: staging-ci
|
||||
cancel-in-progress: false # Let running suites finish
|
||||
|
||||
jobs:
|
||||
# ── Check for new commits ──────────────────────────────────────
|
||||
check-changes:
|
||||
name: Check for new commits
|
||||
runs-on: ubuntu-latest
|
||||
outputs:
|
||||
has_changes: ${{ steps.check.outputs.has_changes }}
|
||||
current_head: ${{ steps.check.outputs.current_head }}
|
||||
diff_range: ${{ steps.check.outputs.diff_range }}
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: staging
|
||||
fetch-depth: 0
|
||||
fetch-tags: true
|
||||
|
||||
- name: Check for changes since last tested
|
||||
id: check
|
||||
env:
|
||||
FORCE_RUN: ${{ inputs.force }}
|
||||
run: |
|
||||
CURRENT_HEAD=$(git rev-parse HEAD)
|
||||
echo "current_head=${CURRENT_HEAD}" >> "$GITHUB_OUTPUT"
|
||||
|
||||
if git rev-parse staging-tested >/dev/null 2>&1; then
|
||||
LAST_TESTED=$(git rev-parse staging-tested)
|
||||
else
|
||||
LAST_TESTED=""
|
||||
fi
|
||||
|
||||
DIFF_RANGE=""
|
||||
if [ -n "$LAST_TESTED" ] && [ "$LAST_TESTED" = "$CURRENT_HEAD" ]; then
|
||||
echo "No new commits since last tested (${CURRENT_HEAD})"
|
||||
HAS_CHANGES=false
|
||||
else
|
||||
HAS_CHANGES=true
|
||||
if [ -n "$LAST_TESTED" ]; then
|
||||
COMMIT_COUNT=$(git rev-list --count "${LAST_TESTED}..HEAD")
|
||||
echo "Found ${COMMIT_COUNT} new commit(s) since last tested"
|
||||
DIFF_RANGE="${LAST_TESTED}..${CURRENT_HEAD}"
|
||||
else
|
||||
git fetch origin main
|
||||
MERGE_BASE=$(git merge-base origin/main HEAD)
|
||||
echo "First run -- reviewing from merge-base ${MERGE_BASE}"
|
||||
DIFF_RANGE="${MERGE_BASE}..${CURRENT_HEAD}"
|
||||
fi
|
||||
fi
|
||||
|
||||
# Force override from workflow_dispatch
|
||||
if [ "$FORCE_RUN" = "true" ]; then
|
||||
echo "Force run requested"
|
||||
HAS_CHANGES=true
|
||||
if [ -z "$DIFF_RANGE" ]; then
|
||||
DIFF_RANGE="${CURRENT_HEAD}..${CURRENT_HEAD}"
|
||||
fi
|
||||
fi
|
||||
|
||||
echo "has_changes=${HAS_CHANGES}" >> "$GITHUB_OUTPUT"
|
||||
echo "diff_range=${DIFF_RANGE}" >> "$GITHUB_OUTPUT"
|
||||
|
||||
# ── Run full test suite ──────────────────────────────────────────
|
||||
tests:
|
||||
name: Test Suite
|
||||
needs: check-changes
|
||||
if: needs.check-changes.outputs.has_changes == 'true'
|
||||
uses: ./.github/workflows/test.yml
|
||||
|
||||
# ── Run E2E browser tests ────────────────────────────────────────
|
||||
e2e:
|
||||
name: E2E Browser Tests
|
||||
needs: check-changes
|
||||
if: needs.check-changes.outputs.has_changes == 'true'
|
||||
uses: ./.github/workflows/e2e.yml
|
||||
|
||||
# ── Create promotion PR (triggers claude-review.yml on the PR) ──
|
||||
create-promotion-pr:
|
||||
name: Create Promotion PR
|
||||
needs: check-changes
|
||||
if: needs.check-changes.outputs.has_changes == 'true'
|
||||
runs-on: ubuntu-latest
|
||||
outputs:
|
||||
pr_number: ${{ steps.create-pr.outputs.pr_number }}
|
||||
promotion_branch: ${{ steps.branch.outputs.branch }}
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: staging
|
||||
fetch-depth: 0
|
||||
|
||||
- name: Generate GitHub App token
|
||||
id: app-token
|
||||
uses: actions/create-github-app-token@v2
|
||||
with:
|
||||
app-id: ${{ secrets.GH_RELEASES_MANAGER_APP_ID }}
|
||||
private-key: ${{ secrets.GH_RELEASES_MANAGER_APP_PRIVATE_KEY }}
|
||||
|
||||
- name: Set token
|
||||
id: token
|
||||
run: |
|
||||
if [ -n "${{ steps.app-token.outputs.token }}" ]; then
|
||||
echo "token=${{ steps.app-token.outputs.token }}" >> "$GITHUB_OUTPUT"
|
||||
else
|
||||
echo "token=${{ github.token }}" >> "$GITHUB_OUTPUT"
|
||||
fi
|
||||
|
||||
- name: Check if staging is ahead of main
|
||||
id: ahead-check
|
||||
env:
|
||||
GH_TOKEN: ${{ steps.token.outputs.token }}
|
||||
run: |
|
||||
git fetch origin main
|
||||
AHEAD=$(git rev-list --count origin/main..origin/staging)
|
||||
echo "commits_ahead=${AHEAD}" >> "$GITHUB_OUTPUT"
|
||||
if [ "$AHEAD" -eq 0 ]; then
|
||||
echo "Staging is not ahead of main. Nothing to promote."
|
||||
else
|
||||
echo "Staging is ${AHEAD} commits ahead of main."
|
||||
fi
|
||||
|
||||
- name: Create promotion branch
|
||||
id: branch
|
||||
if: steps.ahead-check.outputs.commits_ahead != '0'
|
||||
run: |
|
||||
SHORT_SHA=$(echo "${{ needs.check-changes.outputs.current_head }}" | cut -c1-8)
|
||||
BRANCH="staging-promote/${SHORT_SHA}-${{ github.run_id }}"
|
||||
git checkout -b "$BRANCH"
|
||||
git push origin "$BRANCH"
|
||||
echo "branch=${BRANCH}" >> "$GITHUB_OUTPUT"
|
||||
echo "Created promotion branch: ${BRANCH}"
|
||||
|
||||
- name: Find base branch
|
||||
id: find-base
|
||||
if: steps.ahead-check.outputs.commits_ahead != '0'
|
||||
env:
|
||||
GH_TOKEN: ${{ steps.token.outputs.token }}
|
||||
run: |
|
||||
# Find the newest open promotion PR with a staging-promote/* head branch
|
||||
LATEST=$(gh pr list --label staging-promotion --state open \
|
||||
--json headRefName,createdAt \
|
||||
--jq '[.[] | select(.headRefName | startswith("staging-promote/"))] | sort_by(.createdAt) | last | .headRefName // empty')
|
||||
if [ -n "$LATEST" ]; then
|
||||
echo "base=${LATEST}" >> "$GITHUB_OUTPUT"
|
||||
echo "Chaining onto existing promotion branch: ${LATEST}"
|
||||
else
|
||||
echo "base=main" >> "$GITHUB_OUTPUT"
|
||||
echo "No existing promotion PR — targeting main"
|
||||
fi
|
||||
|
||||
- name: Create promotion PR
|
||||
id: create-pr
|
||||
if: steps.ahead-check.outputs.commits_ahead != '0'
|
||||
env:
|
||||
GH_TOKEN: ${{ steps.token.outputs.token }}
|
||||
run: |
|
||||
RANGE="${{ needs.check-changes.outputs.diff_range }}"
|
||||
TIMESTAMP=$(date -u +"%Y-%m-%d %H:%M UTC")
|
||||
BRANCH="${{ steps.branch.outputs.branch }}"
|
||||
BASE="${{ steps.find-base.outputs.base }}"
|
||||
|
||||
PR_URL=$(gh pr create \
|
||||
--base "$BASE" \
|
||||
--head "$BRANCH" \
|
||||
--title "chore: promote staging to main (${TIMESTAMP})" \
|
||||
--body "## Auto-promotion from staging CI
|
||||
|
||||
**Batch range:** \`${RANGE}\`
|
||||
**Promotion branch:** \`${BRANCH}\`
|
||||
**Base:** \`${BASE}\`
|
||||
**Triggered by:** Staging CI batch at ${TIMESTAMP}
|
||||
|
||||
Waiting for gates:
|
||||
- Tests: pending
|
||||
- E2E: pending
|
||||
- Claude Code review: pending (will post comments on this PR)
|
||||
|
||||
---
|
||||
*Auto-created by staging-ci workflow*" \
|
||||
--label "staging-promotion")
|
||||
|
||||
PR_NUM=$(echo "$PR_URL" | grep -oE '[0-9]+$')
|
||||
echo "pr_number=${PR_NUM}" >> "$GITHUB_OUTPUT"
|
||||
echo "Created promotion PR #${PR_NUM}"
|
||||
|
||||
# ── Gate: wait for review, process findings, merge or block ─────
|
||||
gate:
|
||||
name: Staging Gate
|
||||
needs: [check-changes, tests, e2e, create-promotion-pr]
|
||||
if: >
|
||||
always() &&
|
||||
needs.check-changes.outputs.has_changes == 'true' &&
|
||||
needs.tests.result == 'success' &&
|
||||
needs.e2e.result == 'success' &&
|
||||
needs.create-promotion-pr.result == 'success'
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 25
|
||||
outputs:
|
||||
gate_passed: ${{ steps.evaluate.outputs.passed }}
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: staging
|
||||
fetch-depth: 1
|
||||
|
||||
- name: Generate GitHub App token
|
||||
id: app-token
|
||||
uses: actions/create-github-app-token@v2
|
||||
with:
|
||||
app-id: ${{ secrets.GH_RELEASES_MANAGER_APP_ID }}
|
||||
private-key: ${{ secrets.GH_RELEASES_MANAGER_APP_PRIVATE_KEY }}
|
||||
|
||||
- name: Set token
|
||||
id: token
|
||||
run: |
|
||||
if [ -n "${{ steps.app-token.outputs.token }}" ]; then
|
||||
echo "token=${{ steps.app-token.outputs.token }}" >> "$GITHUB_OUTPUT"
|
||||
else
|
||||
echo "token=${{ github.token }}" >> "$GITHUB_OUTPUT"
|
||||
fi
|
||||
|
||||
- name: Wait for Claude review job
|
||||
env:
|
||||
GH_TOKEN: ${{ steps.token.outputs.token }}
|
||||
PR_NUMBER: ${{ needs.create-promotion-pr.outputs.pr_number }}
|
||||
REPO: ${{ github.repository }}
|
||||
run: |
|
||||
if [ -z "$PR_NUMBER" ]; then
|
||||
echo "No PR number — skipping wait"
|
||||
exit 0
|
||||
fi
|
||||
|
||||
PR_SHA=$(gh pr view "$PR_NUMBER" --json headRefOid --jq '.headRefOid' || echo "")
|
||||
if [ -z "$PR_SHA" ]; then
|
||||
echo "::warning::Could not get PR head SHA"
|
||||
exit 0
|
||||
fi
|
||||
|
||||
echo "Polling for Claude Code Review job on PR #${PR_NUMBER} (SHA: ${PR_SHA})..."
|
||||
TIMEOUT=1200 # 20 minutes
|
||||
ELAPSED=0
|
||||
INTERVAL=30
|
||||
|
||||
while [ "$ELAPSED" -lt "$TIMEOUT" ]; do
|
||||
STATUS=$(gh api "repos/${REPO}/commits/${PR_SHA}/check-runs" \
|
||||
--jq '[.check_runs[] | select(.name == "Claude Code Review") | .conclusion // .status] | first // "pending"' 2>/dev/null || echo "pending")
|
||||
|
||||
if [ "$STATUS" = "success" ] || [ "$STATUS" = "failure" ] || [ "$STATUS" = "cancelled" ]; then
|
||||
echo "Claude review job completed with status: ${STATUS} (${ELAPSED}s)"
|
||||
exit 0
|
||||
fi
|
||||
|
||||
echo "Claude review status: ${STATUS} (${ELAPSED}s elapsed)"
|
||||
sleep "$INTERVAL"
|
||||
ELAPSED=$((ELAPSED + INTERVAL))
|
||||
done
|
||||
|
||||
echo "::warning::Claude review job not completed after ${TIMEOUT}s"
|
||||
|
||||
- name: Process Claude review comments and create issues
|
||||
id: process-findings
|
||||
env:
|
||||
GH_TOKEN: ${{ steps.token.outputs.token }}
|
||||
PR_NUMBER: ${{ needs.create-promotion-pr.outputs.pr_number }}
|
||||
REPO: ${{ github.repository }}
|
||||
run: |
|
||||
HAS_BLOCKING=false
|
||||
ISSUES_CREATED=0
|
||||
|
||||
if [ -z "$PR_NUMBER" ]; then
|
||||
echo "No PR — skipping finding processing"
|
||||
echo "has_blocking=false" >> "$GITHUB_OUTPUT"
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# Check for "No issues found" first (clean pass)
|
||||
NO_ISSUES=$(gh api "repos/${REPO}/issues/${PR_NUMBER}/comments" \
|
||||
--jq '[.[] | select(.user.login == "claude[bot]") | select(.body | test("No issues found"))] | length' 2>/dev/null || echo "0")
|
||||
if [ "$NO_ISSUES" -gt 0 ]; then
|
||||
echo "Claude review found no issues — gate passes"
|
||||
echo "has_blocking=false" >> "$GITHUB_OUTPUT"
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# Get the last Claude comment that contains findings
|
||||
JQ_FILTER='[.[] | select(.user.login == "claude[bot]") | select(.body | test("Found [0-9]+ issue"))] | last'
|
||||
BODY=$(gh api "repos/${REPO}/issues/${PR_NUMBER}/comments" \
|
||||
--jq "${JQ_FILTER} | .body // empty" 2>/dev/null || echo "")
|
||||
COMMENT_URL=$(gh api "repos/${REPO}/issues/${PR_NUMBER}/comments" \
|
||||
--jq "${JQ_FILTER} | .html_url // empty" 2>/dev/null || echo "")
|
||||
|
||||
if [ -z "$BODY" ]; then
|
||||
echo "::warning::No Claude review comment found for PR #${PR_NUMBER} — treating as blocking"
|
||||
echo "has_blocking=true" >> "$GITHUB_OUTPUT"
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# Parse [SEVERITY:CONFIDENCE] tags from each numbered finding
|
||||
# Matrix: CRITICAL always→issue, ≥80→block. HIGH ≥50→issue. MEDIUM ≥80→issue. LOW ≥80→issue.
|
||||
# Use process substitution so variables propagate to parent shell
|
||||
while read -r line; do
|
||||
TAG=$(echo "$line" | grep -oE '^\[(CRITICAL|HIGH|MEDIUM|LOW):[0-9]+\]')
|
||||
SEVERITY=$(echo "$TAG" | sed 's/\[\(.*\):\(.*\)\]/\1/')
|
||||
CONFIDENCE=$(echo "$TAG" | sed 's/\[\(.*\):\(.*\)\]/\2/')
|
||||
DESC=$(echo "$line" | sed "s/\[${SEVERITY}:${CONFIDENCE}\] *//" | head -1)
|
||||
|
||||
echo "Found: [${SEVERITY}:${CONFIDENCE}] ${DESC}"
|
||||
|
||||
# Check if blocking (CRITICAL ≥80)
|
||||
if [ "$SEVERITY" = "CRITICAL" ] && [ "$CONFIDENCE" -ge 80 ]; then
|
||||
HAS_BLOCKING=true
|
||||
fi
|
||||
|
||||
# Determine if this should create an issue
|
||||
CREATE_ISSUE=false
|
||||
case "$SEVERITY" in
|
||||
CRITICAL) CREATE_ISSUE=true ;;
|
||||
HIGH) [ "$CONFIDENCE" -ge 50 ] && CREATE_ISSUE=true ;;
|
||||
MEDIUM) [ "$CONFIDENCE" -ge 80 ] && CREATE_ISSUE=true ;;
|
||||
LOW) [ "$CONFIDENCE" -ge 80 ] && CREATE_ISSUE=true ;;
|
||||
esac
|
||||
|
||||
if [ "$CREATE_ISSUE" = "true" ]; then
|
||||
case "$SEVERITY" in
|
||||
CRITICAL) LABELS="bug,risk: high,staging-ci-review" ;;
|
||||
HIGH) LABELS="bug,risk: medium,staging-ci-review" ;;
|
||||
MEDIUM) LABELS="risk: medium,staging-ci-review" ;;
|
||||
LOW) LABELS="risk: low,staging-ci-review" ;;
|
||||
esac
|
||||
|
||||
TITLE=$(echo "$DESC" | cut -c1-80)
|
||||
{
|
||||
echo "## [${SEVERITY}:${CONFIDENCE}] Issue Found by Staging CI Review"
|
||||
echo ""
|
||||
echo "**Severity:** ${SEVERITY}"
|
||||
echo "**Confidence:** ${CONFIDENCE}/100"
|
||||
echo "**PR comment:** ${COMMENT_URL}"
|
||||
echo ""
|
||||
echo "### Description"
|
||||
echo "$DESC"
|
||||
echo ""
|
||||
echo "---"
|
||||
echo "*Auto-created by staging-ci Claude Code review*"
|
||||
} > /tmp/issue-body.md
|
||||
|
||||
if gh issue create \
|
||||
--title "[${SEVERITY}] ${TITLE}" \
|
||||
--body-file /tmp/issue-body.md \
|
||||
--label "${LABELS}"; then
|
||||
ISSUES_CREATED=$((ISSUES_CREATED + 1))
|
||||
else
|
||||
echo "::warning::Failed to create issue for ${SEVERITY} finding"
|
||||
fi
|
||||
fi
|
||||
done < <(echo "$BODY" | grep -oE '\[(CRITICAL|HIGH|MEDIUM|LOW):[0-9]+\].*')
|
||||
|
||||
echo "Created ${ISSUES_CREATED} issues"
|
||||
echo "has_blocking=${HAS_BLOCKING}" >> "$GITHUB_OUTPUT"
|
||||
|
||||
- name: Evaluate gate
|
||||
id: evaluate
|
||||
env:
|
||||
PR_NUMBER: ${{ needs.create-promotion-pr.outputs.pr_number }}
|
||||
SKIP_GATE: ${{ inputs.skip_claude_gate }}
|
||||
HAS_BLOCKING: ${{ steps.process-findings.outputs.has_blocking }}
|
||||
run: |
|
||||
SKIP_INPUT="$SKIP_GATE"
|
||||
|
||||
if [ "$HAS_BLOCKING" = "true" ]; then
|
||||
echo "::warning::Claude review found blocking issues (CRITICAL ≥80 confidence)"
|
||||
if [ "$SKIP_INPUT" = "true" ]; then
|
||||
echo "::warning::Gate overridden by skip_claude_gate workflow input"
|
||||
echo "passed=true" >> "$GITHUB_OUTPUT"
|
||||
else
|
||||
echo "::error::Blocking promotion due to CRITICAL findings (≥80 confidence)"
|
||||
echo "::error::PR #${PR_NUMBER} left open with review comments"
|
||||
echo "passed=false" >> "$GITHUB_OUTPUT"
|
||||
exit 1
|
||||
fi
|
||||
else
|
||||
echo "No blocking findings. Gate passed."
|
||||
echo "passed=true" >> "$GITHUB_OUTPUT"
|
||||
fi
|
||||
|
||||
- name: Merge promotion PR
|
||||
id: merge
|
||||
if: steps.evaluate.outputs.passed == 'true'
|
||||
env:
|
||||
GH_TOKEN: ${{ steps.token.outputs.token }}
|
||||
PR_NUMBER: ${{ needs.create-promotion-pr.outputs.pr_number }}
|
||||
run: |
|
||||
if [ -n "$PR_NUMBER" ]; then
|
||||
echo "Merging promotion PR #${PR_NUMBER}"
|
||||
# Do NOT use --delete-branch: deleting a promotion branch closes
|
||||
# any chained PRs that use it as their base (verified in ironclaw-ci-test).
|
||||
# Stale promotion branches are cleaned up separately.
|
||||
gh pr merge "$PR_NUMBER" --merge
|
||||
echo "merged=true" >> "$GITHUB_OUTPUT"
|
||||
fi
|
||||
|
||||
# ── Update tested tag (always, so next batch covers only new commits) ──
|
||||
update-tag:
|
||||
name: Update staging-tested tag
|
||||
needs: [check-changes, tests, e2e, create-promotion-pr, gate]
|
||||
if: >
|
||||
always() &&
|
||||
needs.check-changes.outputs.has_changes == 'true' &&
|
||||
needs.tests.result == 'success' &&
|
||||
needs.e2e.result == 'success' &&
|
||||
needs.create-promotion-pr.result == 'success'
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
with:
|
||||
ref: staging
|
||||
fetch-depth: 1
|
||||
|
||||
- name: Update staging-tested tag
|
||||
run: |
|
||||
git tag -f staging-tested "${{ needs.check-changes.outputs.current_head }}"
|
||||
git push origin staging-tested --force
|
||||
echo "Updated staging-tested tag to ${{ needs.check-changes.outputs.current_head }}"
|
||||
|
||||
# ── Report ───────────────────────────────────────────────────────
|
||||
report:
|
||||
name: Staging CI Summary
|
||||
needs: [check-changes, tests, e2e, create-promotion-pr, gate, update-tag]
|
||||
if: always() && needs.check-changes.outputs.has_changes == 'true'
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Summary
|
||||
run: |
|
||||
echo "## Staging CI Batch Results" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "| Check | Result |" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "|-------|--------|" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "| Tests | ${{ needs.tests.result }} |" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "| E2E | ${{ needs.e2e.result }} |" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "| Promotion PR | ${{ needs.create-promotion-pr.result }} |" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "| Gate | ${{ needs.gate.result }} |" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "| Tag Updated | ${{ needs.update-tag.result }} |" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "" >> "$GITHUB_STEP_SUMMARY"
|
||||
echo "Range: ${{ needs.check-changes.outputs.diff_range }}" >> "$GITHUB_STEP_SUMMARY"
|
||||
PR_NUM="${{ needs.create-promotion-pr.outputs.pr_number }}"
|
||||
if [ -n "$PR_NUM" ]; then
|
||||
echo "Promotion PR: #${PR_NUM}" >> "$GITHUB_STEP_SUMMARY"
|
||||
fi
|
||||
+33
-14
@@ -1,6 +1,9 @@
|
||||
name: Run Tests
|
||||
on:
|
||||
workflow_call:
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
@@ -14,7 +17,7 @@ jobs:
|
||||
matrix:
|
||||
include:
|
||||
- name: all-features
|
||||
flags: "--all-features"
|
||||
flags: "--features postgres,libsql,html-to-markdown"
|
||||
- name: default
|
||||
flags: ""
|
||||
- name: libsql-only
|
||||
@@ -25,7 +28,6 @@ jobs:
|
||||
- name: Install Rust
|
||||
uses: dtolnay/rust-toolchain@stable
|
||||
with:
|
||||
profile: minimal
|
||||
targets: wasm32-wasip2
|
||||
- uses: Swatinem/rust-cache@v2
|
||||
with:
|
||||
@@ -39,20 +41,24 @@ jobs:
|
||||
|
||||
telegram-tests:
|
||||
name: Telegram Channel Tests
|
||||
if: >
|
||||
github.event_name == 'push' ||
|
||||
(github.event_name == 'pull_request' && github.base_ref != 'staging')
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@v6
|
||||
- name: Install Rust
|
||||
uses: dtolnay/rust-toolchain@stable
|
||||
with:
|
||||
profile: minimal
|
||||
- uses: Swatinem/rust-cache@v2
|
||||
- name: Run Telegram Channel Tests
|
||||
run: cargo test --manifest-path channels-src/telegram/Cargo.toml -- --nocapture
|
||||
|
||||
windows-build:
|
||||
name: Windows Build (${{ matrix.name }})
|
||||
if: >
|
||||
github.event_name == 'push' ||
|
||||
(github.event_name == 'pull_request' && github.base_ref != 'staging')
|
||||
runs-on: windows-latest
|
||||
strategy:
|
||||
fail-fast: false
|
||||
@@ -69,8 +75,6 @@ jobs:
|
||||
uses: actions/checkout@v6
|
||||
- name: Install Rust
|
||||
uses: dtolnay/rust-toolchain@stable
|
||||
with:
|
||||
profile: minimal
|
||||
- uses: Swatinem/rust-cache@v2
|
||||
with:
|
||||
key: windows-${{ matrix.name }}
|
||||
@@ -79,6 +83,9 @@ jobs:
|
||||
|
||||
wasm-wit-compat:
|
||||
name: WASM WIT Compatibility
|
||||
if: >
|
||||
github.event_name == 'push' ||
|
||||
(github.event_name == 'pull_request' && github.base_ref != 'staging')
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
@@ -86,7 +93,6 @@ jobs:
|
||||
- name: Install Rust
|
||||
uses: dtolnay/rust-toolchain@stable
|
||||
with:
|
||||
profile: minimal
|
||||
targets: wasm32-wasip2
|
||||
- uses: Swatinem/rust-cache@v2
|
||||
with:
|
||||
@@ -100,6 +106,9 @@ jobs:
|
||||
|
||||
docker-build:
|
||||
name: Docker Build
|
||||
if: >
|
||||
github.event_name == 'push' ||
|
||||
(github.event_name == 'pull_request' && github.base_ref != 'staging')
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
@@ -129,12 +138,22 @@ jobs:
|
||||
needs: [tests, telegram-tests, wasm-wit-compat, docker-build, windows-build, version-check]
|
||||
steps:
|
||||
- run: |
|
||||
if [[ "${{ needs.tests.result }}" != "success" || "${{ needs.telegram-tests.result }}" != "success" || "${{ needs.wasm-wit-compat.result }}" != "success" || "${{ needs.docker-build.result }}" != "success" || "${{ needs.windows-build.result }}" != "success" ]]; then
|
||||
echo "One or more jobs failed"
|
||||
exit 1
|
||||
fi
|
||||
# version-check only runs on PRs, so skip/success are both acceptable
|
||||
if [[ "${{ needs.version-check.result }}" == "failure" ]]; then
|
||||
echo "Version bump check failed"
|
||||
# Unit tests must always pass
|
||||
if [[ "${{ needs.tests.result }}" != "success" ]]; then
|
||||
echo "Unit tests failed"
|
||||
exit 1
|
||||
fi
|
||||
# Gated jobs: must pass on promotion PRs / push, skipped on developer PRs
|
||||
for job in telegram-tests wasm-wit-compat docker-build windows-build version-check; do
|
||||
case "$job" in
|
||||
telegram-tests) result="${{ needs.telegram-tests.result }}" ;;
|
||||
wasm-wit-compat) result="${{ needs.wasm-wit-compat.result }}" ;;
|
||||
docker-build) result="${{ needs.docker-build.result }}" ;;
|
||||
windows-build) result="${{ needs.windows-build.result }}" ;;
|
||||
version-check) result="${{ needs.version-check.result }}" ;;
|
||||
esac
|
||||
if [[ "$result" == "failure" || "$result" == "cancelled" ]]; then
|
||||
echo "$job failed"
|
||||
exit 1
|
||||
fi
|
||||
done
|
||||
|
||||
+7
-1
@@ -4,8 +4,9 @@
|
||||
.env.*
|
||||
!.env.example
|
||||
|
||||
# Claude Code worktrees
|
||||
# Claude Code worktrees and lock files
|
||||
.claude/worktrees/
|
||||
.claude/scheduled_tasks.lock
|
||||
|
||||
# Sidecar tool data
|
||||
.sidecar/
|
||||
@@ -22,3 +23,8 @@ bench-results/
|
||||
# WASM build artifacts (loaded from disk, not bundled)
|
||||
*.wasm
|
||||
|
||||
# Traces
|
||||
trace_*.json
|
||||
|
||||
# Local Claude Code settings (machine-specific, should not be committed)
|
||||
.claude/settings.local.json
|
||||
|
||||
@@ -7,6 +7,10 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
||||
|
||||
## [Unreleased]
|
||||
|
||||
### Added
|
||||
|
||||
- AWS Bedrock LLM provider via native Converse API with IAM and SSO auth support (feature-gated: `--features bedrock`)
|
||||
|
||||
## [0.16.1](https://github.com/nearai/ironclaw/compare/v0.16.0...v0.16.1) - 2026-03-06
|
||||
|
||||
### Fixed
|
||||
|
||||
@@ -1,149 +1,122 @@
|
||||
# IronClaw Development Guide
|
||||
|
||||
## Project Overview
|
||||
|
||||
**IronClaw** is a secure personal AI assistant that protects your data and expands its capabilities on the fly.
|
||||
|
||||
### Core Philosophy
|
||||
- **User-first security** - Your data stays yours, encrypted and local
|
||||
- **Self-expanding** - Build new tools dynamically without vendor dependency
|
||||
- **Defense in depth** - Multiple security layers against prompt injection and data exfiltration
|
||||
- **Always available** - Multi-channel access with proactive background execution
|
||||
|
||||
### Features
|
||||
- **Multi-channel input**: TUI (Ratatui), HTTP webhooks, WASM channels (Telegram, Slack), web gateway
|
||||
- **Parallel job execution** with state machine and self-repair for stuck jobs
|
||||
- **Sandbox execution**: Docker container isolation with network proxy and credential injection
|
||||
- **Claude Code mode**: Delegate jobs to Claude CLI inside containers
|
||||
- **Skills system**: SKILL.md prompt extensions with trust model, tool attenuation, and ClawHub registry
|
||||
- **Routines**: Scheduled (cron) and reactive (event, webhook) task execution
|
||||
- **Web gateway**: Browser UI with SSE/WebSocket real-time streaming
|
||||
- **Extension management**: Install, auth, activate MCP/WASM extensions
|
||||
- **Extensible tools**: Built-in tools, WASM sandbox, MCP client, dynamic builder
|
||||
- **Persistent memory**: Workspace with hybrid search (FTS + vector via RRF)
|
||||
- **Prompt injection defense**: Sanitizer, validator, policy rules, leak detection, shell env scrubbing
|
||||
- **Multi-provider LLM**: NEAR AI, OpenAI, Anthropic, Ollama, OpenAI-compatible, Tinfoil private inference
|
||||
- **Setup wizard**: 7-step interactive onboarding for first-run configuration
|
||||
- **Heartbeat system**: Proactive periodic execution with checklist
|
||||
**IronClaw** is a secure personal AI assistant — user-first security, self-expanding tools, defense in depth, multi-channel access with proactive background execution.
|
||||
|
||||
## Build & Test
|
||||
|
||||
```bash
|
||||
# Format code
|
||||
cargo fmt
|
||||
|
||||
# Lint (fix ALL warnings before committing, including pre-existing ones)
|
||||
cargo clippy --all --benches --tests --examples --all-features
|
||||
|
||||
# Run all tests
|
||||
cargo test
|
||||
|
||||
# Run specific test
|
||||
cargo test test_name
|
||||
|
||||
# Run with logging
|
||||
RUST_LOG=ironclaw=debug cargo run
|
||||
cargo fmt # format
|
||||
cargo clippy --all --benches --tests --examples --all-features # lint (zero warnings)
|
||||
cargo test # unit tests
|
||||
cargo test --features integration # + PostgreSQL tests
|
||||
RUST_LOG=ironclaw=debug cargo run # run with logging
|
||||
```
|
||||
|
||||
E2E tests: see `tests/e2e/CLAUDE.md`.
|
||||
|
||||
## Code Style
|
||||
|
||||
- Prefer `crate::` for cross-module imports; `super::` is fine in tests and intra-module refs
|
||||
- No `pub use` re-exports unless exposing to downstream consumers
|
||||
- No `.unwrap()` or `.expect()` in production code (tests are fine)
|
||||
- Use `thiserror` for error types in `error.rs`
|
||||
- Map errors with context: `.map_err(|e| SomeError::Variant { reason: e.to_string() })?`
|
||||
- Prefer strong types over strings (enums, newtypes)
|
||||
- Keep functions focused, extract helpers when logic is reused
|
||||
- Comments for non-obvious logic only
|
||||
|
||||
## Architecture
|
||||
|
||||
Prefer generic/extensible architectures over hardcoding specific integrations. Ask clarifying questions about the desired abstraction level before implementing.
|
||||
|
||||
Key traits for extensibility: `Database`, `Channel`, `Tool`, `LlmProvider`, `SuccessEvaluator`, `EmbeddingProvider`, `NetworkPolicyDecider`, `Hook`, `Observer`, `Tunnel`.
|
||||
|
||||
All I/O is async with tokio. Use `Arc<T>` for shared state, `RwLock` for concurrent access.
|
||||
|
||||
## Project Structure
|
||||
|
||||
```
|
||||
src/
|
||||
├── lib.rs # Library root, module declarations
|
||||
├── main.rs # Entry point, CLI args, startup
|
||||
├── config.rs # Configuration from env vars
|
||||
├── app.rs # App startup orchestration (channel wiring, DB init)
|
||||
├── bootstrap.rs # Base directory resolution (~/.ironclaw), early .env loading
|
||||
├── settings.rs # User settings persistence (~/.ironclaw/settings.json)
|
||||
├── service.rs # OS service management (launchd/systemd daemon install)
|
||||
├── tracing_fmt.rs # Custom tracing formatter
|
||||
├── util.rs # Shared utilities
|
||||
├── config/ # Configuration from env vars (split by subsystem)
|
||||
│ ├── mod.rs # Re-exports all config types; top-level Config struct
|
||||
│ ├── agent.rs, llm.rs, channels.rs, database.rs, sandbox.rs, skills.rs
|
||||
│ ├── heartbeat.rs, routines.rs, safety.rs, embeddings.rs, wasm.rs
|
||||
│ ├── tunnel.rs # Tunnel provider config (TUNNEL_PROVIDER, TUNNEL_URL, etc.)
|
||||
│ └── secrets.rs, hygiene.rs, builder.rs, helpers.rs
|
||||
├── error.rs # Error types (thiserror)
|
||||
│
|
||||
├── agent/ # Core agent logic
|
||||
│ ├── agent_loop.rs # Main Agent struct, message handling loop
|
||||
│ ├── router.rs # MessageIntent classification
|
||||
│ ├── scheduler.rs # Parallel job scheduling
|
||||
│ ├── worker.rs # Per-job execution with LLM reasoning
|
||||
│ ├── self_repair.rs # Stuck job detection and recovery
|
||||
│ ├── heartbeat.rs # Proactive periodic execution
|
||||
│ ├── session.rs # Session/thread/turn model with state machine
|
||||
│ ├── session_manager.rs # Thread/session lifecycle management
|
||||
│ ├── compaction.rs # Context window management with turn summarization
|
||||
│ ├── context_monitor.rs # Memory pressure detection
|
||||
│ ├── undo.rs # Turn-based undo/redo with checkpoints
|
||||
│ ├── submission.rs # Submission parsing (undo, redo, compact, clear, etc.)
|
||||
│ ├── dispatcher.rs # Skill-aware job dispatching
|
||||
│ ├── task.rs # Sub-task execution framework
|
||||
│ ├── routine.rs # Routine types (Trigger, Action, Guardrails)
|
||||
│ └── routine_engine.rs # Routine execution (cron ticker, event matcher)
|
||||
├── agent/ # Core agent loop, dispatcher, scheduler, sessions — see src/agent/CLAUDE.md
|
||||
│
|
||||
├── channels/ # Multi-channel input
|
||||
│ ├── channel.rs # Channel trait, IncomingMessage, OutgoingResponse
|
||||
│ ├── manager.rs # ChannelManager merges streams
|
||||
│ ├── cli/ # Full TUI with Ratatui
|
||||
│ │ ├── mod.rs # TuiChannel implementation
|
||||
│ │ ├── app.rs # Application state
|
||||
│ │ ├── render.rs # UI rendering
|
||||
│ │ ├── events.rs # Input handling
|
||||
│ │ ├── overlay.rs # Approval overlays
|
||||
│ │ └── composer.rs # Message composition
|
||||
│ ├── http.rs # HTTP webhook (axum) with secret validation
|
||||
│ ├── webhook_server.rs # Unified HTTP server composing all webhook routes
|
||||
│ ├── repl.rs # Simple REPL (for testing)
|
||||
│ ├── web/ # Web gateway (browser UI)
|
||||
│ │ ├── mod.rs # Gateway builder, startup
|
||||
│ │ ├── server.rs # Axum router, 40+ API endpoints
|
||||
│ │ ├── sse.rs # SSE broadcast manager
|
||||
│ │ ├── ws.rs # WebSocket gateway + connection tracking
|
||||
│ │ ├── types.rs # Request/response types, SseEvent enum
|
||||
│ │ ├── auth.rs # Bearer token auth middleware
|
||||
│ │ ├── log_layer.rs # Tracing layer for log streaming
|
||||
│ │ └── static/ # HTML, CSS, JS (single-page app)
|
||||
│ ├── web/ # Web gateway (browser UI) — see src/channels/web/CLAUDE.md
|
||||
│ └── 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)
|
||||
│ ├── mod.rs # Cli struct, Command enum (run/onboard/config/tool/registry/mcp/memory/pairing/service/doctor/status/completion)
|
||||
│ └── config.rs, tool.rs, registry.rs, mcp.rs, memory.rs, pairing.rs, service.rs, doctor.rs, status.rs, completion.rs
|
||||
│
|
||||
├── registry/ # Extension registry catalog
|
||||
│ ├── manifest.rs # ExtensionManifest, ArtifactSpec, BundleDefinition types
|
||||
│ ├── catalog.rs # RegistryCatalog: load from filesystem and embedded JSON
|
||||
│ └── installer.rs # RegistryInstaller: download, verify, install WASM artifacts
|
||||
│
|
||||
├── hooks/ # Lifecycle hooks (6 points: BeforeInbound, BeforeToolCall, BeforeOutbound, OnSessionStart, OnSessionEnd, TransformResponse)
|
||||
│
|
||||
├── tunnel/ # Tunnel abstraction for public internet exposure
|
||||
│ ├── 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)
|
||||
│
|
||||
├── orchestrator/ # Internal HTTP API for sandbox containers
|
||||
│ ├── mod.rs
|
||||
│ ├── api.rs # Axum endpoints (LLM proxy, events, prompts)
|
||||
│ ├── auth.rs # Per-job bearer token store
|
||||
│ └── job_manager.rs # Container lifecycle (create, stop, cleanup)
|
||||
│
|
||||
├── worker/ # Runs inside Docker containers
|
||||
│ ├── mod.rs
|
||||
│ ├── runtime.rs # Worker execution loop (tool calls, LLM)
|
||||
│ ├── claude_bridge.rs # Claude Code bridge (spawns claude CLI)
|
||||
│ ├── api.rs # HTTP client to orchestrator
|
||||
│ └── proxy_llm.rs # LlmProvider that proxies through orchestrator
|
||||
│
|
||||
├── safety/ # Prompt injection defense
|
||||
│ ├── sanitizer.rs # Pattern detection, content escaping
|
||||
│ ├── validator.rs # Input validation (length, encoding, patterns)
|
||||
│ ├── policy.rs # PolicyRule system with severity/actions
|
||||
│ └── leak_detector.rs # Secret detection (API keys, tokens, etc.)
|
||||
│ ├── leak_detector.rs # Secret detection (API keys, tokens, etc.)
|
||||
│ └── credential_detect.rs # HTTP request credential detection
|
||||
│
|
||||
├── llm/ # LLM integration (multi-provider)
|
||||
│ ├── mod.rs # Provider factory, LlmBackend enum
|
||||
│ ├── provider.rs # LlmProvider trait, message types
|
||||
│ ├── nearai_chat.rs # NEAR AI Chat Completions provider (session token + API key auth)
|
||||
│ ├── reasoning.rs # Planning, tool selection, evaluation
|
||||
│ ├── session.rs # Session token management with auto-renewal
|
||||
│ ├── circuit_breaker.rs # Circuit breaker for provider failures
|
||||
│ ├── retry.rs # Retry with exponential backoff
|
||||
│ ├── failover.rs # Multi-provider failover chain
|
||||
│ ├── response_cache.rs # LLM response caching
|
||||
│ ├── costs.rs # Token cost tracking
|
||||
│ └── rig_adapter.rs # Rig framework adapter
|
||||
├── llm/ # Multi-provider LLM integration — see src/llm/CLAUDE.md
|
||||
│
|
||||
├── tools/ # Extensible tool system
|
||||
│ ├── tool.rs # Tool trait, ToolOutput, ToolError
|
||||
│ ├── registry.rs # ToolRegistry for discovery
|
||||
│ ├── sandbox.rs # Process-based sandbox (stub, superseded by wasm/)
|
||||
│ ├── builtin/ # Built-in tools
|
||||
│ │ ├── echo.rs, time.rs, json.rs, http.rs
|
||||
│ │ ├── file.rs # ReadFile, WriteFile, ListDir, ApplyPatch
|
||||
│ │ ├── shell.rs # Shell command execution
|
||||
│ │ ├── memory.rs # Memory tools (search, write, read, tree)
|
||||
│ │ ├── job.rs # CreateJob, ListJobs, JobStatus, CancelJob
|
||||
│ │ ├── routine.rs # routine_create/list/update/delete/history
|
||||
│ │ ├── extension_tools.rs # Extension install/auth/activate/remove
|
||||
│ │ ├── skill_tools.rs # skill_list/search/install/remove tools
|
||||
│ │ └── marketplace.rs, ecommerce.rs, taskrabbit.rs, restaurant.rs (stubs)
|
||||
│ ├── 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)
|
||||
│ ├── builder/ # Dynamic tool building
|
||||
│ │ ├── core.rs # BuildRequirement, SoftwareType, Language
|
||||
│ │ ├── templates.rs # Project scaffolding
|
||||
@@ -151,7 +124,9 @@ src/
|
||||
│ │ └── validation.rs # WASM validation
|
||||
│ ├── mcp/ # Model Context Protocol
|
||||
│ │ ├── client.rs # MCP client over HTTP
|
||||
│ │ └── protocol.rs # JSON-RPC types
|
||||
│ │ ├── 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
|
||||
@@ -161,130 +136,60 @@ src/
|
||||
│ ├── 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/ # Database abstraction layer
|
||||
│ ├── mod.rs # Database trait (~60 async methods)
|
||||
│ ├── postgres.rs # PostgreSQL backend (delegates to Store + Repository)
|
||||
│ ├── libsql_backend.rs # libSQL/Turso backend (embedded SQLite)
|
||||
│ └── libsql_migrations.rs # SQLite-dialect schema (idempotent)
|
||||
├── db/ # Dual-backend persistence (PostgreSQL + libSQL) — see src/db/CLAUDE.md
|
||||
│
|
||||
├── workspace/ # Persistent memory system (OpenClaw-inspired)
|
||||
│ ├── mod.rs # Workspace struct, memory operations
|
||||
│ ├── document.rs # MemoryDocument, MemoryChunk, WorkspaceEntry
|
||||
│ ├── chunker.rs # Document chunking (800 tokens, 15% overlap)
|
||||
│ ├── embeddings.rs # EmbeddingProvider trait, OpenAI implementation
|
||||
│ ├── search.rs # Hybrid search with RRF algorithm
|
||||
│ └── repository.rs # PostgreSQL CRUD and search operations
|
||||
├── workspace/ # Persistent memory system — see src/workspace/README.md
|
||||
│
|
||||
├── context/ # Job context isolation
|
||||
│ ├── state.rs # JobState enum, JobContext, state machine
|
||||
│ ├── memory.rs # ActionRecord, ConversationMemory
|
||||
│ └── manager.rs # ContextManager for concurrent jobs
|
||||
│
|
||||
├── estimation/ # Cost/time/value estimation
|
||||
│ ├── cost.rs # CostEstimator
|
||||
│ ├── time.rs # TimeEstimator
|
||||
│ ├── value.rs # ValueEstimator (profit margins)
|
||||
│ └── learner.rs # Exponential moving average learning
|
||||
│
|
||||
├── evaluation/ # Success evaluation
|
||||
│ ├── success.rs # SuccessEvaluator trait, RuleBasedEvaluator, LlmEvaluator
|
||||
│ └── metrics.rs # MetricsCollector, QualityMetrics
|
||||
├── context/ # Job context isolation (JobState, JobContext, ContextManager)
|
||||
├── estimation/ # Cost/time/value estimation with EMA learning
|
||||
├── evaluation/ # Success evaluation (rule-based, LLM-based)
|
||||
│
|
||||
├── sandbox/ # Docker execution sandbox
|
||||
│ ├── mod.rs # Public API, default allowlist
|
||||
│ ├── config.rs # SandboxConfig, SandboxPolicy enum
|
||||
│ ├── config.rs # SandboxConfig, SandboxPolicy enum (ReadOnly/WorkspaceWrite/FullAccess)
|
||||
│ ├── manager.rs # SandboxManager orchestration
|
||||
│ ├── container.rs # ContainerRunner, Docker lifecycle
|
||||
│ ├── error.rs # SandboxError types
|
||||
│ └── proxy/ # Network proxy for containers
|
||||
│ ├── mod.rs # NetworkProxyBuilder
|
||||
│ ├── http.rs # HttpProxy, CredentialResolver trait
|
||||
│ ├── policy.rs # NetworkPolicyDecider trait
|
||||
│ └── allowlist.rs # DomainAllowlist validation
|
||||
│ └── proxy/ # Network proxy: domain allowlist, credential injection, CONNECT tunnel
|
||||
│
|
||||
├── secrets/ # Secrets management
|
||||
│ ├── crypto.rs # AES-256-GCM encryption
|
||||
│ ├── store.rs # Secret storage
|
||||
│ └── types.rs # Credential types
|
||||
├── secrets/ # Secrets management (AES-256-GCM, OS keychain for master key)
|
||||
│
|
||||
├── setup/ # Onboarding wizard (spec: src/setup/README.md)
|
||||
│ ├── mod.rs # Entry point, check_onboard_needed()
|
||||
│ ├── wizard.rs # 7-step interactive wizard
|
||||
│ ├── channels.rs # Channel setup helpers
|
||||
│ └── prompts.rs # Terminal prompts (select, confirm, secret)
|
||||
├── setup/ # 7-step onboarding wizard — see src/setup/README.md
|
||||
│
|
||||
├── skills/ # SKILL.md prompt extension system
|
||||
│ ├── mod.rs # Core types (SkillTrust, LoadedSkill)
|
||||
│ ├── registry.rs # SkillRegistry: discover, install, remove
|
||||
│ ├── selector.rs # Deterministic scoring prefilter
|
||||
│ ├── attenuation.rs # Trust-based tool ceiling
|
||||
│ ├── gating.rs # Requirement checks (bins, env, config)
|
||||
│ ├── parser.rs # SKILL.md frontmatter + markdown parser
|
||||
│ └── catalog.rs # ClawHub registry client
|
||||
├── skills/ # SKILL.md prompt extension system — see .claude/rules/skills.md
|
||||
│
|
||||
└── history/ # Persistence
|
||||
├── store.rs # PostgreSQL repositories
|
||||
└── analytics.rs # Aggregation queries (JobStats, ToolStats)
|
||||
└── history/ # Persistence (PostgreSQL repositories, analytics)
|
||||
|
||||
tests/
|
||||
├── *.rs # Integration tests (workspace, heartbeat, WS gateway, pairing, etc.)
|
||||
├── test-pages/ # HTML→Markdown conversion fixtures
|
||||
└── e2e/ # Python/Playwright E2E scenarios (see tests/e2e/CLAUDE.md)
|
||||
```
|
||||
|
||||
## Key Patterns
|
||||
## Database
|
||||
|
||||
### Architecture
|
||||
Dual-backend: PostgreSQL + libSQL/Turso. **All new persistence features must support both backends.** See `src/db/CLAUDE.md` and `.claude/rules/database.md`.
|
||||
|
||||
When designing new features or systems, always prefer generic/extensible architectures over hardcoding specific integrations. Ask clarifying questions about the desired abstraction level before implementing.
|
||||
## Module Specs
|
||||
|
||||
### Error Handling
|
||||
- Use `thiserror` for error types in `error.rs`
|
||||
- Never use `.unwrap()` or `.expect()` in production code (tests are fine)
|
||||
- Map errors with context: `.map_err(|e| SomeError::Variant { reason: e.to_string() })?`
|
||||
- Before committing, grep for `.unwrap()` and `.expect(` in changed files to catch violations mechanically
|
||||
When modifying a module with a spec, read the spec first. Code follows spec; spec is the tiebreaker.
|
||||
|
||||
### Async
|
||||
- All I/O is async with tokio
|
||||
- Use `Arc<T>` for shared state across tasks
|
||||
- Use `RwLock` for concurrent read/write access
|
||||
**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.
|
||||
|
||||
### Traits for Extensibility
|
||||
- `Database` - Add new database backends (must implement all ~60 methods)
|
||||
- `Channel` - Add new input sources
|
||||
- `Tool` - Add new capabilities
|
||||
- `LlmProvider` - Add new LLM backends
|
||||
- `SuccessEvaluator` - Custom evaluation logic
|
||||
- `EmbeddingProvider` - Add embedding backends (workspace search)
|
||||
- `NetworkPolicyDecider` - Custom network access policies for sandbox containers
|
||||
| Module | Spec |
|
||||
|--------|------|
|
||||
| `src/agent/` | `src/agent/CLAUDE.md` |
|
||||
| `src/channels/web/` | `src/channels/web/CLAUDE.md` |
|
||||
| `src/db/` | `src/db/CLAUDE.md` |
|
||||
| `src/llm/` | `src/llm/CLAUDE.md` |
|
||||
| `src/setup/` | `src/setup/README.md` |
|
||||
| `src/tools/` | `src/tools/README.md` |
|
||||
| `src/workspace/` | `src/workspace/README.md` |
|
||||
| `tests/e2e/` | `tests/e2e/CLAUDE.md` |
|
||||
|
||||
### Tool Implementation
|
||||
```rust
|
||||
#[async_trait]
|
||||
impl Tool for MyTool {
|
||||
fn name(&self) -> &str { "my_tool" }
|
||||
fn description(&self) -> &str { "Does something useful" }
|
||||
fn parameters_schema(&self) -> serde_json::Value {
|
||||
serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"param": { "type": "string", "description": "A parameter" }
|
||||
},
|
||||
"required": ["param"]
|
||||
})
|
||||
}
|
||||
## Job State Machine
|
||||
|
||||
async fn execute(&self, params: serde_json::Value, ctx: &JobContext)
|
||||
-> Result<ToolOutput, ToolError>
|
||||
{
|
||||
let start = std::time::Instant::now();
|
||||
// ... do work ...
|
||||
Ok(ToolOutput::text("result", start.elapsed()))
|
||||
}
|
||||
|
||||
fn requires_sanitization(&self) -> bool { true } // External data
|
||||
}
|
||||
```
|
||||
|
||||
### State Transitions
|
||||
Job states follow a defined state machine in `context/state.rs`:
|
||||
```
|
||||
Pending -> InProgress -> Completed -> Submitted -> Accepted
|
||||
\-> Failed
|
||||
@@ -292,397 +197,43 @@ Pending -> InProgress -> Completed -> Submitted -> Accepted
|
||||
\-> Failed
|
||||
```
|
||||
|
||||
### Code Style
|
||||
## Skills System
|
||||
|
||||
- Use `crate::` imports, not `super::`
|
||||
- No `pub use` re-exports unless exposing to downstream consumers
|
||||
- Prefer strong types over strings (enums, newtypes)
|
||||
- Keep functions focused, extract helpers when logic is reused
|
||||
- Comments for non-obvious logic only
|
||||
SKILL.md files extend the agent's prompt with domain-specific instructions. See `.claude/rules/skills.md` for full details.
|
||||
|
||||
### Review & Fix Discipline
|
||||
|
||||
Hard-won lessons from code review -- follow these when fixing bugs or addressing review feedback.
|
||||
|
||||
**Fix the pattern, not just the instance:** When a reviewer flags a bug (e.g., TOCTOU race in INSERT + SELECT-back), search the entire codebase for all instances of that same pattern. A fix in `SecretsStore::create()` that doesn't also fix `WasmToolStore::store()` is half a fix.
|
||||
|
||||
**Propagate architectural fixes to satellite types:** If a core type changes its concurrency model (e.g., `LibSqlBackend` switches to connection-per-operation), every type that was handed a resource from the old model (e.g., `LibSqlSecretsStore`, `LibSqlWasmToolStore` holding a single `Connection`) must also be updated. Grep for the old type across the codebase.
|
||||
|
||||
**Schema translation is more than DDL:** When translating a database schema between backends (PostgreSQL to libSQL, etc.), check for:
|
||||
- **Indexes** -- diff `CREATE INDEX` statements between the two schemas
|
||||
- **Seed data** -- check for `INSERT INTO` in migrations (e.g., `leak_detection_patterns`)
|
||||
- **Semantic differences** -- document where SQL functions behave differently (e.g., `json_patch` vs `jsonb_set`)
|
||||
|
||||
**Feature flag testing:** When adding feature-gated code, test compilation with each feature in isolation:
|
||||
```bash
|
||||
cargo check # default features
|
||||
cargo check --no-default-features --features libsql # libsql only
|
||||
cargo check --all-features # all features
|
||||
```
|
||||
Dead code behind the wrong `#[cfg]` gate will only show up when building with a single feature.
|
||||
|
||||
**Regression test with every fix:** Every bug fix must include a test that would have caught the bug. Add a `#[test]` or `#[tokio::test]` that reproduces the original failure. Exempt: changes limited to `src/channels/web/static/` or `.md` files. Use `[skip-regression-check]` in commit message or PR label if genuinely not feasible. The `commit-msg` hook and CI workflow enforce this automatically.
|
||||
|
||||
**Zero clippy warnings policy:** Fix ALL clippy warnings before committing, including pre-existing ones in files you didn't change. Never leave warnings behind — treat `cargo clippy` output as a zero-tolerance gate.
|
||||
|
||||
**Mechanical verification before committing:** Run these checks on changed files before committing:
|
||||
- `cargo clippy --all --benches --tests --examples --all-features` -- zero warnings
|
||||
- `grep -rnE '\.unwrap\(|\.expect\(' <files>` -- no panics in production
|
||||
- `grep -rn 'super::' <files>` -- use `crate::` imports
|
||||
- If you fixed a pattern bug, `grep` for other instances of that pattern across `src/`
|
||||
- Fix commits must include regression tests (enforced by `commit-msg` hook; bypass with `[skip-regression-check]`)
|
||||
- **Trust model**: Trusted (user-placed in `~/.ironclaw/skills/` or workspace `skills/`, full tool access) vs Installed (registry, read-only tools)
|
||||
- **Selection pipeline**: gating (check bin/env/config requirements) -> scoring (keywords/patterns/tags) -> budget (fit within `SKILLS_MAX_TOKENS`) -> attenuation (trust-based tool ceiling)
|
||||
- **Skill tools**: `skill_list`, `skill_search`, `skill_install`, `skill_remove`
|
||||
|
||||
## Configuration
|
||||
|
||||
Environment variables (see `.env.example`):
|
||||
```bash
|
||||
# Database backend (default: postgres)
|
||||
DATABASE_BACKEND=postgres # or "libsql" / "turso"
|
||||
DATABASE_URL=postgres://user:pass@localhost/ironclaw
|
||||
LIBSQL_PATH=~/.ironclaw/ironclaw.db # libSQL local path (default)
|
||||
# LIBSQL_URL=libsql://xxx.turso.io # Turso cloud (optional)
|
||||
# LIBSQL_AUTH_TOKEN=xxx # Required with LIBSQL_URL
|
||||
|
||||
# NEAR AI (when LLM_BACKEND=nearai, the default)
|
||||
# Two auth modes: session token (default) or API key
|
||||
# Session token auth (default): uses browser OAuth on first run
|
||||
NEARAI_SESSION_TOKEN=sess_... # hosting providers: set this
|
||||
NEARAI_BASE_URL=https://private.near.ai
|
||||
# API key auth: set NEARAI_API_KEY, base URL defaults to cloud-api.near.ai
|
||||
# NEARAI_API_KEY=... # API key from cloud.near.ai
|
||||
NEARAI_MODEL=claude-3-5-sonnet-20241022
|
||||
|
||||
# Agent settings
|
||||
AGENT_NAME=ironclaw
|
||||
MAX_PARALLEL_JOBS=5
|
||||
|
||||
# Embeddings (for semantic memory search)
|
||||
OPENAI_API_KEY=sk-... # For OpenAI embeddings
|
||||
# Or use NEAR AI embeddings:
|
||||
# EMBEDDING_PROVIDER=nearai
|
||||
# EMBEDDING_ENABLED=true
|
||||
EMBEDDING_MODEL=text-embedding-3-small # or text-embedding-3-large
|
||||
|
||||
# Heartbeat (proactive periodic execution)
|
||||
HEARTBEAT_ENABLED=true
|
||||
HEARTBEAT_INTERVAL_SECS=1800 # 30 minutes
|
||||
HEARTBEAT_NOTIFY_CHANNEL=tui
|
||||
HEARTBEAT_NOTIFY_USER=default
|
||||
|
||||
# Web gateway
|
||||
GATEWAY_ENABLED=true
|
||||
GATEWAY_HOST=127.0.0.1
|
||||
GATEWAY_PORT=3001
|
||||
GATEWAY_AUTH_TOKEN=changeme # Required for API access
|
||||
GATEWAY_USER_ID=default
|
||||
|
||||
# Docker sandbox
|
||||
SANDBOX_ENABLED=true
|
||||
SANDBOX_IMAGE=ironclaw-worker:latest
|
||||
SANDBOX_MEMORY_LIMIT_MB=512
|
||||
SANDBOX_TIMEOUT_SECS=1800
|
||||
SANDBOX_CPU_LIMIT=1.0 # CPU cores per container
|
||||
SANDBOX_NETWORK_PROXY=true # Enable network proxy for containers
|
||||
SANDBOX_PROXY_PORT=8080 # Proxy listener port
|
||||
SANDBOX_DEFAULT_POLICY=workspace_write # ReadOnly, WorkspaceWrite, FullAccess
|
||||
|
||||
# Claude Code mode (runs inside sandbox containers)
|
||||
CLAUDE_CODE_ENABLED=false
|
||||
CLAUDE_CODE_MODEL=claude-sonnet-4-20250514
|
||||
CLAUDE_CODE_MAX_TURNS=50
|
||||
CLAUDE_CODE_CONFIG_DIR=/home/worker/.claude
|
||||
|
||||
# Routines (scheduled/reactive execution)
|
||||
ROUTINES_ENABLED=true
|
||||
ROUTINES_CRON_INTERVAL=60 # Tick interval in seconds
|
||||
ROUTINES_MAX_CONCURRENT=3
|
||||
|
||||
# Skills system
|
||||
SKILLS_ENABLED=true
|
||||
SKILLS_MAX_TOKENS=4000 # Max prompt budget per turn
|
||||
SKILLS_CATALOG_URL=https://clawhub.dev # ClawHub registry URL
|
||||
SKILLS_AUTO_DISCOVER=true # Scan skill directories on startup
|
||||
|
||||
# Tinfoil private inference
|
||||
TINFOIL_API_KEY=... # Required when LLM_BACKEND=tinfoil
|
||||
TINFOIL_MODEL=kimi-k2-5 # Default model
|
||||
```
|
||||
|
||||
### LLM Providers
|
||||
|
||||
IronClaw supports multiple LLM backends via the `LLM_BACKEND` env var: `nearai` (default), `openai`, `anthropic`, `ollama`, `openai_compatible`, and `tinfoil`.
|
||||
|
||||
**NEAR AI** -- Uses the Chat Completions API with dual auth support. Session token auth (default): authenticates with session tokens (`sess_xxx`) obtained via browser OAuth (GitHub/Google), base URL defaults to `https://private.near.ai`. API key auth: set `NEARAI_API_KEY` (from `cloud.near.ai`), base URL defaults to `https://cloud-api.near.ai`. Both modes use the same Chat Completions endpoint. Tool messages are flattened to plain text for compatibility. Set `NEARAI_SESSION_TOKEN` env var for hosting providers that inject tokens via environment.
|
||||
|
||||
**NEAR AI Cloud** -- Uses the OpenAI-compatible Chat Completions API (`https://cloud-api.near.ai/v1/chat/completions`). Authenticates with API keys from `cloud.near.ai`. Auto-selected when `NEARAI_API_KEY` is set (or explicitly via `NEARAI_API_MODE=chat_completions`). Tool messages are flattened to plain text for compatibility. Configure with `NEARAI_API_KEY` and `NEARAI_BASE_URL` (default: `https://cloud-api.near.ai`).
|
||||
|
||||
**OpenAI-compatible** -- Any endpoint that speaks the OpenAI API (vLLM, LiteLLM, OpenRouter, etc.). Configure with `LLM_BASE_URL`, `LLM_API_KEY` (optional), `LLM_MODEL`. Set `LLM_EXTRA_HEADERS` to inject custom HTTP headers into every request (format: `Key:Value,Key2:Value2`), useful for OpenRouter attribution headers like `HTTP-Referer` and `X-Title`.
|
||||
|
||||
**Tinfoil** -- Private inference via `https://inference.tinfoil.sh/v1`. Runs models inside hardware-attested TEEs so neither Tinfoil nor the cloud provider can see prompts or responses. Uses the OpenAI-compatible Chat Completions API. Configure with `TINFOIL_API_KEY` and `TINFOIL_MODEL` (default: `kimi-k2-5`).
|
||||
|
||||
## Database
|
||||
|
||||
IronClaw supports two database backends, selected at compile time via Cargo feature flags and at runtime via the `DATABASE_BACKEND` environment variable.
|
||||
|
||||
**IMPORTANT: All new features that touch persistence MUST support both backends.** Implement the operation as a method on the `Database` trait in `src/db/mod.rs`, then add the implementation in both `src/db/postgres.rs` (delegate to Store/Repository) and `src/db/libsql_backend.rs` (native SQL).
|
||||
|
||||
### Backends
|
||||
|
||||
| Backend | Feature Flag | Default | Use Case |
|
||||
|---------|-------------|---------|----------|
|
||||
| PostgreSQL | `postgres` (default) | Yes | Production, existing deployments |
|
||||
| libSQL/Turso | `libsql` | No | Zero-dependency local mode, edge, Turso cloud |
|
||||
|
||||
```bash
|
||||
# Build with PostgreSQL only (default)
|
||||
cargo build
|
||||
|
||||
# Build with libSQL only
|
||||
cargo build --no-default-features --features libsql
|
||||
|
||||
# Build with both backends available
|
||||
cargo build --features "postgres,libsql"
|
||||
```
|
||||
|
||||
### Database Trait
|
||||
|
||||
The `Database` trait (`src/db/mod.rs`) defines ~60 async methods covering all persistence:
|
||||
- Conversations, messages, metadata
|
||||
- Jobs, actions, LLM calls, estimation snapshots
|
||||
- Sandbox jobs, job events
|
||||
- Routines, routine runs
|
||||
- Tool failures, settings
|
||||
- Workspace: documents, chunks, hybrid search
|
||||
|
||||
Both backends implement this trait. PostgreSQL delegates to the existing `Store` + `Repository`. libSQL implements native SQLite-dialect SQL.
|
||||
|
||||
### Schema
|
||||
|
||||
**PostgreSQL:** `migrations/V1__initial.sql` (351 lines). Uses pgvector for embeddings, tsvector for FTS, PL/pgSQL functions. Managed by `refinery`.
|
||||
|
||||
**libSQL:** `src/db/libsql_migrations.rs` (consolidated schema, ~480 lines). Translates PG types:
|
||||
- `UUID` -> `TEXT`, `TIMESTAMPTZ` -> `TEXT` (ISO-8601), `JSONB` -> `TEXT`
|
||||
- `VECTOR(1536)` -> `F32_BLOB(1536)` with `libsql_vector_idx`
|
||||
- `tsvector`/`ts_rank_cd` -> FTS5 virtual table with sync triggers
|
||||
- PL/pgSQL functions -> SQLite triggers
|
||||
|
||||
**Tables (both backends):**
|
||||
|
||||
**Core:**
|
||||
- `conversations` - Multi-channel conversation tracking
|
||||
- `agent_jobs` - Job metadata and status
|
||||
- `job_actions` - Event-sourced tool executions
|
||||
- `dynamic_tools` - Agent-built tools
|
||||
- `llm_calls` - Cost tracking
|
||||
- `estimation_snapshots` - Learning data
|
||||
|
||||
**Workspace/Memory:**
|
||||
- `memory_documents` - Flexible path-based files (e.g., "context/vision.md", "daily/2024-01-15.md")
|
||||
- `memory_chunks` - Chunked content with FTS and vector indexes
|
||||
- `heartbeat_state` - Periodic execution tracking
|
||||
|
||||
**Other:**
|
||||
- `routines`, `routine_runs` - Scheduled/reactive execution
|
||||
- `settings` - Per-user key-value settings
|
||||
- `tool_failures` - Self-repair tracking
|
||||
- `secrets`, `wasm_tools`, `tool_capabilities` - Extension infrastructure
|
||||
|
||||
Database configuration: see Configuration section above.
|
||||
|
||||
### Current Limitations (libSQL backend)
|
||||
|
||||
- **Workspace/memory system** not yet wired through Database trait (requires Store migration)
|
||||
- **Secrets store** not yet available (still requires PostgresSecretsStore)
|
||||
- **Hybrid search** uses FTS5 only (vector search via libsql_vector_idx not yet implemented)
|
||||
- **Settings reload from DB** skipped (Config::from_db requires Store)
|
||||
- No incremental migration versioning (schema is CREATE IF NOT EXISTS, no ALTER TABLE support yet)
|
||||
- **No encryption at rest** -- The local SQLite database file stores conversation content, job data, workspace memory, and other application data in plaintext. Only secrets (API tokens, credentials) are encrypted via AES-256-GCM before storage. Users handling sensitive data should use full-disk encryption (FileVault, LUKS, BitLocker) or consider the PostgreSQL backend with TDE/encrypted storage.
|
||||
- **JSON merge patch vs path-targeted update** -- The libSQL backend uses RFC 7396 JSON Merge Patch (`json_patch`) for metadata updates, while PostgreSQL uses path-targeted `jsonb_set`. Merge patch replaces top-level keys entirely, which may drop nested keys not present in the patch. Callers should avoid relying on partial nested object updates in metadata fields.
|
||||
|
||||
## Safety Layer
|
||||
|
||||
All external tool output passes through `SafetyLayer`:
|
||||
1. **Sanitizer** - Detects injection patterns, escapes dangerous content
|
||||
2. **Validator** - Checks length, encoding, forbidden patterns
|
||||
3. **Policy** - Rules with severity (Critical/High/Medium/Low) and actions (Block/Warn/Review/Sanitize)
|
||||
4. **Leak Detector** - Scans for 15+ secret patterns (API keys, tokens, private keys, connection strings) at two points: tool output before it reaches the LLM, and LLM responses before they reach the user. Actions per pattern: Block (reject entirely), Redact (mask the secret), or Warn (flag but allow)
|
||||
|
||||
Tool outputs are wrapped before reaching LLM:
|
||||
```xml
|
||||
<tool_output name="search" sanitized="true">
|
||||
[escaped content]
|
||||
</tool_output>
|
||||
```
|
||||
|
||||
### Shell Environment Scrubbing
|
||||
|
||||
The shell tool (`src/tools/builtin/shell.rs`) scrubs sensitive environment variables before executing commands, preventing secrets from leaking through `env`, `printenv`, or `$VAR` expansion. The sanitizer (`src/safety/sanitizer.rs`) also detects command injection patterns (chained commands, subshells, path traversal) and blocks or escapes them based on policy rules.
|
||||
|
||||
## Skills System
|
||||
|
||||
Skills are SKILL.md files that extend the agent's prompt with domain-specific instructions. Each skill is a YAML frontmatter block (metadata, activation criteria, required tools) followed by a markdown body that gets injected into the LLM context when the skill activates.
|
||||
|
||||
### Trust Model
|
||||
|
||||
| Trust Level | Source | Tool Access |
|
||||
|-------------|--------|-------------|
|
||||
| **Trusted** | User-placed in `~/.ironclaw/skills/` or workspace `skills/` | All tools available to the agent |
|
||||
| **Installed** | Downloaded from ClawHub registry | Read-only tools only (no shell, file write, HTTP) |
|
||||
|
||||
### SKILL.md Format
|
||||
|
||||
```yaml
|
||||
---
|
||||
name: my-skill
|
||||
version: 0.1.0
|
||||
description: Does something useful
|
||||
activation:
|
||||
patterns:
|
||||
- "deploy to.*production"
|
||||
keywords:
|
||||
- "deployment"
|
||||
max_context_tokens: 2000
|
||||
metadata:
|
||||
openclaw:
|
||||
requires:
|
||||
bins: [docker, kubectl]
|
||||
env: [KUBECONFIG]
|
||||
---
|
||||
|
||||
# Deployment Skill
|
||||
|
||||
Instructions for the agent when this skill activates...
|
||||
```
|
||||
|
||||
### Selection Pipeline
|
||||
|
||||
1. **Gating** -- Check binary/env/config requirements; skip skills whose prerequisites are missing
|
||||
2. **Scoring** -- Deterministic scoring against message content using keywords, tags, and regex patterns
|
||||
3. **Budget** -- Select top-scoring skills that fit within `SKILLS_MAX_TOKENS` prompt budget
|
||||
4. **Attenuation** -- Apply trust-based tool ceiling; installed skills lose access to dangerous tools
|
||||
|
||||
### Skill Tools
|
||||
|
||||
Four built-in tools for managing skills at runtime:
|
||||
- **`skill_list`** -- List all discovered skills with trust level and status
|
||||
- **`skill_search`** -- Search ClawHub registry for available skills
|
||||
- **`skill_install`** -- Download and install a skill from ClawHub
|
||||
- **`skill_remove`** -- Remove an installed skill
|
||||
|
||||
### Skill Directories
|
||||
|
||||
- `~/.ironclaw/skills/` -- User's global skills (trusted)
|
||||
- `<workspace>/skills/` -- Per-workspace skills (trusted)
|
||||
- `~/.ironclaw/installed_skills/` -- Registry-installed skills (installed trust)
|
||||
|
||||
### Testing Skills
|
||||
|
||||
- `skills/web-ui-test/` -- Manual test checklist for the web gateway UI via Claude for Chrome extension. Covers connection, chat, skills search/install/remove, and other tabs.
|
||||
|
||||
Skills configuration: see Configuration section above.
|
||||
|
||||
## Docker Sandbox
|
||||
|
||||
The `src/sandbox/` module provides Docker-based isolation for job execution with a network proxy that controls outbound access and injects credentials.
|
||||
|
||||
### Sandbox Policies
|
||||
|
||||
| Policy | Filesystem | Network | Use Case |
|
||||
|--------|-----------|---------|----------|
|
||||
| **ReadOnly** | Read-only workspace mount | Allowlisted domains only | Analysis, code review |
|
||||
| **WorkspaceWrite** | Read-write workspace mount | Allowlisted domains only | Code generation, file edits |
|
||||
| **FullAccess** | Full filesystem | Unrestricted | Trusted admin tasks |
|
||||
|
||||
### Network Proxy
|
||||
|
||||
Containers route all HTTP/HTTPS traffic through a host-side proxy (`src/sandbox/proxy/`):
|
||||
- **Domain allowlist** -- Only allowlisted domains are reachable (default: package registries, docs sites, GitHub, common APIs)
|
||||
- **Credential injection** -- The `CredentialResolver` trait injects auth headers into proxied requests so secrets never enter the container environment
|
||||
- **CONNECT tunnel** -- HTTPS traffic uses CONNECT method; the proxy validates the target domain against the allowlist before establishing the tunnel
|
||||
- **Policy decisions** -- The `NetworkPolicyDecider` trait allows custom logic for allow/deny/inject decisions per request
|
||||
|
||||
### Zero-Exposure Credential Model
|
||||
|
||||
Secrets (API keys, tokens) are stored encrypted on the host and injected into HTTP requests by the proxy at transit time. Container processes never have access to raw credential values, preventing exfiltration even if container code is compromised.
|
||||
|
||||
Sandbox configuration: see Configuration section above.
|
||||
|
||||
## Testing
|
||||
|
||||
Tests are in `mod tests {}` blocks at the bottom of each file. Run specific module tests:
|
||||
```bash
|
||||
cargo test safety::sanitizer::tests
|
||||
cargo test tools::registry::tests
|
||||
```
|
||||
|
||||
Key test patterns:
|
||||
- Unit tests for pure functions
|
||||
- Async tests with `#[tokio::test]`
|
||||
- No mocks, prefer real implementations or stubs
|
||||
|
||||
## Current Limitations / TODOs
|
||||
|
||||
1. **Domain-specific tools** - `marketplace.rs`, `restaurant.rs`, `taskrabbit.rs`, `ecommerce.rs` return placeholder responses; need real API integrations
|
||||
2. **Integration tests** - Need testcontainers setup for PostgreSQL
|
||||
3. **MCP stdio transport** - Only HTTP transport implemented
|
||||
4. **WIT bindgen integration** - Auto-extract tool description/schema from WASM modules (stubbed)
|
||||
5. **Capability granting after tool build** - Built tools get empty capabilities; need UX for granting HTTP/secrets access
|
||||
6. **Tool versioning workflow** - No version tracking or rollback for dynamically built tools
|
||||
7. **Webhook trigger endpoint** - Routines webhook trigger not yet exposed in web gateway
|
||||
8. **Full channel status view** - Gateway status widget exists, but no per-channel connection dashboard
|
||||
|
||||
## Tool Architecture
|
||||
|
||||
**Keep tool-specific logic out of the main agent codebase.** The main agent provides generic infrastructure; tools are self-contained units that declare their requirements through `capabilities.json` files (API endpoints, credentials, rate limits, auth setup). Service-specific auth flows, CLI commands, and configuration do not belong in the main agent.
|
||||
|
||||
Tools can be built as **WASM** (sandboxed, credential-injected, single binary) or **MCP servers** (ecosystem of pre-built servers, any language, but no sandbox). Both are first-class via `ironclaw tool install`. Auth is declared in capabilities files with OAuth and manual token entry support.
|
||||
|
||||
See `src/tools/README.md` for full tool architecture, adding new tools (built-in Rust and WASM), auth JSON examples, and WASM vs MCP decision guide.
|
||||
See `.env.example` for all environment variables. LLM backends (`nearai`, `openai`, `anthropic`, `ollama`, `openai_compatible`, `tinfoil`, `bedrock`) documented in `src/llm/CLAUDE.md`.
|
||||
|
||||
## Adding a New Channel
|
||||
|
||||
1. Create `src/channels/my_channel.rs`
|
||||
2. Implement the `Channel` trait
|
||||
3. Add config in `src/config.rs`
|
||||
4. Wire up in `main.rs` channel setup section
|
||||
3. Add config in `src/config/channels.rs`
|
||||
4. Wire up in `src/app.rs` channel setup section
|
||||
|
||||
## Workspace & Memory
|
||||
|
||||
Persistent memory with hybrid search (FTS + vector via RRF). Four tools: `memory_search`, `memory_write`, `memory_read`, `memory_tree`. Identity files (AGENTS.md, SOUL.md, USER.md, IDENTITY.md) injected into system prompt. Heartbeat system runs proactive periodic execution (default: 30 minutes), reading `HEARTBEAT.md` and notifying via channel if findings. See `src/workspace/README.md`.
|
||||
|
||||
## Debugging
|
||||
|
||||
```bash
|
||||
# Verbose logging
|
||||
RUST_LOG=ironclaw=trace cargo run
|
||||
|
||||
# Just the agent module
|
||||
RUST_LOG=ironclaw::agent=debug cargo run
|
||||
|
||||
# With HTTP request logging
|
||||
RUST_LOG=ironclaw=debug,tower_http=debug cargo run
|
||||
RUST_LOG=ironclaw=trace cargo run # verbose
|
||||
RUST_LOG=ironclaw::agent=debug cargo run # agent module only
|
||||
RUST_LOG=ironclaw=debug,tower_http=debug cargo run # + HTTP request logging
|
||||
```
|
||||
|
||||
## Module Specifications
|
||||
## Current Limitations
|
||||
|
||||
Some modules have a `README.md` that serves as the authoritative specification
|
||||
for that module's behavior. When modifying code in a module that has a spec:
|
||||
|
||||
1. **Read the spec first** before making changes
|
||||
2. **Code follows spec**: if the spec says X, the code must do X
|
||||
3. **Update both sides**: if you change behavior, update the spec to match;
|
||||
if you're implementing a spec change, update the code to match
|
||||
4. **Spec is the tiebreaker**: when code and spec disagree, the spec is correct
|
||||
(unless the spec is clearly outdated, in which case fix the spec first)
|
||||
|
||||
| Module | Spec File |
|
||||
|--------|-----------|
|
||||
| `src/setup/` | `src/setup/README.md` |
|
||||
| `src/workspace/` | `src/workspace/README.md` |
|
||||
| `src/tools/` | `src/tools/README.md` |
|
||||
|
||||
## Workspace & Memory System
|
||||
|
||||
OpenClaw-inspired persistent memory with a flexible filesystem-like structure. Principle: "Memory is database, not RAM" -- if you want to remember something, write it explicitly. Uses hybrid search combining FTS (keyword) + vector (semantic) via Reciprocal Rank Fusion.
|
||||
|
||||
Four memory tools for LLM use: `memory_search` (hybrid search -- call before answering questions about prior work), `memory_write`, `memory_read`, `memory_tree`. Identity files (AGENTS.md, SOUL.md, USER.md, IDENTITY.md) are injected into the LLM system prompt.
|
||||
|
||||
The heartbeat system runs proactive periodic execution (default: 30 minutes), reading `HEARTBEAT.md` and notifying via channel if findings are detected.
|
||||
|
||||
See `src/workspace/README.md` for full API documentation, filesystem structure, hybrid search details, chunking strategy, and heartbeat system.
|
||||
1. Domain-specific tools (`marketplace.rs`, `restaurant.rs`, etc.) are stubs
|
||||
2. Integration tests need testcontainers for PostgreSQL
|
||||
3. MCP: no streaming support; stdio/HTTP/Unix transports all use request-response
|
||||
4. WIT bindgen: auto-extract tool schema from WASM is stubbed
|
||||
5. Built tools get empty capabilities; need UX for granting access
|
||||
6. No tool versioning or rollback
|
||||
7. Observability: only `log` and `noop` backends (no OpenTelemetry)
|
||||
|
||||
@@ -1,5 +1,34 @@
|
||||
# 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
|
||||
|
||||
When your change affects a tracked capability, update `FEATURE_PARITY.md` in the same branch.
|
||||
@@ -9,3 +38,23 @@ When your change affects a tracked capability, update `FEATURE_PARITY.md` in the
|
||||
1. Review the relevant parity rows in `FEATURE_PARITY.md`.
|
||||
2. Update status/notes if behavior changed.
|
||||
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/`, `src/sandbox/`, `src/orchestrator/`), 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 (requires `deny.toml`; see the `cargo-deny` CI job).
|
||||
|
||||
@@ -0,0 +1,862 @@
|
||||
# IronClaw Coverage Plan: 63.3% to 95%
|
||||
|
||||
> Generated 2025-03-06 from [Codecov](https://app.codecov.io/gh/nearai/ironclaw/tree/main/src)
|
||||
|
||||
## Current State
|
||||
|
||||
| Metric | Value |
|
||||
|--------|-------|
|
||||
| **Current coverage** | 48,571 / 76,694 lines = **63.33%** |
|
||||
| **Target** | 72,859 / 76,694 lines = **95.0%** |
|
||||
| **Gap** | **24,288 lines** need coverage |
|
||||
| **Files >= 95%** | 43 / 239 |
|
||||
| **Files < 95%** | 196 (27,872 total misses) |
|
||||
|
||||
## Module Summary
|
||||
|
||||
Sorted by uncovered lines (descending):
|
||||
|
||||
| Module | Lines | Hits | Miss | Coverage | Priority |
|
||||
|--------|------:|-----:|-----:|---------:|----------|
|
||||
| `channels/` | 14,079 | 8,677 | 5,402 | 61.6% | P0 |
|
||||
| `tools/` | 13,445 | 9,407 | 4,038 | 70.0% | P1 |
|
||||
| `agent/` | 9,152 | 6,096 | 3,056 | 66.6% | P0 |
|
||||
| `setup/` | 3,005 | 462 | 2,543 | 15.4% | P1 |
|
||||
| `extensions/` | 3,540 | 1,298 | 2,242 | 36.7% | P0 |
|
||||
| `cli/` | 2,834 | 697 | 2,137 | 24.6% | P1 |
|
||||
| `history/` | 1,626 | 0 | 1,626 | 0.0% | P0 |
|
||||
| `llm/` | 7,029 | 5,776 | 1,253 | 82.2% | P2 |
|
||||
| `(root)` | 4,122 | 3,121 | 1,001 | 75.7% | P2 |
|
||||
| `worker/` | 1,274 | 480 | 794 | 37.7% | P1 |
|
||||
| `sandbox/` | 1,615 | 897 | 718 | 55.5% | P2 |
|
||||
| `registry/` | 1,588 | 1,107 | 481 | 69.7% | P2 |
|
||||
| `db/` | 921 | 441 | 480 | 47.9% | P1 |
|
||||
| `workspace/` | 2,006 | 1,584 | 422 | 79.0% | P2 |
|
||||
| `orchestrator/` | 1,199 | 795 | 404 | 66.3% | P2 |
|
||||
| `config/` | 1,464 | 1,095 | 369 | 74.8% | P2 |
|
||||
| `hooks/` | 1,379 | 1,081 | 298 | 78.4% | P2 |
|
||||
| `secrets/` | 687 | 407 | 280 | 59.2% | P2 |
|
||||
| `skills/` | 1,714 | 1,585 | 129 | 92.5% | P3 |
|
||||
| `context/` | 693 | 586 | 107 | 84.6% | P3 |
|
||||
| `estimation/` | 467 | 369 | 98 | 79.0% | P3 |
|
||||
| `safety/` | 1,424 | 1,337 | 87 | 93.9% | P3 |
|
||||
| `evaluation/` | 226 | 152 | 74 | 67.3% | P3 |
|
||||
| `pairing/` | 498 | 446 | 52 | 89.6% | P3 |
|
||||
| `tunnel/` | 391 | 368 | 23 | 94.1% | P3 |
|
||||
| `observability/` | 316 | 307 | 9 | 97.2% | Done |
|
||||
|
||||
## Top 40 Files by Uncovered Lines
|
||||
|
||||
These files account for the vast majority of the coverage gap:
|
||||
|
||||
| File | Lines | Miss | Coverage | Lines to 95% |
|
||||
|------|------:|-----:|---------:|--------------:|
|
||||
| `src/extensions/manager.rs` | 2,404 | 2,083 | 13.3% | 1,962 |
|
||||
| `src/setup/wizard.rs` | 2,150 | 1,789 | 16.8% | 1,681 |
|
||||
| `src/history/store.rs` | 1,486 | 1,486 | 0.0% | 1,411 |
|
||||
| `src/channels/web/server.rs` | 1,985 | 993 | 50.0% | 893 |
|
||||
| `src/channels/wasm/wrapper.rs` | 2,237 | 934 | 58.2% | 822 |
|
||||
| `src/agent/thread_ops.rs` | 1,044 | 763 | 26.9% | 710 |
|
||||
| `src/cli/tool.rs` | 757 | 735 | 2.9% | 697 |
|
||||
| `src/setup/channels.rs` | 645 | 596 | 7.6% | 563 |
|
||||
| `src/agent/commands.rs` | 587 | 587 | 0.0% | 557 |
|
||||
| `src/main.rs` | 740 | 522 | 29.4% | 485 |
|
||||
| `src/channels/web/handlers/jobs.rs` | 513 | 456 | 11.1% | 430 |
|
||||
| `src/tools/builder/core.rs` | 524 | 456 | 13.0% | 429 |
|
||||
| `src/agent/worker.rs` | 1,078 | 467 | 56.7% | 413 |
|
||||
| `src/channels/web/handlers/chat.rs` | 564 | 417 | 26.1% | 388 |
|
||||
| `src/tools/wasm/wrapper.rs` | 1,005 | 436 | 56.6% | 385 |
|
||||
| `src/channels/signal.rs` | 1,814 | 472 | 74.0% | 381 |
|
||||
| `src/tools/mcp/auth.rs` | 472 | 378 | 19.9% | 354 |
|
||||
| `src/worker/runtime.rs` | 350 | 330 | 5.7% | 312 |
|
||||
| `src/tools/builtin/job.rs` | 1,014 | 359 | 64.6% | 308 |
|
||||
| `src/cli/mcp.rs` | 322 | 319 | 0.9% | 302 |
|
||||
| `src/cli/oauth_defaults.rs` | 730 | 335 | 54.1% | 298 |
|
||||
| `src/llm/nearai_chat.rs` | 854 | 340 | 60.2% | 297 |
|
||||
| `src/sandbox/container.rs` | 407 | 317 | 22.1% | 296 |
|
||||
| `src/tools/mcp/client.rs` | 341 | 291 | 14.7% | 273 |
|
||||
| `src/registry/installer.rs` | 765 | 311 | 59.3% | 272 |
|
||||
| `src/orchestrator/job_manager.rs` | 405 | 270 | 33.3% | 249 |
|
||||
| `src/channels/web/handlers/routines.rs` | 249 | 249 | 0.0% | 236 |
|
||||
| `src/agent/scheduler.rs` | 559 | 263 | 53.0% | 235 |
|
||||
| `src/tools/wasm/storage.rs` | 296 | 243 | 17.9% | 228 |
|
||||
| `src/channels/repl.rs` | 233 | 233 | 0.0% | 221 |
|
||||
| `src/llm/session.rs` | 413 | 242 | 41.4% | 221 |
|
||||
| `src/worker/claude_bridge.rs` | 629 | 247 | 60.7% | 215 |
|
||||
| `src/agent/agent_loop.rs` | 523 | 234 | 55.2% | 207 |
|
||||
| `src/worker/api.rs` | 258 | 207 | 19.8% | 194 |
|
||||
| `src/sandbox/proxy/http.rs` | 307 | 192 | 37.5% | 176 |
|
||||
| `src/channels/wasm/storage.rs` | 182 | 182 | 0.0% | 172 |
|
||||
| `src/cli/registry.rs` | 177 | 177 | 0.0% | 168 |
|
||||
| `src/llm/reasoning.rs` | 1,163 | 219 | 81.2% | 160 |
|
||||
| `src/tools/builder/testing.rs` | 308 | 174 | 43.5% | 158 |
|
||||
| `src/db/postgres.rs` | 166 | 166 | 0.0% | 157 |
|
||||
|
||||
---
|
||||
|
||||
## Tier 1 -- High-Impact Unit Tests (~8,500 lines)
|
||||
|
||||
Pure logic, serialization, and database queries testable in isolation without real
|
||||
infrastructure. Highest coverage gain per unit of effort.
|
||||
|
||||
### `src/history/store.rs` -- 0% -> 95% (+1,411 lines)
|
||||
|
||||
PostgreSQL repository layer (conversations, jobs, actions, LLM calls, estimation
|
||||
snapshots). Test query construction and result mapping. Can use the libSQL backend
|
||||
as a real in-memory database or test doubles for the `Database` trait.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_store_conversation_crud` -- create, read, update, delete conversations
|
||||
- `test_store_job_lifecycle` -- insert job, update status through state machine
|
||||
- `test_store_action_recording` -- record and query job actions
|
||||
- `test_store_llm_call_tracking` -- insert and aggregate LLM call records
|
||||
- `test_store_estimation_snapshots` -- save and retrieve estimation data
|
||||
|
||||
### `src/history/analytics.rs` -- 0% -> 95% (+133 lines)
|
||||
|
||||
Aggregation queries (JobStats, ToolStats). Test the query builders and result
|
||||
deserialization.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_job_stats_aggregation` -- verify counts, durations, success rates
|
||||
- `test_tool_stats_ranking` -- verify tool usage frequency sorting
|
||||
- `test_analytics_empty_db` -- graceful handling of no data
|
||||
|
||||
### `src/extensions/manager.rs` -- 13.3% -> 95% (+1,962 lines)
|
||||
|
||||
Largest single file gap. Extension lifecycle orchestration (install, auth,
|
||||
activate, remove), config parsing, and state transitions.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_extension_install_from_manifest` -- parse manifest, create extension record
|
||||
- `test_extension_auth_flow` -- OAuth token setup, credential storage
|
||||
- `test_extension_activate_deactivate` -- state transitions, tool registration
|
||||
- `test_extension_remove_cleanup` -- remove extension, clean up artifacts
|
||||
- `test_extension_config_validation` -- reject invalid configs, handle defaults
|
||||
- `test_extension_list_filtering` -- filter by status, type, search query
|
||||
- `test_extension_capability_check` -- verify required capabilities before activation
|
||||
|
||||
### `src/extensions/discovery.rs` -- 27.8% -> 95% (+125 lines)
|
||||
|
||||
Extension discovery from filesystem and registry.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_discover_local_extensions` -- scan directory, parse manifests
|
||||
- `test_discover_skip_invalid` -- gracefully skip malformed extension dirs
|
||||
- `test_discover_dedup` -- handle duplicate extensions across paths
|
||||
|
||||
### `src/tools/builder/core.rs` -- 13% -> 95% (+429 lines)
|
||||
|
||||
`BuildRequirement`, `SoftwareType`, `Language` types and project scaffolding.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_build_requirement_parsing` -- deserialize from JSON
|
||||
- `test_scaffold_project_structure` -- verify generated file tree
|
||||
- `test_language_detection` -- detect language from file extensions
|
||||
- `test_software_type_constraints` -- validate type-specific requirements
|
||||
|
||||
### `src/tools/builder/testing.rs` -- 43.5% -> 95% (+158 lines)
|
||||
|
||||
Test harness integration for built tools.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_harness_setup_teardown` -- lifecycle of test environment
|
||||
- `test_harness_run_tests` -- execute tests and capture results
|
||||
- `test_harness_failure_reporting` -- verify error details on test failure
|
||||
|
||||
### `src/tools/mcp/auth.rs` -- 19.9% -> 95% (+354 lines)
|
||||
|
||||
OAuth token management for MCP servers.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_token_refresh_on_expiry` -- auto-refresh when token expires
|
||||
- `test_token_header_injection` -- correct Authorization header format
|
||||
- `test_token_persistence` -- save/load tokens across restarts
|
||||
- `test_oauth_pkce_flow` -- code verifier/challenge generation
|
||||
- `test_auth_config_parsing` -- parse various auth config formats
|
||||
|
||||
### `src/tools/mcp/client.rs` -- 14.7% -> 95% (+273 lines)
|
||||
|
||||
JSON-RPC client for MCP protocol.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_jsonrpc_request_serialization` -- correct JSON-RPC 2.0 format
|
||||
- `test_jsonrpc_response_parsing` -- handle success, error, and batch responses
|
||||
- `test_jsonrpc_error_codes` -- map MCP error codes to ToolError
|
||||
- `test_tool_list_discovery` -- parse tools/list response
|
||||
- `test_tool_call_roundtrip` -- serialize call, parse result
|
||||
|
||||
### `src/tools/wasm/storage.rs` -- 17.9% -> 95% (+228 lines)
|
||||
|
||||
WASM tool persistence (store, load, delete, list).
|
||||
|
||||
**Tests to write:**
|
||||
- `test_wasm_tool_store_roundtrip` -- store and retrieve tool binary + metadata
|
||||
- `test_wasm_tool_delete` -- remove tool and verify gone
|
||||
- `test_wasm_tool_list_filtering` -- filter by name, capability
|
||||
- `test_wasm_tool_update_metadata` -- update without re-uploading binary
|
||||
|
||||
### `src/tools/wasm/wrapper.rs` -- 56.6% -> 95% (+385 lines)
|
||||
|
||||
Tool trait wrapper for WASM modules.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_wasm_param_marshalling` -- JSON params to WASM component model types
|
||||
- `test_wasm_output_conversion` -- WASM return values to ToolOutput
|
||||
- `test_wasm_error_propagation` -- WASM traps to ToolError
|
||||
- `test_wasm_fuel_exhaustion` -- verify fuel limit enforcement
|
||||
- `test_wasm_memory_limit` -- verify memory ceiling
|
||||
|
||||
### `src/tools/wasm/loader.rs` -- 62.4% -> 95% (+156 lines)
|
||||
|
||||
WASM tool discovery from filesystem.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_loader_scan_directory` -- find .wasm files with capabilities.json
|
||||
- `test_loader_skip_invalid` -- skip files without valid WIT exports
|
||||
- `test_loader_cache_invalidation` -- reload when file changes
|
||||
|
||||
### `src/tools/builtin/job.rs` -- 64.6% -> 95% (+308 lines)
|
||||
|
||||
Job management tools (CreateJob, ListJobs, JobStatus, CancelJob).
|
||||
|
||||
**Tests to write:**
|
||||
- `test_create_job_params` -- validate required/optional parameters
|
||||
- `test_list_jobs_formatting` -- verify output structure
|
||||
- `test_job_status_transitions` -- query status at each state
|
||||
- `test_cancel_job_running` -- cancel an in-progress job
|
||||
- `test_cancel_job_completed` -- error on already-completed job
|
||||
|
||||
### `src/secrets/store.rs` -- 48.1% -> 95% (+145 lines)
|
||||
|
||||
Encrypted secret storage.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_secret_store_roundtrip` -- store encrypted, retrieve decrypted
|
||||
- `test_secret_update` -- overwrite existing secret
|
||||
- `test_secret_delete` -- remove and verify inaccessible
|
||||
- `test_secret_list_redacted` -- list shows names but not values
|
||||
|
||||
### `src/llm/session.rs` -- 41.4% -> 95% (+221 lines)
|
||||
|
||||
Session token management with auto-renewal.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_session_token_parsing` -- parse `sess_xxx` format
|
||||
- `test_session_expiry_detection` -- detect expired tokens
|
||||
- `test_session_auto_renewal` -- trigger renewal before expiry
|
||||
- `test_session_concurrent_renewal` -- only one renewal in flight
|
||||
|
||||
### `src/llm/nearai_chat.rs` -- 60.2% -> 95% (+297 lines)
|
||||
|
||||
NEAR AI Chat Completions provider.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_nearai_request_building` -- correct endpoint, headers, body
|
||||
- `test_nearai_response_parsing` -- parse streaming and non-streaming responses
|
||||
- `test_nearai_tool_message_flattening` -- tool messages flattened to text
|
||||
- `test_nearai_auth_modes` -- session token vs API key auth
|
||||
- `test_nearai_error_handling` -- rate limits, auth failures, server errors
|
||||
|
||||
### `src/llm/mod.rs` -- 53.7% -> 95% (+112 lines)
|
||||
|
||||
Provider factory and backend selection.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_provider_factory_nearai` -- select NEAR AI from config
|
||||
- `test_provider_factory_openai` -- select OpenAI from config
|
||||
- `test_provider_factory_ollama` -- select Ollama from config
|
||||
- `test_provider_factory_invalid` -- error on unknown backend
|
||||
|
||||
### `src/llm/reasoning.rs` -- 81.2% -> 95% (+160 lines)
|
||||
|
||||
Planning, tool selection, evaluation logic.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_reasoning_step_parsing` -- parse planning steps from LLM output
|
||||
- `test_tool_selection_scoring` -- rank tools by relevance
|
||||
- `test_evaluation_rubric` -- score completions against criteria
|
||||
- `test_reasoning_with_no_tools` -- handle tool-less responses
|
||||
|
||||
### `src/db/postgres.rs` -- 0% -> 95% (+157 lines)
|
||||
|
||||
PostgreSQL backend delegation to Store + Repository.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_postgres_backend_delegates` -- verify delegation pattern (trait-level)
|
||||
- `test_postgres_connection_config` -- TLS, pool size, timeout parsing
|
||||
|
||||
### `src/workspace/mod.rs` -- 75.9% -> 95% (+109 lines)
|
||||
|
||||
Memory operations (write, read, search, tree).
|
||||
|
||||
**Tests to write:**
|
||||
- `test_workspace_write_read` -- write document, read it back
|
||||
- `test_workspace_search_hybrid` -- FTS + vector search via RRF
|
||||
- `test_workspace_tree` -- directory listing of memory filesystem
|
||||
- `test_workspace_overwrite` -- update existing document
|
||||
|
||||
### `src/workspace/embeddings.rs` -- 35.1% -> 95% (~100 lines)
|
||||
|
||||
Embedding provider abstraction.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_embedding_dimension_handling` -- verify dimension config
|
||||
- `test_embedding_batch_processing` -- batch multiple chunks
|
||||
- `test_embedding_provider_fallback` -- graceful degradation when unavailable
|
||||
|
||||
---
|
||||
|
||||
## Tier 2 -- Trace Tests (~7,000 lines)
|
||||
|
||||
End-to-end tests that exercise the agent loop, worker, scheduler, and dispatcher
|
||||
by replaying LLM traces through `TestRig` (see `tests/support/test_rig.rs`). Each
|
||||
trace test covers multiple modules simultaneously, making them high-leverage.
|
||||
|
||||
Each trace test needs:
|
||||
1. A JSON fixture in `tests/fixtures/llm_traces/`
|
||||
2. A test file in `tests/` using `TestRigBuilder`
|
||||
|
||||
### Trace: Thread Operations
|
||||
|
||||
**Covers:** `agent/thread_ops.rs` (+710 lines)
|
||||
|
||||
Test thread creation, listing, switching, and deletion via trace replay.
|
||||
|
||||
**Fixture:** `thread_operations.json`
|
||||
**Tests:**
|
||||
- `test_thread_create_and_switch` -- create thread, switch to it, verify context
|
||||
- `test_thread_list` -- list all threads, verify metadata
|
||||
- `test_thread_delete` -- delete thread, verify removal
|
||||
- `test_thread_switch_nonexistent` -- error handling for missing thread
|
||||
|
||||
### Trace: Agent Commands
|
||||
|
||||
**Covers:** `agent/commands.rs` (+557 lines)
|
||||
|
||||
Test slash commands through the agent loop.
|
||||
|
||||
**Fixture:** `agent_commands.json`
|
||||
**Tests:**
|
||||
- `test_command_help` -- /help returns command list
|
||||
- `test_command_clear` -- /clear resets conversation
|
||||
- `test_command_compact` -- /compact triggers summarization
|
||||
- `test_command_undo_redo` -- /undo then /redo restores state
|
||||
- `test_command_status` -- /status shows agent state
|
||||
|
||||
### Trace: Worker Multi-Turn Execution
|
||||
|
||||
**Covers:** `agent/worker.rs` (+413 lines), `agent/agent_loop.rs` (+207 lines)
|
||||
|
||||
Test multi-turn tool calling, error recovery, and completion flows.
|
||||
|
||||
**Fixture:** `worker_multi_turn.json`
|
||||
**Tests:**
|
||||
- `test_worker_sequential_tools` -- call tool A, then tool B based on A's result
|
||||
- `test_worker_tool_error_recovery` -- tool fails, agent retries or adapts
|
||||
- `test_worker_max_turns` -- verify turn limit enforcement
|
||||
|
||||
### Trace: Scheduler Parallel Jobs
|
||||
|
||||
**Covers:** `agent/scheduler.rs` (+235 lines)
|
||||
|
||||
Test parallel job dispatch and completion tracking.
|
||||
|
||||
**Fixture:** `scheduler_parallel.json`
|
||||
**Tests:**
|
||||
- `test_scheduler_parallel_dispatch` -- dispatch 3 jobs, all complete
|
||||
- `test_scheduler_job_dependency` -- job B waits for job A
|
||||
- `test_scheduler_stuck_detection` -- detect and recover stuck job
|
||||
|
||||
### Trace: Dispatcher Skill Selection
|
||||
|
||||
**Covers:** `agent/dispatcher.rs` (+153 lines)
|
||||
|
||||
Test skill-aware routing and tool attenuation.
|
||||
|
||||
**Fixture:** `dispatcher_skills.json`
|
||||
**Tests:**
|
||||
- `test_dispatcher_skill_match` -- match message to skill, inject prompt
|
||||
- `test_dispatcher_tool_attenuation` -- installed skill loses dangerous tools
|
||||
- `test_dispatcher_no_skill` -- fallback when no skill matches
|
||||
|
||||
### Trace: Routine Execution
|
||||
|
||||
**Covers:** `agent/routine_engine.rs` (~80 lines), `agent/routine.rs` (~40 lines)
|
||||
|
||||
Test cron tick and event-triggered routine execution.
|
||||
|
||||
**Fixture:** `routine_execution.json`
|
||||
**Tests:**
|
||||
- `test_routine_cron_trigger` -- routine fires on schedule
|
||||
- `test_routine_event_trigger` -- routine fires on matching event
|
||||
- `test_routine_guardrails` -- routine respects policy constraints
|
||||
|
||||
### Trace: Compaction and Context Pressure
|
||||
|
||||
**Covers:** `agent/compaction.rs` (~50 lines), `agent/context_monitor.rs` (~30 lines)
|
||||
|
||||
Test turn summarization and memory pressure detection.
|
||||
|
||||
**Fixture:** `compaction_flow.json`
|
||||
**Tests:**
|
||||
- `test_compaction_triggers_at_threshold` -- summarize when context exceeds limit
|
||||
- `test_compaction_preserves_recent` -- keep recent turns intact
|
||||
- `test_context_pressure_warning` -- emit warning at high usage
|
||||
|
||||
### Trace: Job Tool Coverage
|
||||
|
||||
**Covers:** `tools/builtin/job.rs` (+308 lines), `tools/builtin/skill_tools.rs` (+110 lines)
|
||||
|
||||
Test job and skill management tools through agent execution.
|
||||
|
||||
**Fixture:** `job_and_skill_tools.json`
|
||||
**Tests:**
|
||||
- `test_create_and_list_jobs` -- create job, list shows it
|
||||
- `test_job_status_query` -- query status of running job
|
||||
- `test_skill_list_and_search` -- list local skills, search registry
|
||||
|
||||
### Trace: Memory Tools
|
||||
|
||||
**Covers:** `tools/builtin/memory.rs` (~20 lines), `workspace/` (+109 lines)
|
||||
|
||||
Test memory operations through agent tool calls.
|
||||
|
||||
**Fixture:** `memory_tools.json`
|
||||
**Tests:**
|
||||
- `test_memory_write_and_search` -- write doc, search finds it
|
||||
- `test_memory_read_by_path` -- read specific document
|
||||
- `test_memory_tree` -- list memory filesystem structure
|
||||
|
||||
### Trace: Extension Management
|
||||
|
||||
**Covers:** `tools/builtin/extension_tools.rs` (~40 lines)
|
||||
|
||||
Test extension lifecycle via agent tool calls.
|
||||
|
||||
**Fixture:** `extension_management.json`
|
||||
**Tests:**
|
||||
- `test_extension_install_via_tool` -- agent installs an extension
|
||||
- `test_extension_auth_via_tool` -- agent configures auth
|
||||
- `test_extension_activate_via_tool` -- agent activates extension
|
||||
|
||||
### Trace: Self-Repair
|
||||
|
||||
**Covers:** `agent/self_repair.rs` (~40 lines)
|
||||
|
||||
Test stuck job detection and recovery.
|
||||
|
||||
**Fixture:** `self_repair.json`
|
||||
**Tests:**
|
||||
- `test_stuck_job_detected` -- job stuck for > threshold triggers repair
|
||||
- `test_stuck_job_recovered` -- recovery restarts job successfully
|
||||
- `test_stuck_job_fails_permanently` -- recovery fails, job marked failed
|
||||
|
||||
### Trace: Heartbeat
|
||||
|
||||
**Covers:** `agent/heartbeat.rs` (+80 lines)
|
||||
|
||||
Test periodic proactive execution.
|
||||
|
||||
**Fixture:** `heartbeat.json`
|
||||
**Tests:**
|
||||
- `test_heartbeat_periodic_fire` -- heartbeat triggers at interval
|
||||
- `test_heartbeat_reads_checklist` -- reads HEARTBEAT.md, processes items
|
||||
- `test_heartbeat_notification` -- sends notification on findings
|
||||
|
||||
---
|
||||
|
||||
## Tier 3 -- Web/Channel Handler Tests (~4,500 lines)
|
||||
|
||||
Test HTTP handlers and SSE/WS endpoints using `axum_test` or
|
||||
`tower::ServiceExt::oneshot` with a real router and in-memory database.
|
||||
|
||||
### `src/channels/web/server.rs` -- 50% -> 95% (+893 lines)
|
||||
|
||||
The single biggest web gap. 40+ API endpoints.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_api_health` -- GET /health returns 200
|
||||
- `test_api_chat_submit` -- POST /api/chat sends message
|
||||
- `test_api_jobs_list` -- GET /api/jobs returns job list
|
||||
- `test_api_jobs_create` -- POST /api/jobs creates job
|
||||
- `test_api_routines_crud` -- full CRUD cycle for routines
|
||||
- `test_api_settings_get_set` -- GET/PUT settings
|
||||
- `test_api_memory_search` -- POST /api/memory/search
|
||||
- `test_api_extensions_list` -- GET /api/extensions
|
||||
- `test_api_skills_list` -- GET /api/skills
|
||||
- `test_api_sse_connect` -- SSE stream connects and receives events
|
||||
- `test_api_auth_required` -- endpoints reject missing/bad tokens
|
||||
- `test_api_cors_headers` -- verify CORS configuration
|
||||
|
||||
### `src/channels/web/handlers/chat.rs` -- 26.1% -> 95% (+388 lines)
|
||||
|
||||
Chat message submission and SSE streaming.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_chat_submit_message` -- submit message, receive response
|
||||
- `test_chat_sse_stream` -- verify SSE event format
|
||||
- `test_chat_thread_context` -- messages scoped to thread
|
||||
- `test_chat_invalid_payload` -- reject malformed requests
|
||||
|
||||
### `src/channels/web/handlers/jobs.rs` -- 11.1% -> 95% (+430 lines)
|
||||
|
||||
Job CRUD endpoints.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_jobs_list_empty` -- empty list returns []
|
||||
- `test_jobs_create_and_get` -- create, then GET by ID
|
||||
- `test_jobs_cancel` -- cancel running job
|
||||
- `test_jobs_filter_by_status` -- filter by pending/running/completed
|
||||
- `test_jobs_pagination` -- limit/offset parameters
|
||||
|
||||
### `src/channels/web/handlers/routines.rs` -- 0% -> 95% (+236 lines)
|
||||
|
||||
Routine CRUD endpoints.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_routines_create` -- POST creates routine
|
||||
- `test_routines_list` -- GET lists all routines
|
||||
- `test_routines_update` -- PUT updates routine config
|
||||
- `test_routines_delete` -- DELETE removes routine
|
||||
- `test_routines_history` -- GET history for a routine
|
||||
|
||||
### `src/channels/web/handlers/extensions.rs` -- 0% -> 95% (+129 lines)
|
||||
|
||||
Extension management endpoints.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_extensions_list` -- list installed extensions
|
||||
- `test_extensions_install` -- install from manifest URL
|
||||
- `test_extensions_activate` -- activate/deactivate toggle
|
||||
- `test_extensions_remove` -- remove installed extension
|
||||
|
||||
### `src/channels/web/handlers/memory.rs` -- 0% -> 95% (+110 lines)
|
||||
|
||||
Memory/workspace endpoints.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_memory_search` -- search returns ranked results
|
||||
- `test_memory_write` -- write a document
|
||||
- `test_memory_read` -- read by path
|
||||
- `test_memory_tree` -- tree returns filesystem structure
|
||||
|
||||
### `src/channels/web/handlers/settings.rs` -- 0% -> 95% (+103 lines)
|
||||
|
||||
Settings endpoints.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_settings_get` -- retrieve current settings
|
||||
- `test_settings_update` -- update individual setting
|
||||
- `test_settings_validation` -- reject invalid setting values
|
||||
|
||||
### `src/channels/web/handlers/static_files.rs` -- 0% -> 95% (+97 lines)
|
||||
|
||||
Static file serving.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_static_index_html` -- GET / serves index.html
|
||||
- `test_static_css_js` -- serve CSS/JS with correct content types
|
||||
- `test_static_404` -- missing file returns 404
|
||||
|
||||
### `src/channels/wasm/wrapper.rs` -- 58.2% -> 95% (+822 lines)
|
||||
|
||||
WASM channel wrapper (message routing, lifecycle).
|
||||
|
||||
**Tests to write:**
|
||||
- `test_wasm_channel_start` -- initialize WASM channel module
|
||||
- `test_wasm_channel_message_routing` -- route incoming message to WASM
|
||||
- `test_wasm_channel_response` -- return WASM response to caller
|
||||
- `test_wasm_channel_error_handling` -- handle WASM trap gracefully
|
||||
- `test_wasm_channel_lifecycle` -- start, process, shutdown
|
||||
|
||||
### `src/channels/wasm/loader.rs` -- 38.1% -> 95% (+141 lines)
|
||||
|
||||
WASM channel discovery.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_channel_loader_scan` -- find channel WASM modules
|
||||
- `test_channel_loader_validation` -- reject invalid modules
|
||||
- `test_channel_loader_manifest` -- parse channel capabilities
|
||||
|
||||
### `src/channels/wasm/storage.rs` -- 0% -> 95% (+172 lines)
|
||||
|
||||
WASM channel state persistence.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_channel_storage_save_load` -- persist and restore channel state
|
||||
- `test_channel_storage_isolation` -- per-channel state isolation
|
||||
- `test_channel_storage_cleanup` -- remove state on channel uninstall
|
||||
|
||||
### `src/channels/signal.rs` -- 74% -> 95% (+381 lines)
|
||||
|
||||
Signal protocol channel.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_signal_message_send` -- send encrypted message
|
||||
- `test_signal_message_receive` -- decrypt incoming message
|
||||
- `test_signal_attachment_handling` -- handle media attachments
|
||||
- `test_signal_group_message` -- group chat routing
|
||||
- `test_signal_error_handling` -- handle connection failures
|
||||
|
||||
### `src/channels/repl.rs` -- 0% -> 95% (+221 lines)
|
||||
|
||||
Simple REPL channel.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_repl_input_parsing` -- parse user input lines
|
||||
- `test_repl_output_formatting` -- format agent responses
|
||||
- `test_repl_multiline` -- handle multi-line input
|
||||
- `test_repl_special_commands` -- handle /quit, /help
|
||||
|
||||
---
|
||||
|
||||
## Tier 4 -- CLI Tests (~2,100 lines)
|
||||
|
||||
CLI subcommands can be tested by invoking clap-parsed command structs directly
|
||||
or by calling the handler functions with constructed arguments.
|
||||
|
||||
### `src/cli/tool.rs` -- 2.9% -> 95% (+697 lines)
|
||||
|
||||
Tool CLI (install, list, remove, build).
|
||||
|
||||
**Tests to write:**
|
||||
- `test_cli_tool_list` -- list installed tools
|
||||
- `test_cli_tool_install_local` -- install from local .wasm file
|
||||
- `test_cli_tool_install_registry` -- install from registry
|
||||
- `test_cli_tool_remove` -- remove installed tool
|
||||
- `test_cli_tool_build` -- scaffold and build tool project
|
||||
- `test_cli_tool_info` -- display tool details
|
||||
|
||||
### `src/cli/mcp.rs` -- 0.9% -> 95% (+302 lines)
|
||||
|
||||
MCP server management CLI.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_cli_mcp_list` -- list configured MCP servers
|
||||
- `test_cli_mcp_add` -- add MCP server config
|
||||
- `test_cli_mcp_remove` -- remove MCP server config
|
||||
- `test_cli_mcp_tools` -- list tools from MCP server
|
||||
- `test_cli_mcp_test_connection` -- verify MCP server reachable
|
||||
|
||||
### `src/cli/oauth_defaults.rs` -- 54.1% -> 95% (+298 lines)
|
||||
|
||||
OAuth default configurations.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_oauth_defaults_loading` -- load default OAuth configs
|
||||
- `test_oauth_url_construction` -- build auth/token URLs
|
||||
- `test_oauth_scope_merging` -- merge requested scopes with defaults
|
||||
- `test_oauth_provider_lookup` -- lookup by provider name
|
||||
|
||||
### `src/cli/registry.rs` -- 0% -> 95% (+168 lines)
|
||||
|
||||
Registry CLI commands.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_cli_registry_search` -- search for packages
|
||||
- `test_cli_registry_install` -- install package from registry
|
||||
- `test_cli_registry_info` -- display package details
|
||||
|
||||
### `src/cli/status.rs` -- 0% -> 95% (+142 lines)
|
||||
|
||||
Status display commands.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_cli_status_gathering` -- collect system status info
|
||||
- `test_cli_status_formatting` -- render status output
|
||||
- `test_cli_status_components` -- check individual components
|
||||
|
||||
### `src/cli/memory.rs` -- 15.5% -> 95% (+138 lines)
|
||||
|
||||
Memory CLI subcommands.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_cli_memory_search` -- search workspace from CLI
|
||||
- `test_cli_memory_write` -- write document from CLI
|
||||
- `test_cli_memory_read` -- read document from CLI
|
||||
- `test_cli_memory_tree` -- display memory tree
|
||||
|
||||
### `src/cli/doctor.rs` -- 28.7% -> 95% (+115 lines)
|
||||
|
||||
Diagnostic checks.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_doctor_check_database` -- verify DB connectivity check
|
||||
- `test_doctor_check_llm` -- verify LLM provider check
|
||||
- `test_doctor_check_tools` -- verify tool availability check
|
||||
- `test_doctor_report_format` -- verify output format
|
||||
|
||||
### `src/cli/config.rs` -- 36.5% -> 95% (~100 lines)
|
||||
|
||||
Config CLI subcommands.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_cli_config_get` -- read config value
|
||||
- `test_cli_config_set` -- write config value
|
||||
- `test_cli_config_list` -- list all config keys
|
||||
- `test_cli_config_reset` -- reset to defaults
|
||||
|
||||
---
|
||||
|
||||
## Tier 5 -- Setup/Infra Tests (~2,400 lines)
|
||||
|
||||
Hardest to test: interactive wizards, Docker, process spawning. Strategy: extract
|
||||
pure logic into testable functions, test the interactive parts by injecting mock
|
||||
input.
|
||||
|
||||
### `src/setup/wizard.rs` -- 16.8% -> 95% (+1,681 lines)
|
||||
|
||||
7-step interactive onboarding wizard. Refactor to extract validation functions,
|
||||
step logic, and config generation into testable units.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_wizard_step_validation` -- each step validates input correctly
|
||||
- `test_wizard_config_generation` -- generate config from wizard answers
|
||||
- `test_wizard_default_values` -- verify sensible defaults
|
||||
- `test_wizard_skip_completed` -- skip already-configured steps
|
||||
- `test_wizard_llm_backend_selection` -- provider-specific config paths
|
||||
- `test_wizard_channel_setup` -- channel configuration logic
|
||||
|
||||
### `src/setup/channels.rs` -- 7.6% -> 95% (+563 lines)
|
||||
|
||||
Channel setup helpers.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_channel_setup_defaults` -- default channel configuration
|
||||
- `test_channel_setup_validation` -- reject invalid channel configs
|
||||
- `test_channel_setup_telegram` -- Telegram-specific setup logic
|
||||
- `test_channel_setup_signal` -- Signal-specific setup logic
|
||||
- `test_channel_setup_webhook` -- webhook URL validation
|
||||
|
||||
### `src/setup/prompts.rs` -- 24.8% -> 95% (+147 lines)
|
||||
|
||||
Terminal prompt utilities.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_prompt_select` -- selection from list
|
||||
- `test_prompt_confirm` -- yes/no confirmation
|
||||
- `test_prompt_secret` -- masked input
|
||||
- `test_prompt_validation` -- input validation rules
|
||||
|
||||
### `src/sandbox/container.rs` -- 22.1% -> 95% (+296 lines)
|
||||
|
||||
Docker container lifecycle. Test command construction without actual Docker.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_container_config_to_docker_args` -- generate correct docker run args
|
||||
- `test_container_volume_mounts` -- workspace mount configuration
|
||||
- `test_container_env_scrubbing` -- sensitive env vars removed
|
||||
- `test_container_resource_limits` -- CPU/memory limit args
|
||||
- `test_container_network_config` -- proxy network setup
|
||||
|
||||
### `src/sandbox/manager.rs` -- 59% -> 95% (+114 lines)
|
||||
|
||||
Sandbox orchestration.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_sandbox_policy_enforcement` -- policy to container config mapping
|
||||
- `test_sandbox_cleanup` -- cleanup on job completion
|
||||
- `test_sandbox_concurrent_limit` -- enforce max concurrent containers
|
||||
|
||||
### `src/sandbox/proxy/http.rs` -- 37.5% -> 95% (+176 lines)
|
||||
|
||||
HTTP proxy for container network access.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_proxy_allowlist_enforcement` -- block disallowed domains
|
||||
- `test_proxy_credential_injection` -- inject auth headers
|
||||
- `test_proxy_connect_tunnel` -- HTTPS CONNECT method handling
|
||||
- `test_proxy_logging` -- request/response logging
|
||||
|
||||
### `src/worker/runtime.rs` -- 5.7% -> 95% (+312 lines)
|
||||
|
||||
Worker execution loop (runs inside containers).
|
||||
|
||||
**Tests to write:**
|
||||
- `test_worker_tool_dispatch` -- dispatch tool call, return result
|
||||
- `test_worker_llm_interaction` -- send prompt, receive response
|
||||
- `test_worker_turn_limit` -- enforce max turns
|
||||
- `test_worker_error_propagation` -- tool error surfaces to agent
|
||||
|
||||
### `src/worker/claude_bridge.rs` -- 60.7% -> 95% (+215 lines)
|
||||
|
||||
Claude CLI bridge.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_claude_command_construction` -- build claude CLI command
|
||||
- `test_claude_output_parsing` -- parse claude CLI JSON output
|
||||
- `test_claude_error_handling` -- handle CLI crashes gracefully
|
||||
- `test_claude_config_injection` -- inject config dir and model
|
||||
|
||||
### `src/worker/api.rs` -- 19.8% -> 95% (+194 lines)
|
||||
|
||||
Worker HTTP client to orchestrator.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_worker_api_request_building` -- correct endpoint URLs and headers
|
||||
- `test_worker_api_response_parsing` -- parse orchestrator responses
|
||||
- `test_worker_api_auth_token` -- bearer token injection
|
||||
- `test_worker_api_retry` -- retry on transient failures
|
||||
|
||||
### `src/main.rs` -- 29.4% -> 95% (+485 lines)
|
||||
|
||||
Entry point and startup. Extract startup logic into testable functions.
|
||||
|
||||
**Tests to write:**
|
||||
- `test_cli_arg_parsing` -- verify clap argument parsing
|
||||
- `test_startup_config_loading` -- config from env + file
|
||||
- `test_startup_channel_selection` -- select channels from config
|
||||
- `test_startup_feature_flags` -- feature-gated code paths
|
||||
|
||||
---
|
||||
|
||||
## Tier 6 -- Remaining Files to 95% (~2,000 lines)
|
||||
|
||||
Smaller files that each need a handful of additional tests.
|
||||
|
||||
| File | Lines Needed | Test Focus |
|
||||
|------|-------------:|------------|
|
||||
| `src/tools/builtin/skill_tools.rs` | 110 | skill_list, skill_search, skill_install, skill_remove |
|
||||
| `src/hooks/bundled.rs` | 115 | bundled hook execution, hook discovery |
|
||||
| `src/registry/installer.rs` | 272 | package download, verification, installation |
|
||||
| `src/registry/artifacts.rs` | 72 | artifact packaging, checksums |
|
||||
| `src/orchestrator/job_manager.rs` | 249 | container lifecycle, job routing |
|
||||
| `src/orchestrator/api.rs` | 125 | LLM proxy, event dispatch endpoints |
|
||||
| `src/app.rs` | 137 | AppBuilder configuration, startup sequence |
|
||||
| `src/service.rs` | 120 | service lifecycle, signal handling |
|
||||
| `src/config/channels.rs` | 55 | channel config parsing |
|
||||
| `src/config/sandbox.rs` | 61 | sandbox config parsing |
|
||||
| `src/config/tunnel.rs` | 43 | tunnel config parsing |
|
||||
| `src/config/mod.rs` | 63 | config merging, env override |
|
||||
| `src/config/database.rs` | 38 | database URL parsing |
|
||||
| `src/evaluation/success.rs` | 34 | success evaluator logic |
|
||||
| `src/evaluation/metrics.rs` | 40 | metrics collection |
|
||||
| `src/context/manager.rs` | 57 | concurrent job context isolation |
|
||||
| `src/context/memory.rs` | 36 | action recording, conversation memory |
|
||||
|
||||
---
|
||||
|
||||
## Execution Priority
|
||||
|
||||
Maximize coverage gain per unit of effort:
|
||||
|
||||
| Order | Category | Lines Gained | Effort |
|
||||
|------:|----------|-------------:|--------|
|
||||
| 1 | Trace tests (Tier 2) | ~7,000 | Medium (high leverage, each test covers many modules) |
|
||||
| 2 | Unit tests for 0% files (Tier 1 subset) | ~3,500 | Low (pure logic, no infrastructure) |
|
||||
| 3 | Web handler tests (Tier 3) | ~4,500 | Medium (axum_test + in-memory DB) |
|
||||
| 4 | Extension/MCP/WASM unit tests (Tier 1 remainder) | ~3,500 | Medium |
|
||||
| 5 | CLI subcommand tests (Tier 4) | ~2,100 | Low-Medium |
|
||||
| 6 | Setup wizard extraction + tests (Tier 5) | ~2,400 | High (requires refactoring) |
|
||||
| 7 | LLM provider tests (Tier 1 subset) | ~800 | Medium |
|
||||
| 8 | Remaining small files (Tier 6) | ~2,000 | Low |
|
||||
|
||||
## Notes
|
||||
|
||||
- All trace tests require `--features libsql` and use `TestRigBuilder` from `tests/support/`
|
||||
- Web handler tests can use `axum::test` helpers or build the router directly
|
||||
- CLI tests should call handler functions directly, not shell out to the binary
|
||||
- Setup wizard tests require extracting pure logic from interactive prompts first
|
||||
- Sandbox/container tests should verify command construction, not run Docker
|
||||
- Worker tests can use `TraceLlm` for the LLM provider, same as trace tests
|
||||
Generated
+836
-23
File diff suppressed because it is too large
Load Diff
+14
-2
@@ -40,7 +40,7 @@ tokio-stream = { version = "0.1", features = ["sync"] }
|
||||
futures = "0.3"
|
||||
|
||||
# HTTP client
|
||||
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls-native-roots", "stream"] }
|
||||
reqwest = { version = "0.12", default-features = false, features = ["json", "multipart", "rustls-tls-native-roots", "stream"] }
|
||||
|
||||
# Serialization
|
||||
serde = { version = "1", features = ["derive"] }
|
||||
@@ -56,7 +56,7 @@ rustls = { version = "0.23", optional = true, default-features = false }
|
||||
rustls-native-certs = { version = "0.8", optional = true }
|
||||
|
||||
# Database - libSQL/Turso (optional embedded database)
|
||||
libsql = { version = "0.6", optional = true, default-features = false, features = ["core", "replication"] }
|
||||
libsql = { version = "0.6", optional = true, default-features = false, features = ["core", "replication", "remote", "tls"] }
|
||||
|
||||
# Error handling
|
||||
thiserror = "2"
|
||||
@@ -73,6 +73,8 @@ toml = "0.8"
|
||||
# Core types
|
||||
uuid = { version = "1", features = ["v4", "v5", "serde"] }
|
||||
chrono = { version = "0.4", features = ["serde"] }
|
||||
chrono-tz = "0.10"
|
||||
iana-time-zone = "0.1"
|
||||
rust_decimal = { version = "1", features = ["serde", "serde-with-str", "maths"] }
|
||||
rust_decimal_macros = "1"
|
||||
|
||||
@@ -140,6 +142,11 @@ subtle = "2" # Constant-time comparisons for token validation
|
||||
# Multi-provider LLM support
|
||||
rig-core = "0.30"
|
||||
|
||||
# AWS Bedrock (native Converse API, opt-in via --features bedrock)
|
||||
aws-config = { version = "1", features = ["behavior-version-latest"], optional = true }
|
||||
aws-sdk-bedrockruntime = { version = "1", optional = true }
|
||||
aws-smithy-types = { version = "1", optional = true }
|
||||
|
||||
# Docker sandbox
|
||||
bollard = "0.18"
|
||||
|
||||
@@ -147,6 +154,10 @@ bollard = "0.18"
|
||||
flate2 = "1"
|
||||
tar = "0.4"
|
||||
|
||||
# Document text extraction
|
||||
pdf-extract = "0.7"
|
||||
zip = { version = "2", default-features = false, features = ["deflate"] }
|
||||
|
||||
# HTTP proxy for sandboxed network access
|
||||
hyper = { version = "1.5", features = ["server", "http1", "http2"] }
|
||||
hyper-util = { version = "0.1", features = ["server", "tokio", "http1", "http2"] }
|
||||
@@ -197,6 +208,7 @@ postgres = [
|
||||
libsql = ["dep:libsql"]
|
||||
integration = []
|
||||
html-to-markdown = ["dep:html-to-markdown-rs", "dep:readabilityrs"]
|
||||
bedrock = ["dep:aws-config", "dep:aws-sdk-bedrockruntime", "dep:aws-smithy-types"]
|
||||
|
||||
[[test]]
|
||||
name = "html_to_markdown"
|
||||
|
||||
+45
-21
@@ -10,6 +10,8 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
||||
- 🚫 Out of scope (intentionally skipped)
|
||||
- ➖ N/A (not applicable to Rust implementation)
|
||||
|
||||
**Last reviewed against OpenClaw PRs:** 2026-03-10 (merged 2026-02-24 through 2026-03-10)
|
||||
|
||||
---
|
||||
|
||||
## 1. Architecture
|
||||
@@ -43,7 +45,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
||||
| launchd/systemd integration | ✅ | ❌ | |
|
||||
| Bonjour/mDNS discovery | ✅ | ❌ | |
|
||||
| Tailscale integration | ✅ | ❌ | |
|
||||
| Health check endpoints | ✅ | ✅ | /api/health + /api/gateway/status |
|
||||
| Health check endpoints | ✅ | ✅ | /api/health + /api/gateway/status + /healthz + /readyz, with channel-backed readiness probes |
|
||||
| `doctor` diagnostics | ✅ | ❌ | |
|
||||
| Agent event broadcast | ✅ | 🚧 | SSE broadcast manager exists (SseManager) but tool/job-state events not fully wired |
|
||||
| Channel health monitor | ✅ | ❌ | Auto-restart with configurable interval |
|
||||
@@ -66,17 +68,17 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
||||
| REPL (simple) | ✅ | ✅ | - | For testing |
|
||||
| WASM channels | ❌ | ✅ | - | IronClaw innovation |
|
||||
| WhatsApp | ✅ | ❌ | P1 | Baileys (Web), same-phone mode with echo detection |
|
||||
| Telegram | ✅ | ✅ | - | WASM channel(MTProto), DM pairing, caption, /start, bot_username |
|
||||
| Telegram | ✅ | ✅ | - | WASM channel(MTProto), DM pairing, caption, /start, bot_username, DM topics |
|
||||
| Discord | ✅ | ❌ | P2 | discord.js, thread parent binding inheritance |
|
||||
| Signal | ✅ | ✅ | P2 | signal-cli daemonPC, SSE listener HTTP/JSON-R, user/group allowlists, DM pairing |
|
||||
| Slack | ✅ | ✅ | - | WASM tool |
|
||||
| iMessage | ✅ | ❌ | P3 | BlueBubbles or Linq recommended |
|
||||
| Linq | ✅ | ❌ | P3 | Real iMessage via API, no Mac required |
|
||||
| Feishu/Lark | ✅ | ❌ | P3 | Bitable create app/field tools |
|
||||
| Feishu/Lark | ✅ | ❌ | P3 | Bitable create app/field tools, Docx table/image/file actions, rich-text media extraction |
|
||||
| LINE | ✅ | ❌ | P3 | |
|
||||
| WebChat | ✅ | ✅ | - | Web gateway chat |
|
||||
| Matrix | ✅ | ❌ | P3 | E2EE support |
|
||||
| Mattermost | ✅ | ❌ | P3 | Emoji reactions |
|
||||
| Mattermost | ✅ | ❌ | P3 | Emoji reactions, interactive buttons, model picker |
|
||||
| Google Chat | ✅ | ❌ | P3 | |
|
||||
| MS Teams | ✅ | ❌ | P3 | |
|
||||
| Twitch | ✅ | ❌ | P3 | |
|
||||
@@ -92,6 +94,8 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
||||
| User message reactions | ✅ | ❌ | Surface inbound reactions |
|
||||
| sendPoll | ✅ | ❌ | Poll creation via agent |
|
||||
| Cron/heartbeat topic targeting | ✅ | ❌ | Messages land in correct topic |
|
||||
| DM topics support | ✅ | ❌ | Agent/topic bindings in DMs and agent-scoped SessionKeys |
|
||||
| Persistent ACP topic binding | ✅ | ❌ | ACP harness sessions can pin to Telegram forum or DM topics |
|
||||
|
||||
### Discord-Specific Features (since Feb 2025)
|
||||
|
||||
@@ -107,21 +111,36 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
||||
|---------|----------|----------|-------|
|
||||
| Streaming draft replies | ✅ | ❌ | Partial replies via draft message updates |
|
||||
| Configurable stream modes | ✅ | ❌ | Per-channel stream behavior |
|
||||
| Thread ownership | ✅ | ❌ | Thread-level ownership tracking |
|
||||
| Thread ownership | ✅ | ❌ | Thread-level ownership tracking plus reply participation memory |
|
||||
| Download-file action | ✅ | ❌ | On-demand attachment downloads via message actions |
|
||||
|
||||
### Mattermost-Specific Features (since Mar 2026)
|
||||
|
||||
| Feature | OpenClaw | IronClaw | Notes |
|
||||
|---------|----------|----------|-------|
|
||||
| Interactive buttons | ✅ | ❌ | Clickable message buttons with signed callback flow |
|
||||
| Interactive model picker | ✅ | ❌ | In-channel provider/model chooser |
|
||||
|
||||
### Feishu/Lark-Specific Features (since Mar 2026)
|
||||
|
||||
| Feature | OpenClaw | IronClaw | Notes |
|
||||
|---------|----------|----------|-------|
|
||||
| Doc/table actions | ✅ | ❌ | `feishu_doc` supports tables, positional insert, color_text, image upload, and file upload |
|
||||
| Rich-text embedded media extraction | ✅ | ❌ | Pull video/media attachments from post messages |
|
||||
|
||||
### Channel Features
|
||||
|
||||
| Feature | OpenClaw | IronClaw | Notes |
|
||||
|---------|----------|----------|-------|
|
||||
| DM pairing codes | ✅ | ✅ | `ironclaw pairing list/approve`, host APIs |
|
||||
| Allowlist/blocklist | ✅ | 🚧 | allow_from + pairing store |
|
||||
| Allowlist/blocklist | ✅ | 🚧 | `allow_from` + pairing store + hardened command/group allowlists |
|
||||
| Self-message bypass | ✅ | ❌ | Own messages skip pairing |
|
||||
| Mention-based activation | ✅ | ✅ | bot_username + respond_to_all_group_messages |
|
||||
| Per-group tool policies | ✅ | ❌ | Allow/deny specific tools |
|
||||
| Thread isolation | ✅ | ✅ | Separate sessions per thread |
|
||||
| Per-channel media limits | ✅ | 🚧 | Caption support for media; no size limits |
|
||||
| Typing indicators | ✅ | 🚧 | TUI + Telegram typing/actionable status prompts; richer parity pending |
|
||||
| Per-channel ackReaction config | ✅ | ❌ | Customizable acknowledgement reactions |
|
||||
| Thread isolation | ✅ | ✅ | Separate sessions per thread/topic |
|
||||
| Per-channel media limits | ✅ | 🚧 | Caption support plus `mediaMaxMb` enforcement for WhatsApp, Telegram, and Discord |
|
||||
| Typing indicators | ✅ | 🚧 | TUI + channel typing, with configurable silence timeout; richer parity pending |
|
||||
| Per-channel ackReaction config | ✅ | ❌ | Customizable acknowledgement reactions/scopes |
|
||||
| Group session priming | ✅ | ❌ | Member roster injected for context |
|
||||
| Sender_id in trusted metadata | ✅ | ❌ | Exposed in system metadata |
|
||||
|
||||
@@ -138,7 +157,8 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
||||
| `gateway start/stop` | ✅ | ❌ | P2 | |
|
||||
| `onboard` (wizard) | ✅ | ✅ | - | Interactive setup |
|
||||
| `tui` | ✅ | ✅ | - | Ratatui TUI |
|
||||
| `config` | ✅ | ✅ | - | Read/write config |
|
||||
| `config` | ✅ | ✅ | - | Read/write config plus validate/path helpers |
|
||||
| `backup` | ✅ | ❌ | P3 | Create/verify local backup archives |
|
||||
| `channels` | ✅ | ❌ | P2 | Channel management |
|
||||
| `models` | ✅ | 🚧 | - | Model selector in TUI |
|
||||
| `status` | ✅ | ✅ | - | System status (enriched session details) |
|
||||
@@ -177,14 +197,15 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
||||
| Global sessions | ✅ | ❌ | Optional shared context |
|
||||
| Session pruning | ✅ | ❌ | Auto cleanup old sessions |
|
||||
| Context compaction | ✅ | ✅ | Auto summarization |
|
||||
| Compaction model override | ✅ | ❌ | Use a dedicated provider/model for summarization only |
|
||||
| Post-compaction read audit | ✅ | ❌ | Layer 3: workspace rules appended to summaries |
|
||||
| Post-compaction context injection | ✅ | ❌ | Workspace context as system event |
|
||||
| Custom system prompts | ✅ | ✅ | Template variables, safety guardrails |
|
||||
| Skills (modular capabilities) | ✅ | ✅ | Prompt-based skills with trust gating, attenuation, activation criteria, catalog, selector |
|
||||
| Skill routing blocks | ✅ | 🚧 | ActivationCriteria (keywords, patterns, tags) but no "Use when / Don't use when" blocks |
|
||||
| Skill path compaction | ✅ | ❌ | ~ prefix to reduce prompt tokens |
|
||||
| Thinking modes (low/med/high) | ✅ | ❌ | Configurable reasoning depth |
|
||||
| Per-model thinkingDefault override | ✅ | ❌ | Override thinking level per model |
|
||||
| Thinking modes (off/minimal/low/medium/high/xhigh/adaptive) | ✅ | ❌ | Configurable reasoning depth |
|
||||
| Per-model thinkingDefault override | ✅ | ❌ | Override thinking level per model; Anthropic Claude 4.6 defaults to adaptive |
|
||||
| Block-level streaming | ✅ | ❌ | |
|
||||
| Tool-level streaming | ✅ | ❌ | |
|
||||
| Z.AI tool_stream | ✅ | ❌ | Real-time tool call streaming |
|
||||
@@ -213,8 +234,8 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
||||
| Provider | OpenClaw | IronClaw | Priority | Notes |
|
||||
|----------|----------|----------|----------|-------|
|
||||
| NEAR AI | ✅ | ✅ | - | Primary provider |
|
||||
| Anthropic (Claude) | ✅ | 🚧 | - | Via NEAR AI proxy; Opus 4.5, Sonnet 4, Sonnet 4.6 |
|
||||
| OpenAI | ✅ | 🚧 | - | Via NEAR AI proxy |
|
||||
| Anthropic (Claude) | ✅ | 🚧 | - | Via NEAR AI proxy; Opus 4.5, Sonnet 4, Sonnet 4.6, adaptive thinking default |
|
||||
| OpenAI | ✅ | 🚧 | - | Via NEAR AI proxy; GPT-5.4 + Codex OAuth |
|
||||
| AWS Bedrock | ✅ | ❌ | P3 | |
|
||||
| Google Gemini | ✅ | ❌ | P3 | |
|
||||
| NVIDIA API | ✅ | ❌ | P3 | New provider |
|
||||
@@ -238,7 +259,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
||||
| Per-session model override | ✅ | ✅ | Model selector in TUI |
|
||||
| Model selection UI | ✅ | ✅ | TUI keyboard shortcut |
|
||||
| Per-model thinkingDefault | ✅ | ❌ | Override thinking level per model in config |
|
||||
| 1M context beta header | ✅ | ❌ | Anthropic extended context support |
|
||||
| 1M context support | ✅ | ❌ | Anthropic extended context beta + OpenAI Codex GPT-5.4 1M context |
|
||||
|
||||
### Owner: _Unassigned_
|
||||
|
||||
@@ -253,7 +274,8 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
||||
| Multiple images per tool call | ✅ | ❌ | P2 | Single tool invocation, multiple images |
|
||||
| Audio transcription | ✅ | ❌ | P2 | |
|
||||
| Video support | ✅ | ❌ | P3 | |
|
||||
| PDF parsing | ✅ | ❌ | P2 | pdfjs-dist |
|
||||
| PDF analysis tool | ✅ | ❌ | P2 | Native Anthropic/Gemini path with text/image extraction fallback |
|
||||
| PDF parsing | ✅ | ❌ | P2 | `pdfjs-dist` fallback path |
|
||||
| MIME detection | ✅ | ❌ | P2 | |
|
||||
| Media caching | ✅ | ❌ | P3 | |
|
||||
| Vision model integration | ✅ | ❌ | P2 | Image understanding |
|
||||
@@ -276,7 +298,8 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
||||
| Workspace-relative install | ✅ | ✅ | ~/.ironclaw/tools/ |
|
||||
| Channel plugins | ✅ | ✅ | WASM channels |
|
||||
| Auth plugins | ✅ | ❌ | |
|
||||
| Memory plugins | ✅ | ❌ | Custom backends |
|
||||
| Memory plugins | ✅ | ❌ | Custom backends + selectable memory slot |
|
||||
| Context-engine plugins | ✅ | ❌ | Custom context management + subagent/context hooks |
|
||||
| Tool plugins | ✅ | ✅ | WASM tools |
|
||||
| Hook plugins | ✅ | ✅ | Declarative hooks from extension capabilities |
|
||||
| Provider plugins | ✅ | ❌ | |
|
||||
@@ -298,7 +321,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
||||
| JSON5 support | ✅ | ❌ | Comments, trailing commas |
|
||||
| YAML alternative | ✅ | ❌ | |
|
||||
| Environment variable interpolation | ✅ | ✅ | `${VAR}` |
|
||||
| Config validation/schema | ✅ | ✅ | Type-safe Config struct |
|
||||
| Config validation/schema | ✅ | ✅ | Type-safe Config struct + `openclaw config validate` |
|
||||
| Hot-reload | ✅ | ❌ | |
|
||||
| Legacy migration | ✅ | ➖ | |
|
||||
| State directory | ✅ `~/.openclaw-state/` | ✅ `~/.ironclaw/` | |
|
||||
@@ -405,6 +428,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
||||
| Feature | OpenClaw | IronClaw | Priority | Notes |
|
||||
|---------|----------|----------|----------|-------|
|
||||
| Cron jobs | ✅ | ✅ | - | Routines with cron trigger |
|
||||
| Per-job model fallback override | ✅ | ❌ | P2 | `payload.fallbacks` overrides agent-level fallbacks |
|
||||
| Cron stagger controls | ✅ | ❌ | P3 | Default stagger for scheduled jobs |
|
||||
| Cron finished-run webhook | ✅ | ❌ | P3 | Webhook on job completion |
|
||||
| Timezone support | ✅ | ✅ | - | Via cron expressions |
|
||||
@@ -458,10 +482,10 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
||||
| Elevated mode | ✅ | ❌ | |
|
||||
| Safe bins allowlist | ✅ | ❌ | Hardened path trust |
|
||||
| LD*/DYLD* validation | ✅ | ❌ | |
|
||||
| Path traversal prevention | ✅ | ✅ | Including config includes (OC-06) |
|
||||
| Path traversal prevention | ✅ | ✅ | Including config includes (OC-06) + workspace-only tool mounts |
|
||||
| Credential theft via env injection | ✅ | 🚧 | Shell env scrubbing + command injection detection; no full OC-09 defense |
|
||||
| Session file permissions (0o600) | ✅ | ✅ | Session token file set to 0o600 in llm/session.rs |
|
||||
| Skill download path restriction | ✅ | ❌ | Prevent arbitrary write targets |
|
||||
| Skill download path restriction | ✅ | ❌ | Validated download roots prevent arbitrary write targets |
|
||||
| Webhook signature verification | ✅ | ✅ | |
|
||||
| Media URL validation | ✅ | ❌ | |
|
||||
| Prompt injection defense | ✅ | ✅ | Pattern detection, sanitization |
|
||||
|
||||
@@ -14,6 +14,11 @@
|
||||
<a href="https://www.reddit.com/r/ironclawAI/"><img src="https://img.shields.io/badge/Reddit-r%2FironclawAI-FF4500?style=flat&logo=reddit&logoColor=white" alt="Reddit: r/ironclawAI" /></a>
|
||||
</p>
|
||||
|
||||
<p align="center">
|
||||
<a href="README.md">English</a> |
|
||||
<a href="README.zh-CN.md">简体中文</a>
|
||||
</p>
|
||||
|
||||
<p align="center">
|
||||
<a href="#philosophy">Philosophy</a> •
|
||||
<a href="#features">Features</a> •
|
||||
|
||||
+319
@@ -0,0 +1,319 @@
|
||||
<p align="center">
|
||||
<img src="ironclaw.png?v=2" alt="IronClaw" width="200"/>
|
||||
</p>
|
||||
|
||||
<h1 align="center">IronClaw</h1>
|
||||
|
||||
<p align="center">
|
||||
<strong>安全可靠的个人 AI 助手,始终站在你这边</strong>
|
||||
</p>
|
||||
|
||||
<p align="center">
|
||||
<a href="#license"><img src="https://img.shields.io/badge/license-MIT%20OR%20Apache%202.0-blue.svg" alt="License: MIT OR Apache-2.0" /></a>
|
||||
<a href="https://t.me/ironclawAI"><img src="https://img.shields.io/badge/Telegram-%40ironclawAI-26A5E4?style=flat&logo=telegram&logoColor=white" alt="Telegram: @ironclawAI" /></a>
|
||||
<a href="https://www.reddit.com/r/ironclawAI/"><img src="https://img.shields.io/badge/Reddit-r%2FironclawAI-FF4500?style=flat&logo=reddit&logoColor=white" alt="Reddit: r/ironclawAI" /></a>
|
||||
</p>
|
||||
|
||||
<p align="center">
|
||||
<a href="README.md">English</a> |
|
||||
<a href="README.zh-CN.md">简体中文</a>
|
||||
</p>
|
||||
|
||||
<p align="center">
|
||||
<a href="#设计理念">设计理念</a> •
|
||||
<a href="#功能特性">功能特性</a> •
|
||||
<a href="#安装">安装</a> •
|
||||
<a href="#配置">配置</a> •
|
||||
<a href="#安全机制">安全机制</a> •
|
||||
<a href="#系统架构">系统架构</a>
|
||||
</p>
|
||||
|
||||
---
|
||||
|
||||
## 设计理念
|
||||
|
||||
IronClaw 基于一个简单的原则:**你的 AI 助手应该为你服务,而不是与你为敌。**
|
||||
|
||||
在 AI 系统对数据处理日益不透明、与企业利益捆绑的今天,IronClaw 选择了一条不同的路:
|
||||
|
||||
- **数据归你所有** — 所有信息存储在本地,加密保护,始终在你掌控之下
|
||||
- **透明至上** — 完全开源,可审计,没有隐藏的遥测或数据收集
|
||||
- **自主扩展** — 随时构建新工具,无需等待供应商更新
|
||||
- **纵深防御** — 多层安全机制抵御提示注入和数据泄露
|
||||
|
||||
IronClaw 是一个你真正可以信赖的 AI 助手,无论是个人生活还是工作。
|
||||
|
||||
## 功能特性
|
||||
|
||||
### 安全优先
|
||||
|
||||
- **WASM 沙箱** — 不受信任的工具在隔离的 WebAssembly 容器中运行,采用基于能力的权限模型
|
||||
- **凭据保护** — 密钥永远不会暴露给工具;在宿主边界注入并进行泄露检测
|
||||
- **提示注入防御** — 模式检测、内容清理和策略执行
|
||||
- **端点白名单** — HTTP 请求仅限于明确批准的主机和路径
|
||||
|
||||
### 随时可用
|
||||
|
||||
- **多渠道接入** — REPL、HTTP webhook、WASM 渠道(Telegram、Slack)和 Web 网关
|
||||
- **Docker 沙箱** — 隔离的容器执行,支持每任务令牌和编排器/工作器模式
|
||||
- **Web 网关** — 浏览器 UI,支持实时 SSE/WebSocket 流式传输
|
||||
- **定时任务** — Cron 调度、事件触发器、Webhook 处理器,实现后台自动化
|
||||
- **心跳系统** — 主动后台执行,用于监控和维护任务
|
||||
- **并行任务** — 使用隔离上下文同时处理多个请求
|
||||
- **自修复** — 自动检测并恢复卡住的操作
|
||||
|
||||
### 自主扩展
|
||||
|
||||
- **动态工具构建** — 描述你的需求,IronClaw 会将其构建为 WASM 工具
|
||||
- **MCP 协议** — 连接模型上下文协议(Model Context Protocol)服务器以获取额外能力
|
||||
- **插件架构** — 无需重启即可加载新的 WASM 工具和渠道
|
||||
|
||||
### 持久记忆
|
||||
|
||||
- **混合搜索** — 全文搜索 + 向量搜索,采用倒数排名融合(Reciprocal Rank Fusion)
|
||||
- **工作空间文件系统** — 灵活的基于路径的存储,用于笔记、日志和上下文
|
||||
- **身份文件** — 跨会话保持一致的个性和偏好设置
|
||||
|
||||
## 安装
|
||||
|
||||
### 前置要求
|
||||
|
||||
- Rust 1.85+
|
||||
- PostgreSQL 15+,需安装 [pgvector](https://github.com/pgvector/pgvector) 扩展
|
||||
- NEAR AI 账户(通过设置向导进行身份验证)
|
||||
|
||||
## 下载或编译
|
||||
|
||||
访问 [Releases 页面](https://github.com/nearai/ironclaw/releases/) 查看最新版本。
|
||||
|
||||
<details>
|
||||
<summary>通过 Windows 安装程序安装 (Windows)</summary>
|
||||
|
||||
下载 [Windows 安装程序](https://github.com/nearai/ironclaw/releases/latest/download/ironclaw-x86_64-pc-windows-msvc.msi) 并运行。
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary>通过 PowerShell 脚本安装 (Windows)</summary>
|
||||
|
||||
```sh
|
||||
irm https://github.com/nearai/ironclaw/releases/latest/download/ironclaw-installer.ps1 | iex
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary>通过 Shell 脚本安装 (macOS、Linux、Windows/WSL)</summary>
|
||||
|
||||
```sh
|
||||
curl --proto '=https' --tlsv1.2 -LsSf https://github.com/nearai/ironclaw/releases/latest/download/ironclaw-installer.sh | sh
|
||||
```
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary>通过 Homebrew 安装 (macOS/Linux)</summary>
|
||||
|
||||
```sh
|
||||
brew install ironclaw
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary>从源码编译 (Windows、Linux、macOS 上使用 Cargo)</summary>
|
||||
|
||||
确保你已安装 [Rust](https://rustup.rs)。
|
||||
|
||||
```bash
|
||||
# 克隆仓库
|
||||
git clone https://github.com/nearai/ironclaw.git
|
||||
cd ironclaw
|
||||
|
||||
# 编译
|
||||
cargo build --release
|
||||
|
||||
# 运行测试
|
||||
cargo test
|
||||
```
|
||||
|
||||
如需进行**完整发布构建**(修改了渠道源码后),先运行 `./scripts/build-all.sh` 重新编译渠道。
|
||||
|
||||
</details>
|
||||
|
||||
### 数据库设置
|
||||
|
||||
```bash
|
||||
# 创建数据库
|
||||
createdb ironclaw
|
||||
|
||||
# 启用 pgvector 扩展
|
||||
psql ironclaw -c "CREATE EXTENSION IF NOT EXISTS vector;"
|
||||
```
|
||||
|
||||
## 配置
|
||||
|
||||
运行设置向导来配置 IronClaw:
|
||||
|
||||
```bash
|
||||
ironclaw onboard
|
||||
```
|
||||
|
||||
向导将引导你完成数据库连接、NEAR AI 身份验证(通过浏览器 OAuth)和密钥加密(使用系统钥匙串)。设置会保存在数据库中;引导变量(如 `DATABASE_URL`、`LLM_BACKEND`)写入 `~/.ironclaw/.env`,以便在数据库连接前可用。
|
||||
|
||||
### 替代 LLM 提供商
|
||||
|
||||
IronClaw 默认使用 NEAR AI,但兼容任何 OpenAI 兼容的端点。
|
||||
常用选项包括 **OpenRouter**(300+ 模型)、**Together AI**、**Fireworks AI**、**Ollama**(本地部署)以及自托管服务器如 **vLLM** 或 **LiteLLM**。
|
||||
|
||||
在向导中选择 *"OpenAI-compatible"*,或直接设置环境变量:
|
||||
|
||||
```env
|
||||
LLM_BACKEND=openai_compatible
|
||||
LLM_BASE_URL=https://openrouter.ai/api/v1
|
||||
LLM_API_KEY=sk-or-...
|
||||
LLM_MODEL=anthropic/claude-sonnet-4
|
||||
```
|
||||
|
||||
详见 [docs/LLM_PROVIDERS.md](docs/LLM_PROVIDERS.md) 获取完整的提供商指南。
|
||||
|
||||
## 安全机制
|
||||
|
||||
IronClaw 实现了纵深防御策略来保护你的数据并防止滥用。
|
||||
|
||||
### WASM 沙箱
|
||||
|
||||
所有不受信任的工具都在隔离的 WebAssembly 容器中运行:
|
||||
|
||||
- **基于能力的权限** — 明确授权 HTTP、密钥、工具调用等能力
|
||||
- **端点白名单** — HTTP 请求仅限已批准的主机和路径
|
||||
- **凭据注入** — 密钥在宿主边界注入,永远不会暴露给 WASM 代码
|
||||
- **泄露检测** — 扫描请求和响应以防止密钥外泄
|
||||
- **速率限制** — 每个工具独立的请求限制,防止滥用
|
||||
- **资源限制** — 内存、CPU 和执行时间约束
|
||||
|
||||
```
|
||||
WASM ──► 白名单 ──► 泄露扫描 ──► 凭据 ──► 执行 ──► 泄露扫描 ──► WASM
|
||||
验证器 (请求) 注入器 请求 (响应)
|
||||
```
|
||||
|
||||
### 提示注入防御
|
||||
|
||||
外部内容需通过多个安全层:
|
||||
|
||||
- 基于模式的注入尝试检测
|
||||
- 内容清理和转义
|
||||
- 带严重级别的策略规则(阻止/警告/审核/清理)
|
||||
- 工具输出包装,确保安全的 LLM 上下文注入
|
||||
|
||||
### 数据保护
|
||||
|
||||
- 所有数据存储在本地 PostgreSQL 数据库中
|
||||
- 密钥使用 AES-256-GCM 加密
|
||||
- 无遥测、无分析、无数据共享
|
||||
- 所有工具执行的完整审计日志
|
||||
|
||||
## 系统架构
|
||||
|
||||
```
|
||||
┌────────────────────────────────────────────────────────────────┐
|
||||
│ 渠道 │
|
||||
│ ┌──────┐ ┌──────┐ ┌─────────────┐ ┌─────────────┐ │
|
||||
│ │ REPL │ │ HTTP │ │ WASM 渠道 │ │ Web 网关 │ │
|
||||
│ └──┬───┘ └──┬───┘ └──────┬──────┘ │ (SSE + WS) │ │
|
||||
│ │ │ │ └──────┬──────┘ │
|
||||
│ └─────────┴──────────────┴────────────────┘ │
|
||||
│ │ │
|
||||
│ ┌─────────▼─────────┐ │
|
||||
│ │ 代理循环 │ 意图路由 │
|
||||
│ └────┬──────────┬───┘ │
|
||||
│ │ │ │
|
||||
│ ┌──────────▼────┐ ┌──▼───────────────┐ │
|
||||
│ │ 调度器 │ │ 定时任务引擎 │ │
|
||||
│ │ (并行任务) │ │(cron, 事件, wh) │ │
|
||||
│ └──────┬────────┘ └────────┬─────────┘ │
|
||||
│ │ │ │
|
||||
│ ┌─────────────┼────────────────────┘ │
|
||||
│ │ │ │
|
||||
│ ┌───▼─────┐ ┌────▼────────────────┐ │
|
||||
│ │ 本地 │ │ 编排器 │ │
|
||||
│ │ 工作器 │ │ ┌───────────────┐ │ │
|
||||
│ │(进程内) │ │ │ Docker 沙箱 │ │ │
|
||||
│ └───┬─────┘ │ │ 容器 │ │ │
|
||||
│ │ │ │ ┌───────────┐ │ │ │
|
||||
│ │ │ │ │工作器/CC │ │ │ │
|
||||
│ │ │ │ └───────────┘ │ │ │
|
||||
│ │ │ └───────────────┘ │ │
|
||||
│ │ └─────────┬───────────┘ │
|
||||
│ └──────────────────┤ │
|
||||
│ │ │
|
||||
│ ┌───────────▼──────────┐ │
|
||||
│ │ 工具注册表 │ │
|
||||
│ │ 内置、MCP、WASM │ │
|
||||
│ └──────────────────────┘ │
|
||||
└────────────────────────────────────────────────────────────────┘
|
||||
```
|
||||
|
||||
### 核心组件
|
||||
|
||||
| 组件 | 用途 |
|
||||
|------|------|
|
||||
| **代理循环** | 主消息处理和任务协调 |
|
||||
| **路由器** | 分类用户意图(命令、查询、任务) |
|
||||
| **调度器** | 管理带优先级的并行任务执行 |
|
||||
| **工作器** | 执行包含 LLM 推理和工具调用的任务 |
|
||||
| **编排器** | 容器生命周期、LLM 代理、每任务认证 |
|
||||
| **Web 网关** | 浏览器 UI,含聊天、记忆、任务、日志、扩展、定时任务 |
|
||||
| **定时任务引擎** | 定时(cron)和响应式(事件、webhook)后台任务 |
|
||||
| **工作空间** | 带混合搜索的持久记忆 |
|
||||
| **安全层** | 提示注入防御和内容清理 |
|
||||
|
||||
## 使用方式
|
||||
|
||||
```bash
|
||||
# 首次设置(配置数据库、认证等)
|
||||
ironclaw onboard
|
||||
|
||||
# 启动交互式 REPL
|
||||
cargo run
|
||||
|
||||
# 启用调试日志
|
||||
RUST_LOG=ironclaw=debug cargo run
|
||||
```
|
||||
|
||||
## 开发
|
||||
|
||||
```bash
|
||||
# 格式化代码
|
||||
cargo fmt
|
||||
|
||||
# 代码检查
|
||||
cargo clippy --all --benches --tests --examples --all-features
|
||||
|
||||
# 运行测试
|
||||
createdb ironclaw_test
|
||||
cargo test
|
||||
|
||||
# 运行指定测试
|
||||
cargo test test_name
|
||||
```
|
||||
|
||||
- **Telegram 渠道**:参见 [docs/TELEGRAM_SETUP.md](docs/TELEGRAM_SETUP.md) 了解设置和私信配对。
|
||||
- **修改渠道源码**:在 `cargo build` 之前运行 `./channels-src/telegram/build.sh` 以便打包更新后的 WASM。
|
||||
|
||||
## OpenClaw 传承
|
||||
|
||||
IronClaw 是受 [OpenClaw](https://github.com/openclaw/openclaw) 启发的 Rust 重新实现。参见 [FEATURE_PARITY.md](FEATURE_PARITY.md) 了解完整的功能追踪矩阵。
|
||||
|
||||
主要差异:
|
||||
|
||||
- **Rust vs TypeScript** — 原生性能、内存安全、单一二进制文件
|
||||
- **WASM 沙箱 vs Docker** — 轻量级、基于能力的安全机制
|
||||
- **PostgreSQL vs SQLite** — 生产级持久化存储
|
||||
- **安全优先设计** — 多层防御、凭据保护
|
||||
|
||||
## 许可证
|
||||
|
||||
可选择以下任一许可证:
|
||||
|
||||
- Apache License, Version 2.0 ([LICENSE-APACHE](LICENSE-APACHE))
|
||||
- MIT License ([LICENSE-MIT](LICENSE-MIT))
|
||||
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "discord-channel"
|
||||
version = "0.1.0"
|
||||
version = "0.2.0"
|
||||
edition = "2021"
|
||||
description = "Discord channel for IronClaw"
|
||||
license = "MIT OR Apache-2.0"
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"version": "0.2.0",
|
||||
"wit_version": "0.3.0",
|
||||
"type": "channel",
|
||||
"name": "discord",
|
||||
"description": "Discord Gateway/Webhook channel for handling slash commands, buttons, and messages",
|
||||
|
||||
@@ -312,6 +312,10 @@ impl Guest for DiscordChannel {
|
||||
|
||||
fn on_status(_update: StatusUpdate) {}
|
||||
|
||||
fn on_broadcast(_user_id: String, _response: AgentResponse) -> Result<(), String> {
|
||||
Err("broadcast not yet implemented for Discord channel".to_string())
|
||||
}
|
||||
|
||||
fn on_shutdown() {
|
||||
channel_host::log(
|
||||
channel_host::LogLevel::Info,
|
||||
@@ -414,6 +418,7 @@ fn handle_slash_command(interaction: &DiscordInteraction) -> bool {
|
||||
content,
|
||||
thread_id: None,
|
||||
metadata_json,
|
||||
attachments: vec![],
|
||||
});
|
||||
true
|
||||
}
|
||||
@@ -467,6 +472,7 @@ fn handle_message_component(interaction: &DiscordInteraction, message: &DiscordM
|
||||
content: format!("[Button clicked] {}", message.content),
|
||||
thread_id: None,
|
||||
metadata_json,
|
||||
attachments: vec![],
|
||||
});
|
||||
}
|
||||
|
||||
@@ -683,4 +689,34 @@ mod tests {
|
||||
assert_eq!(parsed.channel_id, "123");
|
||||
assert_eq!(parsed.interaction_id, "456");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_slash_command_interaction() {
|
||||
// Verify that a slash command interaction deserializes correctly.
|
||||
let json = r#"{
|
||||
"type": 2,
|
||||
"id": "int_1",
|
||||
"application_id": "app_1",
|
||||
"channel_id": "ch_1",
|
||||
"member": {
|
||||
"user": {
|
||||
"id": "user_1",
|
||||
"username": "testuser",
|
||||
"global_name": "Test User"
|
||||
}
|
||||
},
|
||||
"data": {
|
||||
"id": "cmd_1",
|
||||
"name": "ask",
|
||||
"options": [
|
||||
{"name": "question", "value": "What is rust?"}
|
||||
]
|
||||
},
|
||||
"token": "token_abc"
|
||||
}"#;
|
||||
|
||||
let interaction: DiscordInteraction = serde_json::from_str(json).unwrap();
|
||||
assert_eq!(interaction.interaction_type, 2);
|
||||
assert!(interaction.data.is_some());
|
||||
}
|
||||
}
|
||||
|
||||
Generated
+1
-1
@@ -267,7 +267,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "slack-channel"
|
||||
version = "0.1.0"
|
||||
version = "0.2.1"
|
||||
dependencies = [
|
||||
"hex",
|
||||
"hmac",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "slack-channel"
|
||||
version = "0.1.0"
|
||||
version = "0.2.1"
|
||||
edition = "2021"
|
||||
description = "Slack Events API channel for IronClaw"
|
||||
license = "MIT OR Apache-2.0"
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"version": "0.2.0",
|
||||
"wit_version": "0.3.0",
|
||||
"type": "channel",
|
||||
"name": "slack",
|
||||
"description": "Slack Events API channel for receiving and responding to Slack messages",
|
||||
|
||||
@@ -29,7 +29,7 @@ use exports::near::agent::channel::{
|
||||
AgentResponse, ChannelConfig, Guest, HttpEndpointConfig, IncomingHttpRequest,
|
||||
OutgoingHttpResponse, StatusUpdate,
|
||||
};
|
||||
use near::agent::channel_host::{self, EmittedMessage};
|
||||
use near::agent::channel_host::{self, EmittedMessage, InboundAttachment};
|
||||
|
||||
/// Slack event wrapper.
|
||||
#[derive(Debug, Deserialize)]
|
||||
@@ -78,6 +78,25 @@ struct SlackEvent {
|
||||
|
||||
/// Subtype (bot_message, etc.)
|
||||
subtype: Option<String>,
|
||||
|
||||
/// File attachments shared in the message.
|
||||
#[serde(default)]
|
||||
files: Option<Vec<SlackFile>>,
|
||||
}
|
||||
|
||||
/// Slack file attachment.
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct SlackFile {
|
||||
/// File ID.
|
||||
id: String,
|
||||
/// MIME type.
|
||||
mimetype: Option<String>,
|
||||
/// Original filename.
|
||||
name: Option<String>,
|
||||
/// File size in bytes.
|
||||
size: Option<u64>,
|
||||
/// URL to download the file (requires auth).
|
||||
url_private: Option<String>,
|
||||
}
|
||||
|
||||
/// Metadata stored with emitted messages for response routing.
|
||||
@@ -306,13 +325,140 @@ impl Guest for SlackChannel {
|
||||
|
||||
fn on_status(_update: StatusUpdate) {}
|
||||
|
||||
fn on_broadcast(_user_id: String, _response: AgentResponse) -> Result<(), String> {
|
||||
Err("broadcast not yet implemented for Slack channel".to_string())
|
||||
}
|
||||
|
||||
fn on_shutdown() {
|
||||
channel_host::log(channel_host::LogLevel::Info, "Slack channel shutting down");
|
||||
}
|
||||
}
|
||||
|
||||
/// Extract attachments from Slack file objects.
|
||||
fn extract_slack_attachments(files: &Option<Vec<SlackFile>>) -> Vec<InboundAttachment> {
|
||||
let Some(files) = files else {
|
||||
return Vec::new();
|
||||
};
|
||||
files
|
||||
.iter()
|
||||
.map(|f| InboundAttachment {
|
||||
id: f.id.clone(),
|
||||
mime_type: f
|
||||
.mimetype
|
||||
.clone()
|
||||
.unwrap_or_else(|| "application/octet-stream".to_string()),
|
||||
filename: f.name.clone(),
|
||||
size_bytes: f.size,
|
||||
source_url: f.url_private.clone(),
|
||||
storage_key: None,
|
||||
extracted_text: None,
|
||||
extras_json: String::new(),
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Download a file from Slack using the url_private endpoint.
|
||||
///
|
||||
/// Slack file downloads require Bearer auth with the bot token, which is
|
||||
/// injected by the host credential system via `channel_host::http_request`.
|
||||
fn download_slack_file(url: &str) -> Result<Vec<u8>, String> {
|
||||
let headers = serde_json::json!({});
|
||||
|
||||
let result = channel_host::http_request("GET", url, &headers.to_string(), None, None);
|
||||
|
||||
let response = result.map_err(|e| format!("Slack file download failed: {}", e))?;
|
||||
|
||||
if response.status != 200 {
|
||||
let body_str = String::from_utf8_lossy(&response.body);
|
||||
return Err(format!(
|
||||
"Slack file download returned {}: {}",
|
||||
response.status, body_str
|
||||
));
|
||||
}
|
||||
|
||||
Ok(response.body)
|
||||
}
|
||||
|
||||
/// Download file bytes and store them via the host for processing.
|
||||
///
|
||||
/// Downloads all file types (images, documents, etc.) so the host-side
|
||||
/// middleware can process them (vision pipeline for images, text extraction
|
||||
/// for documents, transcription for audio, etc.).
|
||||
/// Maximum file size to download (20 MB). Files larger than this are skipped
|
||||
/// to avoid excessive memory use and slow downloads in the WASM runtime.
|
||||
const MAX_DOWNLOAD_SIZE_BYTES: u64 = 20 * 1024 * 1024;
|
||||
|
||||
fn download_and_store_slack_files(attachments: &[InboundAttachment]) {
|
||||
for att in attachments {
|
||||
let Some(ref url) = att.source_url else {
|
||||
continue;
|
||||
};
|
||||
|
||||
// Skip files that exceed the size limit
|
||||
if let Some(size) = att.size_bytes {
|
||||
if size > MAX_DOWNLOAD_SIZE_BYTES {
|
||||
channel_host::log(
|
||||
channel_host::LogLevel::Warn,
|
||||
&format!(
|
||||
"Skipping Slack file download: {} bytes exceeds {} MB limit (id={})",
|
||||
size,
|
||||
MAX_DOWNLOAD_SIZE_BYTES / (1024 * 1024),
|
||||
att.id
|
||||
),
|
||||
);
|
||||
continue;
|
||||
}
|
||||
}
|
||||
|
||||
match download_slack_file(url) {
|
||||
Ok(bytes) => {
|
||||
// Post-download size guard: metadata size_bytes is optional,
|
||||
// so a file with no size info could bypass the pre-download check.
|
||||
if bytes.len() as u64 > MAX_DOWNLOAD_SIZE_BYTES {
|
||||
channel_host::log(
|
||||
channel_host::LogLevel::Warn,
|
||||
&format!(
|
||||
"Discarding Slack file after download: {} bytes exceeds {} MB limit (id={})",
|
||||
bytes.len(),
|
||||
MAX_DOWNLOAD_SIZE_BYTES / (1024 * 1024),
|
||||
att.id
|
||||
),
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
channel_host::log(
|
||||
channel_host::LogLevel::Info,
|
||||
&format!(
|
||||
"Downloaded Slack file: {} bytes, mime={}",
|
||||
bytes.len(),
|
||||
att.mime_type
|
||||
),
|
||||
);
|
||||
if let Err(e) = channel_host::store_attachment_data(&att.id, &bytes) {
|
||||
channel_host::log(
|
||||
channel_host::LogLevel::Error,
|
||||
&format!("Failed to store Slack file data: {}", e),
|
||||
);
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
channel_host::log(
|
||||
channel_host::LogLevel::Error,
|
||||
&format!("Failed to download Slack file: {}", e),
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Handle a Slack event and emit message if applicable.
|
||||
fn handle_slack_event(event: SlackEvent, team_id: Option<String>, _event_id: Option<String>) {
|
||||
let attachments = extract_slack_attachments(&event.files);
|
||||
|
||||
// Download and store file attachments for host-side processing
|
||||
download_and_store_slack_files(&attachments);
|
||||
|
||||
match event.event_type.as_str() {
|
||||
// Direct mention of the bot (always in a channel, not a DM)
|
||||
"app_mention" => {
|
||||
@@ -326,7 +472,14 @@ fn handle_slack_event(event: SlackEvent, team_id: Option<String>, _event_id: Opt
|
||||
if !check_sender_permission(&user, &channel, false) {
|
||||
return;
|
||||
}
|
||||
emit_message(user, text, channel, event.thread_ts.or(Some(ts)), team_id);
|
||||
emit_message(
|
||||
user,
|
||||
text,
|
||||
channel,
|
||||
event.thread_ts.or(Some(ts)),
|
||||
team_id,
|
||||
attachments,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -348,7 +501,14 @@ fn handle_slack_event(event: SlackEvent, team_id: Option<String>, _event_id: Opt
|
||||
if !check_sender_permission(&user, &channel, true) {
|
||||
return;
|
||||
}
|
||||
emit_message(user, text, channel, event.thread_ts.or(Some(ts)), team_id);
|
||||
emit_message(
|
||||
user,
|
||||
text,
|
||||
channel,
|
||||
event.thread_ts.or(Some(ts)),
|
||||
team_id,
|
||||
attachments,
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -369,6 +529,7 @@ fn emit_message(
|
||||
channel: String,
|
||||
thread_ts: Option<String>,
|
||||
team_id: Option<String>,
|
||||
attachments: Vec<InboundAttachment>,
|
||||
) {
|
||||
let message_ts = thread_ts.clone().unwrap_or_default();
|
||||
|
||||
@@ -396,6 +557,7 @@ fn emit_message(
|
||||
content: cleaned_text,
|
||||
thread_id: thread_ts,
|
||||
metadata_json,
|
||||
attachments,
|
||||
});
|
||||
}
|
||||
|
||||
@@ -551,3 +713,117 @@ fn json_response(status: u16, value: serde_json::Value) -> OutgoingHttpResponse
|
||||
|
||||
// Export the component
|
||||
export!(SlackChannel);
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_extract_slack_attachments_with_files() {
|
||||
let files = Some(vec![
|
||||
SlackFile {
|
||||
id: "F123".to_string(),
|
||||
mimetype: Some("image/png".to_string()),
|
||||
name: Some("screenshot.png".to_string()),
|
||||
size: Some(50000),
|
||||
url_private: Some("https://files.slack.com/F123".to_string()),
|
||||
},
|
||||
SlackFile {
|
||||
id: "F456".to_string(),
|
||||
mimetype: Some("application/pdf".to_string()),
|
||||
name: Some("doc.pdf".to_string()),
|
||||
size: Some(120000),
|
||||
url_private: None,
|
||||
},
|
||||
]);
|
||||
|
||||
let attachments = extract_slack_attachments(&files);
|
||||
assert_eq!(attachments.len(), 2);
|
||||
|
||||
assert_eq!(attachments[0].id, "F123");
|
||||
assert_eq!(attachments[0].mime_type, "image/png");
|
||||
assert_eq!(attachments[0].filename, Some("screenshot.png".to_string()));
|
||||
assert_eq!(attachments[0].size_bytes, Some(50000));
|
||||
assert_eq!(
|
||||
attachments[0].source_url,
|
||||
Some("https://files.slack.com/F123".to_string())
|
||||
);
|
||||
|
||||
assert_eq!(attachments[1].id, "F456");
|
||||
assert_eq!(attachments[1].mime_type, "application/pdf");
|
||||
assert!(attachments[1].source_url.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_slack_attachments_none() {
|
||||
let attachments = extract_slack_attachments(&None);
|
||||
assert!(attachments.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_slack_attachments_empty() {
|
||||
let attachments = extract_slack_attachments(&Some(vec![]));
|
||||
assert!(attachments.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_slack_attachments_missing_mime() {
|
||||
let files = Some(vec![SlackFile {
|
||||
id: "F789".to_string(),
|
||||
mimetype: None,
|
||||
name: Some("unknown".to_string()),
|
||||
size: None,
|
||||
url_private: None,
|
||||
}]);
|
||||
|
||||
let attachments = extract_slack_attachments(&files);
|
||||
assert_eq!(attachments.len(), 1);
|
||||
assert_eq!(attachments[0].mime_type, "application/octet-stream");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_slack_event_with_files() {
|
||||
let json = r#"{
|
||||
"type": "message",
|
||||
"user": "U123",
|
||||
"channel": "D456",
|
||||
"text": "Check this file",
|
||||
"ts": "1234567890.000001",
|
||||
"files": [
|
||||
{
|
||||
"id": "F001",
|
||||
"mimetype": "image/jpeg",
|
||||
"name": "photo.jpg",
|
||||
"size": 30000,
|
||||
"url_private": "https://files.slack.com/F001"
|
||||
}
|
||||
]
|
||||
}"#;
|
||||
|
||||
let event: SlackEvent = serde_json::from_str(json).unwrap();
|
||||
assert!(event.files.is_some());
|
||||
let files = event.files.unwrap();
|
||||
assert_eq!(files.len(), 1);
|
||||
assert_eq!(files[0].id, "F001");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_slack_event_without_files() {
|
||||
let json = r#"{
|
||||
"type": "message",
|
||||
"user": "U123",
|
||||
"channel": "D456",
|
||||
"text": "Just text",
|
||||
"ts": "1234567890.000001"
|
||||
}"#;
|
||||
|
||||
let event: SlackEvent = serde_json::from_str(json).unwrap();
|
||||
assert!(event.files.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_max_download_size_constant() {
|
||||
// Verify the constant is 20 MB
|
||||
assert_eq!(MAX_DOWNLOAD_SIZE_BYTES, 20 * 1024 * 1024);
|
||||
}
|
||||
}
|
||||
|
||||
Generated
+1
-1
@@ -212,7 +212,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "telegram-channel"
|
||||
version = "0.1.0"
|
||||
version = "0.2.1"
|
||||
dependencies = [
|
||||
"serde",
|
||||
"serde_json",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "telegram-channel"
|
||||
version = "0.1.0"
|
||||
version = "0.2.1"
|
||||
edition = "2021"
|
||||
description = "Telegram Bot API channel for IronClaw"
|
||||
license = "MIT OR Apache-2.0"
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,9 +1,17 @@
|
||||
{
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"version": "0.2.2",
|
||||
"wit_version": "0.3.0",
|
||||
"type": "channel",
|
||||
"name": "telegram",
|
||||
"description": "Telegram Bot API channel for receiving and responding to Telegram messages",
|
||||
"auth": {
|
||||
"secret_name": "telegram_bot_token",
|
||||
"display_name": "Telegram",
|
||||
"instructions": "Get your bot token from @BotFather on Telegram (https://t.me/BotFather). Send /newbot or /token to get it.",
|
||||
"setup_url": "https://t.me/BotFather",
|
||||
"token_hint": "Looks like 123456789:AABBccDDeeFFgg...",
|
||||
"env_var": "TELEGRAM_BOT_TOKEN"
|
||||
},
|
||||
"setup": {
|
||||
"required_secrets": [
|
||||
{
|
||||
@@ -17,7 +25,8 @@
|
||||
"capabilities": {
|
||||
"http": {
|
||||
"allowlist": [
|
||||
{ "host": "api.telegram.org", "path_prefix": "/bot" }
|
||||
{ "host": "api.telegram.org", "path_prefix": "/bot" },
|
||||
{ "host": "api.telegram.org", "path_prefix": "/file/bot" }
|
||||
],
|
||||
"credentials": {
|
||||
"telegram_bot": {
|
||||
@@ -26,6 +35,7 @@
|
||||
"host_patterns": ["api.telegram.org"]
|
||||
}
|
||||
},
|
||||
"max_response_bytes": 52428800,
|
||||
"rate_limit": {
|
||||
"requests_per_minute": 30,
|
||||
"requests_per_hour": 1000
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "whatsapp-channel"
|
||||
version = "0.1.0"
|
||||
version = "0.2.0"
|
||||
edition = "2021"
|
||||
description = "WhatsApp Cloud API channel for IronClaw"
|
||||
|
||||
|
||||
@@ -32,7 +32,7 @@ use exports::near::agent::channel::{
|
||||
AgentResponse, ChannelConfig, Guest, HttpEndpointConfig, IncomingHttpRequest,
|
||||
OutgoingHttpResponse, StatusUpdate,
|
||||
};
|
||||
use near::agent::channel_host::{self, EmittedMessage};
|
||||
use near::agent::channel_host::{self, EmittedMessage, InboundAttachment};
|
||||
|
||||
// ============================================================================
|
||||
// WhatsApp Cloud API Types
|
||||
@@ -137,10 +137,46 @@ struct WhatsAppMessage {
|
||||
/// Text content (if type is "text")
|
||||
text: Option<TextContent>,
|
||||
|
||||
/// Image content
|
||||
image: Option<WhatsAppMedia>,
|
||||
|
||||
/// Audio content
|
||||
audio: Option<WhatsAppMedia>,
|
||||
|
||||
/// Video content
|
||||
video: Option<WhatsAppMedia>,
|
||||
|
||||
/// Document content
|
||||
document: Option<WhatsAppDocument>,
|
||||
|
||||
/// Context for replies
|
||||
context: Option<MessageContext>,
|
||||
}
|
||||
|
||||
/// WhatsApp media attachment (image, audio, video).
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct WhatsAppMedia {
|
||||
/// Media ID (use to download via Graph API)
|
||||
id: String,
|
||||
/// MIME type
|
||||
mime_type: Option<String>,
|
||||
/// Caption text
|
||||
caption: Option<String>,
|
||||
}
|
||||
|
||||
/// WhatsApp document attachment.
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct WhatsAppDocument {
|
||||
/// Media ID
|
||||
id: String,
|
||||
/// MIME type
|
||||
mime_type: Option<String>,
|
||||
/// Filename
|
||||
filename: Option<String>,
|
||||
/// Caption text
|
||||
caption: Option<String>,
|
||||
}
|
||||
|
||||
/// Text message content.
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct TextContent {
|
||||
@@ -476,6 +512,10 @@ impl Guest for WhatsAppChannel {
|
||||
|
||||
fn on_status(_update: StatusUpdate) {}
|
||||
|
||||
fn on_broadcast(_user_id: String, _response: AgentResponse) -> Result<(), String> {
|
||||
Err("broadcast not yet implemented for WhatsApp channel".to_string())
|
||||
}
|
||||
|
||||
fn on_shutdown() {
|
||||
channel_host::log(
|
||||
channel_host::LogLevel::Info,
|
||||
@@ -618,26 +658,102 @@ fn handle_incoming_message(req: &IncomingHttpRequest) -> OutgoingHttpResponse {
|
||||
json_response(200, serde_json::json!({"status": "ok"}))
|
||||
}
|
||||
|
||||
/// Extract attachments from a WhatsApp message.
|
||||
fn extract_whatsapp_attachments(message: &WhatsAppMessage) -> Vec<InboundAttachment> {
|
||||
let mut attachments = Vec::new();
|
||||
|
||||
if let Some(ref img) = message.image {
|
||||
attachments.push(InboundAttachment {
|
||||
id: img.id.clone(),
|
||||
mime_type: img
|
||||
.mime_type
|
||||
.clone()
|
||||
.unwrap_or_else(|| "image/jpeg".to_string()),
|
||||
filename: None,
|
||||
size_bytes: None,
|
||||
source_url: None, // WhatsApp requires Graph API call with media ID to get URL
|
||||
storage_key: None,
|
||||
extracted_text: img.caption.clone(),
|
||||
extras_json: String::new(),
|
||||
});
|
||||
}
|
||||
|
||||
if let Some(ref audio) = message.audio {
|
||||
attachments.push(InboundAttachment {
|
||||
id: audio.id.clone(),
|
||||
mime_type: audio
|
||||
.mime_type
|
||||
.clone()
|
||||
.unwrap_or_else(|| "audio/ogg".to_string()),
|
||||
filename: None,
|
||||
size_bytes: None,
|
||||
source_url: None,
|
||||
storage_key: None,
|
||||
extracted_text: audio.caption.clone(),
|
||||
extras_json: String::new(),
|
||||
});
|
||||
}
|
||||
|
||||
if let Some(ref video) = message.video {
|
||||
attachments.push(InboundAttachment {
|
||||
id: video.id.clone(),
|
||||
mime_type: video
|
||||
.mime_type
|
||||
.clone()
|
||||
.unwrap_or_else(|| "video/mp4".to_string()),
|
||||
filename: None,
|
||||
size_bytes: None,
|
||||
source_url: None,
|
||||
storage_key: None,
|
||||
extracted_text: video.caption.clone(),
|
||||
extras_json: String::new(),
|
||||
});
|
||||
}
|
||||
|
||||
if let Some(ref doc) = message.document {
|
||||
attachments.push(InboundAttachment {
|
||||
id: doc.id.clone(),
|
||||
mime_type: doc
|
||||
.mime_type
|
||||
.clone()
|
||||
.unwrap_or_else(|| "application/octet-stream".to_string()),
|
||||
filename: doc.filename.clone(),
|
||||
size_bytes: None,
|
||||
source_url: None,
|
||||
storage_key: None,
|
||||
extracted_text: doc.caption.clone(),
|
||||
extras_json: String::new(),
|
||||
});
|
||||
}
|
||||
|
||||
attachments
|
||||
}
|
||||
|
||||
/// Process a single WhatsApp message.
|
||||
fn handle_message(
|
||||
message: &WhatsAppMessage,
|
||||
phone_number_id: &str,
|
||||
contact_names: &std::collections::HashMap<String, String>,
|
||||
) {
|
||||
// Only handle text messages for now
|
||||
// TODO: Add support for image, audio, video, document, etc.
|
||||
if message.message_type != "text" {
|
||||
channel_host::log(
|
||||
channel_host::LogLevel::Debug,
|
||||
&format!("Skipping non-text message type: {}", message.message_type),
|
||||
);
|
||||
return;
|
||||
}
|
||||
let attachments = extract_whatsapp_attachments(message);
|
||||
|
||||
// Extract text content
|
||||
// Extract text content (from text body or media captions)
|
||||
let text = match &message.text {
|
||||
Some(t) if !t.body.is_empty() => t.body.clone(),
|
||||
_ => return,
|
||||
_ => {
|
||||
// Try to use caption from media messages as content
|
||||
let caption = message
|
||||
.image
|
||||
.as_ref()
|
||||
.and_then(|m| m.caption.clone())
|
||||
.or_else(|| message.video.as_ref().and_then(|m| m.caption.clone()))
|
||||
.or_else(|| message.document.as_ref().and_then(|m| m.caption.clone()));
|
||||
match caption {
|
||||
Some(c) if !c.is_empty() => c,
|
||||
_ if !attachments.is_empty() => String::new(),
|
||||
_ => return,
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
// Look up sender's name from contacts
|
||||
@@ -670,6 +786,7 @@ fn handle_message(
|
||||
content: text,
|
||||
thread_id: None, // WhatsApp doesn't have threads like Slack/Discord
|
||||
metadata_json,
|
||||
attachments,
|
||||
});
|
||||
|
||||
channel_host::log(
|
||||
@@ -947,4 +1064,138 @@ mod tests {
|
||||
assert_eq!(parsed.phone_number_id, "123456");
|
||||
assert_eq!(parsed.sender_phone, "15551234567");
|
||||
}
|
||||
|
||||
// === Attachment extraction fixture tests ===
|
||||
|
||||
#[test]
|
||||
fn test_extract_whatsapp_image_attachment() {
|
||||
let msg = WhatsAppMessage {
|
||||
id: "msg1".to_string(),
|
||||
from: "15551234567".to_string(),
|
||||
timestamp: "1234567890".to_string(),
|
||||
message_type: "image".to_string(),
|
||||
text: None,
|
||||
image: Some(WhatsAppMedia {
|
||||
id: "media_img_1".to_string(),
|
||||
mime_type: Some("image/jpeg".to_string()),
|
||||
caption: Some("Look at this".to_string()),
|
||||
}),
|
||||
audio: None,
|
||||
video: None,
|
||||
document: None,
|
||||
context: None,
|
||||
};
|
||||
|
||||
let attachments = extract_whatsapp_attachments(&msg);
|
||||
assert_eq!(attachments.len(), 1);
|
||||
assert_eq!(attachments[0].id, "media_img_1");
|
||||
assert_eq!(attachments[0].mime_type, "image/jpeg");
|
||||
assert_eq!(
|
||||
attachments[0].extracted_text,
|
||||
Some("Look at this".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_whatsapp_document_attachment() {
|
||||
let msg = WhatsAppMessage {
|
||||
id: "msg2".to_string(),
|
||||
from: "15551234567".to_string(),
|
||||
timestamp: "1234567890".to_string(),
|
||||
message_type: "document".to_string(),
|
||||
text: None,
|
||||
image: None,
|
||||
audio: None,
|
||||
video: None,
|
||||
document: Some(WhatsAppDocument {
|
||||
id: "media_doc_1".to_string(),
|
||||
mime_type: Some("application/pdf".to_string()),
|
||||
filename: Some("report.pdf".to_string()),
|
||||
caption: None,
|
||||
}),
|
||||
context: None,
|
||||
};
|
||||
|
||||
let attachments = extract_whatsapp_attachments(&msg);
|
||||
assert_eq!(attachments.len(), 1);
|
||||
assert_eq!(attachments[0].id, "media_doc_1");
|
||||
assert_eq!(attachments[0].mime_type, "application/pdf");
|
||||
assert_eq!(
|
||||
attachments[0].filename,
|
||||
Some("report.pdf".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_whatsapp_audio_video_attachments() {
|
||||
let msg = WhatsAppMessage {
|
||||
id: "msg3".to_string(),
|
||||
from: "15551234567".to_string(),
|
||||
timestamp: "1234567890".to_string(),
|
||||
message_type: "audio".to_string(),
|
||||
text: None,
|
||||
image: None,
|
||||
audio: Some(WhatsAppMedia {
|
||||
id: "media_audio_1".to_string(),
|
||||
mime_type: Some("audio/ogg".to_string()),
|
||||
caption: None,
|
||||
}),
|
||||
video: Some(WhatsAppMedia {
|
||||
id: "media_video_1".to_string(),
|
||||
mime_type: Some("video/mp4".to_string()),
|
||||
caption: None,
|
||||
}),
|
||||
document: None,
|
||||
context: None,
|
||||
};
|
||||
|
||||
let attachments = extract_whatsapp_attachments(&msg);
|
||||
assert_eq!(attachments.len(), 2);
|
||||
assert_eq!(attachments[0].id, "media_audio_1");
|
||||
assert_eq!(attachments[1].id, "media_video_1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_whatsapp_text_only_no_attachments() {
|
||||
let msg = WhatsAppMessage {
|
||||
id: "msg4".to_string(),
|
||||
from: "15551234567".to_string(),
|
||||
timestamp: "1234567890".to_string(),
|
||||
message_type: "text".to_string(),
|
||||
text: Some(TextContent {
|
||||
body: "Hello".to_string(),
|
||||
}),
|
||||
image: None,
|
||||
audio: None,
|
||||
video: None,
|
||||
document: None,
|
||||
context: None,
|
||||
};
|
||||
|
||||
let attachments = extract_whatsapp_attachments(&msg);
|
||||
assert!(attachments.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_whatsapp_image_message() {
|
||||
let json = r#"{
|
||||
"id": "wamid.123",
|
||||
"from": "15551234567",
|
||||
"timestamp": "1234567890",
|
||||
"type": "image",
|
||||
"image": {
|
||||
"id": "media_img_abc",
|
||||
"mime_type": "image/jpeg",
|
||||
"caption": "Check this"
|
||||
}
|
||||
}"#;
|
||||
|
||||
let msg: WhatsAppMessage = serde_json::from_str(json).unwrap();
|
||||
assert_eq!(msg.message_type, "image");
|
||||
assert!(msg.image.is_some());
|
||||
|
||||
let attachments = extract_whatsapp_attachments(&msg);
|
||||
assert_eq!(attachments.len(), 1);
|
||||
assert_eq!(attachments[0].id, "media_img_abc");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"version": "0.2.0",
|
||||
"wit_version": "0.3.0",
|
||||
"type": "channel",
|
||||
"name": "whatsapp",
|
||||
"description": "WhatsApp Cloud API channel for receiving and responding to WhatsApp messages",
|
||||
|
||||
+1
-1
@@ -3,7 +3,7 @@ services:
|
||||
postgres:
|
||||
image: pgvector/pgvector:pg16
|
||||
ports:
|
||||
- "5432:5432"
|
||||
- "127.0.0.1:5432:5432"
|
||||
environment:
|
||||
POSTGRES_DB: ironclaw
|
||||
POSTGRES_USER: ironclaw
|
||||
|
||||
@@ -11,7 +11,13 @@ configurations.
|
||||
| NEAR AI | `nearai` | OAuth (browser) | Default; multi-model |
|
||||
| Anthropic | `anthropic` | `ANTHROPIC_API_KEY` | Claude models |
|
||||
| OpenAI | `openai` | `OPENAI_API_KEY` | GPT models |
|
||||
| Google Gemini | `gemini` | `GEMINI_API_KEY` | Gemini models |
|
||||
| io.net | `ionet` | `IONET_API_KEY` | Intelligence API |
|
||||
| Mistral | `mistral` | `MISTRAL_API_KEY` | Mistral models |
|
||||
| Yandex AI Studio | `yandex` | `YANDEX_API_KEY` | YandexGPT models |
|
||||
| Cloudflare Workers AI | `cloudflare` | `CLOUDFLARE_API_KEY` | Access to Workers AI |
|
||||
| Ollama | `ollama` | No | Local inference |
|
||||
| AWS Bedrock | `bedrock` | AWS credentials | Native Converse API |
|
||||
| OpenRouter | `openai_compatible` | `LLM_API_KEY` | 300+ models |
|
||||
| Together AI | `openai_compatible` | `LLM_API_KEY` | Fast inference |
|
||||
| Fireworks AI | `openai_compatible` | `LLM_API_KEY` | Fast inference |
|
||||
@@ -68,6 +74,55 @@ Pull a model first: `ollama pull llama3.2`
|
||||
|
||||
---
|
||||
|
||||
## AWS Bedrock (requires `--features bedrock`)
|
||||
|
||||
Uses the native AWS Converse API via `aws-sdk-bedrockruntime`. Supports standard AWS
|
||||
authentication methods: IAM credentials, SSO profiles, and instance roles.
|
||||
|
||||
> **Build prerequisite:** The `aws-lc-sys` crate (transitive dependency via AWS SDK)
|
||||
> requires **CMake** to compile. Install it before building with `--features bedrock`:
|
||||
> - macOS: `brew install cmake`
|
||||
> - Ubuntu/Debian: `sudo apt install cmake`
|
||||
> - Fedora: `sudo dnf install cmake`
|
||||
|
||||
### With AWS credentials (IAM, SSO, instance roles)
|
||||
|
||||
```env
|
||||
LLM_BACKEND=bedrock
|
||||
BEDROCK_MODEL=anthropic.claude-opus-4-6-v1
|
||||
BEDROCK_REGION=us-east-1
|
||||
BEDROCK_CROSS_REGION=us
|
||||
# AWS_PROFILE=my-sso-profile # optional, for named profiles
|
||||
```
|
||||
|
||||
The AWS SDK credential chain automatically resolves credentials from environment
|
||||
variables (`AWS_ACCESS_KEY_ID`, `AWS_SECRET_ACCESS_KEY`), shared credentials file
|
||||
(`~/.aws/credentials`), SSO profiles, and EC2/ECS instance roles.
|
||||
|
||||
### Cross-region inference
|
||||
|
||||
Set `BEDROCK_CROSS_REGION` to route requests across AWS regions for capacity:
|
||||
|
||||
| Prefix | Routing |
|
||||
|---|---|
|
||||
| `us` | US regions (us-east-1, us-east-2, us-west-2) |
|
||||
| `eu` | European regions |
|
||||
| `apac` | Asia-Pacific regions |
|
||||
| `global` | All commercial AWS regions |
|
||||
| _(unset)_ | Single-region only |
|
||||
|
||||
### Popular Bedrock model IDs
|
||||
|
||||
| Model | ID |
|
||||
|---|---|
|
||||
| Claude Opus 4.6 | `anthropic.claude-opus-4-6-v1` |
|
||||
| Claude Sonnet 4.5 | `anthropic.claude-sonnet-4-5-20250929-v1:0` |
|
||||
| Claude Haiku 4.5 | `anthropic.claude-haiku-4-5-20251001-v1:0` |
|
||||
| Amazon Nova Pro | `amazon.nova-pro-v1:0` |
|
||||
| Llama 4 Maverick | `meta.llama4-maverick-17b-instruct-v1:0` |
|
||||
|
||||
---
|
||||
|
||||
## OpenAI-Compatible Endpoints
|
||||
|
||||
All providers below use `LLM_BACKEND=openai_compatible`. Set `LLM_BASE_URL` to the
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
-- Partial unique indexes to prevent duplicate singleton conversations.
|
||||
-- These guard against TOCTOU races in get_or_create_routine_conversation
|
||||
-- and get_or_create_heartbeat_conversation.
|
||||
|
||||
-- One routine conversation per user per routine_id.
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS uq_conv_routine
|
||||
ON conversations (user_id, (metadata->>'routine_id'))
|
||||
WHERE metadata->>'routine_id' IS NOT NULL;
|
||||
|
||||
-- One heartbeat conversation per user.
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS uq_conv_heartbeat
|
||||
ON conversations (user_id)
|
||||
WHERE metadata->>'thread_type' = 'heartbeat';
|
||||
+141
-11
@@ -1,7 +1,9 @@
|
||||
[
|
||||
{
|
||||
"id": "openai",
|
||||
"aliases": ["open_ai"],
|
||||
"aliases": [
|
||||
"open_ai"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"api_key_env": "OPENAI_API_KEY",
|
||||
"api_key_required": true,
|
||||
@@ -19,7 +21,9 @@
|
||||
},
|
||||
{
|
||||
"id": "anthropic",
|
||||
"aliases": ["claude"],
|
||||
"aliases": [
|
||||
"claude"
|
||||
],
|
||||
"protocol": "anthropic",
|
||||
"api_key_env": "ANTHROPIC_API_KEY",
|
||||
"api_key_required": true,
|
||||
@@ -52,7 +56,10 @@
|
||||
},
|
||||
{
|
||||
"id": "openai_compatible",
|
||||
"aliases": ["openai-compatible", "compatible"],
|
||||
"aliases": [
|
||||
"openai-compatible",
|
||||
"compatible"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"base_url_env": "LLM_BASE_URL",
|
||||
"base_url_required": true,
|
||||
@@ -89,7 +96,9 @@
|
||||
},
|
||||
{
|
||||
"id": "openrouter",
|
||||
"aliases": ["open_router"],
|
||||
"aliases": [
|
||||
"open_router"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://openrouter.ai/api/v1",
|
||||
"api_key_env": "OPENROUTER_API_KEY",
|
||||
@@ -126,7 +135,10 @@
|
||||
},
|
||||
{
|
||||
"id": "nvidia",
|
||||
"aliases": ["nvidia_nim", "nim"],
|
||||
"aliases": [
|
||||
"nvidia_nim",
|
||||
"nim"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://integrate.api.nvidia.com/v1",
|
||||
"api_key_env": "NVIDIA_API_KEY",
|
||||
@@ -144,7 +156,10 @@
|
||||
},
|
||||
{
|
||||
"id": "venice",
|
||||
"aliases": ["venice_ai", "veniceai"],
|
||||
"aliases": [
|
||||
"venice_ai",
|
||||
"veniceai"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://api.venice.ai/api/v1",
|
||||
"api_key_env": "VENICE_API_KEY",
|
||||
@@ -162,7 +177,10 @@
|
||||
},
|
||||
{
|
||||
"id": "together",
|
||||
"aliases": ["together_ai", "togetherai"],
|
||||
"aliases": [
|
||||
"together_ai",
|
||||
"togetherai"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://api.together.xyz/v1",
|
||||
"api_key_env": "TOGETHER_API_KEY",
|
||||
@@ -180,7 +198,9 @@
|
||||
},
|
||||
{
|
||||
"id": "fireworks",
|
||||
"aliases": ["fireworks_ai"],
|
||||
"aliases": [
|
||||
"fireworks_ai"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://api.fireworks.ai/inference/v1",
|
||||
"api_key_env": "FIREWORKS_API_KEY",
|
||||
@@ -198,7 +218,9 @@
|
||||
},
|
||||
{
|
||||
"id": "deepseek",
|
||||
"aliases": ["deep_seek"],
|
||||
"aliases": [
|
||||
"deep_seek"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://api.deepseek.com/v1",
|
||||
"api_key_env": "DEEPSEEK_API_KEY",
|
||||
@@ -234,7 +256,9 @@
|
||||
},
|
||||
{
|
||||
"id": "sambanova",
|
||||
"aliases": ["samba_nova"],
|
||||
"aliases": [
|
||||
"samba_nova"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://api.sambanova.ai/v1",
|
||||
"api_key_env": "SAMBANOVA_API_KEY",
|
||||
@@ -249,5 +273,111 @@
|
||||
"display_name": "SambaNova",
|
||||
"can_list_models": false
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": "gemini",
|
||||
"aliases": [
|
||||
"google_gemini",
|
||||
"google"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://generativelanguage.googleapis.com/v1beta/openai",
|
||||
"api_key_env": "GEMINI_API_KEY",
|
||||
"api_key_required": true,
|
||||
"model_env": "GEMINI_MODEL",
|
||||
"default_model": "gemini-2.5-flash",
|
||||
"description": "Google Gemini (via OpenAI-compatible endpoint)",
|
||||
"setup": {
|
||||
"kind": "api_key",
|
||||
"secret_name": "llm_gemini_api_key",
|
||||
"key_url": "https://aistudio.google.com/app/apikey",
|
||||
"display_name": "Google Gemini",
|
||||
"can_list_models": true
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": "ionet",
|
||||
"aliases": [
|
||||
"io_net",
|
||||
"io.net"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://api.intelligence.io.solutions/api/v1",
|
||||
"api_key_env": "IONET_API_KEY",
|
||||
"api_key_required": true,
|
||||
"model_env": "IONET_MODEL",
|
||||
"default_model": "deepseek-coder-v2-instruct",
|
||||
"description": "io.net Intelligence API",
|
||||
"setup": {
|
||||
"kind": "api_key",
|
||||
"secret_name": "llm_ionet_api_key",
|
||||
"key_url": "https://cloud.io.net/intelligence",
|
||||
"display_name": "io.net",
|
||||
"can_list_models": true
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": "mistral",
|
||||
"aliases": [
|
||||
"mistral_ai",
|
||||
"mistralai"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://api.mistral.ai/v1",
|
||||
"api_key_env": "MISTRAL_API_KEY",
|
||||
"api_key_required": true,
|
||||
"model_env": "MISTRAL_MODEL",
|
||||
"default_model": "mistral-large-latest",
|
||||
"description": "Mistral AI API",
|
||||
"setup": {
|
||||
"kind": "api_key",
|
||||
"secret_name": "llm_mistral_api_key",
|
||||
"key_url": "https://console.mistral.ai/api-keys",
|
||||
"display_name": "Mistral",
|
||||
"can_list_models": true
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": "yandex",
|
||||
"aliases": [
|
||||
"yandex_ai_studio",
|
||||
"yandexgpt",
|
||||
"yandex_gpt"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://ai.api.cloud.yandex.net/v1",
|
||||
"api_key_env": "YANDEX_API_KEY",
|
||||
"api_key_required": true,
|
||||
"model_env": "YANDEX_MODEL",
|
||||
"extra_headers_env": "YANDEX_EXTRA_HEADERS",
|
||||
"default_model": "yandexgpt-lite",
|
||||
"description": "Yandex AI Studio (YandexGPT)",
|
||||
"setup": {
|
||||
"kind": "api_key",
|
||||
"secret_name": "llm_yandex_api_key",
|
||||
"key_url": "https://aistudio.yandex.ru/platform/folders/",
|
||||
"display_name": "Yandex AI Studio",
|
||||
"can_list_models": true
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": "cloudflare",
|
||||
"aliases": [
|
||||
"cloudflare_ai",
|
||||
"cf_ai"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"api_key_env": "CLOUDFLARE_API_KEY",
|
||||
"api_key_required": true,
|
||||
"base_url_env": "CLOUDFLARE_BASE_URL",
|
||||
"model_env": "CLOUDFLARE_MODEL",
|
||||
"default_model": "@cf/meta/llama-3.3-70b-instruct-fp8-fast",
|
||||
"description": "Cloudflare Workers AI",
|
||||
"setup": {
|
||||
"kind": "open_ai_compatible",
|
||||
"secret_name": "llm_cloudflare_api_key",
|
||||
"display_name": "Cloudflare Workers AI",
|
||||
"can_list_models": false
|
||||
}
|
||||
}
|
||||
]
|
||||
]
|
||||
@@ -2,8 +2,8 @@
|
||||
"name": "discord",
|
||||
"display_name": "Discord Channel",
|
||||
"kind": "channel",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"version": "0.2.0",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Talk to your agent in Discord",
|
||||
"keywords": [
|
||||
"messaging",
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/discord-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "030707431717bca3411a48f311c6ab5f92a45c747de26cafe4f6e3e23a8b3b2d"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,8 +2,8 @@
|
||||
"name": "slack",
|
||||
"display_name": "Slack Channel",
|
||||
"kind": "channel",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"version": "0.2.1",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Talk to your agent in Slack",
|
||||
"keywords": [
|
||||
"messaging",
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/slack-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "6ed36077b67ac70a041f06f760f93ba79b33269885413c3c3f2c8c87ee60807e"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,8 +2,8 @@
|
||||
"name": "telegram",
|
||||
"display_name": "Telegram Channel",
|
||||
"kind": "channel",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"version": "0.2.2",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Talk to your agent through a Telegram bot",
|
||||
"keywords": [
|
||||
"messaging",
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/telegram-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "98c86895a9c4b0a1e19fe8a47f1ccbfe7e972e112b05e584bc897130dc32283a"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,8 +2,8 @@
|
||||
"name": "whatsapp",
|
||||
"display_name": "WhatsApp Channel",
|
||||
"kind": "channel",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"version": "0.2.0",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Talk to your agent through WhatsApp",
|
||||
"keywords": [
|
||||
"messaging",
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/whatsapp-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "bd35cad18d87292ea8d2f52db9b514ed9f814a414de910f59073d475c26c4c14"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,8 +2,8 @@
|
||||
"name": "github",
|
||||
"display_name": "GitHub",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"version": "0.2.0",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "GitHub integration for issues, PRs, repos, and code search",
|
||||
"keywords": [
|
||||
"git",
|
||||
@@ -20,7 +20,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/github-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "6fcd32719a4ff15641a4b50fff8984686550f0c491dce60518f4126857d0c544"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,8 +2,8 @@
|
||||
"name": "gmail",
|
||||
"display_name": "Gmail",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"version": "0.2.0",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Read, send, and manage Gmail messages and threads",
|
||||
"keywords": [
|
||||
"email",
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/gmail-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "023da7000b17568bf0e64b2e5013c8a042b2f323c85f1632339231c73d500e39"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,8 +2,8 @@
|
||||
"name": "google-calendar",
|
||||
"display_name": "Google Calendar",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"version": "0.2.0",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Create, read, update, and delete Google Calendar events",
|
||||
"keywords": [
|
||||
"calendar",
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-calendar-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "fc42277b65881d6e9bcc5403dc54c7f5b3ddeaaaf04617fce2c5da05d76325f0"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,8 +2,8 @@
|
||||
"name": "google-docs",
|
||||
"display_name": "Google Docs",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"version": "0.2.0",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Create and edit Google Docs documents",
|
||||
"keywords": [
|
||||
"documents",
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-docs-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "385c04abd1e6b8011ccc330e1f4bd7ce58577e488959b51594aa04eb26cbe7cc"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,8 +2,8 @@
|
||||
"name": "google-drive",
|
||||
"display_name": "Google Drive",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"version": "0.2.0",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Upload, download, search, and manage Google Drive files and folders",
|
||||
"keywords": [
|
||||
"storage",
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-drive-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "1b107d575a5d52cc8c76d9a681802190f4373fb485f7f54f445533f097fa37c0"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,8 +2,8 @@
|
||||
"name": "google-sheets",
|
||||
"display_name": "Google Sheets",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"version": "0.2.0",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Read and write Google Sheets spreadsheet data",
|
||||
"keywords": [
|
||||
"spreadsheets",
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-sheets-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "c4f6b1e8c5126ac2c8a4b98e4283a3afa32223d2488fc3c3a609758c0c9beb90"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,8 +2,8 @@
|
||||
"name": "google-slides",
|
||||
"display_name": "Google Slides",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"version": "0.2.0",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Create and edit Google Slides presentations",
|
||||
"keywords": [
|
||||
"presentations",
|
||||
@@ -18,7 +18,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-slides-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "7110b8565340c888e51f99e9c013bf4de8f8a7f7b33bace00eb8fc47831ff20b"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,8 +2,8 @@
|
||||
"name": "slack-tool",
|
||||
"display_name": "Slack Tool",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"version": "0.2.0",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Your agent uses Slack to post and read messages in your workspace",
|
||||
"keywords": [
|
||||
"messaging",
|
||||
@@ -18,7 +18,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/slack-tool-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "6ed36077b67ac70a041f06f760f93ba79b33269885413c3c3f2c8c87ee60807e"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,8 +2,8 @@
|
||||
"name": "telegram-mtproto",
|
||||
"display_name": "Telegram Tool",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"version": "0.2.0",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Your agent uses your Telegram account to read and send messages",
|
||||
"keywords": [
|
||||
"messaging",
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/telegram-mtproto-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "98c86895a9c4b0a1e19fe8a47f1ccbfe7e972e112b05e584bc897130dc32283a"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -2,8 +2,8 @@
|
||||
"name": "web-search",
|
||||
"display_name": "Web Search",
|
||||
"kind": "tool",
|
||||
"version": "0.1.0",
|
||||
"wit_version": "0.2.0",
|
||||
"version": "0.2.0",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Search the web using Brave Search API",
|
||||
"keywords": [
|
||||
"search",
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/web-search-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "66cb2b9b00652385e9f30f17c74902b9222c17c53e9d3bd1ef42f5cab705bcf6"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
Executable
+273
@@ -0,0 +1,273 @@
|
||||
#!/usr/bin/env bash
|
||||
# Architecture boundary checks for IronClaw.
|
||||
# Run as: bash scripts/check-boundaries.sh
|
||||
# Returns non-zero if hard violations are found.
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)"
|
||||
cd "$REPO_ROOT"
|
||||
|
||||
violations=0
|
||||
|
||||
echo "=== Architecture Boundary Checks ==="
|
||||
echo
|
||||
|
||||
# --------------------------------------------------------------------------
|
||||
# Check 1: Direct database driver usage outside the db layer
|
||||
# --------------------------------------------------------------------------
|
||||
# tokio_postgres:: and libsql:: types should only appear in:
|
||||
# - src/db/ (the database abstraction layer)
|
||||
# - src/workspace/repository.rs (workspace's own DB layer)
|
||||
# - src/error.rs (needs From impls for driver error types)
|
||||
# - src/app.rs (bootstraps/initialises the database)
|
||||
# - src/testing.rs (test infrastructure)
|
||||
# - src/cli/ (CLI commands that bootstrap DB connections)
|
||||
# - src/setup/ (onboarding wizard bootstraps DB)
|
||||
# - src/main.rs (entry point)
|
||||
#
|
||||
# Everything else is a boundary violation -- those modules should go through
|
||||
# the Database trait, not touch driver types directly.
|
||||
# --------------------------------------------------------------------------
|
||||
|
||||
echo "--- Check 1: Direct database driver usage outside db layer ---"
|
||||
|
||||
results=$(grep -rn 'tokio_postgres::\|libsql::' src/ \
|
||||
--include='*.rs' \
|
||||
| grep -v 'src/db/' \
|
||||
| grep -v 'src/workspace/repository.rs' \
|
||||
| grep -v 'src/error.rs' \
|
||||
| grep -v 'src/app.rs' \
|
||||
| grep -v 'src/testing.rs' \
|
||||
| grep -v 'src/cli/' \
|
||||
| grep -v 'src/setup/' \
|
||||
| grep -v 'src/main.rs' \
|
||||
| grep -v '^\s*//' \
|
||||
| grep -v '//.*tokio_postgres\|//.*libsql' \
|
||||
|| true)
|
||||
|
||||
if [ -n "$results" ]; then
|
||||
echo "VIOLATION: Direct database driver usage found outside db layer:"
|
||||
echo "$results"
|
||||
echo
|
||||
count=$(echo "$results" | wc -l | tr -d ' ')
|
||||
echo "($count occurrence(s) -- these modules should use the Database trait)"
|
||||
violations=$((violations + 1))
|
||||
else
|
||||
echo "OK"
|
||||
fi
|
||||
echo
|
||||
|
||||
# --------------------------------------------------------------------------
|
||||
# Check 2: .unwrap() / .expect() in production code (heuristic)
|
||||
# --------------------------------------------------------------------------
|
||||
# We cannot perfectly distinguish test vs production code with grep alone
|
||||
# (test modules span many lines). Instead we:
|
||||
# 1. Exclude files that are entirely test infrastructure
|
||||
# 2. Exclude lines that are clearly in test code (assert, #[test], etc.)
|
||||
# 3. Report a per-file summary so reviewers can focus on the worst files
|
||||
#
|
||||
# This is a WARNING, not a hard violation.
|
||||
# --------------------------------------------------------------------------
|
||||
|
||||
echo "--- Check 2: .unwrap() / .expect() in production code ---"
|
||||
|
||||
# Collect raw matches excluding obvious test-only files and lines
|
||||
raw_results=$(grep -rn '\.unwrap()\|\.expect(' src/ \
|
||||
--include='*.rs' \
|
||||
| grep -v 'src/main.rs' \
|
||||
| grep -v 'src/testing.rs' \
|
||||
| grep -v 'src/setup/' \
|
||||
|| true)
|
||||
|
||||
if [ -n "$raw_results" ]; then
|
||||
total=$(echo "$raw_results" | wc -l | tr -d ' ')
|
||||
echo "WARNING: ~$total .unwrap()/.expect() calls found in src/ (excluding main/testing/setup)."
|
||||
echo "Many are in test modules; a per-file breakdown helps triage:"
|
||||
echo
|
||||
# Show per-file counts, sorted by count descending, top 15
|
||||
file_counts=$(echo "$raw_results" | cut -d: -f1 | sort | uniq -c | sort -rn)
|
||||
echo "$file_counts" | head -15
|
||||
fc_total=$(echo "$file_counts" | wc -l | tr -d ' ')
|
||||
if [ "$fc_total" -gt 15 ]; then
|
||||
echo " ... and $((fc_total - 15)) more files"
|
||||
fi
|
||||
echo
|
||||
echo "(This is a warning for gradual cleanup, not a blocking violation.)"
|
||||
echo "(Many of these are inside #[cfg(test)] modules which is acceptable.)"
|
||||
else
|
||||
echo "OK"
|
||||
fi
|
||||
echo
|
||||
|
||||
# --------------------------------------------------------------------------
|
||||
# Check 3: std::env::var reads outside config/bootstrap layers
|
||||
# --------------------------------------------------------------------------
|
||||
# Sensitive values should come through Config or the secrets module.
|
||||
# Direct std::env::var / env::var() reads are allowed in:
|
||||
# - src/config/ (the config layer itself)
|
||||
# - src/main.rs (entry point)
|
||||
# - src/setup/ (onboarding wizard)
|
||||
# - src/testing.rs (test infrastructure)
|
||||
# - src/cli/ (CLI commands that read env for bootstrap)
|
||||
# - src/bootstrap.rs (bootstrap logic)
|
||||
# --------------------------------------------------------------------------
|
||||
|
||||
echo "--- Check 3: Direct env var reads outside config layer ---"
|
||||
|
||||
results=$(grep -rn 'std::env::var\|env::var(' src/ \
|
||||
--include='*.rs' \
|
||||
| grep -v 'src/config/' \
|
||||
| grep -v 'src/main.rs' \
|
||||
| grep -v 'src/setup/' \
|
||||
| grep -v 'src/testing.rs' \
|
||||
| grep -v 'src/cli/' \
|
||||
| grep -v 'src/bootstrap.rs' \
|
||||
| grep -v '#\[cfg(test)\]' \
|
||||
| grep -v '#\[test\]' \
|
||||
| grep -v 'mod tests' \
|
||||
| grep -v 'fn test_' \
|
||||
| grep -v '//.*env::var' \
|
||||
|| true)
|
||||
|
||||
if [ -n "$results" ]; then
|
||||
count=$(echo "$results" | wc -l | tr -d ' ')
|
||||
echo "WARNING: Direct env var reads found outside config layer ($count occurrences):"
|
||||
echo "$results"
|
||||
echo
|
||||
echo "(Review these -- secrets/config should come through Config or the secrets module)"
|
||||
else
|
||||
echo "OK"
|
||||
fi
|
||||
echo
|
||||
|
||||
# --------------------------------------------------------------------------
|
||||
# Check 4: Test tier gating — integration tests must use feature flags
|
||||
# --------------------------------------------------------------------------
|
||||
# Files in tests/ that connect to PostgreSQL or use DATABASE_URL must be
|
||||
# gated behind #![cfg(all(feature = "postgres", feature = "integration"))].
|
||||
# This ensures `cargo test` (no flags) never requires external services.
|
||||
#
|
||||
# Heuristic: any test file referencing DATABASE_URL, connect(), PgPool,
|
||||
# or tokio_postgres should have the cfg gate on the first few lines.
|
||||
# --------------------------------------------------------------------------
|
||||
|
||||
echo "--- Check 4: Test tier gating for integration tests ---"
|
||||
|
||||
tier_violations=()
|
||||
for test_file in tests/*.rs; do
|
||||
[ -f "$test_file" ] || continue
|
||||
|
||||
# Check if the file actually connects to a database (imports DB types
|
||||
# or calls pool/connect). Mere string references like "DATABASE_URL"
|
||||
# in config tests don't count.
|
||||
needs_gate=false
|
||||
if grep -q 'PgPool\|tokio_postgres::\|create_pool\|\.connect(' "$test_file" 2>/dev/null; then
|
||||
needs_gate=true
|
||||
fi
|
||||
|
||||
if [ "$needs_gate" = true ]; then
|
||||
# Check first 5 lines for the cfg gate
|
||||
if ! head -5 "$test_file" | grep -q 'cfg.*feature.*integration' 2>/dev/null; then
|
||||
tier_violations+=(" $test_file: needs '#![cfg(all(feature = \"postgres\", feature = \"integration\"))]'")
|
||||
fi
|
||||
fi
|
||||
done
|
||||
|
||||
if [ ${#tier_violations[@]} -gt 0 ]; then
|
||||
echo "VIOLATION: Integration tests missing feature gate:"
|
||||
printf '%s\n' "${tier_violations[@]}"
|
||||
echo
|
||||
echo "(Tests requiring external services must be gated behind the 'integration' feature)"
|
||||
violations=$((violations + 1))
|
||||
else
|
||||
echo "OK"
|
||||
fi
|
||||
echo
|
||||
|
||||
# --------------------------------------------------------------------------
|
||||
# Check 5: No silent test-skip patterns (try_connect, is_available, etc.)
|
||||
# --------------------------------------------------------------------------
|
||||
# Tests must fail loudly when prerequisites are missing, not silently skip.
|
||||
# The correct approach is feature-flag gating (#![cfg(feature = "integration")]).
|
||||
# Patterns like try_connect().is_none() { return; } hide broken tests.
|
||||
# --------------------------------------------------------------------------
|
||||
|
||||
echo "--- Check 5: No silent test-skip patterns ---"
|
||||
|
||||
skip_results=$(grep -rn 'try_connect\|is_available.*return\|is_none.*return\|is_err.*return.*//.*skip' tests/ \
|
||||
--include='*.rs' \
|
||||
|| true)
|
||||
|
||||
if [ -n "$skip_results" ]; then
|
||||
echo "VIOLATION: Silent test-skip patterns found (use feature gates instead):"
|
||||
echo "$skip_results"
|
||||
echo
|
||||
violations=$((violations + 1))
|
||||
else
|
||||
echo "OK"
|
||||
fi
|
||||
echo
|
||||
|
||||
# --------------------------------------------------------------------------
|
||||
# Check 6: LLM module isolation — no imports from other crate modules
|
||||
# --------------------------------------------------------------------------
|
||||
# src/llm/ should only import from:
|
||||
# - crate::llm (self-references)
|
||||
# - external crates (no crate:: prefix)
|
||||
# It must NOT import from crate::agent, crate::tools, crate::channels,
|
||||
# crate::safety, crate::config, crate::bootstrap, crate::cli, crate::db,
|
||||
# crate::workspace, crate::worker, crate::orchestrator, crate::skills,
|
||||
# crate::hooks, crate::setup, crate::context, etc.
|
||||
#
|
||||
# Test-only imports (crate::testing) are excluded since they don't affect
|
||||
# the runtime dependency graph and won't exist in the extracted crate.
|
||||
# --------------------------------------------------------------------------
|
||||
|
||||
echo "--- Check 6: LLM module isolation ---"
|
||||
|
||||
# Match any `crate::` reference (use-imports AND inline paths) that isn't
|
||||
# crate::llm or crate::testing. Filter out comments.
|
||||
# We strip inline comments (everything after //) with sed before checking,
|
||||
# so a line like `real_code(crate::foo); // crate::llm` is still caught.
|
||||
results=$(grep -rn 'crate::' src/llm/ \
|
||||
--include='*.rs' \
|
||||
| grep -v '^\s*//' \
|
||||
| sed 's|//.*||' \
|
||||
| grep 'crate::' \
|
||||
| grep -v 'crate::llm' \
|
||||
| grep -v 'crate::testing' \
|
||||
|| true)
|
||||
|
||||
if [ -n "$results" ]; then
|
||||
count=$(echo "$results" | wc -l | tr -d ' ')
|
||||
echo "WARNING: src/llm/ has $count reference(s) to modules outside crate::llm:"
|
||||
echo "$results"
|
||||
echo
|
||||
echo "(These are pre-existing; fix them before extracting the crate.)"
|
||||
echo "(New 'use crate::' imports are hard violations — see below.)"
|
||||
echo
|
||||
# Hard-fail only on new `use crate::` imports (easy to avoid in new code).
|
||||
use_imports=$(echo "$results" | grep '^[^:]*:.*use crate::' || true)
|
||||
if [ -n "$use_imports" ]; then
|
||||
echo "HARD VIOLATION: new 'use crate::' imports in src/llm/:"
|
||||
echo "$use_imports"
|
||||
violations=$((violations + 1))
|
||||
fi
|
||||
else
|
||||
echo "OK"
|
||||
fi
|
||||
echo
|
||||
|
||||
# --------------------------------------------------------------------------
|
||||
# Summary
|
||||
# --------------------------------------------------------------------------
|
||||
|
||||
echo "=== Summary ==="
|
||||
if [ "$violations" -gt 0 ]; then
|
||||
echo "FAILED: $violations hard violation(s) found"
|
||||
exit 1
|
||||
else
|
||||
echo "PASSED: No hard violations found (review warnings above)"
|
||||
exit 0
|
||||
fi
|
||||
@@ -51,9 +51,11 @@ echo "[6/6] Installing git hooks..."
|
||||
HOOKS_DIR=$(git rev-parse --git-path hooks 2>/dev/null) || true
|
||||
if [ -n "$HOOKS_DIR" ]; then
|
||||
mkdir -p "$HOOKS_DIR"
|
||||
SCRIPT_ABS="$(cd "$(dirname "$0")" && pwd)/commit-msg-regression.sh"
|
||||
ln -sf "$SCRIPT_ABS" "$HOOKS_DIR/commit-msg"
|
||||
SCRIPTS_ABS="$(cd "$(dirname "$0")" && pwd)"
|
||||
ln -sf "$SCRIPTS_ABS/commit-msg-regression.sh" "$HOOKS_DIR/commit-msg"
|
||||
echo " commit-msg hook installed (regression test enforcement)"
|
||||
ln -sf "$SCRIPTS_ABS/pre-commit-safety.sh" "$HOOKS_DIR/pre-commit"
|
||||
echo " pre-commit hook installed (UTF-8, case-sensitivity, /tmp, redaction checks)"
|
||||
else
|
||||
echo " Skipped: not a git repository"
|
||||
fi
|
||||
|
||||
Executable
+136
@@ -0,0 +1,136 @@
|
||||
#!/usr/bin/env bash
|
||||
# Pre-commit safety checks for common issues caught by AI code reviewers.
|
||||
#
|
||||
# Can be run standalone: bash scripts/pre-commit-safety.sh
|
||||
# Or installed as a git pre-commit hook via dev-setup.sh.
|
||||
#
|
||||
# Checks staged .rs files for:
|
||||
# 1. Unsafe UTF-8 byte slicing (panics on multi-byte chars)
|
||||
# 2. Case-sensitive file extension comparisons
|
||||
# 3. Hardcoded /tmp paths in tests (flaky in parallel runs)
|
||||
# 4. Tool parameters logged without redaction (secret leaks)
|
||||
# 5. Multi-step DB operations without transaction wrapping
|
||||
#
|
||||
# Suppress individual lines with an inline "// safety: <reason>" comment.
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
# Determine a suitable base ref for standalone diffs.
|
||||
resolve_base_ref() {
|
||||
local candidates=(
|
||||
"@{upstream}"
|
||||
"origin/HEAD"
|
||||
"origin/main"
|
||||
"origin/master"
|
||||
"main"
|
||||
"master"
|
||||
)
|
||||
|
||||
for ref in "${candidates[@]}"; do
|
||||
if git rev-parse --verify --quiet "$ref" >/dev/null 2>&1; then
|
||||
echo "$ref"
|
||||
return 0
|
||||
fi
|
||||
done
|
||||
|
||||
echo "pre-commit-safety: could not determine a base Git ref for diff (tried: ${candidates[*]})." >&2
|
||||
echo "pre-commit-safety: ensure your repository has an upstream or a local main/master branch." >&2
|
||||
exit 1
|
||||
}
|
||||
|
||||
# Support both pre-commit hook (staged files) and standalone (all changed vs base)
|
||||
if git diff --cached --quiet 2>/dev/null; then
|
||||
# No staged changes -- compare working tree against a resolved base ref
|
||||
BASE_REF="$(resolve_base_ref)"
|
||||
DIFF_OUTPUT=$(git diff "$BASE_REF" -- '*.rs' 2>/dev/null || true)
|
||||
else
|
||||
DIFF_OUTPUT=$(git diff --cached -U0 -- '*.rs' 2>/dev/null || true)
|
||||
fi
|
||||
|
||||
# Early exit if there are no relevant .rs changes
|
||||
if [ -z "$DIFF_OUTPUT" ]; then
|
||||
exit 0
|
||||
fi
|
||||
|
||||
WARNINGS=0
|
||||
|
||||
warn() {
|
||||
if [ "$WARNINGS" -eq 0 ]; then
|
||||
echo ""
|
||||
echo "=== Pre-commit Safety Checks ==="
|
||||
echo ""
|
||||
fi
|
||||
WARNINGS=$((WARNINGS + 1))
|
||||
echo " [$1] $2"
|
||||
}
|
||||
|
||||
# 1. Unsafe UTF-8 byte slicing: &s[..N] or &s[..some_var] on strings
|
||||
# Safe patterns: is_char_boundary, char_indices, // safety:
|
||||
if echo "$DIFF_OUTPUT" | grep -nE '^\+' | grep -E '\[\.\..*\]' | grep -vE 'is_char_boundary|char_indices|// safety:|as_bytes|Vec<|&\[u8\]|\[u8\]|bytes\(\)|&bytes' | head -3 | grep -q .; then
|
||||
warn "UTF8" "Possible unsafe byte-index string slicing. Use is_char_boundary() or char_indices()."
|
||||
echo "$DIFF_OUTPUT" | grep -nE '^\+' | grep -E '\[\.\..*\]' | grep -vE 'is_char_boundary|char_indices|// safety:|as_bytes|Vec<|&\[u8\]|\[u8\]|bytes\(\)|&bytes' | head -3 | sed 's/^/ /'
|
||||
fi
|
||||
|
||||
# 2. Case-sensitive file extension checks
|
||||
# Match: .ends_with(".png") without prior to_lowercase
|
||||
if echo "$DIFF_OUTPUT" | grep -nE '^\+.*ends_with\("\.([pP][nN][gG]|[jJ][pP][eE]?[gG]|[gG][iI][fF]|[wW][eE][bB][pP]|[mM][dD])"\)' | grep -vE 'to_lowercase|to_ascii_lowercase|// safety:' | head -3 | grep -q .; then
|
||||
warn "CASE" "Case-sensitive file extension comparison. Normalize to lowercase first."
|
||||
echo "$DIFF_OUTPUT" | grep -nE '^\+.*ends_with\("\.([pP][nN][gG]|[jJ][pP][eE]?[gG]|[gG][iI][fF]|[wW][eE][bB][pP]|[mM][dD])"\)' | grep -vE 'to_lowercase|to_ascii_lowercase|// safety:' | head -3 | sed 's/^/ /'
|
||||
fi
|
||||
|
||||
# 3. Hardcoded /tmp paths in test files
|
||||
if echo "$DIFF_OUTPUT" | grep -nE '^\+.*"/tmp/' | grep -vE 'tempfile|tempdir|// safety:' | head -3 | grep -q .; then
|
||||
warn "TMPDIR" "Hardcoded /tmp path. Use tempfile::tempdir() for parallel-safe tests."
|
||||
echo "$DIFF_OUTPUT" | grep -nE '^\+.*"/tmp/' | grep -vE 'tempfile|tempdir|// safety:' | head -3 | sed 's/^/ /'
|
||||
fi
|
||||
|
||||
# 4. Logging tool parameters without redaction
|
||||
if echo "$DIFF_OUTPUT" | grep -nE '^\+.*tracing::(info|debug|warn|error).*param' | grep -vE 'redact|// safety:' | head -3 | grep -q .; then
|
||||
warn "REDACT" "Logging tool parameters without redaction. Use redact_params() first."
|
||||
echo "$DIFF_OUTPUT" | grep -nE '^\+.*tracing::(info|debug|warn|error).*param' | grep -vE 'redact|// safety:' | head -3 | sed 's/^/ /'
|
||||
fi
|
||||
|
||||
# 5. Multi-step DB operations without transaction
|
||||
# Uses -W (function context) to reduce false positives from existing transactions.
|
||||
# Suppressible with "// safety:" in the hunk.
|
||||
DIFF_W_OUTPUT=$(git diff --cached -W -- '*.rs' 2>/dev/null || git diff "$(resolve_base_ref)" -W -- '*.rs' 2>/dev/null || true)
|
||||
if [ -n "$DIFF_W_OUTPUT" ]; then
|
||||
HUNK_COUNT=$(echo "$DIFF_W_OUTPUT" | awk '
|
||||
/^@@/ {
|
||||
if (count >= 2 && !has_tx && !has_safety) found++
|
||||
count=0; has_tx=0; has_safety=0
|
||||
}
|
||||
/^\+.*\.(execute|query)\(/ { count++ }
|
||||
/^\+.*(transaction|\.tx\.|\.begin\()/ { has_tx=1 }
|
||||
/ .*(transaction|\.tx\.|\.begin\()/ { has_tx=1 }
|
||||
/\/\/ safety:/ { has_safety=1 }
|
||||
END {
|
||||
if (count >= 2 && !has_tx && !has_safety) found++
|
||||
print found+0
|
||||
}
|
||||
')
|
||||
if [ "$HUNK_COUNT" -gt 0 ]; then
|
||||
warn "TX" "Multiple DB operations in same function without transaction. Wrap in a transaction for atomicity."
|
||||
echo "$DIFF_W_OUTPUT" | awk '
|
||||
/^@@/ {
|
||||
if (count >= 2 && !has_tx && !has_safety) { print buf }
|
||||
buf=""; count=0; has_tx=0; has_safety=0
|
||||
}
|
||||
/^\+.*\.(execute|query)\(/ { count++ }
|
||||
/^\+.*(transaction|\.tx\.|\.begin\()/ { has_tx=1 }
|
||||
/ .*(transaction|\.tx\.|\.begin\()/ { has_tx=1 }
|
||||
/\/\/ safety:/ { has_safety=1 }
|
||||
{ buf = buf "\n" $0 }
|
||||
END {
|
||||
if (count >= 2 && !has_tx && !has_safety) { print buf }
|
||||
}
|
||||
' | grep -E '^\+.*\.(execute|query)\(' | head -4 | sed 's/^/ /'
|
||||
fi
|
||||
fi
|
||||
|
||||
if [ "$WARNINGS" -gt 0 ]; then
|
||||
echo ""
|
||||
echo "Found $WARNINGS potential issue(s). Fix them or add '// safety: <reason>' to suppress."
|
||||
echo ""
|
||||
exit 1
|
||||
fi
|
||||
@@ -0,0 +1,54 @@
|
||||
---
|
||||
name: review-checklist
|
||||
version: 0.1.0
|
||||
description: Pre-merge review checklist based on recurring AI reviewer feedback patterns
|
||||
activation:
|
||||
patterns:
|
||||
- "review.*checklist"
|
||||
- "ready to merge"
|
||||
- "pre-merge check"
|
||||
- "check.*before.*merge"
|
||||
keywords:
|
||||
- review
|
||||
- checklist
|
||||
- merge
|
||||
- pre-merge
|
||||
max_context_tokens: 1500
|
||||
---
|
||||
|
||||
# Pre-Merge Review Checklist
|
||||
|
||||
Before merging, verify these items. They represent the most common issues caught by automated code reviewers (Copilot, Gemini) on IronClaw PRs.
|
||||
|
||||
## Database Operations
|
||||
- [ ] Multi-step DB operations are wrapped in transactions (INSERT+INSERT, UPDATE+DELETE, read-modify-write)
|
||||
- [ ] Both postgres AND libsql backends updated for any new Database trait methods
|
||||
- [ ] Migrations are atomic (SQL execution + version recording in same transaction)
|
||||
|
||||
## Security & Data Safety
|
||||
- [ ] Tool parameters are redacted via `redact_params()` before logging or SSE/WebSocket broadcast
|
||||
- [ ] URL validation resolves DNS before checking for private/loopback IPs (anti-SSRF via DNS rebinding)
|
||||
- [ ] Destructive tools have `requires_approval()` returning `Always` or `UnlessAutoApproved`
|
||||
- [ ] Data from worker containers is treated as untrusted (tool domain checks, server-side nesting depth)
|
||||
- [ ] No secrets or credentials in error messages, logs, or SSE events
|
||||
|
||||
## String Safety
|
||||
- [ ] No byte-index slicing (`&s[..n]`) on external/user strings -- use `is_char_boundary()` or `char_indices()`
|
||||
- [ ] File extension and media type comparisons are case-insensitive (`.to_ascii_lowercase()` before matching)
|
||||
- [ ] Path comparisons are case-insensitive where needed (macOS/Windows filesystems)
|
||||
|
||||
## Trait Wrappers & Decorator Chain
|
||||
- [ ] New `LlmProvider` trait methods are delegated in ALL wrapper types (grep `impl LlmProvider for`)
|
||||
- [ ] New trait methods are tested through the full decorator/provider chain, not just the base impl
|
||||
- [ ] Default trait method implementations are intentional -- wrappers that silently return defaults are bugs
|
||||
|
||||
## Tests
|
||||
- [ ] Temporary files/dirs use `tempfile` crate, no hardcoded `/tmp/` paths
|
||||
- [ ] Tests don't mutate global statics without synchronization (use per-test state or `serial_test`)
|
||||
- [ ] Tests don't make real network requests (use mocks, stubs, or RFC 5737 TEST-NET IPs like 192.0.2.1)
|
||||
- [ ] Test names and comments match actual test behavior and assertions
|
||||
|
||||
## Comments & Documentation
|
||||
- [ ] Code comments match actual behavior (especially route paths, tool names, function semantics)
|
||||
- [ ] Spec/README files updated if module behavior changed
|
||||
- [ ] Error messages are clear and non-redundant (don't nest tool name inside tool error that already contains it)
|
||||
@@ -0,0 +1,171 @@
|
||||
# Agent Module
|
||||
|
||||
Core agent logic. This is the most complex subsystem — read this before working in `src/agent/`.
|
||||
|
||||
## Module Map
|
||||
|
||||
| File | Role |
|
||||
|------|------|
|
||||
| `agent_loop.rs` | `Agent` struct, `AgentDeps`, main `run()` event loop. Delegates to siblings. |
|
||||
| `dispatcher.rs` | Agentic loop for conversational turns: LLM call → tool execution → repeat. Injects skill context. Returns `Response` or `NeedApproval`. |
|
||||
| `thread_ops.rs` | Thread/session operations: `process_user_input`, undo/redo, approval, auth-mode interception, DB hydration, compaction. |
|
||||
| `commands.rs` | System command handlers (`/help`, `/model`, `/status`, `/skills`, etc.) and job intent handlers. |
|
||||
| `session.rs` | Data model: `Session` → `Thread` → `Turn`. State machines for threads and turns. |
|
||||
| `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. |
|
||||
| `scheduler.rs` | Parallel job scheduling. Maintains `jobs` map (full LLM-driven) and `subtasks` map (tool-exec/background). |
|
||||
| `worker.rs` | Per-job execution for background scheduler jobs: calls LLM, runs tools, handles the reasoning loop. Distinct from `dispatcher.rs`. |
|
||||
| `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. |
|
||||
| `self_repair.rs` | Detects stuck jobs and broken tools, attempts recovery. |
|
||||
| `heartbeat.rs` | Proactive periodic execution. Reads `HEARTBEAT.md`, notifies via channel if findings. |
|
||||
| `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). |
|
||||
| `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`. |
|
||||
| `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`. |
|
||||
| `job_monitor.rs` | Subscribes to SSE broadcast and injects Claude Code (container) output back into the agent loop as `IncomingMessage`. |
|
||||
|
||||
## Session / Thread / Turn Model
|
||||
|
||||
```
|
||||
Session (per user)
|
||||
└── Thread (per conversation — can have many)
|
||||
└── Turn (per request/response pair)
|
||||
├── user_input: String
|
||||
├── response: Option<String>
|
||||
├── tool_calls: Vec<ToolCall>
|
||||
└── state: TurnState (Pending | Running | Complete | Failed)
|
||||
```
|
||||
|
||||
- A session has one **active thread** at a time; threads can be switched.
|
||||
- Turns are append-only. Undo rolls back by restoring a prior checkpoint (message list, not a full thread snapshot).
|
||||
- `UndoManager` is per-thread, stored in `SessionManager`, not on `Session` itself. Max 20 checkpoints (oldest dropped when exceeded).
|
||||
- Group chat detection: if `metadata.chat_type` is `group`/`channel`/`supergroup`, `MEMORY.md` is excluded from the system prompt to prevent leaking personal context.
|
||||
- **Auth mode**: if a thread has `pending_auth` set (e.g. from `tool_auth` returning `awaiting_token`), the next user message is intercepted before any turn creation, logging, or safety validation and sent directly to the credential store. Any control submission (undo, interrupt, etc.) cancels auth mode.
|
||||
- `ThreadState` values: `Idle`, `Processing`, `AwaitingApproval`, `Completed`, `Interrupted`.
|
||||
- `SessionManager` maps `(user_id, channel, external_thread_id)` → internal UUID. Prunes idle sessions every 10 minutes (warns at 1000 sessions).
|
||||
|
||||
## Agentic Loop (dispatcher.rs)
|
||||
|
||||
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.
|
||||
|
||||
```
|
||||
run_agentic_loop() [dispatcher.rs — conversational turns]
|
||||
1. Load workspace system prompt (identity files: AGENTS.md, SOUL.md, etc.)
|
||||
2. Detect group chat from metadata; exclude MEMORY.md if group chat
|
||||
3. Select active skills (keyword/pattern scoring against message content)
|
||||
4. Build skill context block (injected before user message)
|
||||
5. LLM call → text response OR tool calls
|
||||
6. If tool calls:
|
||||
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 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.
|
||||
|
||||
**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).
|
||||
|
||||
## Command Routing (router.rs)
|
||||
|
||||
The `Router` handles explicit `/commands` (prefix `/`). It parses them into `MessageIntent` variants: `CreateJob`, `CheckJobStatus`, `CancelJob`, `ListJobs`, `HelpJob`, `Command`. Natural language messages bypass the router entirely — they go directly to `dispatcher.rs` via `process_user_input`. Note: most user-facing commands (undo, compact, etc.) are handled by `SubmissionParser` before the router runs, so `Router` only sees unrecognized `/xxx` patterns that haven't already been claimed by `submission.rs`.
|
||||
|
||||
## Compaction
|
||||
|
||||
Triggered by `ContextMonitor` when token usage approaches the model's context limit.
|
||||
|
||||
**Token estimation**: Word-count × 1.3 + 4 overhead per message. Default context limit: 100,000 tokens. Compaction threshold: 80% (configurable).
|
||||
|
||||
Three strategies, chosen by `ContextMonitor.suggest_compaction()` based on usage ratio:
|
||||
- **MoveToWorkspace** — Writes full turn transcript to workspace daily log, keeps 10 recent turns. Used when usage is 80–85% (moderate). Falls back to `Truncate(5)` if no workspace.
|
||||
- **Summarize** (`keep_recent: N`) — LLM generates a summary of old turns, writes it to workspace daily log (`daily/YYYY-MM-DD.md`), removes old turns. Used when usage is 85–95%.
|
||||
- **Truncate** (`keep_recent: N`) — Removes oldest turns without summarization (fast path). Used when usage >95% (critical).
|
||||
|
||||
If the LLM call for summarization fails, the error propagates — turns are **not** truncated on failure.
|
||||
|
||||
Manual trigger: user sends `/compact` (parsed by `submission.rs`).
|
||||
|
||||
## Scheduler
|
||||
|
||||
`Scheduler` maintains two maps under `Arc<RwLock<HashMap>>`:
|
||||
- `jobs` — full LLM-driven jobs, each with a `Worker` and an `mpsc` channel for `WorkerMessage` (`Start`, `Stop`, `Ping`, `UserMessage`).
|
||||
- `subtasks` — lightweight `ToolExec` or `Background` tasks spawned via `spawn_subtask()` / `spawn_batch()`.
|
||||
|
||||
**Preferred entry point**: `dispatch_job()` — creates context, optionally sets metadata, persists to DB (so FK references from `job_actions`/`llm_calls` are valid immediately), then calls `schedule()`. Don't call `schedule()` directly unless you've already persisted.
|
||||
|
||||
Check-insert is done under a single write lock to prevent TOCTOU races. A cleanup task polls every second for job completion and removes the entry from the map.
|
||||
|
||||
`spawn_subtask()` returns a `oneshot::Receiver` — callers must await it to get the result. `spawn_batch()` runs all tasks concurrently and returns results in input order.
|
||||
|
||||
## Self-Repair
|
||||
|
||||
`DefaultSelfRepair` runs on `repair_check_interval` (from `AgentConfig`). It:
|
||||
1. Calls `ContextManager::find_stuck_jobs()` to find jobs in `JobState::Stuck`.
|
||||
2. Attempts `ctx.attempt_recovery()` (transitions back to `InProgress`).
|
||||
3. Returns `ManualRequired` if `repair_attempts >= max_repair_attempts`.
|
||||
4. Detects broken tools via `store.get_broken_tools(5)` (threshold: 5 failures). Requires `with_store()` to be called; returns empty without a store.
|
||||
5. Attempts to rebuild broken tools via `SoftwareBuilder`. Requires `with_builder()` to be called; returns `ManualRequired` without a builder.
|
||||
|
||||
Note: the `stuck_threshold` duration is stored but currently unused (marked `#[allow(dead_code)]`). Stuck detection relies on `JobState::Stuck` being set by the state machine, not wall-clock time comparison.
|
||||
|
||||
Repair results: `Success`, `Retry`, `Failed`, `ManualRequired`. `Retry` does NOT notify the user (to avoid spam).
|
||||
|
||||
## Key Invariants
|
||||
|
||||
- Never call `.unwrap()` or `.expect()` — use `?` with proper error mapping.
|
||||
- All state mutations on `Session`/`Thread` happen under `Arc<Mutex<Session>>` lock.
|
||||
- The agent loop is single-threaded per thread; parallel execution happens at the job/scheduler level.
|
||||
- Skills are selected **deterministically** (no LLM call) — see `skills/selector.rs`.
|
||||
- Tool results pass through `SafetyLayer` before returning to LLM (sanitizer → validator → policy → leak detector).
|
||||
- `SessionManager` uses double-checked locking for session creation. Read lock first (fast path), then write lock with re-check to prevent duplicate sessions.
|
||||
- `Scheduler.schedule()` holds the write lock for the entire check-insert sequence — don't hold any other locks when calling it.
|
||||
- `cheap_llm` in `AgentDeps` is used for heartbeat and other lightweight tasks. Falls back to main `llm` if `None`. Use `agent.cheap_llm()` accessor, not `deps.cheap_llm` directly.
|
||||
- `CostGuard.check_allowed()` must be called **before** LLM calls; `record_llm_call()` must be called **after**. Both calls are separate — the guard does not auto-record.
|
||||
- `BeforeInbound` and `BeforeOutbound` hooks run for every user message and agent response respectively. Hooks can modify content or reject. Hook errors are logged but **fail-open** (processing continues).
|
||||
|
||||
## Complete Submission Command Reference
|
||||
|
||||
All commands parsed by `SubmissionParser::parse()`:
|
||||
|
||||
| Input | Variant | Notes |
|
||||
|-------|---------|-------|
|
||||
| `/undo` | `Undo` | |
|
||||
| `/redo` | `Redo` | |
|
||||
| `/interrupt`, `/stop` | `Interrupt` | |
|
||||
| `/compact` | `Compact` | |
|
||||
| `/clear` | `Clear` | |
|
||||
| `/heartbeat` | `Heartbeat` | |
|
||||
| `/summarize`, `/summary` | `Summarize` | |
|
||||
| `/suggest` | `Suggest` | |
|
||||
| `/new`, `/thread new` | `NewThread` | |
|
||||
| `/thread <uuid>` | `SwitchThread` | Must be valid UUID |
|
||||
| `/resume <uuid>` | `Resume` | Must be valid UUID |
|
||||
| `/status [id]`, `/progress [id]`, `/list` | `JobStatus` | `/list` = all jobs |
|
||||
| `/cancel <id>` | `JobCancel` | |
|
||||
| `/quit`, `/exit`, `/shutdown` | `Quit` | |
|
||||
| `yes/y/approve/ok` and aliases | `ApprovalResponse { approved: true, always: false }` | |
|
||||
| `always/a` and aliases | `ApprovalResponse { approved: true, always: true }` | |
|
||||
| `no/n/deny/reject/cancel` and aliases | `ApprovalResponse { approved: false }` | |
|
||||
| JSON `ExecApproval{...}` | `ExecApproval` | From web gateway approval endpoint |
|
||||
| `/help`, `/?` | `SystemCommand { "help" }` | Bypasses thread-state checks |
|
||||
| `/version` | `SystemCommand { "version" }` | |
|
||||
| `/tools` | `SystemCommand { "tools" }` | |
|
||||
| `/skills [search <q>]` | `SystemCommand { "skills" }` | |
|
||||
| `/ping` | `SystemCommand { "ping" }` | |
|
||||
| `/debug` | `SystemCommand { "debug" }` | |
|
||||
| `/model [name]` | `SystemCommand { "model" }` | |
|
||||
| Everything else | `UserInput` | Starts a new agentic turn |
|
||||
|
||||
**`SystemCommand` vs control**: `SystemCommand` variants bypass thread-state checks entirely (no session lock, no turn creation). `Quit` returns `Ok(None)` from `handle_message` which breaks the main loop.
|
||||
|
||||
## Adding a New Submission Command
|
||||
|
||||
Submissions are special messages parsed in `submission.rs` before the agentic loop runs. To add a new one:
|
||||
1. Add a variant to `Submission` enum in `submission.rs`
|
||||
2. Add parsing in `SubmissionParser::parse()`
|
||||
3. Handle in `agent_loop.rs` where `SubmissionResult` is matched (the `match submission { ... }` block in `handle_message`)
|
||||
4. Implement the handler method (usually in `thread_ops.rs` for session operations, or `commands.rs` for system commands)
|
||||
+126
-8
@@ -77,6 +77,10 @@ pub struct AgentDeps {
|
||||
pub sse_tx: Option<tokio::sync::broadcast::Sender<crate::channels::web::types::SseEvent>>,
|
||||
/// HTTP interceptor for trace recording/replay.
|
||||
pub http_interceptor: Option<Arc<dyn crate::llm::recording::HttpInterceptor>>,
|
||||
/// Audio transcription middleware for voice messages.
|
||||
pub transcription: Option<Arc<crate::transcription::TranscriptionMiddleware>>,
|
||||
/// Document text extraction middleware for PDF, DOCX, PPTX, etc.
|
||||
pub document_extraction: Option<Arc<crate::document_extraction::DocumentExtractionMiddleware>>,
|
||||
}
|
||||
|
||||
/// The main agent that coordinates all components.
|
||||
@@ -92,6 +96,9 @@ pub struct Agent {
|
||||
pub(super) heartbeat_config: Option<HeartbeatConfig>,
|
||||
pub(super) hygiene_config: Option<crate::config::HygieneConfig>,
|
||||
pub(super) routine_config: Option<RoutineConfig>,
|
||||
/// Optional slot to expose the routine engine to the gateway for manual triggering.
|
||||
pub(super) routine_engine_slot:
|
||||
Option<Arc<tokio::sync::RwLock<Option<Arc<crate::agent::routine_engine::RoutineEngine>>>>>,
|
||||
}
|
||||
|
||||
impl Agent {
|
||||
@@ -127,6 +134,9 @@ impl Agent {
|
||||
if let Some(ref tx) = deps.sse_tx {
|
||||
scheduler.set_sse_sender(tx.clone());
|
||||
}
|
||||
if let Some(ref interceptor) = deps.http_interceptor {
|
||||
scheduler.set_http_interceptor(Arc::clone(interceptor));
|
||||
}
|
||||
let scheduler = Arc::new(scheduler);
|
||||
|
||||
Self {
|
||||
@@ -141,9 +151,18 @@ impl Agent {
|
||||
heartbeat_config,
|
||||
hygiene_config,
|
||||
routine_config,
|
||||
routine_engine_slot: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Set the routine engine slot for exposing the engine to the gateway.
|
||||
pub fn set_routine_engine_slot(
|
||||
&mut self,
|
||||
slot: Arc<tokio::sync::RwLock<Option<Arc<crate::agent::routine_engine::RoutineEngine>>>>,
|
||||
) {
|
||||
self.routine_engine_slot = Some(slot);
|
||||
}
|
||||
|
||||
// Convenience accessors
|
||||
|
||||
/// Get the scheduler (for external wiring, e.g. CreateJobTool).
|
||||
@@ -335,8 +354,19 @@ impl Agent {
|
||||
let heartbeat_handle = if let Some(ref hb_config) = self.heartbeat_config {
|
||||
if hb_config.enabled {
|
||||
if let Some(workspace) = self.workspace() {
|
||||
let config = AgentHeartbeatConfig::default()
|
||||
let mut config = AgentHeartbeatConfig::default()
|
||||
.with_interval(std::time::Duration::from_secs(hb_config.interval_secs));
|
||||
config.quiet_hours_start = hb_config.quiet_hours_start;
|
||||
config.quiet_hours_end = hb_config.quiet_hours_end;
|
||||
config.timezone = hb_config
|
||||
.timezone
|
||||
.clone()
|
||||
.or_else(|| Some(self.config.default_timezone.clone()));
|
||||
if let (Some(user), Some(channel)) =
|
||||
(&hb_config.notify_user, &hb_config.notify_channel)
|
||||
{
|
||||
config = config.with_notify(user, channel);
|
||||
}
|
||||
|
||||
// Set up notification channel
|
||||
let (notify_tx, mut notify_rx) =
|
||||
@@ -387,8 +417,8 @@ impl Agent {
|
||||
hygiene,
|
||||
workspace.clone(),
|
||||
self.cheap_llm().clone(),
|
||||
self.safety().clone(),
|
||||
Some(notify_tx),
|
||||
self.store().map(Arc::clone),
|
||||
))
|
||||
} else {
|
||||
tracing::warn!("Heartbeat enabled but no workspace available");
|
||||
@@ -416,6 +446,8 @@ impl Agent {
|
||||
Arc::clone(workspace),
|
||||
notify_tx,
|
||||
Some(self.scheduler.clone()),
|
||||
self.tools().clone(),
|
||||
self.safety().clone(),
|
||||
));
|
||||
|
||||
// Register routine tools
|
||||
@@ -479,7 +511,12 @@ impl Agent {
|
||||
// SAFETY: self is consumed by run(), we can smuggle the engine in
|
||||
// via a local to use in the message loop below.
|
||||
|
||||
tracing::info!(
|
||||
// Expose engine to gateway for manual triggering
|
||||
if let Some(ref slot) = self.routine_engine_slot {
|
||||
*slot.write().await = Some(Arc::clone(&engine));
|
||||
}
|
||||
|
||||
tracing::debug!(
|
||||
"Routines enabled: cron ticker every {}s, max {} concurrent",
|
||||
rt_config.cron_check_interval_secs,
|
||||
rt_config.max_concurrent_routines
|
||||
@@ -501,26 +538,40 @@ impl Agent {
|
||||
let routine_engine_for_loop = routine_handle.as_ref().map(|(_, e)| Arc::clone(e));
|
||||
|
||||
// Main message loop
|
||||
tracing::info!("Agent {} ready and listening", self.config.name);
|
||||
tracing::debug!("Agent {} ready and listening", self.config.name);
|
||||
|
||||
loop {
|
||||
let message = tokio::select! {
|
||||
biased;
|
||||
_ = tokio::signal::ctrl_c() => {
|
||||
tracing::info!("Ctrl+C received, shutting down...");
|
||||
tracing::debug!("Ctrl+C received, shutting down...");
|
||||
break;
|
||||
}
|
||||
msg = message_stream.next() => {
|
||||
match msg {
|
||||
Some(m) => m,
|
||||
None => {
|
||||
tracing::info!("All channel streams ended, shutting down...");
|
||||
tracing::debug!("All channel streams ended, shutting down...");
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
// Apply transcription middleware to audio attachments
|
||||
let mut message = message;
|
||||
if let Some(ref transcription) = self.deps.transcription {
|
||||
transcription.process(&mut message).await;
|
||||
}
|
||||
|
||||
// Apply document extraction middleware to document attachments
|
||||
if let Some(ref doc_extraction) = self.deps.document_extraction {
|
||||
doc_extraction.process(&mut message).await;
|
||||
}
|
||||
|
||||
// Store successfully extracted document text in workspace for indexing
|
||||
self.store_extracted_documents(&message).await;
|
||||
|
||||
match self.handle_message(&message).await {
|
||||
Ok(Some(response)) if !response.is_empty() => {
|
||||
// Hook: BeforeOutbound — allow hooks to modify or suppress outbound
|
||||
@@ -575,7 +626,7 @@ impl Agent {
|
||||
}
|
||||
Ok(None) => {
|
||||
// Shutdown signal received (/quit, /exit, /shutdown)
|
||||
tracing::info!("Shutdown command received, exiting...");
|
||||
tracing::debug!("Shutdown command received, exiting...");
|
||||
break;
|
||||
}
|
||||
Err(e) => {
|
||||
@@ -604,7 +655,7 @@ impl Agent {
|
||||
}
|
||||
|
||||
// Cleanup
|
||||
tracing::info!("Agent shutting down...");
|
||||
tracing::debug!("Agent shutting down...");
|
||||
repair_handle.abort();
|
||||
pruning_handle.abort();
|
||||
if let Some(handle) = heartbeat_handle {
|
||||
@@ -619,6 +670,73 @@ impl Agent {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Store extracted document text in workspace memory for future search/recall.
|
||||
async fn store_extracted_documents(&self, message: &IncomingMessage) {
|
||||
let workspace = match self.workspace() {
|
||||
Some(ws) => ws,
|
||||
None => return,
|
||||
};
|
||||
|
||||
for attachment in &message.attachments {
|
||||
if attachment.kind != crate::channels::AttachmentKind::Document {
|
||||
continue;
|
||||
}
|
||||
let text = match &attachment.extracted_text {
|
||||
Some(t) if !t.starts_with('[') => t, // skip error messages like "[Failed to..."
|
||||
_ => continue,
|
||||
};
|
||||
|
||||
// Sanitize filename: strip path separators to prevent directory traversal
|
||||
let raw_name = attachment.filename.as_deref().unwrap_or("unnamed_document");
|
||||
let filename: String = raw_name
|
||||
.chars()
|
||||
.map(|c| {
|
||||
if c == '/' || c == '\\' || c == '\0' {
|
||||
'_'
|
||||
} else {
|
||||
c
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
let filename = filename.trim_start_matches('.');
|
||||
let filename = if filename.is_empty() {
|
||||
"unnamed_document"
|
||||
} else {
|
||||
filename
|
||||
};
|
||||
let date = chrono::Utc::now().format("%Y-%m-%d");
|
||||
let path = format!("documents/{date}/{filename}");
|
||||
|
||||
let header = format!(
|
||||
"# {filename}\n\n\
|
||||
> Uploaded by **{}** via **{}** on {date}\n\
|
||||
> MIME: {} | Size: {} bytes\n\n---\n\n",
|
||||
message.user_id,
|
||||
message.channel,
|
||||
attachment.mime_type,
|
||||
attachment.size_bytes.unwrap_or(0),
|
||||
);
|
||||
let content = format!("{header}{text}");
|
||||
|
||||
match workspace.write(&path, &content).await {
|
||||
Ok(_) => {
|
||||
tracing::info!(
|
||||
path = %path,
|
||||
text_len = text.len(),
|
||||
"Stored extracted document in workspace memory"
|
||||
);
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
path = %path,
|
||||
error = %e,
|
||||
"Failed to store extracted document in workspace"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn handle_message(&self, message: &IncomingMessage) -> Result<Option<String>, Error> {
|
||||
// Set message tool context for this turn (current channel and target)
|
||||
// For Signal, use signal_target from metadata (group:ID or phone number),
|
||||
|
||||
@@ -0,0 +1,307 @@
|
||||
//! Augment user message content with structured attachment context.
|
||||
|
||||
use base64::Engine;
|
||||
|
||||
use crate::channels::{AttachmentKind, IncomingAttachment};
|
||||
use crate::llm::{ContentPart, ImageUrl};
|
||||
|
||||
/// Result of processing attachments for the LLM pipeline.
|
||||
pub struct AugmentResult {
|
||||
/// Augmented text content with attachment metadata appended.
|
||||
pub text: String,
|
||||
/// Image content parts to include as multimodal input.
|
||||
pub image_parts: Vec<ContentPart>,
|
||||
}
|
||||
|
||||
/// Process attachments into augmented text and multimodal image parts.
|
||||
///
|
||||
/// Returns `None` if `attachments` is empty (caller should use original content).
|
||||
/// Returns `Some(AugmentResult)` with:
|
||||
/// - `text`: original content + `<attachments>` block (metadata, transcripts, etc.)
|
||||
/// - `image_parts`: `ContentPart::ImageUrl` entries for images with data
|
||||
pub fn augment_with_attachments(
|
||||
content: &str,
|
||||
attachments: &[IncomingAttachment],
|
||||
) -> Option<AugmentResult> {
|
||||
if attachments.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let mut text = content.to_string();
|
||||
text.push_str("\n\n<attachments>");
|
||||
|
||||
let mut image_parts = Vec::new();
|
||||
|
||||
for (i, att) in attachments.iter().enumerate() {
|
||||
text.push('\n');
|
||||
text.push_str(&format_attachment(i + 1, att));
|
||||
|
||||
// Build multimodal image part when image data is available
|
||||
if att.kind == AttachmentKind::Image && !att.data.is_empty() {
|
||||
let b64 = base64::engine::general_purpose::STANDARD.encode(&att.data);
|
||||
let data_url = format!("data:{};base64,{}", att.mime_type, b64);
|
||||
image_parts.push(ContentPart::ImageUrl {
|
||||
image_url: ImageUrl {
|
||||
url: data_url,
|
||||
detail: None,
|
||||
},
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
text.push_str("\n</attachments>");
|
||||
Some(AugmentResult { text, image_parts })
|
||||
}
|
||||
|
||||
/// Escape a string for use as an XML attribute value.
|
||||
fn escape_xml_attr(s: &str) -> String {
|
||||
s.replace('&', "&")
|
||||
.replace('"', """)
|
||||
.replace('<', "<")
|
||||
.replace('>', ">")
|
||||
}
|
||||
|
||||
/// Escape a string for use as XML text content.
|
||||
fn escape_xml_text(s: &str) -> String {
|
||||
s.replace('&', "&")
|
||||
.replace('<', "<")
|
||||
.replace('>', ">")
|
||||
}
|
||||
|
||||
fn format_attachment(index: usize, att: &IncomingAttachment) -> String {
|
||||
let filename = escape_xml_attr(att.filename.as_deref().unwrap_or("unknown"));
|
||||
let mime = escape_xml_attr(&att.mime_type);
|
||||
|
||||
match &att.kind {
|
||||
AttachmentKind::Audio => {
|
||||
let duration_attr = att
|
||||
.duration_secs
|
||||
.map(|d| format!(" duration=\"{d}s\""))
|
||||
.unwrap_or_default();
|
||||
|
||||
let body = match &att.extracted_text {
|
||||
Some(text) => format!("Transcript: {}", escape_xml_text(text)),
|
||||
None => "Audio transcript unavailable.".to_string(),
|
||||
};
|
||||
|
||||
format!(
|
||||
"<attachment index=\"{index}\" type=\"audio\" filename=\"{filename}\"{duration_attr}>\n\
|
||||
{body}\n\
|
||||
</attachment>"
|
||||
)
|
||||
}
|
||||
AttachmentKind::Image => {
|
||||
let size_attr = att
|
||||
.size_bytes
|
||||
.map(|s| format!(" size=\"{}\"", format_size(s)))
|
||||
.unwrap_or_default();
|
||||
|
||||
let body = if att.data.is_empty() {
|
||||
"[Image attached — visual content not available in this conversation]"
|
||||
} else {
|
||||
"[Image attached — sent as visual content]"
|
||||
};
|
||||
|
||||
format!(
|
||||
"<attachment index=\"{index}\" type=\"image\" filename=\"{filename}\" mime=\"{mime}\"{size_attr}>\n\
|
||||
{body}\n\
|
||||
</attachment>"
|
||||
)
|
||||
}
|
||||
AttachmentKind::Document => {
|
||||
let body: String = match &att.extracted_text {
|
||||
Some(text) => escape_xml_text(text),
|
||||
None => {
|
||||
let size_info = att
|
||||
.size_bytes
|
||||
.map(|s| format!(" size=\"{}\"", format_size(s)))
|
||||
.unwrap_or_default();
|
||||
return format!(
|
||||
"<attachment index=\"{index}\" type=\"document\" filename=\"{filename}\" mime=\"{mime}\"{size_info}>\n\
|
||||
[Document attached — text extraction unavailable]\n\
|
||||
</attachment>"
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
let size_attr = att
|
||||
.size_bytes
|
||||
.map(|s| format!(" size=\"{}\"", format_size(s)))
|
||||
.unwrap_or_default();
|
||||
|
||||
format!(
|
||||
"<attachment index=\"{index}\" type=\"document\" filename=\"{filename}\" mime=\"{mime}\"{size_attr}>\n\
|
||||
{body}\n\
|
||||
</attachment>"
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn format_size(bytes: u64) -> String {
|
||||
if bytes < 1024 {
|
||||
format!("{bytes}B")
|
||||
} else if bytes < 1024 * 1024 {
|
||||
format!("{}KB", bytes / 1024)
|
||||
} else {
|
||||
format!("{:.1}MB", bytes as f64 / (1024.0 * 1024.0))
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn make_attachment(kind: AttachmentKind) -> IncomingAttachment {
|
||||
IncomingAttachment {
|
||||
id: "test-id".to_string(),
|
||||
kind,
|
||||
mime_type: "application/octet-stream".to_string(),
|
||||
filename: None,
|
||||
size_bytes: None,
|
||||
source_url: None,
|
||||
storage_key: None,
|
||||
extracted_text: None,
|
||||
data: vec![],
|
||||
duration_secs: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn empty_attachments_returns_none() {
|
||||
assert!(augment_with_attachments("hello", &[]).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn audio_with_transcript() {
|
||||
let mut att = make_attachment(AttachmentKind::Audio);
|
||||
att.filename = Some("voice.ogg".to_string());
|
||||
att.extracted_text = Some("Hello, can you help me?".to_string());
|
||||
att.duration_secs = Some(5);
|
||||
|
||||
let result = augment_with_attachments("hi", &[att]).unwrap();
|
||||
assert!(result.text.starts_with("hi\n\n<attachments>"));
|
||||
assert!(result.text.contains("type=\"audio\""));
|
||||
assert!(result.text.contains("filename=\"voice.ogg\""));
|
||||
assert!(result.text.contains("duration=\"5s\""));
|
||||
assert!(result.text.contains("Transcript: Hello, can you help me?"));
|
||||
assert!(result.text.ends_with("</attachments>"));
|
||||
assert!(result.image_parts.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn audio_without_transcript() {
|
||||
let mut att = make_attachment(AttachmentKind::Audio);
|
||||
att.filename = Some("voice.ogg".to_string());
|
||||
att.duration_secs = Some(10);
|
||||
|
||||
let result = augment_with_attachments("hi", &[att]).unwrap();
|
||||
assert!(result.text.contains("Audio transcript unavailable."));
|
||||
assert!(result.text.contains("duration=\"10s\""));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn image_without_data_no_visual() {
|
||||
let mut att = make_attachment(AttachmentKind::Image);
|
||||
att.filename = Some("screenshot.png".to_string());
|
||||
att.mime_type = "image/png".to_string();
|
||||
att.size_bytes = Some(245_000);
|
||||
|
||||
let result = augment_with_attachments("check this", &[att]).unwrap();
|
||||
assert!(result.text.contains("type=\"image\""));
|
||||
assert!(result.text.contains("filename=\"screenshot.png\""));
|
||||
assert!(result.text.contains("mime=\"image/png\""));
|
||||
assert!(result.text.contains("size=\"239KB\""));
|
||||
assert!(
|
||||
result
|
||||
.text
|
||||
.contains("[Image attached — visual content not available in this conversation]")
|
||||
);
|
||||
assert!(result.image_parts.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn image_with_data_produces_content_part() {
|
||||
let mut att = make_attachment(AttachmentKind::Image);
|
||||
att.filename = Some("photo.jpg".to_string());
|
||||
att.mime_type = "image/jpeg".to_string();
|
||||
att.data = vec![0xFF, 0xD8, 0xFF]; // fake JPEG header
|
||||
|
||||
let result = augment_with_attachments("look", &[att]).unwrap();
|
||||
assert!(
|
||||
result
|
||||
.text
|
||||
.contains("[Image attached — sent as visual content]")
|
||||
);
|
||||
assert_eq!(result.image_parts.len(), 1);
|
||||
match &result.image_parts[0] {
|
||||
ContentPart::ImageUrl { image_url } => {
|
||||
assert!(image_url.url.starts_with("data:image/jpeg;base64,"));
|
||||
}
|
||||
other => panic!("Expected ImageUrl, got: {:?}", other),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn document_with_extracted_text() {
|
||||
let mut att = make_attachment(AttachmentKind::Document);
|
||||
att.filename = Some("report.pdf".to_string());
|
||||
att.extracted_text = Some("Executive summary: Q3 results".to_string());
|
||||
|
||||
let result = augment_with_attachments("review", &[att]).unwrap();
|
||||
assert!(result.text.contains("type=\"document\""));
|
||||
assert!(result.text.contains("filename=\"report.pdf\""));
|
||||
assert!(result.text.contains("Executive summary: Q3 results"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn document_without_extracted_text() {
|
||||
let mut att = make_attachment(AttachmentKind::Document);
|
||||
att.filename = Some("data.csv".to_string());
|
||||
att.mime_type = "text/csv".to_string();
|
||||
att.size_bytes = Some(1024);
|
||||
|
||||
let result = augment_with_attachments("analyze", &[att]).unwrap();
|
||||
assert!(result.text.contains("type=\"document\""));
|
||||
assert!(result.text.contains("mime=\"text/csv\""));
|
||||
assert!(
|
||||
result
|
||||
.text
|
||||
.contains("[Document attached — text extraction unavailable]")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn multiple_attachments_with_mixed_images() {
|
||||
let mut audio = make_attachment(AttachmentKind::Audio);
|
||||
audio.filename = Some("voice.ogg".to_string());
|
||||
audio.extracted_text = Some("Hello".to_string());
|
||||
|
||||
let mut image_with_data = make_attachment(AttachmentKind::Image);
|
||||
image_with_data.filename = Some("photo.jpg".to_string());
|
||||
image_with_data.mime_type = "image/jpeg".to_string();
|
||||
image_with_data.data = vec![0xFF, 0xD8];
|
||||
|
||||
let mut image_no_data = make_attachment(AttachmentKind::Image);
|
||||
image_no_data.filename = Some("remote.png".to_string());
|
||||
image_no_data.mime_type = "image/png".to_string();
|
||||
|
||||
let result =
|
||||
augment_with_attachments("msg", &[audio, image_with_data, image_no_data]).unwrap();
|
||||
assert!(result.text.contains("index=\"1\""));
|
||||
assert!(result.text.contains("index=\"2\""));
|
||||
assert!(result.text.contains("index=\"3\""));
|
||||
// Only the image with data produces a content part
|
||||
assert_eq!(result.image_parts.len(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn original_content_preserved() {
|
||||
let original = "Please help me with this task";
|
||||
let mut att = make_attachment(AttachmentKind::Audio);
|
||||
att.extracted_text = Some("transcript".to_string());
|
||||
|
||||
let result = augment_with_attachments(original, &[att]).unwrap();
|
||||
assert!(result.text.starts_with(original));
|
||||
}
|
||||
}
|
||||
+48
-7
@@ -345,7 +345,6 @@ impl Agent {
|
||||
crate::workspace::hygiene::HygieneConfig::default(),
|
||||
workspace.clone(),
|
||||
self.llm().clone(),
|
||||
self.safety().clone(),
|
||||
);
|
||||
|
||||
match runner.check_heartbeat().await {
|
||||
@@ -406,7 +405,7 @@ impl Agent {
|
||||
.with_max_tokens(512)
|
||||
.with_temperature(0.3);
|
||||
|
||||
let reasoning = Reasoning::new(self.llm().clone(), self.safety().clone());
|
||||
let reasoning = Reasoning::new(self.llm().clone());
|
||||
match reasoning.complete(request).await {
|
||||
Ok((text, _usage)) => Ok(SubmissionResult::response(format!(
|
||||
"Thread Summary:\n\n{}",
|
||||
@@ -454,7 +453,7 @@ impl Agent {
|
||||
.with_max_tokens(512)
|
||||
.with_temperature(0.5);
|
||||
|
||||
let reasoning = Reasoning::new(self.llm().clone(), self.safety().clone());
|
||||
let reasoning = Reasoning::new(self.llm().clone());
|
||||
match reasoning.complete(request).await {
|
||||
Ok((text, _usage)) => Ok(SubmissionResult::response(format!(
|
||||
"Suggested Next Steps:\n\n{}",
|
||||
@@ -663,10 +662,14 @@ impl Agent {
|
||||
}
|
||||
|
||||
match self.llm().set_model(requested) {
|
||||
Ok(()) => Ok(SubmissionResult::response(format!(
|
||||
"Switched model to: {}",
|
||||
requested
|
||||
))),
|
||||
Ok(()) => {
|
||||
// Persist the model choice so it survives restarts.
|
||||
self.persist_selected_model(requested).await;
|
||||
Ok(SubmissionResult::response(format!(
|
||||
"Switched model to: {}",
|
||||
requested
|
||||
)))
|
||||
}
|
||||
Err(e) => Ok(SubmissionResult::error(format!(
|
||||
"Failed to switch model: {}",
|
||||
e
|
||||
@@ -822,4 +825,42 @@ impl Agent {
|
||||
_ => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
/// Persist the selected model to the settings store (DB and/or TOML config).
|
||||
///
|
||||
/// Best-effort: logs warnings on failure but does not propagate errors,
|
||||
/// since the in-memory model switch already succeeded.
|
||||
async fn persist_selected_model(&self, model: &str) {
|
||||
// 1. Persist to DB if available.
|
||||
if let Some(store) = self.store() {
|
||||
let value = serde_json::Value::String(model.to_string());
|
||||
if let Err(e) = store.set_setting("default", "selected_model", &value).await {
|
||||
tracing::warn!("Failed to persist model to DB: {}", e);
|
||||
}
|
||||
}
|
||||
|
||||
// 2. Update TOML config file if it exists (sync I/O in spawn_blocking).
|
||||
let model_owned = model.to_string();
|
||||
if let Err(e) = tokio::task::spawn_blocking(move || {
|
||||
let toml_path = crate::settings::Settings::default_toml_path();
|
||||
match crate::settings::Settings::load_toml(&toml_path) {
|
||||
Ok(Some(mut settings)) => {
|
||||
settings.selected_model = Some(model_owned);
|
||||
if let Err(e) = settings.save_toml(&toml_path) {
|
||||
tracing::warn!("Failed to persist model to config.toml: {}", e);
|
||||
}
|
||||
}
|
||||
Ok(None) => {
|
||||
// No config file on disk; nothing to update.
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::warn!("Failed to load config.toml for model persistence: {}", e);
|
||||
}
|
||||
}
|
||||
})
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Model TOML persistence task failed: {}", e);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+112
-37
@@ -13,7 +13,6 @@ use crate::agent::context_monitor::{CompactionStrategy, ContextBreakdown};
|
||||
use crate::agent::session::Thread;
|
||||
use crate::error::Error;
|
||||
use crate::llm::{ChatMessage, CompletionRequest, LlmProvider, Reasoning};
|
||||
use crate::safety::SafetyLayer;
|
||||
use crate::workspace::Workspace;
|
||||
|
||||
/// Result of a compaction operation.
|
||||
@@ -34,13 +33,12 @@ pub struct CompactionResult {
|
||||
/// Compacts conversation context to stay within limits.
|
||||
pub struct ContextCompactor {
|
||||
llm: Arc<dyn LlmProvider>,
|
||||
safety: Arc<SafetyLayer>,
|
||||
}
|
||||
|
||||
impl ContextCompactor {
|
||||
/// Create a new context compactor.
|
||||
pub fn new(llm: Arc<dyn LlmProvider>, safety: Arc<SafetyLayer>) -> Self {
|
||||
Self { llm, safety }
|
||||
pub fn new(llm: Arc<dyn LlmProvider>) -> Self {
|
||||
Self { llm }
|
||||
}
|
||||
|
||||
/// Compact a thread's context using the given strategy.
|
||||
@@ -105,27 +103,26 @@ impl ContextCompactor {
|
||||
// Generate summary
|
||||
let summary = self.generate_summary(&to_summarize).await?;
|
||||
|
||||
// Write to workspace if available
|
||||
let summary_written = if let Some(ws) = workspace {
|
||||
// Write to workspace if available.
|
||||
// If archival fails, preserve turns to avoid context loss.
|
||||
let (summary_written, turns_removed) = if let Some(ws) = workspace {
|
||||
match self.write_summary_to_workspace(ws, &summary).await {
|
||||
Ok(()) => true,
|
||||
Ok(()) => {
|
||||
thread.truncate_turns(keep_recent);
|
||||
(true, turns_to_remove)
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
"Compaction summary write failed (turns will still be truncated): {}",
|
||||
e
|
||||
);
|
||||
false
|
||||
tracing::warn!("Compaction summary write failed (turns preserved): {}", e);
|
||||
(false, 0)
|
||||
}
|
||||
}
|
||||
} else {
|
||||
false
|
||||
thread.truncate_turns(keep_recent);
|
||||
(false, turns_to_remove)
|
||||
};
|
||||
|
||||
// Truncate thread
|
||||
thread.truncate_turns(keep_recent);
|
||||
|
||||
Ok(CompactionPartial {
|
||||
turns_removed: turns_to_remove,
|
||||
turns_removed,
|
||||
summary_written,
|
||||
summary: Some(summary),
|
||||
})
|
||||
@@ -167,23 +164,20 @@ impl ContextCompactor {
|
||||
// Format turns for storage
|
||||
let content = format_turns_for_storage(old_turns);
|
||||
|
||||
// Write to workspace
|
||||
let written = match self.write_context_to_workspace(ws, &content).await {
|
||||
Ok(()) => true,
|
||||
// Write to workspace. If archival fails, preserve turns.
|
||||
let (written, turns_removed) = match self.write_context_to_workspace(ws, &content).await {
|
||||
Ok(()) => {
|
||||
thread.truncate_turns(keep_recent);
|
||||
(true, turns_to_remove)
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
"Compaction context write failed (turns will still be truncated): {}",
|
||||
e
|
||||
);
|
||||
false
|
||||
tracing::warn!("Compaction context write failed (turns preserved): {}", e);
|
||||
(false, 0)
|
||||
}
|
||||
};
|
||||
|
||||
// Truncate
|
||||
thread.truncate_turns(keep_recent);
|
||||
|
||||
Ok(CompactionPartial {
|
||||
turns_removed: turns_to_remove,
|
||||
turns_removed,
|
||||
summary_written: written,
|
||||
summary: None,
|
||||
})
|
||||
@@ -233,7 +227,7 @@ Be brief but capture all important details. Use bullet points."#,
|
||||
.with_max_tokens(1024)
|
||||
.with_temperature(0.3);
|
||||
|
||||
let reasoning = Reasoning::new(self.llm.clone(), self.safety.clone());
|
||||
let reasoning = Reasoning::new(self.llm.clone());
|
||||
let (text, _) = reasoning.complete(request).await?;
|
||||
Ok(text)
|
||||
}
|
||||
@@ -346,17 +340,11 @@ mod tests {
|
||||
// === QA Plan - Compaction strategy tests ===
|
||||
|
||||
use crate::agent::context_monitor::CompactionStrategy;
|
||||
use crate::config::SafetyConfig;
|
||||
use crate::safety::SafetyLayer;
|
||||
use crate::testing::StubLlm;
|
||||
|
||||
/// Helper: build a `ContextCompactor` with the given `StubLlm`.
|
||||
fn make_compactor(llm: Arc<StubLlm>) -> ContextCompactor {
|
||||
let safety = Arc::new(SafetyLayer::new(&SafetyConfig {
|
||||
max_output_length: 100_000,
|
||||
injection_check_enabled: false,
|
||||
}));
|
||||
ContextCompactor::new(llm, safety)
|
||||
ContextCompactor::new(llm)
|
||||
}
|
||||
|
||||
/// Helper: build a thread with `n` completed turns.
|
||||
@@ -370,6 +358,19 @@ mod tests {
|
||||
thread
|
||||
}
|
||||
|
||||
#[cfg(feature = "libsql")]
|
||||
async fn make_unmigrated_workspace() -> crate::workspace::Workspace {
|
||||
use crate::db::Database;
|
||||
use crate::db::libsql::LibSqlBackend;
|
||||
|
||||
// Intentionally skip migrations so workspace append operations fail.
|
||||
let backend = LibSqlBackend::new_memory()
|
||||
.await
|
||||
.expect("should create in-memory libsql backend");
|
||||
let db: Arc<dyn Database> = Arc::new(backend);
|
||||
crate::workspace::Workspace::new_with_db("compaction-test", db)
|
||||
}
|
||||
|
||||
// ------------------------------------------------------------------
|
||||
// 1. compact_truncate keeps last N turns
|
||||
// ------------------------------------------------------------------
|
||||
@@ -568,6 +569,43 @@ mod tests {
|
||||
assert_eq!(llm.calls(), 0);
|
||||
}
|
||||
|
||||
#[cfg(feature = "libsql")]
|
||||
#[tokio::test]
|
||||
async fn test_compact_with_summary_preserves_turns_when_workspace_write_fails() {
|
||||
let llm = Arc::new(StubLlm::new("summary"));
|
||||
let compactor = make_compactor(llm.clone());
|
||||
let mut thread = make_thread(8);
|
||||
let original_inputs: Vec<String> =
|
||||
thread.turns.iter().map(|t| t.user_input.clone()).collect();
|
||||
let workspace = make_unmigrated_workspace().await;
|
||||
|
||||
let result = compactor
|
||||
.compact(
|
||||
&mut thread,
|
||||
CompactionStrategy::Summarize { keep_recent: 3 },
|
||||
Some(&workspace),
|
||||
)
|
||||
.await
|
||||
.expect("compact should succeed even when workspace write fails");
|
||||
|
||||
// On archival failure, no turns should be removed.
|
||||
assert_eq!(thread.turns.len(), 8);
|
||||
assert_eq!(
|
||||
thread
|
||||
.turns
|
||||
.iter()
|
||||
.map(|t| t.user_input.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
original_inputs
|
||||
.iter()
|
||||
.map(|s| s.as_str())
|
||||
.collect::<Vec<_>>()
|
||||
);
|
||||
assert_eq!(result.turns_removed, 0);
|
||||
assert!(!result.summary_written);
|
||||
assert_eq!(llm.calls(), 1);
|
||||
}
|
||||
|
||||
// ------------------------------------------------------------------
|
||||
// 7. compact_to_workspace without workspace falls back to truncation
|
||||
// ------------------------------------------------------------------
|
||||
@@ -616,6 +654,43 @@ mod tests {
|
||||
assert_eq!(result.turns_removed, 0);
|
||||
}
|
||||
|
||||
#[cfg(feature = "libsql")]
|
||||
#[tokio::test]
|
||||
async fn test_compact_to_workspace_preserves_turns_when_workspace_write_fails() {
|
||||
let llm = Arc::new(StubLlm::new("unused"));
|
||||
let compactor = make_compactor(llm.clone());
|
||||
let mut thread = make_thread(20);
|
||||
let original_inputs: Vec<String> =
|
||||
thread.turns.iter().map(|t| t.user_input.clone()).collect();
|
||||
let workspace = make_unmigrated_workspace().await;
|
||||
|
||||
let result = compactor
|
||||
.compact(
|
||||
&mut thread,
|
||||
CompactionStrategy::MoveToWorkspace,
|
||||
Some(&workspace),
|
||||
)
|
||||
.await
|
||||
.expect("compact should succeed even when workspace write fails");
|
||||
|
||||
// On archival failure, no turns should be removed.
|
||||
assert_eq!(thread.turns.len(), 20);
|
||||
assert_eq!(
|
||||
thread
|
||||
.turns
|
||||
.iter()
|
||||
.map(|t| t.user_input.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
original_inputs
|
||||
.iter()
|
||||
.map(|s| s.as_str())
|
||||
.collect::<Vec<_>>()
|
||||
);
|
||||
assert_eq!(result.turns_removed, 0);
|
||||
assert!(!result.summary_written);
|
||||
assert_eq!(llm.calls(), 0);
|
||||
}
|
||||
|
||||
// ------------------------------------------------------------------
|
||||
// 9. format_turns_for_storage includes tool calls
|
||||
// ------------------------------------------------------------------
|
||||
|
||||
+272
-18
@@ -131,10 +131,12 @@ impl CostGuard {
|
||||
// Check hourly rate
|
||||
if let Some(limit) = self.config.max_actions_per_hour {
|
||||
let mut window = self.action_window.lock().await;
|
||||
let cutoff = Instant::now() - std::time::Duration::from_secs(3600);
|
||||
// Drain expired entries
|
||||
while window.front().is_some_and(|t| *t < cutoff) {
|
||||
window.pop_front();
|
||||
// checked_sub avoids panic when system uptime < 1 hour (Windows)
|
||||
if let Some(cutoff) = Instant::now().checked_sub(std::time::Duration::from_secs(3600)) {
|
||||
// Drain expired entries
|
||||
while window.front().is_some_and(|t| *t < cutoff) {
|
||||
window.pop_front();
|
||||
}
|
||||
}
|
||||
let count = window.len() as u64;
|
||||
if count >= limit {
|
||||
@@ -151,21 +153,46 @@ impl CostGuard {
|
||||
/// Record a completed LLM action: its token costs and the action timestamp.
|
||||
///
|
||||
/// Call this AFTER an LLM call completes so that costs are tracked.
|
||||
/// - `cache_read_input_tokens`: tokens served from cache.
|
||||
/// - `cache_creation_input_tokens`: tokens written to cache.
|
||||
/// - `cache_read_discount`: divisor for cache-read cost (e.g. 10 for Anthropic 90% off, 2 for OpenAI 50% off).
|
||||
/// - `cache_write_multiplier`: cost multiplier for cache writes (1.25 for 5m, 2.0 for 1h).
|
||||
///
|
||||
/// When `cost_per_token` is `Some`, those rates are used directly (provider-
|
||||
/// sourced pricing). When `None`, falls back to the static `costs::model_cost`
|
||||
/// lookup table, then `costs::default_cost`.
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub async fn record_llm_call(
|
||||
&self,
|
||||
model: &str,
|
||||
input_tokens: u32,
|
||||
output_tokens: u32,
|
||||
cache_read_input_tokens: u32,
|
||||
cache_creation_input_tokens: u32,
|
||||
cache_read_discount: Decimal,
|
||||
cache_write_multiplier: Decimal,
|
||||
cost_per_token: Option<(Decimal, Decimal)>,
|
||||
) -> Decimal {
|
||||
let (input_rate, output_rate) = cost_per_token
|
||||
.unwrap_or_else(|| costs::model_cost(model).unwrap_or_else(costs::default_cost));
|
||||
let cost =
|
||||
input_rate * Decimal::from(input_tokens) + output_rate * Decimal::from(output_tokens);
|
||||
// Cached read tokens cost input_rate / cache_read_discount (provider-specific).
|
||||
// Cached write tokens cost write_multiplier × input_rate (e.g. 1.25× for 5m, 2× for 1h).
|
||||
// Uncached tokens = total input - cache reads - cache writes.
|
||||
let cached_total = cache_read_input_tokens.saturating_add(cache_creation_input_tokens);
|
||||
let uncached_input = input_tokens.saturating_sub(cached_total);
|
||||
let effective_discount = if cache_read_discount.is_zero() {
|
||||
Decimal::ONE
|
||||
} else {
|
||||
cache_read_discount
|
||||
};
|
||||
let cache_read_cost =
|
||||
input_rate * Decimal::from(cache_read_input_tokens) / effective_discount;
|
||||
let cache_write_cost =
|
||||
input_rate * Decimal::from(cache_creation_input_tokens) * cache_write_multiplier;
|
||||
let cost = input_rate * Decimal::from(uncached_input)
|
||||
+ cache_read_cost
|
||||
+ cache_write_cost
|
||||
+ output_rate * Decimal::from(output_tokens);
|
||||
|
||||
// Update daily cost (reset if new day)
|
||||
{
|
||||
@@ -235,9 +262,11 @@ impl CostGuard {
|
||||
/// Number of actions in the current hourly window.
|
||||
pub async fn actions_this_hour(&self) -> u64 {
|
||||
let mut window = self.action_window.lock().await;
|
||||
let cutoff = Instant::now() - std::time::Duration::from_secs(3600);
|
||||
while window.front().is_some_and(|t| *t < cutoff) {
|
||||
window.pop_front();
|
||||
// checked_sub avoids panic when system uptime < 1 hour (Windows)
|
||||
if let Some(cutoff) = Instant::now().checked_sub(std::time::Duration::from_secs(3600)) {
|
||||
while window.front().is_some_and(|t| *t < cutoff) {
|
||||
window.pop_front();
|
||||
}
|
||||
}
|
||||
window.len() as u64
|
||||
}
|
||||
@@ -267,7 +296,16 @@ mod tests {
|
||||
|
||||
// Record a big call, still allowed
|
||||
guard
|
||||
.record_llm_call("gpt-4o", 100_000, 100_000, None)
|
||||
.record_llm_call(
|
||||
"gpt-4o",
|
||||
100_000,
|
||||
100_000,
|
||||
0,
|
||||
0,
|
||||
Decimal::ONE,
|
||||
Decimal::ONE,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
assert!(guard.check_allowed().await.is_ok());
|
||||
}
|
||||
@@ -285,7 +323,18 @@ mod tests {
|
||||
// Record a call that costs more than $0.01
|
||||
// gpt-4o: input=$0.0000025/tok, output=$0.00001/tok
|
||||
// 10000 input + 10000 output = $0.025 + $0.10 = $0.125
|
||||
guard.record_llm_call("gpt-4o", 10_000, 10_000, None).await;
|
||||
guard
|
||||
.record_llm_call(
|
||||
"gpt-4o",
|
||||
10_000,
|
||||
10_000,
|
||||
0,
|
||||
0,
|
||||
Decimal::ONE,
|
||||
Decimal::ONE,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
// Now should be blocked
|
||||
let result = guard.check_allowed().await;
|
||||
@@ -308,7 +357,9 @@ mod tests {
|
||||
// First 3 actions allowed
|
||||
for _ in 0..3 {
|
||||
assert!(guard.check_allowed().await.is_ok());
|
||||
guard.record_llm_call("gpt-4o", 10, 10, None).await;
|
||||
guard
|
||||
.record_llm_call("gpt-4o", 10, 10, 0, 0, Decimal::ONE, Decimal::ONE, None)
|
||||
.await;
|
||||
}
|
||||
|
||||
// 4th should be blocked
|
||||
@@ -329,7 +380,9 @@ mod tests {
|
||||
|
||||
assert_eq!(guard.daily_spend().await, Decimal::ZERO);
|
||||
|
||||
let cost = guard.record_llm_call("gpt-4o", 1000, 500, None).await;
|
||||
let cost = guard
|
||||
.record_llm_call("gpt-4o", 1000, 500, 0, 0, Decimal::ONE, Decimal::ONE, None)
|
||||
.await;
|
||||
assert!(cost > Decimal::ZERO);
|
||||
assert_eq!(guard.daily_spend().await, cost);
|
||||
}
|
||||
@@ -340,8 +393,12 @@ mod tests {
|
||||
|
||||
assert_eq!(guard.actions_this_hour().await, 0);
|
||||
|
||||
guard.record_llm_call("gpt-4o", 10, 10, None).await;
|
||||
guard.record_llm_call("gpt-4o", 10, 10, None).await;
|
||||
guard
|
||||
.record_llm_call("gpt-4o", 10, 10, 0, 0, Decimal::ONE, Decimal::ONE, None)
|
||||
.await;
|
||||
guard
|
||||
.record_llm_call("gpt-4o", 10, 10, 0, 0, Decimal::ONE, Decimal::ONE, None)
|
||||
.await;
|
||||
|
||||
assert_eq!(guard.actions_this_hour().await, 2);
|
||||
}
|
||||
@@ -378,10 +435,23 @@ mod tests {
|
||||
assert!(guard.model_usage().await.is_empty());
|
||||
|
||||
// Record calls for two different models
|
||||
guard.record_llm_call("gpt-4o", 1000, 500, None).await;
|
||||
guard.record_llm_call("gpt-4o", 2000, 1000, None).await;
|
||||
guard
|
||||
.record_llm_call("claude-3-5-sonnet-20241022", 500, 200, None)
|
||||
.record_llm_call("gpt-4o", 1000, 500, 0, 0, Decimal::ONE, Decimal::ONE, None)
|
||||
.await;
|
||||
guard
|
||||
.record_llm_call("gpt-4o", 2000, 1000, 0, 0, Decimal::ONE, Decimal::ONE, None)
|
||||
.await;
|
||||
guard
|
||||
.record_llm_call(
|
||||
"claude-3-5-sonnet-20241022",
|
||||
500,
|
||||
200,
|
||||
0,
|
||||
0,
|
||||
Decimal::ONE,
|
||||
Decimal::ONE,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
let usage = guard.model_usage().await;
|
||||
@@ -402,4 +472,188 @@ mod tests {
|
||||
// Costs should differ since models have different pricing
|
||||
assert_ne!(gpt.cost, claude.cost);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_cache_discount_reduces_cost() {
|
||||
let guard = CostGuard::new(CostGuardConfig::default());
|
||||
|
||||
// Full price: 1000 input + 500 output, no cache
|
||||
let full_cost = guard
|
||||
.record_llm_call(
|
||||
"claude-opus-4-6",
|
||||
1000,
|
||||
500,
|
||||
0,
|
||||
0,
|
||||
Decimal::ONE,
|
||||
Decimal::ONE,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
let guard2 = CostGuard::new(CostGuardConfig::default());
|
||||
|
||||
// Same tokens but all input cached (90% discount on input)
|
||||
let cached_cost = guard2
|
||||
.record_llm_call(
|
||||
"claude-opus-4-6",
|
||||
1000,
|
||||
500,
|
||||
1000,
|
||||
0,
|
||||
dec!(10),
|
||||
Decimal::ONE,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
// Cached cost must be strictly less than full cost
|
||||
assert!(
|
||||
cached_cost < full_cost,
|
||||
"cached_cost ({}) should be less than full_cost ({})",
|
||||
cached_cost,
|
||||
full_cost
|
||||
);
|
||||
|
||||
// The difference should be exactly 90% of the input cost
|
||||
let (input_rate, _) = costs::model_cost("claude-opus-4-6").unwrap();
|
||||
let expected_savings = input_rate * Decimal::from(1000u32) * dec!(9) / dec!(10);
|
||||
let actual_savings = full_cost - cached_cost;
|
||||
assert_eq!(
|
||||
actual_savings, expected_savings,
|
||||
"savings should be 90% of input cost for fully-cached request"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_cache_write_surcharge_increases_cost() {
|
||||
let guard = CostGuard::new(CostGuardConfig::default());
|
||||
|
||||
// Full price: 1000 input + 500 output, no cache activity
|
||||
let full_cost = guard
|
||||
.record_llm_call(
|
||||
"claude-opus-4-6",
|
||||
1000,
|
||||
500,
|
||||
0,
|
||||
0,
|
||||
Decimal::ONE,
|
||||
Decimal::ONE,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
let guard2 = CostGuard::new(CostGuardConfig::default());
|
||||
|
||||
// Same tokens, but all input tokens are cache writes (1.25x surcharge for 5m TTL)
|
||||
let short_multiplier = Decimal::new(125, 2); // 1.25
|
||||
let write_cost = guard2
|
||||
.record_llm_call(
|
||||
"claude-opus-4-6",
|
||||
1000,
|
||||
500,
|
||||
0,
|
||||
1000,
|
||||
Decimal::ONE,
|
||||
short_multiplier,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
// Write cost must be strictly greater than full cost
|
||||
assert!(
|
||||
write_cost > full_cost,
|
||||
"write_cost ({}) should be greater than full_cost ({})",
|
||||
write_cost,
|
||||
full_cost
|
||||
);
|
||||
|
||||
// The difference should be exactly 25% of the input cost
|
||||
let (input_rate, _) = costs::model_cost("claude-opus-4-6").unwrap();
|
||||
let expected_surcharge = input_rate * Decimal::from(1000u32) * dec!(0.25);
|
||||
let actual_surcharge = write_cost - full_cost;
|
||||
assert_eq!(
|
||||
actual_surcharge, expected_surcharge,
|
||||
"surcharge should be 25% of input cost for 5m cache writes"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_cache_write_surcharge_long_ttl() {
|
||||
let guard = CostGuard::new(CostGuardConfig::default());
|
||||
|
||||
// Full price: 1000 input + 500 output
|
||||
let full_cost = guard
|
||||
.record_llm_call(
|
||||
"claude-opus-4-6",
|
||||
1000,
|
||||
500,
|
||||
0,
|
||||
0,
|
||||
Decimal::ONE,
|
||||
Decimal::ONE,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
let guard2 = CostGuard::new(CostGuardConfig::default());
|
||||
|
||||
// All input tokens are cache writes with 2.0x multiplier (1h TTL)
|
||||
let long_multiplier = Decimal::TWO;
|
||||
let write_cost = guard2
|
||||
.record_llm_call(
|
||||
"claude-opus-4-6",
|
||||
1000,
|
||||
500,
|
||||
0,
|
||||
1000,
|
||||
Decimal::ONE,
|
||||
long_multiplier,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
// Write cost > full cost
|
||||
assert!(write_cost > full_cost);
|
||||
|
||||
// Surcharge should be 100% of input cost (2.0x - 1.0x = 1.0x)
|
||||
let (input_rate, _) = costs::model_cost("claude-opus-4-6").unwrap();
|
||||
let expected_surcharge = input_rate * Decimal::from(1000u32);
|
||||
let actual_surcharge = write_cost - full_cost;
|
||||
assert_eq!(
|
||||
actual_surcharge, expected_surcharge,
|
||||
"surcharge should be 100% of input cost for 1h cache writes"
|
||||
);
|
||||
}
|
||||
|
||||
/// Regression test for #657: Instant::now() - Duration panics on Windows
|
||||
/// when system uptime is less than the subtracted duration.
|
||||
#[tokio::test]
|
||||
async fn test_checked_sub_no_panic_on_fresh_guard() {
|
||||
// A fresh CostGuard with rate limits should not panic even if
|
||||
// checked_sub returns None (simulating short uptime).
|
||||
let guard = CostGuard::new(CostGuardConfig {
|
||||
max_cost_per_day_cents: None,
|
||||
max_actions_per_hour: Some(100),
|
||||
});
|
||||
|
||||
// These must not panic regardless of system uptime
|
||||
assert!(guard.check_allowed().await.is_ok());
|
||||
assert_eq!(guard.actions_this_hour().await, 0);
|
||||
|
||||
// Record some actions and verify again
|
||||
guard
|
||||
.record_llm_call("gpt-4o", 10, 10, 0, 0, Decimal::ONE, Decimal::ONE, None)
|
||||
.await;
|
||||
assert!(guard.check_allowed().await.is_ok());
|
||||
assert_eq!(guard.actions_this_hour().await, 1);
|
||||
}
|
||||
|
||||
/// Verify that checked_sub itself behaves as expected for the pattern we use.
|
||||
#[test]
|
||||
fn test_instant_checked_sub_returns_none_for_overflow() {
|
||||
// Duration::MAX will always exceed uptime, so checked_sub must return None
|
||||
let result = Instant::now().checked_sub(std::time::Duration::MAX);
|
||||
assert!(result.is_none());
|
||||
}
|
||||
}
|
||||
|
||||
+312
-34
@@ -50,8 +50,18 @@ impl Agent {
|
||||
|
||||
// Load workspace system prompt (identity files: AGENTS.md, SOUL.md, etc.)
|
||||
// In group chats, MEMORY.md is excluded to prevent leaking personal context.
|
||||
// Resolve the user's timezone
|
||||
let user_tz = crate::timezone::resolve_timezone(
|
||||
message.timezone.as_deref(),
|
||||
None, // user setting lookup can be added later
|
||||
&self.config.default_timezone,
|
||||
);
|
||||
|
||||
let system_prompt = if let Some(ws) = self.workspace() {
|
||||
match ws.system_prompt_for_context(is_group_chat).await {
|
||||
match ws
|
||||
.system_prompt_for_context_tz(is_group_chat, user_tz)
|
||||
.await
|
||||
{
|
||||
Ok(prompt) if !prompt.is_empty() => Some(prompt),
|
||||
Ok(_) => None,
|
||||
Err(e) => {
|
||||
@@ -103,7 +113,7 @@ impl Agent {
|
||||
None
|
||||
};
|
||||
|
||||
let mut reasoning = Reasoning::new(self.llm().clone(), self.safety().clone())
|
||||
let mut reasoning = Reasoning::new(self.llm().clone())
|
||||
.with_channel(message.channel.clone())
|
||||
.with_model_name(self.llm().active_model_name())
|
||||
.with_group_chat(is_group_chat);
|
||||
@@ -130,6 +140,18 @@ impl Agent {
|
||||
let mut job_ctx =
|
||||
JobContext::with_user(&message.user_id, "chat", "Interactive chat session");
|
||||
job_ctx.http_interceptor = self.deps.http_interceptor.clone();
|
||||
job_ctx.user_timezone = user_tz.name().to_string();
|
||||
|
||||
// Build system prompts once for this turn. Two variants: with tools
|
||||
// (normal iterations) and without (force_text final iteration).
|
||||
let initial_tool_defs = self.tools().tool_definitions().await;
|
||||
let initial_tool_defs = if !active_skills.is_empty() {
|
||||
crate::skills::attenuate_tools(&initial_tool_defs, &active_skills).tools
|
||||
} else {
|
||||
initial_tool_defs
|
||||
};
|
||||
let cached_prompt = reasoning.build_system_prompt_with_tools(&initial_tool_defs);
|
||||
let cached_prompt_no_tools = reasoning.build_system_prompt_with_tools(&[]);
|
||||
|
||||
let max_tool_iterations = self.config.max_tool_iterations;
|
||||
// Force a text-only response on the last iteration to guarantee termination
|
||||
@@ -138,6 +160,8 @@ impl Agent {
|
||||
let force_text_at = max_tool_iterations;
|
||||
let nudge_at = max_tool_iterations.saturating_sub(1);
|
||||
let mut iteration = 0;
|
||||
const MAX_TOOL_INTENT_NUDGES: u32 = 2;
|
||||
let mut consecutive_tool_intent_nudges: u32 = 0;
|
||||
loop {
|
||||
iteration += 1;
|
||||
// Hard ceiling one past the forced-text iteration (should never be reached
|
||||
@@ -206,10 +230,16 @@ impl Agent {
|
||||
};
|
||||
|
||||
// Call LLM with current context; force_text drops tools to guarantee a
|
||||
// text response on the final iteration.
|
||||
// text response on the final iteration. The pre-built system prompt
|
||||
// avoids rebuilding the same ~1,500-token string each iteration.
|
||||
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 {
|
||||
cached_prompt.clone()
|
||||
})
|
||||
.with_metadata({
|
||||
let mut m = std::collections::HashMap::new();
|
||||
m.insert("thread_id".to_string(), thread_id.to_string());
|
||||
@@ -246,7 +276,7 @@ impl Agent {
|
||||
// Compact: keep system messages + last user message + current turn
|
||||
context_messages = compact_messages_for_retry(&context_messages);
|
||||
|
||||
// Rebuild context with compacted messages
|
||||
// Rebuild context with compacted messages, reusing cached prompt
|
||||
let mut retry_context = ReasoningContext::new()
|
||||
.with_messages(context_messages.clone())
|
||||
.with_tools(if force_text {
|
||||
@@ -256,6 +286,7 @@ impl Agent {
|
||||
})
|
||||
.with_metadata(context.metadata.clone());
|
||||
retry_context.force_text = force_text;
|
||||
retry_context.system_prompt = context.system_prompt.clone();
|
||||
|
||||
reasoning
|
||||
.respond_with_tools(&retry_context)
|
||||
@@ -276,12 +307,18 @@ impl Agent {
|
||||
|
||||
// Record cost and track token usage
|
||||
let model_name = self.llm().active_model_name();
|
||||
let read_discount = self.llm().cache_read_discount();
|
||||
let write_multiplier = self.llm().cache_write_multiplier();
|
||||
let call_cost = self
|
||||
.cost_guard()
|
||||
.record_llm_call(
|
||||
&model_name,
|
||||
output.usage.input_tokens,
|
||||
output.usage.output_tokens,
|
||||
output.usage.cache_read_input_tokens,
|
||||
output.usage.cache_creation_input_tokens,
|
||||
read_discount,
|
||||
write_multiplier,
|
||||
Some(self.llm().cost_per_token()),
|
||||
)
|
||||
.await;
|
||||
@@ -294,6 +331,24 @@ impl Agent {
|
||||
|
||||
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;
|
||||
}
|
||||
|
||||
// Strip internal "[Called tool ...]" text that can leak when
|
||||
// provider flattening (e.g. NEAR AI) converts tool_calls to
|
||||
// plain text and the LLM echoes it back.
|
||||
@@ -304,6 +359,7 @@ impl Agent {
|
||||
tool_calls,
|
||||
content,
|
||||
} => {
|
||||
consecutive_tool_intent_nudges = 0;
|
||||
// Add the assistant message with tool_calls to context.
|
||||
// OpenAI protocol requires this before tool-result messages.
|
||||
context_messages.push(ChatMessage::assistant_with_tool_calls(
|
||||
@@ -625,8 +681,53 @@ impl Agent {
|
||||
.into())
|
||||
});
|
||||
|
||||
// Send ToolResult preview
|
||||
if let Ok(ref output) = tool_result
|
||||
// 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
|
||||
&& matches!(tc.name.as_str(), "image_generate" | "image_edit")
|
||||
{
|
||||
if let Ok(sentinel) =
|
||||
serde_json::from_str::<serde_json::Value>(output)
|
||||
&& sentinel.get("type").and_then(|v| v.as_str())
|
||||
== Some("image_generated")
|
||||
{
|
||||
let data_url = sentinel
|
||||
.get("data")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or_default()
|
||||
.to_string();
|
||||
let path = sentinel
|
||||
.get("path")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(String::from);
|
||||
// Skip broadcasting if data_url is empty to avoid
|
||||
// sending a broken ImageGenerated SSE event.
|
||||
if data_url.is_empty() {
|
||||
tracing::warn!(
|
||||
"Image generation sentinel has empty data URL, skipping broadcast"
|
||||
);
|
||||
} else {
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::ImageGenerated { data_url, path },
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
true
|
||||
} else {
|
||||
false
|
||||
}
|
||||
} else {
|
||||
false
|
||||
};
|
||||
|
||||
// Send ToolResult preview (skip for image sentinels to avoid
|
||||
// broadcasting multi-MB base64 data as a preview)
|
||||
if !is_image_sentinel
|
||||
&& let Ok(ref output) = tool_result
|
||||
&& !output.is_empty()
|
||||
{
|
||||
let _ = self
|
||||
@@ -642,23 +743,6 @@ impl Agent {
|
||||
.await;
|
||||
}
|
||||
|
||||
// Record result in thread
|
||||
{
|
||||
let mut sess = session.lock().await;
|
||||
if let Some(thread) = sess.threads.get_mut(&thread_id)
|
||||
&& let Some(turn) = thread.last_turn_mut()
|
||||
{
|
||||
match &tool_result {
|
||||
Ok(output) => {
|
||||
turn.record_tool_result(serde_json::json!(output));
|
||||
}
|
||||
Err(e) => {
|
||||
turn.record_tool_error(e.to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Check for auth awaiting — defer the return
|
||||
// until all results are recorded.
|
||||
if deferred_auth.is_none()
|
||||
@@ -698,6 +782,7 @@ impl Agent {
|
||||
}
|
||||
|
||||
// Sanitize and add tool result to context
|
||||
let is_tool_error = tool_result.is_err();
|
||||
let result_content = match tool_result {
|
||||
Ok(output) => {
|
||||
let sanitized =
|
||||
@@ -708,9 +793,26 @@ impl Agent {
|
||||
sanitized.was_modified,
|
||||
)
|
||||
}
|
||||
Err(e) => format!("Error: {}", e),
|
||||
Err(e) => format!("Tool '{}' failed: {}", tc.name, e),
|
||||
};
|
||||
|
||||
// Record sanitized result in thread so messages()
|
||||
// and persist_tool_calls() use cleaned content.
|
||||
{
|
||||
let mut sess = session.lock().await;
|
||||
if let Some(thread) = sess.threads.get_mut(&thread_id)
|
||||
&& let Some(turn) = thread.last_turn_mut()
|
||||
{
|
||||
if is_tool_error {
|
||||
turn.record_tool_error(result_content.clone());
|
||||
} else {
|
||||
turn.record_tool_result(serde_json::json!(
|
||||
result_content
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
context_messages.push(ChatMessage::tool_result(
|
||||
&tc.id,
|
||||
&tc.name,
|
||||
@@ -740,6 +842,7 @@ impl Agent {
|
||||
tool_call_id: tc.id.clone(),
|
||||
context_messages: context_messages.clone(),
|
||||
deferred_tool_calls: tool_calls[approval_idx + 1..].to_vec(),
|
||||
user_timezone: Some(user_tz.name().to_string()),
|
||||
};
|
||||
|
||||
return Ok(AgenticLoopResult::NeedApproval { pending });
|
||||
@@ -1041,6 +1144,8 @@ mod tests {
|
||||
input_tokens: 0,
|
||||
output_tokens: 0,
|
||||
finish_reason: FinishReason::Stop,
|
||||
cache_read_input_tokens: 0,
|
||||
cache_creation_input_tokens: 0,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1054,6 +1159,8 @@ mod tests {
|
||||
input_tokens: 0,
|
||||
output_tokens: 0,
|
||||
finish_reason: FinishReason::Stop,
|
||||
cache_read_input_tokens: 0,
|
||||
cache_creation_input_tokens: 0,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1078,6 +1185,8 @@ mod tests {
|
||||
cost_guard: Arc::new(CostGuard::new(CostGuardConfig::default())),
|
||||
sse_tx: None,
|
||||
http_interceptor: None,
|
||||
transcription: None,
|
||||
document_extraction: None,
|
||||
};
|
||||
|
||||
Agent::new(
|
||||
@@ -1095,6 +1204,8 @@ mod tests {
|
||||
max_actions_per_hour: None,
|
||||
max_tool_iterations: 50,
|
||||
auto_approve_tools: false,
|
||||
default_timezone: "UTC".to_string(),
|
||||
max_tokens_per_job: 0,
|
||||
},
|
||||
deps,
|
||||
Arc::new(ChannelManager::new()),
|
||||
@@ -1153,6 +1264,96 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_always_approval_requirement_bypasses_session_auto_approve() {
|
||||
// Regression test: even if tool is auto-approved in session,
|
||||
// ApprovalRequirement::Always must still trigger approval.
|
||||
use crate::tools::ApprovalRequirement;
|
||||
|
||||
let mut session = Session::new("user-1");
|
||||
let tool_name = "tool_remove";
|
||||
|
||||
// Manually auto-approve tool_remove in this session
|
||||
session.auto_approve_tool(tool_name);
|
||||
assert!(
|
||||
session.is_tool_auto_approved(tool_name),
|
||||
"tool should be auto-approved"
|
||||
);
|
||||
|
||||
// However, ApprovalRequirement::Always should always require approval
|
||||
// This is verified by the dispatcher logic: Always => true (ignores session state)
|
||||
let always_req = ApprovalRequirement::Always;
|
||||
let requires_approval = match always_req {
|
||||
ApprovalRequirement::Never => false,
|
||||
ApprovalRequirement::UnlessAutoApproved => !session.is_tool_auto_approved(tool_name),
|
||||
ApprovalRequirement::Always => true,
|
||||
};
|
||||
|
||||
assert!(
|
||||
requires_approval,
|
||||
"ApprovalRequirement::Always must require approval even when tool is auto-approved"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_always_approval_requirement_vs_unless_auto_approved() {
|
||||
// Verify the two requirements behave differently
|
||||
use crate::tools::ApprovalRequirement;
|
||||
|
||||
let mut session = Session::new("user-2");
|
||||
let tool_name = "http";
|
||||
|
||||
// Scenario 1: Tool is auto-approved
|
||||
session.auto_approve_tool(tool_name);
|
||||
|
||||
// UnlessAutoApproved → doesn't require approval if auto-approved
|
||||
let unless_req = ApprovalRequirement::UnlessAutoApproved;
|
||||
let unless_needs = match unless_req {
|
||||
ApprovalRequirement::Never => false,
|
||||
ApprovalRequirement::UnlessAutoApproved => !session.is_tool_auto_approved(tool_name),
|
||||
ApprovalRequirement::Always => true,
|
||||
};
|
||||
assert!(
|
||||
!unless_needs,
|
||||
"UnlessAutoApproved should not need approval when auto-approved"
|
||||
);
|
||||
|
||||
// Always → always requires approval
|
||||
let always_req = ApprovalRequirement::Always;
|
||||
let always_needs = match always_req {
|
||||
ApprovalRequirement::Never => false,
|
||||
ApprovalRequirement::UnlessAutoApproved => !session.is_tool_auto_approved(tool_name),
|
||||
ApprovalRequirement::Always => true,
|
||||
};
|
||||
assert!(
|
||||
always_needs,
|
||||
"Always must always require approval, even when auto-approved"
|
||||
);
|
||||
|
||||
// Scenario 2: Tool is NOT auto-approved
|
||||
let new_tool = "new_tool";
|
||||
assert!(!session.is_tool_auto_approved(new_tool));
|
||||
|
||||
// UnlessAutoApproved → requires approval
|
||||
let unless_needs = match unless_req {
|
||||
ApprovalRequirement::Never => false,
|
||||
ApprovalRequirement::UnlessAutoApproved => !session.is_tool_auto_approved(new_tool),
|
||||
ApprovalRequirement::Always => true,
|
||||
};
|
||||
assert!(
|
||||
unless_needs,
|
||||
"UnlessAutoApproved should need approval when not auto-approved"
|
||||
);
|
||||
|
||||
// Always → always requires approval
|
||||
let always_needs = match always_req {
|
||||
ApprovalRequirement::Never => false,
|
||||
ApprovalRequirement::UnlessAutoApproved => !session.is_tool_auto_approved(new_tool),
|
||||
ApprovalRequirement::Always => true,
|
||||
};
|
||||
assert!(always_needs, "Always must always require approval");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_pending_approval_serialization_backcompat_without_deferred_calls() {
|
||||
// PendingApproval from before the deferred_tool_calls field was added
|
||||
@@ -1197,6 +1398,7 @@ mod tests {
|
||||
arguments: serde_json::json!({"message": "done"}),
|
||||
},
|
||||
],
|
||||
user_timezone: None,
|
||||
};
|
||||
|
||||
let json = serde_json::to_string(&pending).expect("serialize");
|
||||
@@ -1544,12 +1746,8 @@ mod tests {
|
||||
use crate::testing::StubLlm;
|
||||
|
||||
let stub = Arc::new(StubLlm::failing_non_transient("ctx-bomb"));
|
||||
let safety = Arc::new(SafetyLayer::new(&SafetyConfig {
|
||||
max_output_length: 100_000,
|
||||
injection_check_enabled: false,
|
||||
}));
|
||||
|
||||
let reasoning = Reasoning::new(stub.clone(), safety);
|
||||
let reasoning = Reasoning::new(stub.clone());
|
||||
|
||||
// Build a fat context with lots of history.
|
||||
let messages = vec![
|
||||
@@ -1614,6 +1812,8 @@ mod tests {
|
||||
input_tokens: 0,
|
||||
output_tokens: 5,
|
||||
finish_reason: FinishReason::Stop,
|
||||
cache_read_input_tokens: 0,
|
||||
cache_creation_input_tokens: 0,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1629,6 +1829,8 @@ mod tests {
|
||||
input_tokens: 0,
|
||||
output_tokens: 5,
|
||||
finish_reason: FinishReason::Stop,
|
||||
cache_read_input_tokens: 0,
|
||||
cache_creation_input_tokens: 0,
|
||||
});
|
||||
}
|
||||
// Tools available: always call one.
|
||||
@@ -1642,6 +1844,8 @@ mod tests {
|
||||
input_tokens: 0,
|
||||
output_tokens: 5,
|
||||
finish_reason: FinishReason::ToolUse,
|
||||
cache_read_input_tokens: 0,
|
||||
cache_creation_input_tokens: 0,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1653,11 +1857,7 @@ mod tests {
|
||||
use crate::llm::{Reasoning, ReasoningContext, RespondResult, ToolDefinition};
|
||||
|
||||
let provider = Arc::new(AlwaysToolCallProvider);
|
||||
let safety = Arc::new(SafetyLayer::new(&SafetyConfig {
|
||||
max_output_length: 100_000,
|
||||
injection_check_enabled: false,
|
||||
}));
|
||||
let reasoning = Reasoning::new(provider, safety);
|
||||
let reasoning = Reasoning::new(provider);
|
||||
|
||||
let tool_def = ToolDefinition {
|
||||
name: "echo".to_string(),
|
||||
@@ -1766,6 +1966,8 @@ mod tests {
|
||||
input_tokens: 0,
|
||||
output_tokens: 2,
|
||||
finish_reason: FinishReason::Stop,
|
||||
cache_read_input_tokens: 0,
|
||||
cache_creation_input_tokens: 0,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1780,6 +1982,8 @@ mod tests {
|
||||
input_tokens: 0,
|
||||
output_tokens: 2,
|
||||
finish_reason: FinishReason::Stop,
|
||||
cache_read_input_tokens: 0,
|
||||
cache_creation_input_tokens: 0,
|
||||
});
|
||||
}
|
||||
// Always call a tool that does not exist in the registry.
|
||||
@@ -1793,6 +1997,8 @@ mod tests {
|
||||
input_tokens: 0,
|
||||
output_tokens: 5,
|
||||
finish_reason: FinishReason::ToolUse,
|
||||
cache_read_input_tokens: 0,
|
||||
cache_creation_input_tokens: 0,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1818,6 +2024,8 @@ mod tests {
|
||||
cost_guard: Arc::new(CostGuard::new(CostGuardConfig::default())),
|
||||
sse_tx: None,
|
||||
http_interceptor: None,
|
||||
transcription: None,
|
||||
document_extraction: None,
|
||||
};
|
||||
|
||||
Agent::new(
|
||||
@@ -1835,6 +2043,8 @@ mod tests {
|
||||
max_actions_per_hour: None,
|
||||
max_tool_iterations,
|
||||
auto_approve_tools: true,
|
||||
default_timezone: "UTC".to_string(),
|
||||
max_tokens_per_job: 0,
|
||||
},
|
||||
deps,
|
||||
Arc::new(ChannelManager::new()),
|
||||
@@ -1931,6 +2141,8 @@ mod tests {
|
||||
cost_guard: Arc::new(CostGuard::new(CostGuardConfig::default())),
|
||||
sse_tx: None,
|
||||
http_interceptor: None,
|
||||
transcription: None,
|
||||
document_extraction: None,
|
||||
};
|
||||
|
||||
Agent::new(
|
||||
@@ -1948,6 +2160,8 @@ mod tests {
|
||||
max_actions_per_hour: None,
|
||||
max_tool_iterations: max_iter,
|
||||
auto_approve_tools: true,
|
||||
default_timezone: "UTC".to_string(),
|
||||
max_tokens_per_job: 0,
|
||||
},
|
||||
deps,
|
||||
Arc::new(ChannelManager::new()),
|
||||
@@ -2028,4 +2242,68 @@ mod tests {
|
||||
let result = super::strip_internal_tool_call_text(input);
|
||||
assert_eq!(result, input);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tool_error_format_includes_tool_name() {
|
||||
// Regression test for issue #487: tool errors sent to the LLM should
|
||||
// include the tool name so the model can reason about which tool failed
|
||||
// and try alternatives.
|
||||
let tool_name = "http";
|
||||
let err = crate::error::ToolError::ExecutionFailed {
|
||||
name: tool_name.to_string(),
|
||||
reason: "connection refused".to_string(),
|
||||
};
|
||||
let formatted = format!("Tool '{}' failed: {}", tool_name, err);
|
||||
assert!(
|
||||
formatted.contains("Tool 'http' failed:"),
|
||||
"Error should identify the tool by name, got: {formatted}"
|
||||
);
|
||||
assert!(
|
||||
formatted.contains("connection refused"),
|
||||
"Error should include the underlying reason, got: {formatted}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_image_sentinel_empty_data_url_should_be_skipped() {
|
||||
// Regression: unwrap_or_default() on missing "data" field produces an empty
|
||||
// string. Broadcasting an empty data_url would send a broken SSE event.
|
||||
let sentinel = serde_json::json!({
|
||||
"type": "image_generated",
|
||||
"path": "/tmp/image.png"
|
||||
// "data" field is missing
|
||||
});
|
||||
|
||||
let data_url = sentinel
|
||||
.get("data")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or_default()
|
||||
.to_string();
|
||||
|
||||
assert!(
|
||||
data_url.is_empty(),
|
||||
"Missing 'data' field should produce empty string"
|
||||
);
|
||||
// The fix: empty data_url means we skip broadcasting
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_image_sentinel_present_data_url_is_valid() {
|
||||
let sentinel = serde_json::json!({
|
||||
"type": "image_generated",
|
||||
"data": "data:image/png;base64,abc123",
|
||||
"path": "/tmp/image.png"
|
||||
});
|
||||
|
||||
let data_url = sentinel
|
||||
.get("data")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or_default()
|
||||
.to_string();
|
||||
|
||||
assert!(
|
||||
!data_url.is_empty(),
|
||||
"Present 'data' field should produce non-empty string"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
+168
-8
@@ -29,8 +29,8 @@ use std::time::Duration;
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
use crate::channels::OutgoingResponse;
|
||||
use crate::db::Database;
|
||||
use crate::llm::{ChatMessage, CompletionRequest, LlmProvider, Reasoning};
|
||||
use crate::safety::SafetyLayer;
|
||||
use crate::workspace::Workspace;
|
||||
use crate::workspace::hygiene::HygieneConfig;
|
||||
|
||||
@@ -47,6 +47,12 @@ pub struct HeartbeatConfig {
|
||||
pub notify_user_id: Option<String>,
|
||||
/// Channel to notify on heartbeat findings.
|
||||
pub notify_channel: Option<String>,
|
||||
/// Hour (0-23) when quiet hours start.
|
||||
pub quiet_hours_start: Option<u32>,
|
||||
/// Hour (0-23) when quiet hours end.
|
||||
pub quiet_hours_end: Option<u32>,
|
||||
/// Timezone for quiet hours evaluation (IANA name).
|
||||
pub timezone: Option<String>,
|
||||
}
|
||||
|
||||
impl Default for HeartbeatConfig {
|
||||
@@ -57,6 +63,9 @@ impl Default for HeartbeatConfig {
|
||||
max_failures: 3,
|
||||
notify_user_id: None,
|
||||
notify_channel: None,
|
||||
quiet_hours_start: None,
|
||||
quiet_hours_end: None,
|
||||
timezone: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -74,6 +83,26 @@ impl HeartbeatConfig {
|
||||
self
|
||||
}
|
||||
|
||||
/// Check whether the current time falls within configured quiet hours.
|
||||
pub fn is_quiet_hours(&self) -> bool {
|
||||
use chrono::Timelike;
|
||||
let (Some(start), Some(end)) = (self.quiet_hours_start, self.quiet_hours_end) else {
|
||||
return false;
|
||||
};
|
||||
let tz = self
|
||||
.timezone
|
||||
.as_deref()
|
||||
.and_then(crate::timezone::parse_timezone)
|
||||
.unwrap_or(chrono_tz::UTC);
|
||||
let now_hour = crate::timezone::now_in_tz(tz).hour();
|
||||
if start <= end {
|
||||
now_hour >= start && now_hour < end
|
||||
} else {
|
||||
// Wraps midnight, e.g. 22..06
|
||||
now_hour >= start || now_hour < end
|
||||
}
|
||||
}
|
||||
|
||||
/// Set the notification target.
|
||||
pub fn with_notify(mut self, user_id: impl Into<String>, channel: impl Into<String>) -> Self {
|
||||
self.notify_user_id = Some(user_id.into());
|
||||
@@ -101,8 +130,8 @@ pub struct HeartbeatRunner {
|
||||
hygiene_config: HygieneConfig,
|
||||
workspace: Arc<Workspace>,
|
||||
llm: Arc<dyn LlmProvider>,
|
||||
safety: Arc<SafetyLayer>,
|
||||
response_tx: Option<mpsc::Sender<OutgoingResponse>>,
|
||||
store: Option<Arc<dyn Database>>,
|
||||
consecutive_failures: u32,
|
||||
}
|
||||
|
||||
@@ -113,15 +142,14 @@ impl HeartbeatRunner {
|
||||
hygiene_config: HygieneConfig,
|
||||
workspace: Arc<Workspace>,
|
||||
llm: Arc<dyn LlmProvider>,
|
||||
safety: Arc<SafetyLayer>,
|
||||
) -> Self {
|
||||
Self {
|
||||
config,
|
||||
hygiene_config,
|
||||
workspace,
|
||||
llm,
|
||||
safety,
|
||||
response_tx: None,
|
||||
store: None,
|
||||
consecutive_failures: 0,
|
||||
}
|
||||
}
|
||||
@@ -132,6 +160,12 @@ impl HeartbeatRunner {
|
||||
self
|
||||
}
|
||||
|
||||
/// Set the database store for persistent heartbeat conversations.
|
||||
pub fn with_store(mut self, store: Arc<dyn Database>) -> Self {
|
||||
self.store = Some(store);
|
||||
self
|
||||
}
|
||||
|
||||
/// Run the heartbeat loop.
|
||||
///
|
||||
/// This runs forever, checking periodically based on the configured interval.
|
||||
@@ -153,6 +187,12 @@ impl HeartbeatRunner {
|
||||
loop {
|
||||
interval.tick().await;
|
||||
|
||||
// Skip during quiet hours
|
||||
if self.config.is_quiet_hours() {
|
||||
tracing::debug!("Heartbeat skipped: quiet hours");
|
||||
continue;
|
||||
}
|
||||
|
||||
// Run memory hygiene in the background so it never delays the
|
||||
// heartbeat checklist. Failures are logged inside run_if_due.
|
||||
let hygiene_workspace = Arc::clone(&self.workspace);
|
||||
@@ -263,7 +303,7 @@ impl HeartbeatRunner {
|
||||
.with_max_tokens(max_tokens)
|
||||
.with_temperature(0.3);
|
||||
|
||||
let reasoning = Reasoning::new(self.llm.clone(), self.safety.clone());
|
||||
let reasoning = Reasoning::new(self.llm.clone());
|
||||
let (content, _usage) = match reasoning.complete(request).await {
|
||||
Ok(r) => r,
|
||||
Err(e) => return HeartbeatResult::Failed(format!("LLM call failed: {}", e)),
|
||||
@@ -292,9 +332,32 @@ impl HeartbeatRunner {
|
||||
return;
|
||||
};
|
||||
|
||||
let user_id = self.config.notify_user_id.as_deref().unwrap_or("default");
|
||||
|
||||
// Persist to heartbeat conversation and get thread_id
|
||||
let thread_id = if let Some(ref store) = self.store {
|
||||
match store.get_or_create_heartbeat_conversation(user_id).await {
|
||||
Ok(conv_id) => {
|
||||
if let Err(e) = store
|
||||
.add_conversation_message(conv_id, "assistant", message)
|
||||
.await
|
||||
{
|
||||
tracing::error!("Failed to persist heartbeat message: {}", e);
|
||||
}
|
||||
Some(conv_id.to_string())
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!("Failed to get heartbeat conversation: {}", e);
|
||||
None
|
||||
}
|
||||
}
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let response = OutgoingResponse {
|
||||
content: format!("🔔 *Heartbeat Alert*\n\n{}", message),
|
||||
thread_id: None,
|
||||
thread_id,
|
||||
attachments: Vec::new(),
|
||||
metadata: serde_json::json!({
|
||||
"source": "heartbeat",
|
||||
@@ -354,13 +417,16 @@ pub fn spawn_heartbeat(
|
||||
hygiene_config: HygieneConfig,
|
||||
workspace: Arc<Workspace>,
|
||||
llm: Arc<dyn LlmProvider>,
|
||||
safety: Arc<SafetyLayer>,
|
||||
response_tx: Option<mpsc::Sender<OutgoingResponse>>,
|
||||
store: Option<Arc<dyn Database>>,
|
||||
) -> tokio::task::JoinHandle<()> {
|
||||
let mut runner = HeartbeatRunner::new(config, hygiene_config, workspace, llm, safety);
|
||||
let mut runner = HeartbeatRunner::new(config, hygiene_config, workspace, llm);
|
||||
if let Some(tx) = response_tx {
|
||||
runner = runner.with_response_channel(tx);
|
||||
}
|
||||
if let Some(s) = store {
|
||||
runner = runner.with_store(s);
|
||||
}
|
||||
|
||||
tokio::spawn(async move {
|
||||
runner.run().await;
|
||||
@@ -495,4 +561,98 @@ mod tests {
|
||||
let content = "<!-- comment -->\nActual task here";
|
||||
assert!(!is_effectively_empty(content));
|
||||
}
|
||||
|
||||
// ==================== quiet hours ====================
|
||||
|
||||
#[test]
|
||||
fn test_quiet_hours_inside() {
|
||||
use chrono::{Timelike, Utc};
|
||||
|
||||
let now_utc = Utc::now();
|
||||
let hour = now_utc.hour();
|
||||
let start = hour;
|
||||
let end = (hour + 1) % 24;
|
||||
|
||||
let config = HeartbeatConfig {
|
||||
quiet_hours_start: Some(start),
|
||||
quiet_hours_end: Some(end),
|
||||
timezone: Some("UTC".to_string()),
|
||||
..HeartbeatConfig::default()
|
||||
};
|
||||
// Current UTC hour is inside [start, end) by construction
|
||||
assert!(config.is_quiet_hours());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_quiet_hours_outside() {
|
||||
use chrono::{Timelike, Utc};
|
||||
|
||||
let now_utc = Utc::now();
|
||||
let hour = now_utc.hour();
|
||||
let start = (hour + 1) % 24;
|
||||
let end = (hour + 2) % 24;
|
||||
|
||||
let config = HeartbeatConfig {
|
||||
quiet_hours_start: Some(start),
|
||||
quiet_hours_end: Some(end),
|
||||
timezone: Some("UTC".to_string()),
|
||||
..HeartbeatConfig::default()
|
||||
};
|
||||
// Current UTC hour is outside [start, end) by construction
|
||||
assert!(!config.is_quiet_hours());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_quiet_hours_wraparound_excludes_now() {
|
||||
use chrono::{Timelike, Utc};
|
||||
|
||||
let now_utc = Utc::now();
|
||||
let hour = now_utc.hour();
|
||||
// Window covers all hours except the current one
|
||||
let start = (hour + 1) % 24;
|
||||
let end = hour;
|
||||
|
||||
let config = HeartbeatConfig {
|
||||
quiet_hours_start: Some(start),
|
||||
quiet_hours_end: Some(end),
|
||||
timezone: Some("UTC".to_string()),
|
||||
..HeartbeatConfig::default()
|
||||
};
|
||||
assert!(!config.is_quiet_hours());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_quiet_hours_none_configured() {
|
||||
let config = HeartbeatConfig::default();
|
||||
assert!(!config.is_quiet_hours());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_quiet_hours_same_start_end() {
|
||||
let config = HeartbeatConfig {
|
||||
quiet_hours_start: Some(10),
|
||||
quiet_hours_end: Some(10),
|
||||
timezone: Some("UTC".to_string()),
|
||||
..HeartbeatConfig::default()
|
||||
};
|
||||
// start == end means zero-width window, should be false
|
||||
assert!(!config.is_quiet_hours());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_spawn_heartbeat_accepts_store_param() {
|
||||
// Regression: spawn_heartbeat must accept an optional Database store
|
||||
// for persisting heartbeat notifications to a dedicated conversation.
|
||||
// Compile-time check: the 7th parameter is `Option<Arc<dyn Database>>`.
|
||||
#[allow(clippy::type_complexity)]
|
||||
let _fn_ptr: fn(
|
||||
HeartbeatConfig,
|
||||
HygieneConfig,
|
||||
Arc<crate::workspace::Workspace>,
|
||||
Arc<dyn crate::llm::LlmProvider>,
|
||||
Option<tokio::sync::mpsc::Sender<crate::channels::OutgoingResponse>>,
|
||||
Option<Arc<dyn crate::db::Database>>,
|
||||
) -> tokio::task::JoinHandle<()> = spawn_heartbeat;
|
||||
let _ = _fn_ptr;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -11,6 +11,7 @@
|
||||
//! - Context compaction for long conversations
|
||||
|
||||
mod agent_loop;
|
||||
mod attachments;
|
||||
mod commands;
|
||||
pub mod compaction;
|
||||
pub mod context_monitor;
|
||||
|
||||
+112
-11
@@ -57,7 +57,11 @@ pub struct Routine {
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum Trigger {
|
||||
/// Fire on a cron schedule (e.g. "0 9 * * MON-FRI" or "every 2h").
|
||||
Cron { schedule: String },
|
||||
Cron {
|
||||
schedule: String,
|
||||
#[serde(default)]
|
||||
timezone: Option<String>,
|
||||
},
|
||||
/// Fire when a channel message matches a pattern.
|
||||
Event {
|
||||
/// Optional channel filter (e.g. "telegram", "slack").
|
||||
@@ -99,7 +103,21 @@ impl Trigger {
|
||||
field: "schedule".into(),
|
||||
})?
|
||||
.to_string();
|
||||
Ok(Trigger::Cron { schedule })
|
||||
let timezone = config
|
||||
.get("timezone")
|
||||
.and_then(|v| v.as_str())
|
||||
.and_then(|tz| {
|
||||
if crate::timezone::parse_timezone(tz).is_some() {
|
||||
Some(tz.to_string())
|
||||
} else {
|
||||
tracing::warn!(
|
||||
"Ignoring invalid timezone '{}' from DB for cron trigger",
|
||||
tz
|
||||
);
|
||||
None
|
||||
}
|
||||
});
|
||||
Ok(Trigger::Cron { schedule, timezone })
|
||||
}
|
||||
"event" => {
|
||||
let pattern = config
|
||||
@@ -137,7 +155,10 @@ impl Trigger {
|
||||
/// Serialize trigger-specific config to JSON for DB storage.
|
||||
pub fn to_config_json(&self) -> serde_json::Value {
|
||||
match self {
|
||||
Trigger::Cron { schedule } => serde_json::json!({ "schedule": schedule }),
|
||||
Trigger::Cron { schedule, timezone } => serde_json::json!({
|
||||
"schedule": schedule,
|
||||
"timezone": timezone,
|
||||
}),
|
||||
Trigger::Event { channel, pattern } => serde_json::json!({
|
||||
"pattern": pattern,
|
||||
"channel": channel,
|
||||
@@ -175,6 +196,11 @@ pub enum RoutineAction {
|
||||
/// Max reasoning iterations (default: 10).
|
||||
#[serde(default = "default_max_iterations")]
|
||||
max_iterations: u32,
|
||||
/// Tool names pre-authorized for `Always`-approval tools (e.g. destructive
|
||||
/// shell commands, cross-channel messaging). `UnlessAutoApproved` tools are
|
||||
/// automatically permitted in routine jobs without listing them here.
|
||||
#[serde(default)]
|
||||
tool_permissions: Vec<String>,
|
||||
},
|
||||
}
|
||||
|
||||
@@ -186,6 +212,19 @@ fn default_max_iterations() -> u32 {
|
||||
10
|
||||
}
|
||||
|
||||
/// Parse a `tool_permissions` JSON array into a `Vec<String>`.
|
||||
pub fn parse_tool_permissions(value: &serde_json::Value) -> Vec<String> {
|
||||
value
|
||||
.get("tool_permissions")
|
||||
.and_then(|v| v.as_array())
|
||||
.map(|arr| {
|
||||
arr.iter()
|
||||
.filter_map(|v| v.as_str().map(String::from))
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
impl RoutineAction {
|
||||
/// The string tag stored in the DB action_type column.
|
||||
pub fn type_tag(&self) -> &'static str {
|
||||
@@ -248,10 +287,12 @@ impl RoutineAction {
|
||||
.and_then(|v| v.as_u64())
|
||||
.unwrap_or(default_max_iterations() as u64)
|
||||
as u32;
|
||||
let tool_permissions = parse_tool_permissions(&config);
|
||||
Ok(RoutineAction::FullJob {
|
||||
title,
|
||||
description,
|
||||
max_iterations,
|
||||
tool_permissions,
|
||||
})
|
||||
}
|
||||
other => Err(RoutineError::UnknownActionType {
|
||||
@@ -276,10 +317,12 @@ impl RoutineAction {
|
||||
title,
|
||||
description,
|
||||
max_iterations,
|
||||
tool_permissions,
|
||||
} => serde_json::json!({
|
||||
"title": title,
|
||||
"description": description,
|
||||
"max_iterations": max_iterations,
|
||||
"tool_permissions": tool_permissions,
|
||||
}),
|
||||
}
|
||||
}
|
||||
@@ -393,12 +436,25 @@ pub fn content_hash(content: &str) -> u64 {
|
||||
}
|
||||
|
||||
/// Parse a cron expression and compute the next fire time from now.
|
||||
pub fn next_cron_fire(schedule: &str) -> Result<Option<DateTime<Utc>>, RoutineError> {
|
||||
///
|
||||
/// When `timezone` is provided and valid, the schedule is evaluated in that
|
||||
/// timezone and the result is converted back to UTC. Otherwise UTC is used.
|
||||
pub fn next_cron_fire(
|
||||
schedule: &str,
|
||||
timezone: Option<&str>,
|
||||
) -> Result<Option<DateTime<Utc>>, RoutineError> {
|
||||
let cron_schedule =
|
||||
cron::Schedule::from_str(schedule).map_err(|e| RoutineError::InvalidCron {
|
||||
reason: e.to_string(),
|
||||
})?;
|
||||
Ok(cron_schedule.upcoming(Utc).next())
|
||||
if let Some(tz) = timezone.and_then(crate::timezone::parse_timezone) {
|
||||
Ok(cron_schedule
|
||||
.upcoming(tz)
|
||||
.next()
|
||||
.map(|dt| dt.with_timezone(&Utc)))
|
||||
} else {
|
||||
Ok(cron_schedule.upcoming(Utc).next())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
@@ -411,10 +467,11 @@ mod tests {
|
||||
fn test_trigger_roundtrip() {
|
||||
let trigger = Trigger::Cron {
|
||||
schedule: "0 9 * * MON-FRI".to_string(),
|
||||
timezone: None,
|
||||
};
|
||||
let json = trigger.to_config_json();
|
||||
let parsed = Trigger::from_db("cron", json).expect("parse cron");
|
||||
assert!(matches!(parsed, Trigger::Cron { schedule } if schedule == "0 9 * * MON-FRI"));
|
||||
assert!(matches!(parsed, Trigger::Cron { schedule, .. } if schedule == "0 9 * * MON-FRI"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -450,12 +507,13 @@ mod tests {
|
||||
title: "Deploy review".to_string(),
|
||||
description: "Review and deploy pending changes".to_string(),
|
||||
max_iterations: 5,
|
||||
tool_permissions: vec!["shell".to_string()],
|
||||
};
|
||||
let json = action.to_config_json();
|
||||
let parsed = RoutineAction::from_db("full_job", json).expect("parse full_job");
|
||||
assert!(
|
||||
matches!(parsed, RoutineAction::FullJob { title, max_iterations, .. }
|
||||
if title == "Deploy review" && max_iterations == 5)
|
||||
matches!(parsed, RoutineAction::FullJob { title, max_iterations, tool_permissions, .. }
|
||||
if title == "Deploy review" && max_iterations == 5 && tool_permissions == vec!["shell".to_string()])
|
||||
);
|
||||
}
|
||||
|
||||
@@ -486,16 +544,58 @@ mod tests {
|
||||
#[test]
|
||||
fn test_next_cron_fire_valid() {
|
||||
// Every minute should always have a next fire
|
||||
let next = next_cron_fire("* * * * * *").expect("valid cron");
|
||||
let next = next_cron_fire("* * * * * *", None).expect("valid cron");
|
||||
assert!(next.is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_next_cron_fire_invalid() {
|
||||
let result = next_cron_fire("not a cron");
|
||||
let result = next_cron_fire("not a cron", None);
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_trigger_cron_timezone_roundtrip() {
|
||||
let trigger = Trigger::Cron {
|
||||
schedule: "0 9 * * MON-FRI".to_string(),
|
||||
timezone: Some("America/New_York".to_string()),
|
||||
};
|
||||
let json = trigger.to_config_json();
|
||||
let parsed = Trigger::from_db("cron", json).expect("parse cron");
|
||||
assert!(matches!(parsed, Trigger::Cron { schedule, timezone }
|
||||
if schedule == "0 9 * * MON-FRI"
|
||||
&& timezone.as_deref() == Some("America/New_York")));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_trigger_cron_no_timezone_backward_compat() {
|
||||
let json = serde_json::json!({"schedule": "0 9 * * *"});
|
||||
let parsed = Trigger::from_db("cron", json).expect("parse cron");
|
||||
assert!(matches!(parsed, Trigger::Cron { timezone, .. } if timezone.is_none()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_trigger_cron_invalid_timezone_coerced_to_none() {
|
||||
let json = serde_json::json!({"schedule": "0 9 * * *", "timezone": "Fake/Zone"});
|
||||
let parsed = Trigger::from_db("cron", json).expect("parse cron");
|
||||
assert!(
|
||||
matches!(parsed, Trigger::Cron { timezone, .. } if timezone.is_none()),
|
||||
"invalid timezone should be coerced to None"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_next_cron_fire_with_timezone() {
|
||||
let next_utc = next_cron_fire("0 0 9 * * * *", None)
|
||||
.expect("valid cron")
|
||||
.expect("has next");
|
||||
let next_est = next_cron_fire("0 0 9 * * * *", Some("America/New_York"))
|
||||
.expect("valid cron")
|
||||
.expect("has next");
|
||||
// EST is UTC-5 (or EDT UTC-4), so the UTC result should differ
|
||||
assert_ne!(next_utc, next_est, "timezone should shift the fire time");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_guardrails_default() {
|
||||
let g = RoutineGuardrails::default();
|
||||
@@ -508,7 +608,8 @@ mod tests {
|
||||
fn test_trigger_type_tag() {
|
||||
assert_eq!(
|
||||
Trigger::Cron {
|
||||
schedule: String::new()
|
||||
schedule: String::new(),
|
||||
timezone: None,
|
||||
}
|
||||
.type_tag(),
|
||||
"cron"
|
||||
|
||||
+549
-21
@@ -25,9 +25,14 @@ use crate::agent::routine::{
|
||||
};
|
||||
use crate::channels::{IncomingMessage, OutgoingResponse};
|
||||
use crate::config::RoutineConfig;
|
||||
use crate::context::JobContext;
|
||||
use crate::db::Database;
|
||||
use crate::error::RoutineError;
|
||||
use crate::llm::{ChatMessage, CompletionRequest, FinishReason, LlmProvider};
|
||||
use crate::llm::{
|
||||
ChatMessage, CompletionRequest, FinishReason, LlmProvider, ToolCall, ToolCompletionRequest,
|
||||
};
|
||||
use crate::safety::SafetyLayer;
|
||||
use crate::tools::{ApprovalContext, ApprovalRequirement, ToolError, ToolRegistry, redact_params};
|
||||
use crate::workspace::Workspace;
|
||||
|
||||
/// The routine execution engine.
|
||||
@@ -44,9 +49,14 @@ pub struct RoutineEngine {
|
||||
event_cache: Arc<RwLock<Vec<(Uuid, Routine, Regex)>>>,
|
||||
/// Scheduler for dispatching jobs (FullJob mode).
|
||||
scheduler: Option<Arc<Scheduler>>,
|
||||
/// Tool registry for lightweight routine tool execution.
|
||||
tools: Arc<ToolRegistry>,
|
||||
/// Safety layer for tool output sanitization.
|
||||
safety: Arc<SafetyLayer>,
|
||||
}
|
||||
|
||||
impl RoutineEngine {
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub fn new(
|
||||
config: RoutineConfig,
|
||||
store: Arc<dyn Database>,
|
||||
@@ -54,6 +64,8 @@ impl RoutineEngine {
|
||||
workspace: Arc<Workspace>,
|
||||
notify_tx: mpsc::Sender<OutgoingResponse>,
|
||||
scheduler: Option<Arc<Scheduler>>,
|
||||
tools: Arc<ToolRegistry>,
|
||||
safety: Arc<SafetyLayer>,
|
||||
) -> Self {
|
||||
Self {
|
||||
config,
|
||||
@@ -64,6 +76,8 @@ impl RoutineEngine {
|
||||
running_count: Arc::new(AtomicUsize::new(0)),
|
||||
event_cache: Arc::new(RwLock::new(Vec::new())),
|
||||
scheduler,
|
||||
tools,
|
||||
safety,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -169,7 +183,7 @@ impl RoutineEngine {
|
||||
continue;
|
||||
}
|
||||
|
||||
let detail = if let Trigger::Cron { ref schedule } = routine.trigger {
|
||||
let detail = if let Trigger::Cron { ref schedule, .. } = routine.trigger {
|
||||
Some(schedule.clone())
|
||||
} else {
|
||||
None
|
||||
@@ -180,7 +194,14 @@ impl RoutineEngine {
|
||||
}
|
||||
|
||||
/// Fire a routine manually (from tool call or CLI).
|
||||
pub async fn fire_manual(&self, routine_id: Uuid) -> Result<Uuid, RoutineError> {
|
||||
///
|
||||
/// Bypasses cooldown checks (those only apply to cron/event triggers).
|
||||
/// Still enforces enabled check and concurrent run limit.
|
||||
pub async fn fire_manual(
|
||||
&self,
|
||||
routine_id: Uuid,
|
||||
user_id: Option<&str>,
|
||||
) -> Result<Uuid, RoutineError> {
|
||||
let routine = self
|
||||
.store
|
||||
.get_routine(routine_id)
|
||||
@@ -190,6 +211,13 @@ impl RoutineEngine {
|
||||
})?
|
||||
.ok_or(RoutineError::NotFound { id: routine_id })?;
|
||||
|
||||
// Enforce ownership when a user_id is provided (gateway calls).
|
||||
if let Some(uid) = user_id
|
||||
&& routine.user_id != uid
|
||||
{
|
||||
return Err(RoutineError::NotAuthorized { id: routine_id });
|
||||
}
|
||||
|
||||
if !routine.enabled {
|
||||
return Err(RoutineError::Disabled {
|
||||
name: routine.name.clone(),
|
||||
@@ -225,12 +253,15 @@ impl RoutineEngine {
|
||||
|
||||
// Execute inline for manual triggers (caller wants to wait)
|
||||
let engine = EngineContext {
|
||||
config: self.config.clone(),
|
||||
store: self.store.clone(),
|
||||
llm: self.llm.clone(),
|
||||
workspace: self.workspace.clone(),
|
||||
notify_tx: self.notify_tx.clone(),
|
||||
running_count: self.running_count.clone(),
|
||||
scheduler: self.scheduler.clone(),
|
||||
tools: self.tools.clone(),
|
||||
safety: self.safety.clone(),
|
||||
};
|
||||
|
||||
tokio::spawn(async move {
|
||||
@@ -257,12 +288,15 @@ impl RoutineEngine {
|
||||
};
|
||||
|
||||
let engine = EngineContext {
|
||||
config: self.config.clone(),
|
||||
store: self.store.clone(),
|
||||
llm: self.llm.clone(),
|
||||
workspace: self.workspace.clone(),
|
||||
notify_tx: self.notify_tx.clone(),
|
||||
running_count: self.running_count.clone(),
|
||||
scheduler: self.scheduler.clone(),
|
||||
tools: self.tools.clone(),
|
||||
safety: self.safety.clone(),
|
||||
};
|
||||
|
||||
// Record the run in DB, then spawn execution
|
||||
@@ -304,12 +338,15 @@ impl RoutineEngine {
|
||||
|
||||
/// Shared context passed to the execution function.
|
||||
struct EngineContext {
|
||||
config: RoutineConfig,
|
||||
store: Arc<dyn Database>,
|
||||
llm: Arc<dyn LlmProvider>,
|
||||
workspace: Arc<Workspace>,
|
||||
notify_tx: mpsc::Sender<OutgoingResponse>,
|
||||
running_count: Arc<AtomicUsize>,
|
||||
scheduler: Option<Arc<Scheduler>>,
|
||||
tools: Arc<ToolRegistry>,
|
||||
safety: Arc<SafetyLayer>,
|
||||
}
|
||||
|
||||
/// Execute a routine run. Handles both lightweight and full_job modes.
|
||||
@@ -327,7 +364,19 @@ async fn execute_routine(ctx: EngineContext, routine: Routine, run: RoutineRun)
|
||||
title,
|
||||
description,
|
||||
max_iterations,
|
||||
} => execute_full_job(&ctx, &routine, &run, title, description, *max_iterations).await,
|
||||
tool_permissions,
|
||||
} => {
|
||||
execute_full_job(
|
||||
&ctx,
|
||||
&routine,
|
||||
&run,
|
||||
title,
|
||||
description,
|
||||
*max_iterations,
|
||||
tool_permissions,
|
||||
)
|
||||
.await
|
||||
}
|
||||
};
|
||||
|
||||
// Decrement running count
|
||||
@@ -353,8 +402,12 @@ async fn execute_routine(ctx: EngineContext, routine: Routine, run: RoutineRun)
|
||||
|
||||
// Update routine runtime state
|
||||
let now = Utc::now();
|
||||
let next_fire = if let Trigger::Cron { ref schedule } = routine.trigger {
|
||||
next_cron_fire(schedule).unwrap_or(None)
|
||||
let next_fire = if let Trigger::Cron {
|
||||
ref schedule,
|
||||
ref timezone,
|
||||
} = routine.trigger
|
||||
{
|
||||
next_cron_fire(schedule, timezone.as_deref()).unwrap_or(None)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
@@ -380,6 +433,39 @@ async fn execute_routine(ctx: EngineContext, routine: Routine, run: RoutineRun)
|
||||
tracing::error!(routine = %routine.name, "Failed to update runtime state: {}", e);
|
||||
}
|
||||
|
||||
// Persist routine result to its dedicated conversation thread
|
||||
let thread_id = match ctx
|
||||
.store
|
||||
.get_or_create_routine_conversation(routine.id, &routine.name, &routine.user_id)
|
||||
.await
|
||||
{
|
||||
Ok(conv_id) => {
|
||||
tracing::debug!(
|
||||
routine = %routine.name,
|
||||
routine_id = %routine.id,
|
||||
conversation_id = %conv_id,
|
||||
"Resolved routine conversation thread"
|
||||
);
|
||||
// Record the run result as a conversation message
|
||||
let msg = match (&summary, status) {
|
||||
(Some(s), _) => format!("[{}] {}: {}", run.trigger_type, status, s),
|
||||
(None, _) => format!("[{}] {}", run.trigger_type, status),
|
||||
};
|
||||
if let Err(e) = ctx
|
||||
.store
|
||||
.add_conversation_message(conv_id, "assistant", &msg)
|
||||
.await
|
||||
{
|
||||
tracing::error!(routine = %routine.name, "Failed to persist routine message: {}", e);
|
||||
}
|
||||
Some(conv_id.to_string())
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!(routine = %routine.name, "Failed to get routine conversation: {}", e);
|
||||
None
|
||||
}
|
||||
};
|
||||
|
||||
// Send notifications based on config
|
||||
send_notification(
|
||||
&ctx.notify_tx,
|
||||
@@ -387,6 +473,7 @@ async fn execute_routine(ctx: EngineContext, routine: Routine, run: RoutineRun)
|
||||
&routine.name,
|
||||
status,
|
||||
summary.as_deref(),
|
||||
thread_id.as_deref(),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
@@ -418,6 +505,7 @@ async fn execute_full_job(
|
||||
title: &str,
|
||||
description: &str,
|
||||
max_iterations: u32,
|
||||
tool_permissions: &[String],
|
||||
) -> Result<(RunStatus, Option<String>, Option<i32>), RoutineError> {
|
||||
let scheduler = ctx
|
||||
.scheduler
|
||||
@@ -426,10 +514,26 @@ async fn execute_full_job(
|
||||
reason: "scheduler not available".to_string(),
|
||||
})?;
|
||||
|
||||
let metadata = serde_json::json!({ "max_iterations": max_iterations });
|
||||
let mut metadata = serde_json::json!({ "max_iterations": max_iterations });
|
||||
// Carry the routine's notify config in job metadata so the message tool
|
||||
// can resolve channel/target per-job without global state mutation.
|
||||
if let Some(channel) = &routine.notify.channel {
|
||||
metadata["notify_channel"] = serde_json::json!(channel);
|
||||
}
|
||||
metadata["notify_user"] = serde_json::json!(&routine.notify.user);
|
||||
|
||||
// Build approval context: UnlessAutoApproved tools are auto-approved for routines;
|
||||
// Always tools require explicit listing in tool_permissions.
|
||||
let approval_context = ApprovalContext::autonomous_with_tools(tool_permissions.iter().cloned());
|
||||
|
||||
let job_id = scheduler
|
||||
.dispatch_job(&routine.user_id, title, description, Some(metadata))
|
||||
.dispatch_job_with_context(
|
||||
&routine.user_id,
|
||||
title,
|
||||
description,
|
||||
Some(metadata),
|
||||
approval_context,
|
||||
)
|
||||
.await
|
||||
.map_err(|e| RoutineError::JobDispatchFailed {
|
||||
reason: format!("failed to dispatch job: {e}"),
|
||||
@@ -456,7 +560,10 @@ async fn execute_full_job(
|
||||
Ok((RunStatus::Ok, Some(summary), None))
|
||||
}
|
||||
|
||||
/// Execute a lightweight routine (single LLM call).
|
||||
/// Execute a lightweight routine with optional tool support.
|
||||
///
|
||||
/// If tools are enabled, this runs a simplified agentic loop (max 3-5 iterations).
|
||||
/// If tools are disabled, this does a single LLM call (original behavior).
|
||||
async fn execute_lightweight(
|
||||
ctx: &EngineContext,
|
||||
routine: &Routine,
|
||||
@@ -488,7 +595,7 @@ async fn execute_lightweight(
|
||||
Err(_) => None,
|
||||
};
|
||||
|
||||
// Build the prompt
|
||||
// Build the user-facing prompt
|
||||
let mut full_prompt = String::new();
|
||||
full_prompt.push_str(prompt);
|
||||
|
||||
@@ -516,15 +623,6 @@ async fn execute_lightweight(
|
||||
}
|
||||
};
|
||||
|
||||
let messages = if system_prompt.is_empty() {
|
||||
vec![ChatMessage::user(&full_prompt)]
|
||||
} else {
|
||||
vec![
|
||||
ChatMessage::system(&system_prompt),
|
||||
ChatMessage::user(&full_prompt),
|
||||
]
|
||||
};
|
||||
|
||||
// Determine max_tokens from model metadata with fallback
|
||||
let effective_max_tokens = match ctx.llm.model_metadata().await {
|
||||
Ok(meta) => {
|
||||
@@ -534,6 +632,45 @@ async fn execute_lightweight(
|
||||
Err(_) => max_tokens,
|
||||
};
|
||||
|
||||
// If tools are enabled, use the tool execution loop; otherwise, single LLM call
|
||||
if ctx.config.lightweight_tools_enabled {
|
||||
execute_lightweight_with_tools(
|
||||
ctx,
|
||||
routine,
|
||||
&system_prompt,
|
||||
&full_prompt,
|
||||
effective_max_tokens,
|
||||
)
|
||||
.await
|
||||
} else {
|
||||
execute_lightweight_no_tools(
|
||||
ctx,
|
||||
routine,
|
||||
&system_prompt,
|
||||
&full_prompt,
|
||||
effective_max_tokens,
|
||||
)
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
/// Execute a lightweight routine without tool support (original single-call behavior).
|
||||
async fn execute_lightweight_no_tools(
|
||||
ctx: &EngineContext,
|
||||
_routine: &Routine,
|
||||
system_prompt: &str,
|
||||
full_prompt: &str,
|
||||
effective_max_tokens: u32,
|
||||
) -> Result<(RunStatus, Option<String>, Option<i32>), RoutineError> {
|
||||
let messages = if system_prompt.is_empty() {
|
||||
vec![ChatMessage::user(full_prompt)]
|
||||
} else {
|
||||
vec![
|
||||
ChatMessage::system(system_prompt),
|
||||
ChatMessage::user(full_prompt),
|
||||
]
|
||||
};
|
||||
|
||||
let request = CompletionRequest::new(messages)
|
||||
.with_max_tokens(effective_max_tokens)
|
||||
.with_temperature(0.3);
|
||||
@@ -549,7 +686,7 @@ async fn execute_lightweight(
|
||||
let content = response.content.trim();
|
||||
let tokens_used = Some((response.input_tokens + response.output_tokens) as i32);
|
||||
|
||||
// Empty content guard (same as heartbeat)
|
||||
// Empty content guard
|
||||
if content.is_empty() {
|
||||
return if response.finish_reason == FinishReason::Length {
|
||||
Err(RoutineError::TruncatedResponse)
|
||||
@@ -566,6 +703,269 @@ async fn execute_lightweight(
|
||||
Ok((RunStatus::Attention, Some(content.to_string()), tokens_used))
|
||||
}
|
||||
|
||||
/// Handle a text-only LLM response in lightweight routine execution.
|
||||
///
|
||||
/// Checks for the ROUTINE_OK sentinel, validates content, and returns appropriate status.
|
||||
fn handle_text_response(
|
||||
content: &str,
|
||||
finish_reason: FinishReason,
|
||||
total_input_tokens: u32,
|
||||
total_output_tokens: u32,
|
||||
) -> Result<(RunStatus, Option<String>, Option<i32>), RoutineError> {
|
||||
let content = content.trim();
|
||||
|
||||
// Empty content guard
|
||||
if content.is_empty() {
|
||||
return if finish_reason == FinishReason::Length {
|
||||
Err(RoutineError::TruncatedResponse)
|
||||
} else {
|
||||
Err(RoutineError::EmptyResponse)
|
||||
};
|
||||
}
|
||||
|
||||
// Check for the "nothing to do" sentinel
|
||||
if content == "ROUTINE_OK" || content.contains("ROUTINE_OK") {
|
||||
let total_tokens = Some((total_input_tokens + total_output_tokens) as i32);
|
||||
return Ok((RunStatus::Ok, None, total_tokens));
|
||||
}
|
||||
|
||||
let total_tokens = Some((total_input_tokens + total_output_tokens) as i32);
|
||||
Ok((
|
||||
RunStatus::Attention,
|
||||
Some(content.to_string()),
|
||||
total_tokens,
|
||||
))
|
||||
}
|
||||
|
||||
/// Execute a lightweight routine with tool execution support (agentic loop).
|
||||
///
|
||||
/// This is a simplified version of the full dispatcher loop:
|
||||
/// - Max 3-5 iterations (configurable)
|
||||
/// - Sequential tool execution (not parallel)
|
||||
/// - Auto-approval of non-Always tools
|
||||
/// - No hooks or approval dialogs
|
||||
async fn execute_lightweight_with_tools(
|
||||
ctx: &EngineContext,
|
||||
routine: &Routine,
|
||||
system_prompt: &str,
|
||||
full_prompt: &str,
|
||||
effective_max_tokens: u32,
|
||||
) -> Result<(RunStatus, Option<String>, Option<i32>), RoutineError> {
|
||||
let mut messages = if system_prompt.is_empty() {
|
||||
vec![ChatMessage::user(full_prompt)]
|
||||
} else {
|
||||
vec![
|
||||
ChatMessage::system(system_prompt),
|
||||
ChatMessage::user(full_prompt),
|
||||
]
|
||||
};
|
||||
|
||||
let max_iterations = ctx.config.lightweight_max_iterations.min(5);
|
||||
let mut iteration = 0;
|
||||
let mut total_input_tokens = 0;
|
||||
let mut total_output_tokens = 0;
|
||||
|
||||
// Create a minimal job context for tool execution with unique run ID
|
||||
let run_id = Uuid::new_v4();
|
||||
let job_ctx = JobContext {
|
||||
job_id: run_id,
|
||||
user_id: routine.user_id.clone(),
|
||||
title: "Lightweight Routine".to_string(),
|
||||
description: routine.name.clone(),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
loop {
|
||||
iteration += 1;
|
||||
|
||||
// Force text-only response at iteration limit
|
||||
let force_text = iteration >= max_iterations;
|
||||
|
||||
if force_text {
|
||||
// Final iteration: no tools, just get text response
|
||||
let request = CompletionRequest::new(messages)
|
||||
.with_max_tokens(effective_max_tokens)
|
||||
.with_temperature(0.3);
|
||||
|
||||
let response =
|
||||
ctx.llm
|
||||
.complete(request)
|
||||
.await
|
||||
.map_err(|e| RoutineError::LlmFailed {
|
||||
reason: e.to_string(),
|
||||
})?;
|
||||
|
||||
total_input_tokens += response.input_tokens;
|
||||
total_output_tokens += response.output_tokens;
|
||||
|
||||
return handle_text_response(
|
||||
&response.content,
|
||||
response.finish_reason,
|
||||
total_input_tokens,
|
||||
total_output_tokens,
|
||||
);
|
||||
} else {
|
||||
// Tool-enabled iteration
|
||||
let tool_defs = ctx.tools.tool_definitions().await;
|
||||
|
||||
let request = ToolCompletionRequest::new(messages.clone(), tool_defs)
|
||||
.with_max_tokens(effective_max_tokens)
|
||||
.with_temperature(0.3);
|
||||
|
||||
let response = ctx.llm.complete_with_tools(request).await.map_err(|e| {
|
||||
RoutineError::LlmFailed {
|
||||
reason: e.to_string(),
|
||||
}
|
||||
})?;
|
||||
|
||||
total_input_tokens += response.input_tokens;
|
||||
total_output_tokens += response.output_tokens;
|
||||
|
||||
// Check if LLM returned text (no tool calls)
|
||||
if response.tool_calls.is_empty() {
|
||||
let content = response.content.unwrap_or_default();
|
||||
return handle_text_response(
|
||||
&content,
|
||||
response.finish_reason,
|
||||
total_input_tokens,
|
||||
total_output_tokens,
|
||||
);
|
||||
}
|
||||
|
||||
// LLM returned tool calls: add assistant message and execute tools
|
||||
messages.push(ChatMessage::assistant_with_tool_calls(
|
||||
response.content.clone(),
|
||||
response.tool_calls.clone(),
|
||||
));
|
||||
|
||||
// Execute tools sequentially
|
||||
for tc in response.tool_calls {
|
||||
let result = execute_routine_tool(ctx, &job_ctx, &tc).await;
|
||||
|
||||
// Sanitize and wrap result (including errors)
|
||||
let result_content = match result {
|
||||
Ok(output) => {
|
||||
let sanitized = ctx.safety.sanitize_tool_output(&tc.name, &output);
|
||||
ctx.safety.wrap_for_llm(
|
||||
&tc.name,
|
||||
&sanitized.content,
|
||||
sanitized.was_modified,
|
||||
)
|
||||
}
|
||||
Err(e) => {
|
||||
let error_msg = format!("Tool '{}' failed: {}", tc.name, e);
|
||||
let sanitized = ctx.safety.sanitize_tool_output(&tc.name, &error_msg);
|
||||
ctx.safety.wrap_for_llm(
|
||||
&tc.name,
|
||||
&sanitized.content,
|
||||
sanitized.was_modified,
|
||||
)
|
||||
}
|
||||
};
|
||||
|
||||
// Add tool result to context
|
||||
messages.push(ChatMessage::tool_result(&tc.id, &tc.name, &result_content));
|
||||
}
|
||||
|
||||
// Continue loop to next LLM call
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Execute a single tool for a lightweight routine.
|
||||
async fn execute_routine_tool(
|
||||
ctx: &EngineContext,
|
||||
job_ctx: &JobContext,
|
||||
tc: &ToolCall,
|
||||
) -> Result<String, Box<dyn std::error::Error + Send + Sync>> {
|
||||
// Check if tool exists
|
||||
let tool = ctx
|
||||
.tools
|
||||
.get(&tc.name)
|
||||
.await
|
||||
.ok_or_else(|| format!("Tool '{}' not found", tc.name))?;
|
||||
|
||||
// Check approval requirement: only allow Never tools in lightweight routines.
|
||||
// UnlessAutoApproved and Always tools are blocked to prevent prompt injection attacks.
|
||||
// Lightweight routines can be triggered by external events and may process untrusted data,
|
||||
// making them vulnerable to prompt injection that could trick the LLM into calling
|
||||
// sensitive tools. Blocking these tools entirely is the safest approach.
|
||||
match tool.requires_approval(&tc.arguments) {
|
||||
ApprovalRequirement::Never => {}
|
||||
ApprovalRequirement::UnlessAutoApproved | ApprovalRequirement::Always => {
|
||||
return Err(format!(
|
||||
"Tool '{}' requires manual approval and cannot be used in lightweight routines",
|
||||
tc.name
|
||||
)
|
||||
.into());
|
||||
}
|
||||
}
|
||||
|
||||
// Validate tool parameters
|
||||
let validation = ctx.safety.validator().validate_tool_params(&tc.arguments);
|
||||
if !validation.is_valid {
|
||||
let details = validation
|
||||
.errors
|
||||
.iter()
|
||||
.map(|e| format!("{}: {}", e.field, e.message))
|
||||
.collect::<Vec<_>>()
|
||||
.join("; ");
|
||||
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
|
||||
let timeout = tool.execution_timeout();
|
||||
let start = std::time::Instant::now();
|
||||
let result = tokio::time::timeout(timeout, async {
|
||||
tool.execute(tc.arguments.clone(), job_ctx).await
|
||||
})
|
||||
.await;
|
||||
let elapsed = start.elapsed();
|
||||
|
||||
match &result {
|
||||
Ok(Ok(_)) => {
|
||||
tracing::debug!(
|
||||
tool = %tc.name,
|
||||
elapsed_ms = elapsed.as_millis() as u64,
|
||||
"Lightweight routine tool call succeeded"
|
||||
);
|
||||
}
|
||||
Ok(Err(e)) => {
|
||||
tracing::debug!(
|
||||
tool = %tc.name,
|
||||
elapsed_ms = elapsed.as_millis() as u64,
|
||||
error = %e,
|
||||
"Lightweight routine tool call failed"
|
||||
);
|
||||
}
|
||||
Err(_) => {
|
||||
tracing::debug!(
|
||||
tool = %tc.name,
|
||||
elapsed_ms = elapsed.as_millis() as u64,
|
||||
timeout_secs = timeout.as_secs(),
|
||||
"Lightweight routine tool call timed out"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
let result = result
|
||||
.map_err(|_| ToolError::Timeout(timeout))
|
||||
.map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)?
|
||||
.map_err(|e| Box::new(e) as Box<dyn std::error::Error + Send + Sync>)?;
|
||||
|
||||
// Serialize result to JSON string
|
||||
let result_str =
|
||||
serde_json::to_string(&result.result).unwrap_or_else(|_| "<serialize error>".to_string());
|
||||
Ok(result_str)
|
||||
}
|
||||
|
||||
/// Send a notification based on the routine's notify config and run status.
|
||||
async fn send_notification(
|
||||
tx: &mpsc::Sender<OutgoingResponse>,
|
||||
@@ -573,6 +973,7 @@ async fn send_notification(
|
||||
routine_name: &str,
|
||||
status: RunStatus,
|
||||
summary: Option<&str>,
|
||||
thread_id: Option<&str>,
|
||||
) {
|
||||
let should_notify = match status {
|
||||
RunStatus::Ok => notify.on_success,
|
||||
@@ -599,7 +1000,7 @@ async fn send_notification(
|
||||
|
||||
let response = OutgoingResponse {
|
||||
content: message,
|
||||
thread_id: None,
|
||||
thread_id: thread_id.map(String::from),
|
||||
attachments: Vec::new(),
|
||||
metadata: serde_json::json!({
|
||||
"source": "routine",
|
||||
@@ -644,6 +1045,7 @@ fn truncate(s: &str, max: usize) -> String {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::agent::routine::{NotifyConfig, RunStatus};
|
||||
use crate::config::RoutineConfig;
|
||||
|
||||
#[test]
|
||||
fn test_notification_gating() {
|
||||
@@ -672,4 +1074,130 @@ mod tests {
|
||||
let _ = status.to_string();
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_routine_config_lightweight_tools_enabled_default() {
|
||||
let config = RoutineConfig::default();
|
||||
assert!(
|
||||
config.lightweight_tools_enabled,
|
||||
"Tools should be enabled by default"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_routine_config_lightweight_max_iterations_default() {
|
||||
let config = RoutineConfig::default();
|
||||
assert_eq!(
|
||||
config.lightweight_max_iterations, 3,
|
||||
"Default should be 3 iterations"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_routine_config_can_hold_uncapped_max_iterations() {
|
||||
// The `RoutineConfig` struct can hold a value greater than the safety cap.
|
||||
let config = RoutineConfig {
|
||||
lightweight_max_iterations: 10, // Set a value higher than the cap.
|
||||
..RoutineConfig::default()
|
||||
};
|
||||
// The actual capping to a maximum of 5 is handled at runtime in
|
||||
// `execute_lightweight_with_tools` and during config resolution from env vars.
|
||||
assert_eq!(
|
||||
config.lightweight_max_iterations, 10,
|
||||
"Config struct should store the provided value"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_sanitize_routine_name_replaces_special_chars() {
|
||||
let test_cases = vec![
|
||||
("valid-routine", "valid-routine"),
|
||||
("routine_with_underscore", "routine_with_underscore"),
|
||||
("Routine With Spaces", "Routine_With_Spaces"),
|
||||
("routine/with/slashes", "routine_with_slashes"),
|
||||
("routine@with#symbols", "routine_with_symbols"),
|
||||
];
|
||||
|
||||
for (input, expected) in test_cases {
|
||||
let result = super::sanitize_routine_name(input);
|
||||
assert_eq!(
|
||||
result, expected,
|
||||
"sanitize_routine_name({}) should be {}",
|
||||
input, expected
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_sanitize_routine_name_preserves_alphanumeric_dash_underscore() {
|
||||
let names = vec!["routine123", "routine-name", "routine_name", "ROUTINE"];
|
||||
for name in names {
|
||||
let result = super::sanitize_routine_name(name);
|
||||
assert_eq!(result, name, "Should preserve {}", name);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_routine_sentinel_detection_exact_match() {
|
||||
// The execute_lightweight_no_tools checks: content == "ROUTINE_OK" || content.contains("ROUTINE_OK")
|
||||
// After trim(), whitespace is removed
|
||||
let test_cases = vec![
|
||||
("ROUTINE_OK", true),
|
||||
(" ROUTINE_OK ", true), // After trim, whitespace is removed so matches
|
||||
("something ROUTINE_OK something", true),
|
||||
("ROUTINE_OK is done", true),
|
||||
("done ROUTINE_OK", true),
|
||||
("no sentinel here", false),
|
||||
];
|
||||
|
||||
for (content, should_match) in test_cases {
|
||||
let trimmed = content.trim();
|
||||
let matches = trimmed == "ROUTINE_OK" || trimmed.contains("ROUTINE_OK");
|
||||
assert_eq!(
|
||||
matches, should_match,
|
||||
"Content '{}' sentinel detection should be {}, got {}",
|
||||
content, should_match, matches
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_approval_requirement_pattern_matching() {
|
||||
// Test the approval requirement logic (Never, UnlessAutoApproved, Always)
|
||||
use crate::tools::ApprovalRequirement;
|
||||
|
||||
let requirements = vec![
|
||||
(ApprovalRequirement::Never, "auto-approved"),
|
||||
(ApprovalRequirement::UnlessAutoApproved, "auto-approved"),
|
||||
(ApprovalRequirement::Always, "blocks"),
|
||||
];
|
||||
|
||||
for (req, expected) in requirements {
|
||||
let can_auto_approve = matches!(
|
||||
req,
|
||||
ApprovalRequirement::Never | ApprovalRequirement::UnlessAutoApproved
|
||||
);
|
||||
let label = if can_auto_approve {
|
||||
"auto-approved"
|
||||
} else {
|
||||
"blocks"
|
||||
};
|
||||
assert_eq!(label, expected, "Approval pattern should match");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_empty_response_handling() {
|
||||
// Simulate the empty content guard logic
|
||||
let empty_content = "";
|
||||
let finish_reason_length = crate::llm::FinishReason::Length;
|
||||
let finish_reason_stop = crate::llm::FinishReason::Stop;
|
||||
|
||||
assert!(
|
||||
empty_content.trim().is_empty(),
|
||||
"Should detect empty content"
|
||||
);
|
||||
assert_eq!(finish_reason_length, crate::llm::FinishReason::Length);
|
||||
assert_eq!(finish_reason_stop, crate::llm::FinishReason::Stop);
|
||||
}
|
||||
}
|
||||
|
||||
+286
-3
@@ -18,7 +18,7 @@ use crate::error::{Error, JobError};
|
||||
use crate::hooks::HookRegistry;
|
||||
use crate::llm::LlmProvider;
|
||||
use crate::safety::SafetyLayer;
|
||||
use crate::tools::ToolRegistry;
|
||||
use crate::tools::{ApprovalContext, ToolRegistry};
|
||||
|
||||
/// Message to send to a worker.
|
||||
#[derive(Debug)]
|
||||
@@ -56,6 +56,8 @@ pub struct Scheduler {
|
||||
hooks: Arc<HookRegistry>,
|
||||
/// SSE broadcast sender for live job event streaming.
|
||||
sse_tx: Option<tokio::sync::broadcast::Sender<SseEvent>>,
|
||||
/// HTTP interceptor for trace recording/replay (propagated to workers).
|
||||
http_interceptor: Option<Arc<dyn crate::llm::recording::HttpInterceptor>>,
|
||||
/// Running jobs (main LLM-driven jobs).
|
||||
jobs: Arc<RwLock<HashMap<Uuid, ScheduledJob>>>,
|
||||
/// Running sub-tasks (tool executions, background tasks).
|
||||
@@ -82,6 +84,7 @@ impl Scheduler {
|
||||
store,
|
||||
hooks,
|
||||
sse_tx: None,
|
||||
http_interceptor: None,
|
||||
jobs: Arc::new(RwLock::new(HashMap::new())),
|
||||
subtasks: Arc::new(RwLock::new(HashMap::new())),
|
||||
}
|
||||
@@ -92,6 +95,14 @@ impl Scheduler {
|
||||
self.sse_tx = Some(tx);
|
||||
}
|
||||
|
||||
/// Set the HTTP interceptor for trace recording/replay.
|
||||
pub fn set_http_interceptor(
|
||||
&mut self,
|
||||
interceptor: Arc<dyn crate::llm::recording::HttpInterceptor>,
|
||||
) {
|
||||
self.http_interceptor = Some(interceptor);
|
||||
}
|
||||
|
||||
/// Create, persist, and schedule a job in one shot.
|
||||
///
|
||||
/// This is the preferred entry point for dispatching new jobs. It:
|
||||
@@ -108,12 +119,54 @@ impl Scheduler {
|
||||
title: &str,
|
||||
description: &str,
|
||||
metadata: Option<serde_json::Value>,
|
||||
) -> Result<Uuid, JobError> {
|
||||
self.dispatch_job_inner(user_id, title, description, metadata, None)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Dispatch a job with an explicit approval context for autonomous execution.
|
||||
///
|
||||
/// Same as `dispatch_job`, but the worker will use the given `ApprovalContext`
|
||||
/// to determine which tools are pre-approved (instead of blocking all non-`Never` tools).
|
||||
pub async fn dispatch_job_with_context(
|
||||
&self,
|
||||
user_id: &str,
|
||||
title: &str,
|
||||
description: &str,
|
||||
metadata: Option<serde_json::Value>,
|
||||
approval_context: ApprovalContext,
|
||||
) -> Result<Uuid, JobError> {
|
||||
self.dispatch_job_inner(
|
||||
user_id,
|
||||
title,
|
||||
description,
|
||||
metadata,
|
||||
Some(approval_context),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Shared implementation for `dispatch_job` and `dispatch_job_with_context`.
|
||||
async fn dispatch_job_inner(
|
||||
&self,
|
||||
user_id: &str,
|
||||
title: &str,
|
||||
description: &str,
|
||||
metadata: Option<serde_json::Value>,
|
||||
approval_context: Option<ApprovalContext>,
|
||||
) -> Result<Uuid, JobError> {
|
||||
let job_id = self
|
||||
.context_manager
|
||||
.create_job_for_user(user_id, title, description)
|
||||
.await?;
|
||||
|
||||
// Apply token budget from config, allowing per-job metadata override.
|
||||
let max_tokens = metadata
|
||||
.as_ref()
|
||||
.and_then(|m| m.get("max_tokens"))
|
||||
.and_then(|v| v.as_u64())
|
||||
.unwrap_or(self.config.max_tokens_per_job);
|
||||
|
||||
// Apply metadata if provided
|
||||
if let Some(meta) = metadata {
|
||||
self.context_manager
|
||||
@@ -123,6 +176,15 @@ impl Scheduler {
|
||||
.await?;
|
||||
}
|
||||
|
||||
// Set token budget (separate update to avoid overwriting metadata)
|
||||
if max_tokens > 0 {
|
||||
self.context_manager
|
||||
.update_context(job_id, |ctx| {
|
||||
ctx.max_tokens = max_tokens;
|
||||
})
|
||||
.await?;
|
||||
}
|
||||
|
||||
// Persist to DB before scheduling so the worker's FK references are valid
|
||||
if let Some(ref store) = self.store {
|
||||
let ctx = self.context_manager.get_context(job_id).await?;
|
||||
@@ -132,12 +194,21 @@ impl Scheduler {
|
||||
})?;
|
||||
}
|
||||
|
||||
self.schedule(job_id).await?;
|
||||
self.schedule_with_context(job_id, approval_context).await?;
|
||||
Ok(job_id)
|
||||
}
|
||||
|
||||
/// Schedule a job for execution.
|
||||
pub async fn schedule(&self, job_id: Uuid) -> Result<(), JobError> {
|
||||
self.schedule_with_context(job_id, None).await
|
||||
}
|
||||
|
||||
/// Schedule a job with an optional approval context.
|
||||
async fn schedule_with_context(
|
||||
&self,
|
||||
job_id: Uuid,
|
||||
approval_context: Option<ApprovalContext>,
|
||||
) -> Result<(), JobError> {
|
||||
// Hold write lock for the entire check-insert sequence to prevent
|
||||
// TOCTOU races where two concurrent calls both pass the checks.
|
||||
{
|
||||
@@ -181,6 +252,8 @@ impl Scheduler {
|
||||
timeout: self.config.job_timeout,
|
||||
use_planning: self.config.use_planning,
|
||||
sse_tx: self.sse_tx.clone(),
|
||||
approval_context,
|
||||
http_interceptor: self.http_interceptor.clone(),
|
||||
};
|
||||
let worker = Worker::new(job_id, deps);
|
||||
|
||||
@@ -257,11 +330,14 @@ impl Scheduler {
|
||||
let context_manager = self.context_manager.clone();
|
||||
let safety = self.safety.clone();
|
||||
|
||||
// TODO: propagate parent job's ApprovalContext here when subtasks
|
||||
// are used in autonomous/routine paths (currently only used in tests).
|
||||
tokio::spawn(async move {
|
||||
let result = Self::execute_tool_task(
|
||||
tools,
|
||||
context_manager,
|
||||
safety,
|
||||
None,
|
||||
tool_parent_id,
|
||||
&tool_name,
|
||||
params,
|
||||
@@ -390,6 +466,7 @@ impl Scheduler {
|
||||
tools: Arc<ToolRegistry>,
|
||||
context_manager: Arc<ContextManager>,
|
||||
safety: Arc<SafetyLayer>,
|
||||
approval_context: Option<ApprovalContext>,
|
||||
job_id: Uuid,
|
||||
tool_name: &str,
|
||||
params: serde_json::Value,
|
||||
@@ -413,7 +490,10 @@ impl Scheduler {
|
||||
.into());
|
||||
}
|
||||
|
||||
if tool.requires_approval(¶ms).is_required() {
|
||||
let requirement = tool.requires_approval(¶ms);
|
||||
let blocked =
|
||||
ApprovalContext::is_blocked_or_default(&approval_context, tool_name, requirement);
|
||||
if blocked {
|
||||
return Err(crate::error::ToolError::AuthRequired {
|
||||
name: tool_name.to_string(),
|
||||
}
|
||||
@@ -617,6 +697,11 @@ impl Scheduler {
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::config::SafetyConfig;
|
||||
use crate::safety::SafetyLayer;
|
||||
use crate::tools::{ApprovalRequirement, Tool, ToolError, ToolOutput};
|
||||
|
||||
#[test]
|
||||
fn test_scheduler_creation() {
|
||||
// Would need to mock dependencies for proper testing
|
||||
@@ -627,4 +712,202 @@ mod tests {
|
||||
// This test would need mock dependencies.
|
||||
// For now just verify the empty case doesn't panic.
|
||||
}
|
||||
|
||||
/// A tool that returns `UnlessAutoApproved`.
|
||||
struct SoftApprovalTool;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl Tool for SoftApprovalTool {
|
||||
fn name(&self) -> &str {
|
||||
"soft_gate"
|
||||
}
|
||||
fn description(&self) -> &str {
|
||||
"needs soft approval"
|
||||
}
|
||||
fn parameters_schema(&self) -> serde_json::Value {
|
||||
serde_json::json!({"type": "object", "properties": {}})
|
||||
}
|
||||
async fn execute(
|
||||
&self,
|
||||
_params: serde_json::Value,
|
||||
_ctx: &JobContext,
|
||||
) -> Result<ToolOutput, ToolError> {
|
||||
Ok(ToolOutput::text(
|
||||
"soft_ok",
|
||||
std::time::Instant::now().elapsed(),
|
||||
))
|
||||
}
|
||||
fn requires_approval(&self, _params: &serde_json::Value) -> ApprovalRequirement {
|
||||
ApprovalRequirement::UnlessAutoApproved
|
||||
}
|
||||
fn requires_sanitization(&self) -> bool {
|
||||
false
|
||||
}
|
||||
}
|
||||
|
||||
/// A tool that returns `Always`.
|
||||
struct HardApprovalTool;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl Tool for HardApprovalTool {
|
||||
fn name(&self) -> &str {
|
||||
"hard_gate"
|
||||
}
|
||||
fn description(&self) -> &str {
|
||||
"needs hard approval"
|
||||
}
|
||||
fn parameters_schema(&self) -> serde_json::Value {
|
||||
serde_json::json!({"type": "object", "properties": {}})
|
||||
}
|
||||
async fn execute(
|
||||
&self,
|
||||
_params: serde_json::Value,
|
||||
_ctx: &JobContext,
|
||||
) -> Result<ToolOutput, ToolError> {
|
||||
Ok(ToolOutput::text(
|
||||
"hard_ok",
|
||||
std::time::Instant::now().elapsed(),
|
||||
))
|
||||
}
|
||||
fn requires_approval(&self, _params: &serde_json::Value) -> ApprovalRequirement {
|
||||
ApprovalRequirement::Always
|
||||
}
|
||||
fn requires_sanitization(&self) -> bool {
|
||||
false
|
||||
}
|
||||
}
|
||||
|
||||
async fn setup_tools_and_job() -> (
|
||||
Arc<ToolRegistry>,
|
||||
Arc<ContextManager>,
|
||||
Arc<SafetyLayer>,
|
||||
Uuid,
|
||||
) {
|
||||
let registry = ToolRegistry::new();
|
||||
registry.register(Arc::new(SoftApprovalTool)).await;
|
||||
registry.register(Arc::new(HardApprovalTool)).await;
|
||||
|
||||
let cm = Arc::new(ContextManager::new(5));
|
||||
let job_id = cm.create_job("test", "approval test").await.unwrap();
|
||||
cm.update_context(job_id, |ctx| ctx.transition_to(JobState::InProgress, None))
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
|
||||
let safety = Arc::new(SafetyLayer::new(&SafetyConfig {
|
||||
max_output_length: 100_000,
|
||||
injection_check_enabled: false,
|
||||
}));
|
||||
|
||||
(Arc::new(registry), cm, safety, job_id)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_execute_tool_task_blocks_without_context() {
|
||||
let (tools, cm, safety, job_id) = setup_tools_and_job().await;
|
||||
|
||||
// Without approval context, UnlessAutoApproved is blocked
|
||||
let result = Scheduler::execute_tool_task(
|
||||
tools.clone(),
|
||||
cm.clone(),
|
||||
safety.clone(),
|
||||
None,
|
||||
job_id,
|
||||
"soft_gate",
|
||||
serde_json::json!({}),
|
||||
)
|
||||
.await;
|
||||
assert!(
|
||||
result.is_err(),
|
||||
"soft_gate should be blocked without context"
|
||||
);
|
||||
|
||||
// Always is also blocked
|
||||
let result = Scheduler::execute_tool_task(
|
||||
tools,
|
||||
cm,
|
||||
safety,
|
||||
None,
|
||||
job_id,
|
||||
"hard_gate",
|
||||
serde_json::json!({}),
|
||||
)
|
||||
.await;
|
||||
assert!(
|
||||
result.is_err(),
|
||||
"hard_gate should be blocked without context"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_execute_tool_task_autonomous_unblocks_soft() {
|
||||
let (tools, cm, safety, job_id) = setup_tools_and_job().await;
|
||||
|
||||
// Autonomous context auto-approves UnlessAutoApproved
|
||||
let result = Scheduler::execute_tool_task(
|
||||
tools.clone(),
|
||||
cm.clone(),
|
||||
safety.clone(),
|
||||
Some(ApprovalContext::autonomous()),
|
||||
job_id,
|
||||
"soft_gate",
|
||||
serde_json::json!({}),
|
||||
)
|
||||
.await;
|
||||
assert!(
|
||||
result.is_ok(),
|
||||
"soft_gate should pass with autonomous context"
|
||||
);
|
||||
|
||||
// But still blocks Always
|
||||
let result = Scheduler::execute_tool_task(
|
||||
tools,
|
||||
cm,
|
||||
safety,
|
||||
Some(ApprovalContext::autonomous()),
|
||||
job_id,
|
||||
"hard_gate",
|
||||
serde_json::json!({}),
|
||||
)
|
||||
.await;
|
||||
assert!(
|
||||
result.is_err(),
|
||||
"hard_gate should still be blocked without explicit permission"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_execute_tool_task_autonomous_with_permissions() {
|
||||
let (tools, cm, safety, job_id) = setup_tools_and_job().await;
|
||||
|
||||
// Autonomous context with explicit permission for hard_gate
|
||||
let ctx = ApprovalContext::autonomous_with_tools(["hard_gate".to_string()]);
|
||||
|
||||
let result = Scheduler::execute_tool_task(
|
||||
tools.clone(),
|
||||
cm.clone(),
|
||||
safety.clone(),
|
||||
Some(ctx.clone()),
|
||||
job_id,
|
||||
"soft_gate",
|
||||
serde_json::json!({}),
|
||||
)
|
||||
.await;
|
||||
assert!(result.is_ok(), "soft_gate should pass");
|
||||
|
||||
let result = Scheduler::execute_tool_task(
|
||||
tools,
|
||||
cm,
|
||||
safety,
|
||||
Some(ctx),
|
||||
job_id,
|
||||
"hard_gate",
|
||||
serde_json::json!({}),
|
||||
)
|
||||
.await;
|
||||
assert!(
|
||||
result.is_ok(),
|
||||
"hard_gate should pass with explicit permission"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
+342
-12
@@ -16,6 +16,7 @@ use chrono::{DateTime, Utc};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::channels::web::util::truncate_preview;
|
||||
use crate::llm::{ChatMessage, ToolCall};
|
||||
|
||||
/// A session containing one or more threads.
|
||||
@@ -164,6 +165,10 @@ pub struct PendingApproval {
|
||||
/// executed yet when approval was requested.
|
||||
#[serde(default)]
|
||||
pub deferred_tool_calls: Vec<ToolCall>,
|
||||
/// User timezone at the time the approval was requested, so it persists
|
||||
/// through the approval flow even if the approval message lacks timezone.
|
||||
#[serde(default)]
|
||||
pub user_timezone: Option<String>,
|
||||
}
|
||||
|
||||
/// A conversation thread within a session.
|
||||
@@ -316,11 +321,60 @@ impl Thread {
|
||||
}
|
||||
}
|
||||
|
||||
/// Get all messages for context building.
|
||||
/// Get all messages for context building, including tool call history.
|
||||
///
|
||||
/// Emits the full LLM-compatible message sequence per turn:
|
||||
/// `user → [assistant_with_tool_calls → tool_result*] → assistant`
|
||||
///
|
||||
/// This ensures the LLM sees prior tool executions and won't re-attempt
|
||||
/// completed actions in subsequent turns.
|
||||
pub fn messages(&self) -> Vec<ChatMessage> {
|
||||
let mut messages = Vec::new();
|
||||
for turn in &self.turns {
|
||||
messages.push(ChatMessage::user(&turn.user_input));
|
||||
if turn.image_content_parts.is_empty() {
|
||||
messages.push(ChatMessage::user(&turn.user_input));
|
||||
} else {
|
||||
messages.push(ChatMessage::user_with_parts(
|
||||
&turn.user_input,
|
||||
turn.image_content_parts.clone(),
|
||||
));
|
||||
}
|
||||
|
||||
if !turn.tool_calls.is_empty() {
|
||||
// Build ToolCall objects with synthetic stable IDs
|
||||
let tool_calls: Vec<ToolCall> = turn
|
||||
.tool_calls
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(i, tc)| ToolCall {
|
||||
id: format!("turn{}_{}", turn.turn_number, i),
|
||||
name: tc.name.clone(),
|
||||
arguments: tc.parameters.clone(),
|
||||
})
|
||||
.collect();
|
||||
|
||||
// Assistant message declaring the tool calls (no text content)
|
||||
messages.push(ChatMessage::assistant_with_tool_calls(None, tool_calls));
|
||||
|
||||
// Individual tool result messages, truncated to limit context size.
|
||||
for (i, tc) in turn.tool_calls.iter().enumerate() {
|
||||
let call_id = format!("turn{}_{}", turn.turn_number, i);
|
||||
let content = if let Some(ref err) = tc.error {
|
||||
// .error already contains the full error text;
|
||||
// pass through without wrapping to avoid double-prefix.
|
||||
truncate_preview(err, 1000)
|
||||
} else if let Some(ref res) = tc.result {
|
||||
let raw = match res {
|
||||
serde_json::Value::String(s) => s.clone(),
|
||||
other => other.to_string(),
|
||||
};
|
||||
truncate_preview(&raw, 1000)
|
||||
} else {
|
||||
"OK".to_string()
|
||||
};
|
||||
messages.push(ChatMessage::tool_result(call_id, &tc.name, content));
|
||||
}
|
||||
}
|
||||
if let Some(ref response) = turn.response {
|
||||
messages.push(ChatMessage::assistant(response));
|
||||
}
|
||||
@@ -342,13 +396,16 @@ impl Thread {
|
||||
|
||||
/// Restore thread state from a checkpoint's messages.
|
||||
///
|
||||
/// Clears existing turns and rebuilds from message pairs.
|
||||
/// Messages should alternate: user, assistant, user, assistant...
|
||||
/// Clears existing turns and rebuilds from the message sequence.
|
||||
/// Handles the full message pattern including tool messages:
|
||||
/// `user → [assistant_with_tool_calls → tool_result*] → assistant`
|
||||
///
|
||||
/// Also supports the legacy pattern (user/assistant pairs only) for
|
||||
/// backward compatibility with old checkpoint data.
|
||||
pub fn restore_from_messages(&mut self, messages: Vec<ChatMessage>) {
|
||||
self.turns.clear();
|
||||
self.state = ThreadState::Idle;
|
||||
|
||||
// Messages alternate: user, assistant, user, assistant...
|
||||
let mut iter = messages.into_iter().peekable();
|
||||
let mut turn_number = 0;
|
||||
|
||||
@@ -356,18 +413,58 @@ impl Thread {
|
||||
if msg.role == crate::llm::Role::User {
|
||||
let mut turn = Turn::new(turn_number, &msg.content);
|
||||
|
||||
// Check if next is assistant response
|
||||
if let Some(next) = iter.peek()
|
||||
&& next.role == crate::llm::Role::Assistant
|
||||
{
|
||||
// iter.next() is guaranteed Some after a successful peek()
|
||||
if let Some(response) = iter.next() {
|
||||
turn.complete(&response.content);
|
||||
// Consume tool call sequences (assistant_with_tool_calls + tool_results).
|
||||
// A single turn may contain multiple rounds of tool calls, so we
|
||||
// track the cumulative base index into turn.tool_calls.
|
||||
while let Some(next) = iter.peek() {
|
||||
if next.role == crate::llm::Role::Assistant && next.tool_calls.is_some() {
|
||||
let call_base_idx = turn.tool_calls.len();
|
||||
|
||||
if let Some(assistant_msg) = iter.next()
|
||||
&& let Some(ref tcs) = assistant_msg.tool_calls
|
||||
{
|
||||
for tc in tcs {
|
||||
turn.record_tool_call(&tc.name, tc.arguments.clone());
|
||||
}
|
||||
}
|
||||
|
||||
// Consume the corresponding tool_result messages,
|
||||
// indexing relative to this batch's base offset.
|
||||
let mut pos = 0;
|
||||
while let Some(tr) = iter.peek() {
|
||||
if tr.role != crate::llm::Role::Tool {
|
||||
break;
|
||||
}
|
||||
if let Some(tool_msg) = iter.next() {
|
||||
let idx = call_base_idx + pos;
|
||||
if idx < turn.tool_calls.len() {
|
||||
// Store as result — the error/success distinction
|
||||
// is for the live turn only; restored context just
|
||||
// needs the content the LLM originally saw.
|
||||
turn.tool_calls[idx].result =
|
||||
Some(serde_json::Value::String(tool_msg.content.clone()));
|
||||
}
|
||||
}
|
||||
pos += 1;
|
||||
}
|
||||
} else {
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
// Check if next is the final assistant response for this turn
|
||||
let is_final_assistant = iter.peek().is_some_and(|n| {
|
||||
n.role == crate::llm::Role::Assistant && n.tool_calls.is_none()
|
||||
});
|
||||
if is_final_assistant && let Some(response) = iter.next() {
|
||||
turn.complete(&response.content);
|
||||
}
|
||||
|
||||
self.turns.push(turn);
|
||||
turn_number += 1;
|
||||
} else {
|
||||
// Skip non-user messages that aren't anchored to a turn
|
||||
continue;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -407,6 +504,11 @@ pub struct Turn {
|
||||
pub completed_at: Option<DateTime<Utc>>,
|
||||
/// Error message (if failed).
|
||||
pub error: Option<String>,
|
||||
/// Transient image content parts for multimodal LLM input.
|
||||
/// Not serialized — images are only needed for the current LLM call.
|
||||
/// The text description in `user_input` persists for compaction/context.
|
||||
#[serde(skip)]
|
||||
pub image_content_parts: Vec<crate::llm::ContentPart>,
|
||||
}
|
||||
|
||||
impl Turn {
|
||||
@@ -421,6 +523,7 @@ impl Turn {
|
||||
started_at: Utc::now(),
|
||||
completed_at: None,
|
||||
error: None,
|
||||
image_content_parts: Vec::new(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -429,6 +532,8 @@ impl Turn {
|
||||
self.response = Some(response.into());
|
||||
self.state = TurnState::Completed;
|
||||
self.completed_at = Some(Utc::now());
|
||||
// Free image data — only needed for the initial LLM call, not subsequent turns
|
||||
self.image_content_parts.clear();
|
||||
}
|
||||
|
||||
/// Fail this turn.
|
||||
@@ -436,12 +541,14 @@ impl Turn {
|
||||
self.error = Some(error.into());
|
||||
self.state = TurnState::Failed;
|
||||
self.completed_at = Some(Utc::now());
|
||||
self.image_content_parts.clear();
|
||||
}
|
||||
|
||||
/// Interrupt this turn.
|
||||
pub fn interrupt(&mut self) {
|
||||
self.state = TurnState::Interrupted;
|
||||
self.completed_at = Some(Utc::now());
|
||||
self.image_content_parts.clear();
|
||||
}
|
||||
|
||||
/// Record a tool call.
|
||||
@@ -959,6 +1066,7 @@ mod tests {
|
||||
tool_call_id: "call_123".to_string(),
|
||||
context_messages: vec![ChatMessage::user("do it")],
|
||||
deferred_tool_calls: vec![],
|
||||
user_timezone: None,
|
||||
};
|
||||
|
||||
thread.await_approval(approval);
|
||||
@@ -984,6 +1092,7 @@ mod tests {
|
||||
tool_call_id: "call_456".to_string(),
|
||||
context_messages: vec![],
|
||||
deferred_tool_calls: vec![],
|
||||
user_timezone: None,
|
||||
};
|
||||
|
||||
thread.await_approval(approval);
|
||||
@@ -1012,4 +1121,225 @@ mod tests {
|
||||
ThreadState::Processing
|
||||
);
|
||||
}
|
||||
|
||||
// Regression tests for #568: tool call history must survive hydration.
|
||||
|
||||
#[test]
|
||||
fn test_messages_includes_tool_calls() {
|
||||
let mut thread = Thread::new(Uuid::new_v4());
|
||||
|
||||
thread.start_turn("Search for X");
|
||||
{
|
||||
let turn = thread.turns.last_mut().unwrap();
|
||||
turn.record_tool_call("memory_search", serde_json::json!({"query": "X"}));
|
||||
turn.record_tool_result(serde_json::json!("Found X in doc.md"));
|
||||
}
|
||||
thread.complete_turn("I found X in doc.md.");
|
||||
|
||||
let messages = thread.messages();
|
||||
// user + assistant_with_tool_calls + tool_result + assistant = 4
|
||||
assert_eq!(messages.len(), 4);
|
||||
|
||||
assert_eq!(messages[0].role, crate::llm::Role::User);
|
||||
assert_eq!(messages[0].content, "Search for X");
|
||||
|
||||
assert_eq!(messages[1].role, crate::llm::Role::Assistant);
|
||||
assert!(messages[1].tool_calls.is_some());
|
||||
let tcs = messages[1].tool_calls.as_ref().unwrap();
|
||||
assert_eq!(tcs.len(), 1);
|
||||
assert_eq!(tcs[0].name, "memory_search");
|
||||
|
||||
assert_eq!(messages[2].role, crate::llm::Role::Tool);
|
||||
assert!(messages[2].content.contains("Found X"));
|
||||
|
||||
assert_eq!(messages[3].role, crate::llm::Role::Assistant);
|
||||
assert_eq!(messages[3].content, "I found X in doc.md.");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_messages_multiple_tool_calls_per_turn() {
|
||||
let mut thread = Thread::new(Uuid::new_v4());
|
||||
|
||||
thread.start_turn("Do two things");
|
||||
{
|
||||
let turn = thread.turns.last_mut().unwrap();
|
||||
turn.record_tool_call("echo", serde_json::json!({"msg": "a"}));
|
||||
turn.record_tool_result(serde_json::json!("a"));
|
||||
turn.record_tool_call("time", serde_json::json!({}));
|
||||
turn.record_tool_error("timeout");
|
||||
}
|
||||
thread.complete_turn("Done.");
|
||||
|
||||
let messages = thread.messages();
|
||||
// user + assistant_with_calls(2) + tool_result + tool_result + assistant = 5
|
||||
assert_eq!(messages.len(), 5);
|
||||
|
||||
let tcs = messages[1].tool_calls.as_ref().unwrap();
|
||||
assert_eq!(tcs.len(), 2);
|
||||
|
||||
// First tool: success
|
||||
assert_eq!(messages[2].content, "a");
|
||||
// Second tool: error (passed through directly, no wrapping)
|
||||
assert!(messages[3].content.contains("timeout"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_restore_from_messages_with_tool_calls() {
|
||||
let mut thread = Thread::new(Uuid::new_v4());
|
||||
|
||||
// Build a message sequence with tool calls
|
||||
let tc = ToolCall {
|
||||
id: "call_0".to_string(),
|
||||
name: "search".to_string(),
|
||||
arguments: serde_json::json!({"q": "test"}),
|
||||
};
|
||||
let messages = vec![
|
||||
ChatMessage::user("Find test"),
|
||||
ChatMessage::assistant_with_tool_calls(None, vec![tc]),
|
||||
ChatMessage::tool_result("call_0", "search", "result: found"),
|
||||
ChatMessage::assistant("Found it."),
|
||||
];
|
||||
|
||||
thread.restore_from_messages(messages);
|
||||
|
||||
assert_eq!(thread.turns.len(), 1);
|
||||
let turn = &thread.turns[0];
|
||||
assert_eq!(turn.user_input, "Find test");
|
||||
assert_eq!(turn.tool_calls.len(), 1);
|
||||
assert_eq!(turn.tool_calls[0].name, "search");
|
||||
assert_eq!(
|
||||
turn.tool_calls[0].result,
|
||||
Some(serde_json::Value::String("result: found".to_string()))
|
||||
);
|
||||
assert_eq!(turn.response, Some("Found it.".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_restore_from_messages_with_tool_error() {
|
||||
let mut thread = Thread::new(Uuid::new_v4());
|
||||
|
||||
let tc = ToolCall {
|
||||
id: "call_0".to_string(),
|
||||
name: "http".to_string(),
|
||||
arguments: serde_json::json!({}),
|
||||
};
|
||||
let messages = vec![
|
||||
ChatMessage::user("Fetch URL"),
|
||||
ChatMessage::assistant_with_tool_calls(None, vec![tc]),
|
||||
ChatMessage::tool_result("call_0", "http", "Error: timeout"),
|
||||
ChatMessage::assistant("The request timed out."),
|
||||
];
|
||||
|
||||
thread.restore_from_messages(messages);
|
||||
|
||||
// restore_from_messages stores all tool content as result (not error),
|
||||
// because it can't reliably distinguish errors from results that happen
|
||||
// to start with "Error: ". The content is preserved for LLM context.
|
||||
let turn = &thread.turns[0];
|
||||
assert_eq!(
|
||||
turn.tool_calls[0].result,
|
||||
Some(serde_json::Value::String("Error: timeout".to_string()))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_messages_round_trip_with_tools() {
|
||||
// Build a thread with tool calls, get messages(), restore, get messages() again
|
||||
// The two message sequences should be equivalent.
|
||||
let mut thread = Thread::new(Uuid::new_v4());
|
||||
|
||||
thread.start_turn("Do search");
|
||||
{
|
||||
let turn = thread.turns.last_mut().unwrap();
|
||||
turn.record_tool_call("search", serde_json::json!({"q": "test"}));
|
||||
turn.record_tool_result(serde_json::json!("found"));
|
||||
}
|
||||
thread.complete_turn("Here are results.");
|
||||
|
||||
let messages_original = thread.messages();
|
||||
|
||||
// Restore into a new thread
|
||||
let mut thread2 = Thread::new(Uuid::new_v4());
|
||||
thread2.restore_from_messages(messages_original.clone());
|
||||
|
||||
let messages_restored = thread2.messages();
|
||||
|
||||
// Same number of messages
|
||||
assert_eq!(messages_original.len(), messages_restored.len());
|
||||
|
||||
// Same roles
|
||||
for (orig, rest) in messages_original.iter().zip(messages_restored.iter()) {
|
||||
assert_eq!(orig.role, rest.role);
|
||||
}
|
||||
|
||||
// Same final response
|
||||
assert_eq!(
|
||||
messages_original.last().unwrap().content,
|
||||
messages_restored.last().unwrap().content
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_restore_multi_stage_tool_calls() {
|
||||
let mut thread = Thread::new(Uuid::new_v4());
|
||||
|
||||
let tc1 = ToolCall {
|
||||
id: "call_a".to_string(),
|
||||
name: "search".to_string(),
|
||||
arguments: serde_json::json!({"q": "data"}),
|
||||
};
|
||||
let tc2 = ToolCall {
|
||||
id: "call_b".to_string(),
|
||||
name: "write".to_string(),
|
||||
arguments: serde_json::json!({"path": "out.txt"}),
|
||||
};
|
||||
let messages = vec![
|
||||
ChatMessage::user("Find and save"),
|
||||
ChatMessage::assistant_with_tool_calls(None, vec![tc1]),
|
||||
ChatMessage::tool_result("call_a", "search", "found data"),
|
||||
ChatMessage::assistant_with_tool_calls(None, vec![tc2]),
|
||||
ChatMessage::tool_result("call_b", "write", "written"),
|
||||
ChatMessage::assistant("Done, saved to out.txt"),
|
||||
];
|
||||
|
||||
thread.restore_from_messages(messages);
|
||||
|
||||
assert_eq!(thread.turns.len(), 1);
|
||||
let turn = &thread.turns[0];
|
||||
assert_eq!(turn.tool_calls.len(), 2);
|
||||
assert_eq!(turn.tool_calls[0].name, "search");
|
||||
assert_eq!(turn.tool_calls[1].name, "write");
|
||||
assert_eq!(
|
||||
turn.tool_calls[0].result,
|
||||
Some(serde_json::Value::String("found data".to_string()))
|
||||
);
|
||||
assert_eq!(
|
||||
turn.tool_calls[1].result,
|
||||
Some(serde_json::Value::String("written".to_string()))
|
||||
);
|
||||
assert_eq!(turn.response, Some("Done, saved to out.txt".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_messages_truncates_large_tool_results() {
|
||||
let mut thread = Thread::new(Uuid::new_v4());
|
||||
|
||||
thread.start_turn("Read big file");
|
||||
{
|
||||
let turn = thread.turns.last_mut().unwrap();
|
||||
turn.record_tool_call("read_file", serde_json::json!({"path": "big.txt"}));
|
||||
let big_result = "x".repeat(2000);
|
||||
turn.record_tool_result(serde_json::json!(big_result));
|
||||
}
|
||||
thread.complete_turn("Here's the file content.");
|
||||
|
||||
let messages = thread.messages();
|
||||
let tool_result_content = &messages[2].content;
|
||||
assert!(
|
||||
tool_result_content.len() <= 1010,
|
||||
"Tool result should be truncated, got {} chars",
|
||||
tool_result_content.len()
|
||||
);
|
||||
assert!(tool_result_content.ends_with("..."));
|
||||
}
|
||||
}
|
||||
|
||||
+316
-62
@@ -20,7 +20,7 @@ use crate::channels::web::util::truncate_preview;
|
||||
use crate::channels::{IncomingMessage, StatusUpdate};
|
||||
use crate::context::JobContext;
|
||||
use crate::error::Error;
|
||||
use crate::llm::ChatMessage;
|
||||
use crate::llm::{ChatMessage, ToolCall};
|
||||
use crate::tools::redact_params;
|
||||
|
||||
impl Agent {
|
||||
@@ -66,16 +66,7 @@ impl Agent {
|
||||
.await
|
||||
.unwrap_or_default();
|
||||
msg_count = db_messages.len();
|
||||
chat_messages = db_messages
|
||||
.iter()
|
||||
.filter_map(|m| match m.role.as_str() {
|
||||
"user" => Some(ChatMessage::user(&m.content)),
|
||||
"assistant" => Some(ChatMessage::assistant(&m.content)),
|
||||
// tool_calls rows are UI metadata (tool name + preview),
|
||||
// not part of the LLM conversation context.
|
||||
_ => None,
|
||||
})
|
||||
.collect();
|
||||
chat_messages = rebuild_chat_messages_from_db(&db_messages);
|
||||
} else {
|
||||
msg_count = 0;
|
||||
}
|
||||
@@ -230,7 +221,7 @@ impl Agent {
|
||||
)
|
||||
.await;
|
||||
|
||||
let compactor = ContextCompactor::new(self.llm().clone(), self.safety().clone());
|
||||
let compactor = ContextCompactor::new(self.llm().clone());
|
||||
if let Err(e) = compactor
|
||||
.compact(thread, strategy, self.workspace().map(|w| w.as_ref()))
|
||||
.await
|
||||
@@ -257,6 +248,14 @@ impl Agent {
|
||||
);
|
||||
}
|
||||
|
||||
// Augment content with attachment context (transcripts, metadata, images)
|
||||
let augmented =
|
||||
crate::agent::attachments::augment_with_attachments(content, &message.attachments);
|
||||
let (effective_content, image_parts) = match &augmented {
|
||||
Some(result) => (result.text.as_str(), result.image_parts.clone()),
|
||||
None => (content, Vec::new()),
|
||||
};
|
||||
|
||||
// Start the turn and get messages
|
||||
let turn_messages = {
|
||||
let mut sess = session.lock().await;
|
||||
@@ -264,12 +263,13 @@ impl Agent {
|
||||
.threads
|
||||
.get_mut(&thread_id)
|
||||
.ok_or_else(|| Error::from(crate::error::JobError::NotFound { id: thread_id }))?;
|
||||
thread.start_turn(content);
|
||||
let turn = thread.start_turn(effective_content);
|
||||
turn.image_content_parts = image_parts;
|
||||
thread.messages()
|
||||
};
|
||||
|
||||
// Persist user message to DB immediately so it survives crashes
|
||||
self.persist_user_message(thread_id, &message.user_id, content)
|
||||
self.persist_user_message(thread_id, &message.user_id, effective_content)
|
||||
.await;
|
||||
|
||||
// Send thinking status
|
||||
@@ -331,10 +331,10 @@ impl Agent {
|
||||
};
|
||||
|
||||
thread.complete_turn(&response);
|
||||
let tool_calls = thread
|
||||
let (turn_number, tool_calls) = thread
|
||||
.turns
|
||||
.last()
|
||||
.map(|t| t.tool_calls.clone())
|
||||
.map(|t| (t.turn_number, t.tool_calls.clone()))
|
||||
.unwrap_or_default();
|
||||
let _ = self
|
||||
.channels
|
||||
@@ -346,7 +346,7 @@ impl Agent {
|
||||
.await;
|
||||
|
||||
// Persist tool calls then assistant response (user message already persisted at turn start)
|
||||
self.persist_tool_calls(thread_id, &message.user_id, &tool_calls)
|
||||
self.persist_tool_calls(thread_id, &message.user_id, turn_number, &tool_calls)
|
||||
.await;
|
||||
self.persist_assistant_response(thread_id, &message.user_id, &response)
|
||||
.await;
|
||||
@@ -455,6 +455,7 @@ impl Agent {
|
||||
&self,
|
||||
thread_id: Uuid,
|
||||
user_id: &str,
|
||||
turn_number: usize,
|
||||
tool_calls: &[crate::agent::session::TurnToolCall],
|
||||
) {
|
||||
if tool_calls.is_empty() {
|
||||
@@ -468,14 +469,24 @@ impl Agent {
|
||||
|
||||
let summaries: Vec<serde_json::Value> = tool_calls
|
||||
.iter()
|
||||
.map(|tc| {
|
||||
let mut obj = serde_json::json!({ "name": tc.name });
|
||||
.enumerate()
|
||||
.map(|(i, tc)| {
|
||||
let mut obj = serde_json::json!({
|
||||
"name": tc.name,
|
||||
"call_id": format!("turn{}_{}", turn_number, i),
|
||||
});
|
||||
if let Some(ref result) = tc.result {
|
||||
let preview = match result {
|
||||
serde_json::Value::String(s) => truncate_preview(s, 500),
|
||||
other => truncate_preview(&other.to_string(), 500),
|
||||
};
|
||||
obj["result_preview"] = serde_json::Value::String(preview);
|
||||
// Store full result (truncated to ~1000 chars) for LLM context rebuild
|
||||
let full_result = match result {
|
||||
serde_json::Value::String(s) => truncate_preview(s, 1000),
|
||||
other => truncate_preview(&other.to_string(), 1000),
|
||||
};
|
||||
obj["result"] = serde_json::Value::String(full_result);
|
||||
}
|
||||
if let Some(ref error) = tc.error {
|
||||
obj["error"] = serde_json::Value::String(truncate_preview(error, 200));
|
||||
@@ -618,7 +629,7 @@ impl Agent {
|
||||
crate::agent::context_monitor::CompactionStrategy::Summarize { keep_recent: 5 },
|
||||
);
|
||||
|
||||
let compactor = ContextCompactor::new(self.llm().clone(), self.safety().clone());
|
||||
let compactor = ContextCompactor::new(self.llm().clone());
|
||||
match compactor
|
||||
.compact(thread, strategy, self.workspace().map(|w| w.as_ref()))
|
||||
.await
|
||||
@@ -737,6 +748,16 @@ impl Agent {
|
||||
let mut job_ctx =
|
||||
JobContext::with_user(&message.user_id, "chat", "Interactive chat session");
|
||||
job_ctx.http_interceptor = self.deps.http_interceptor.clone();
|
||||
// Prefer a valid timezone from the approval message, fall back to the
|
||||
// resolved timezone stored when the approval was originally requested.
|
||||
let tz_candidate = message
|
||||
.timezone
|
||||
.as_deref()
|
||||
.filter(|tz| crate::timezone::parse_timezone(tz).is_some())
|
||||
.or(pending.user_timezone.as_deref());
|
||||
if let Some(tz) = tz_candidate {
|
||||
job_ctx.user_timezone = tz.to_string();
|
||||
}
|
||||
|
||||
let _ = self
|
||||
.channels
|
||||
@@ -788,19 +809,33 @@ impl Agent {
|
||||
let mut context_messages = pending.context_messages;
|
||||
let deferred_tool_calls = pending.deferred_tool_calls;
|
||||
|
||||
// Record result in thread
|
||||
// Sanitize tool result, then record the cleaned version in the
|
||||
// thread. Must happen before auth intercept check which may return early.
|
||||
let is_tool_error = tool_result.is_err();
|
||||
let result_content = match &tool_result {
|
||||
Ok(output) => {
|
||||
let sanitized = self
|
||||
.safety()
|
||||
.sanitize_tool_output(&pending.tool_name, output);
|
||||
self.safety().wrap_for_llm(
|
||||
&pending.tool_name,
|
||||
&sanitized.content,
|
||||
sanitized.was_modified,
|
||||
)
|
||||
}
|
||||
Err(e) => format!("Error: {}", e),
|
||||
};
|
||||
|
||||
// Record sanitized result in thread
|
||||
{
|
||||
let mut sess = session.lock().await;
|
||||
if let Some(thread) = sess.threads.get_mut(&thread_id)
|
||||
&& let Some(turn) = thread.last_turn_mut()
|
||||
{
|
||||
match &tool_result {
|
||||
Ok(output) => {
|
||||
turn.record_tool_result(serde_json::json!(output));
|
||||
}
|
||||
Err(e) => {
|
||||
turn.record_tool_error(e.to_string());
|
||||
}
|
||||
if is_tool_error {
|
||||
turn.record_tool_error(result_content.clone());
|
||||
} else {
|
||||
turn.record_tool_result(serde_json::json!(result_content));
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -822,21 +857,6 @@ impl Agent {
|
||||
return Ok(SubmissionResult::response(instructions));
|
||||
}
|
||||
|
||||
// Add tool result to context
|
||||
let result_content = match tool_result {
|
||||
Ok(output) => {
|
||||
let sanitized = self
|
||||
.safety()
|
||||
.sanitize_tool_output(&pending.tool_name, &output);
|
||||
self.safety().wrap_for_llm(
|
||||
&pending.tool_name,
|
||||
&sanitized.content,
|
||||
sanitized.was_modified,
|
||||
)
|
||||
}
|
||||
Err(e) => format!("Error: {}", e),
|
||||
};
|
||||
|
||||
context_messages.push(ChatMessage::tool_result(
|
||||
&pending.tool_call_id,
|
||||
&pending.tool_name,
|
||||
@@ -1041,15 +1061,31 @@ impl Agent {
|
||||
.await;
|
||||
}
|
||||
|
||||
// Record in thread
|
||||
// Sanitize first, then record the cleaned version in thread.
|
||||
// Must happen before auth detection which may set deferred_auth.
|
||||
let is_deferred_error = deferred_result.is_err();
|
||||
let deferred_content = match &deferred_result {
|
||||
Ok(output) => {
|
||||
let sanitized = self.safety().sanitize_tool_output(&tc.name, output);
|
||||
self.safety().wrap_for_llm(
|
||||
&tc.name,
|
||||
&sanitized.content,
|
||||
sanitized.was_modified,
|
||||
)
|
||||
}
|
||||
Err(e) => format!("Error: {}", e),
|
||||
};
|
||||
|
||||
// Record sanitized result in thread
|
||||
{
|
||||
let mut sess = session.lock().await;
|
||||
if let Some(thread) = sess.threads.get_mut(&thread_id)
|
||||
&& let Some(turn) = thread.last_turn_mut()
|
||||
{
|
||||
match &deferred_result {
|
||||
Ok(output) => turn.record_tool_result(serde_json::json!(output)),
|
||||
Err(e) => turn.record_tool_error(e.to_string()),
|
||||
if is_deferred_error {
|
||||
turn.record_tool_error(deferred_content.clone());
|
||||
} else {
|
||||
turn.record_tool_result(serde_json::json!(deferred_content));
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1071,18 +1107,6 @@ impl Agent {
|
||||
deferred_auth = Some(instructions);
|
||||
}
|
||||
|
||||
let deferred_content = match deferred_result {
|
||||
Ok(output) => {
|
||||
let sanitized = self.safety().sanitize_tool_output(&tc.name, &output);
|
||||
self.safety().wrap_for_llm(
|
||||
&tc.name,
|
||||
&sanitized.content,
|
||||
sanitized.was_modified,
|
||||
)
|
||||
}
|
||||
Err(e) => format!("Error: {}", e),
|
||||
};
|
||||
|
||||
context_messages.push(ChatMessage::tool_result(&tc.id, &tc.name, deferred_content));
|
||||
}
|
||||
|
||||
@@ -1102,6 +1126,8 @@ impl Agent {
|
||||
tool_call_id: tc.id.clone(),
|
||||
context_messages: context_messages.clone(),
|
||||
deferred_tool_calls: deferred_tool_calls[approval_idx + 1..].to_vec(),
|
||||
// Carry forward the resolved timezone from the original pending approval
|
||||
user_timezone: pending.user_timezone.clone(),
|
||||
};
|
||||
|
||||
let request_id = new_pending.request_id;
|
||||
@@ -1148,13 +1174,13 @@ impl Agent {
|
||||
match result {
|
||||
Ok(AgenticLoopResult::Response(response)) => {
|
||||
thread.complete_turn(&response);
|
||||
let tool_calls = thread
|
||||
let (turn_number, tool_calls) = thread
|
||||
.turns
|
||||
.last()
|
||||
.map(|t| t.tool_calls.clone())
|
||||
.map(|t| (t.turn_number, t.tool_calls.clone()))
|
||||
.unwrap_or_default();
|
||||
// User message already persisted at turn start; save tool calls then assistant response
|
||||
self.persist_tool_calls(thread_id, &message.user_id, &tool_calls)
|
||||
self.persist_tool_calls(thread_id, &message.user_id, turn_number, &tool_calls)
|
||||
.await;
|
||||
self.persist_assistant_response(thread_id, &message.user_id, &response)
|
||||
.await;
|
||||
@@ -1469,3 +1495,231 @@ impl Agent {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Rebuild full LLM-compatible `ChatMessage` sequence from DB messages.
|
||||
///
|
||||
/// Parses `role="tool_calls"` rows to reconstruct `assistant_with_tool_calls`
|
||||
/// and `tool_result` messages so that the LLM sees the complete tool execution
|
||||
/// history on thread hydration. Falls back gracefully for legacy rows that
|
||||
/// lack the enriched fields (`call_id`, `parameters`, `result`).
|
||||
fn rebuild_chat_messages_from_db(
|
||||
db_messages: &[crate::history::ConversationMessage],
|
||||
) -> Vec<ChatMessage> {
|
||||
let mut result = Vec::new();
|
||||
|
||||
for msg in db_messages {
|
||||
match msg.role.as_str() {
|
||||
"user" => result.push(ChatMessage::user(&msg.content)),
|
||||
"assistant" => result.push(ChatMessage::assistant(&msg.content)),
|
||||
"tool_calls" => {
|
||||
// Try to parse the enriched JSON and rebuild tool messages.
|
||||
if let Ok(calls) = serde_json::from_str::<Vec<serde_json::Value>>(&msg.content) {
|
||||
if calls.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Check if this is an enriched row (has call_id) or legacy
|
||||
let has_call_id = calls
|
||||
.first()
|
||||
.and_then(|c| c.get("call_id"))
|
||||
.and_then(|v| v.as_str())
|
||||
.is_some();
|
||||
|
||||
if has_call_id {
|
||||
// Build assistant_with_tool_calls + tool_result messages
|
||||
let tool_calls: Vec<ToolCall> = calls
|
||||
.iter()
|
||||
.map(|c| ToolCall {
|
||||
id: c["call_id"].as_str().unwrap_or("call_0").to_string(),
|
||||
name: c["name"].as_str().unwrap_or("unknown").to_string(),
|
||||
arguments: c
|
||||
.get("parameters")
|
||||
.cloned()
|
||||
.unwrap_or(serde_json::json!({})),
|
||||
})
|
||||
.collect();
|
||||
|
||||
// The assistant text for tool_calls is always None here;
|
||||
// the final assistant response comes as a separate
|
||||
// "assistant" row after this tool_calls row.
|
||||
result.push(ChatMessage::assistant_with_tool_calls(None, tool_calls));
|
||||
|
||||
// Emit tool_result messages for each call
|
||||
for c in &calls {
|
||||
let call_id = c["call_id"].as_str().unwrap_or("call_0").to_string();
|
||||
let name = c["name"].as_str().unwrap_or("unknown").to_string();
|
||||
let content = if let Some(err) = c.get("error").and_then(|v| v.as_str())
|
||||
{
|
||||
format!("Error: {}", err)
|
||||
} else if let Some(res) = c.get("result").and_then(|v| v.as_str()) {
|
||||
res.to_string()
|
||||
} else if let Some(preview) =
|
||||
c.get("result_preview").and_then(|v| v.as_str())
|
||||
{
|
||||
preview.to_string()
|
||||
} else {
|
||||
"OK".to_string()
|
||||
};
|
||||
result.push(ChatMessage::tool_result(call_id, name, content));
|
||||
}
|
||||
}
|
||||
// Legacy rows without call_id: skip (will appear as
|
||||
// simple user/assistant pairs, same as before this fix).
|
||||
}
|
||||
}
|
||||
_ => {} // Skip unknown roles
|
||||
}
|
||||
}
|
||||
|
||||
result
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_rebuild_chat_messages_user_assistant_only() {
|
||||
let messages = vec![
|
||||
make_db_msg("user", "Hello"),
|
||||
make_db_msg("assistant", "Hi there!"),
|
||||
];
|
||||
let result = rebuild_chat_messages_from_db(&messages);
|
||||
assert_eq!(result.len(), 2);
|
||||
assert_eq!(result[0].role, crate::llm::Role::User);
|
||||
assert_eq!(result[1].role, crate::llm::Role::Assistant);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_rebuild_chat_messages_with_enriched_tool_calls() {
|
||||
let tool_json = serde_json::json!([
|
||||
{
|
||||
"name": "memory_search",
|
||||
"call_id": "call_0",
|
||||
"parameters": {"query": "test"},
|
||||
"result": "Found 3 results",
|
||||
"result_preview": "Found 3 re..."
|
||||
},
|
||||
{
|
||||
"name": "echo",
|
||||
"call_id": "call_1",
|
||||
"parameters": {"message": "hi"},
|
||||
"error": "timeout"
|
||||
}
|
||||
]);
|
||||
let messages = vec![
|
||||
make_db_msg("user", "Search for test"),
|
||||
make_db_msg("tool_calls", &tool_json.to_string()),
|
||||
make_db_msg("assistant", "I found some results."),
|
||||
];
|
||||
let result = rebuild_chat_messages_from_db(&messages);
|
||||
|
||||
// user + assistant_with_tool_calls + tool_result*2 + assistant
|
||||
assert_eq!(result.len(), 5);
|
||||
|
||||
// user
|
||||
assert_eq!(result[0].role, crate::llm::Role::User);
|
||||
|
||||
// assistant with tool_calls
|
||||
assert_eq!(result[1].role, crate::llm::Role::Assistant);
|
||||
assert!(result[1].tool_calls.is_some());
|
||||
let tcs = result[1].tool_calls.as_ref().unwrap();
|
||||
assert_eq!(tcs.len(), 2);
|
||||
assert_eq!(tcs[0].name, "memory_search");
|
||||
assert_eq!(tcs[0].id, "call_0");
|
||||
assert_eq!(tcs[1].name, "echo");
|
||||
|
||||
// tool results
|
||||
assert_eq!(result[2].role, crate::llm::Role::Tool);
|
||||
assert_eq!(result[2].tool_call_id, Some("call_0".to_string()));
|
||||
assert!(result[2].content.contains("Found 3 results"));
|
||||
|
||||
assert_eq!(result[3].role, crate::llm::Role::Tool);
|
||||
assert_eq!(result[3].tool_call_id, Some("call_1".to_string()));
|
||||
assert!(result[3].content.contains("Error: timeout"));
|
||||
|
||||
// final assistant
|
||||
assert_eq!(result[4].role, crate::llm::Role::Assistant);
|
||||
assert_eq!(result[4].content, "I found some results.");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_rebuild_chat_messages_legacy_tool_calls_skipped() {
|
||||
// Legacy format: no call_id field
|
||||
let tool_json = serde_json::json!([
|
||||
{"name": "echo", "result_preview": "hello"}
|
||||
]);
|
||||
let messages = vec![
|
||||
make_db_msg("user", "Hi"),
|
||||
make_db_msg("tool_calls", &tool_json.to_string()),
|
||||
make_db_msg("assistant", "Done"),
|
||||
];
|
||||
let result = rebuild_chat_messages_from_db(&messages);
|
||||
|
||||
// Legacy rows are skipped, only user + assistant
|
||||
assert_eq!(result.len(), 2);
|
||||
assert_eq!(result[0].role, crate::llm::Role::User);
|
||||
assert_eq!(result[1].role, crate::llm::Role::Assistant);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_rebuild_chat_messages_empty() {
|
||||
let result = rebuild_chat_messages_from_db(&[]);
|
||||
assert!(result.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_rebuild_chat_messages_malformed_tool_calls_json() {
|
||||
let messages = vec![
|
||||
make_db_msg("user", "Hi"),
|
||||
make_db_msg("tool_calls", "not valid json"),
|
||||
make_db_msg("assistant", "Done"),
|
||||
];
|
||||
let result = rebuild_chat_messages_from_db(&messages);
|
||||
// Malformed JSON is silently skipped
|
||||
assert_eq!(result.len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_rebuild_chat_messages_multi_turn_with_tools() {
|
||||
let tool_json_1 = serde_json::json!([
|
||||
{"name": "search", "call_id": "call_0", "parameters": {}, "result": "found it"}
|
||||
]);
|
||||
let tool_json_2 = serde_json::json!([
|
||||
{"name": "write", "call_id": "call_0", "parameters": {"path": "a.txt"}, "result": "ok"}
|
||||
]);
|
||||
let messages = vec![
|
||||
make_db_msg("user", "Find X"),
|
||||
make_db_msg("tool_calls", &tool_json_1.to_string()),
|
||||
make_db_msg("assistant", "Found X"),
|
||||
make_db_msg("user", "Write it"),
|
||||
make_db_msg("tool_calls", &tool_json_2.to_string()),
|
||||
make_db_msg("assistant", "Written"),
|
||||
];
|
||||
let result = rebuild_chat_messages_from_db(&messages);
|
||||
|
||||
// Turn 1: user + assistant_with_calls + tool_result + assistant = 4
|
||||
// Turn 2: user + assistant_with_calls + tool_result + assistant = 4
|
||||
assert_eq!(result.len(), 8);
|
||||
|
||||
// Verify turn boundaries
|
||||
assert_eq!(result[0].content, "Find X");
|
||||
assert!(result[1].tool_calls.is_some());
|
||||
assert_eq!(result[2].role, crate::llm::Role::Tool);
|
||||
assert_eq!(result[3].content, "Found X");
|
||||
|
||||
assert_eq!(result[4].content, "Write it");
|
||||
assert!(result[5].tool_calls.is_some());
|
||||
assert_eq!(result[6].role, crate::llm::Role::Tool);
|
||||
assert_eq!(result[7].content, "Written");
|
||||
}
|
||||
|
||||
fn make_db_msg(role: &str, content: &str) -> crate::history::ConversationMessage {
|
||||
crate::history::ConversationMessage {
|
||||
id: uuid::Uuid::new_v4(),
|
||||
role: role.to_string(),
|
||||
content: content.to_string(),
|
||||
created_at: chrono::Utc::now(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+412
-49
@@ -15,11 +15,12 @@ use crate::db::Database;
|
||||
use crate::error::Error;
|
||||
use crate::hooks::HookRegistry;
|
||||
use crate::llm::{
|
||||
ActionPlan, ChatMessage, LlmProvider, Reasoning, ReasoningContext, RespondResult, ToolSelection,
|
||||
ActionPlan, ChatMessage, LlmProvider, Reasoning, ReasoningContext, RespondResult, ToolCall,
|
||||
ToolSelection,
|
||||
};
|
||||
use crate::safety::SafetyLayer;
|
||||
use crate::tools::rate_limiter::RateLimitResult;
|
||||
use crate::tools::{ToolRegistry, redact_params};
|
||||
use crate::tools::{ApprovalContext, ToolRegistry, redact_params};
|
||||
|
||||
/// Shared dependencies for worker execution.
|
||||
///
|
||||
@@ -37,6 +38,12 @@ pub struct WorkerDeps {
|
||||
pub use_planning: bool,
|
||||
/// SSE broadcast sender for live job event streaming to the web gateway.
|
||||
pub sse_tx: Option<tokio::sync::broadcast::Sender<SseEvent>>,
|
||||
/// Approval context for tool execution. When `None`, all non-`Never` tools are
|
||||
/// blocked (legacy behavior). When `Some`, the context determines which tools
|
||||
/// are pre-approved for autonomous execution.
|
||||
pub approval_context: Option<ApprovalContext>,
|
||||
/// HTTP interceptor for trace recording/replay (propagated to JobContext).
|
||||
pub http_interceptor: Option<Arc<dyn crate::llm::recording::HttpInterceptor>>,
|
||||
}
|
||||
|
||||
/// Worker that executes a single job.
|
||||
@@ -205,7 +212,7 @@ impl Worker {
|
||||
let job_ctx = self.context_manager().get_context(self.job_id).await?;
|
||||
|
||||
// Create reasoning engine
|
||||
let reasoning = Reasoning::new(self.llm().clone(), self.safety().clone());
|
||||
let reasoning = Reasoning::new(self.llm().clone());
|
||||
|
||||
// Build initial reasoning context (tool definitions refreshed each iteration in execution_loop)
|
||||
let mut reason_ctx = ReasoningContext::new().with_job(&job_ctx.description);
|
||||
@@ -246,6 +253,9 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
// Already in a terminal state (e.g. execution_loop
|
||||
// called mark_completed itself).
|
||||
}
|
||||
Ok(JobState::Completed) => {
|
||||
// execution_loop already called mark_completed.
|
||||
}
|
||||
Ok(JobState::Stuck) => {
|
||||
// execution_loop marked this as stuck (e.g. "plan
|
||||
// completed but work remains"); leave for self-repair.
|
||||
@@ -296,6 +306,8 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
let mut iteration = 0;
|
||||
const MAX_CONSECUTIVE_RATE_LIMITS: usize = 10;
|
||||
let mut consecutive_rate_limits = 0usize;
|
||||
const MAX_TOOL_INTENT_NUDGES: u32 = 2;
|
||||
let mut consecutive_tool_intent_nudges: u32 = 0;
|
||||
|
||||
// Initial tool definitions for planning (will be refreshed in loop)
|
||||
reason_ctx.available_tools = self.tools().tool_definitions().await;
|
||||
@@ -353,11 +365,13 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
if let Some(ref plan) = plan {
|
||||
self.execute_plan(rx, reasoning, reason_ctx, plan).await?;
|
||||
|
||||
// If the plan marked the job terminal, we're done. Only fall
|
||||
// through to the direct selection loop if the plan was
|
||||
// interrupted or explicitly left the job in-progress.
|
||||
// If the plan marked the job completed, terminal, or stuck, we're
|
||||
// done. Only fall through to the direct selection loop if the
|
||||
// plan was interrupted or explicitly left the job in-progress.
|
||||
if let Ok(ctx) = self.context_manager().get_context(self.job_id).await
|
||||
&& (ctx.state.is_terminal() || ctx.state == JobState::Stuck)
|
||||
&& (ctx.state.is_terminal()
|
||||
|| ctx.state == JobState::Stuck
|
||||
|| ctx.state == JobState::Completed)
|
||||
{
|
||||
return Ok(());
|
||||
}
|
||||
@@ -403,7 +417,8 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
|
||||
iteration += 1;
|
||||
if iteration > max_iterations {
|
||||
self.mark_stuck("Maximum iterations exceeded").await?;
|
||||
self.mark_failed("Maximum iterations exceeded: job hit the iteration cap")
|
||||
.await?;
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
@@ -423,7 +438,8 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
"LLM rate limited during tool selection, backing off"
|
||||
);
|
||||
if consecutive_rate_limits >= MAX_CONSECUTIVE_RATE_LIMITS {
|
||||
self.mark_stuck("Persistent rate limiting").await?;
|
||||
self.mark_failed("Persistent rate limiting: exceeded retry limit")
|
||||
.await?;
|
||||
return Ok(());
|
||||
}
|
||||
self.log_event(
|
||||
@@ -453,7 +469,8 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
"LLM rate limited during respond_with_tools, backing off"
|
||||
);
|
||||
if consecutive_rate_limits >= MAX_CONSECUTIVE_RATE_LIMITS {
|
||||
self.mark_stuck("Persistent rate limiting").await?;
|
||||
self.mark_failed("Persistent rate limiting: exceeded retry limit")
|
||||
.await?;
|
||||
return Ok(());
|
||||
}
|
||||
self.log_event(
|
||||
@@ -469,6 +486,20 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
Err(e) => return Err(e.into()),
|
||||
};
|
||||
|
||||
// Track token usage from LLM call against the job budget.
|
||||
// NOTE: select_tools() also makes LLM calls but doesn't expose
|
||||
// TokenUsage; only respond_with_tools() usage is tracked here.
|
||||
let total_tokens = respond_output.usage.total() as u64;
|
||||
if total_tokens > 0
|
||||
&& let Err(msg) = self
|
||||
.context_manager()
|
||||
.update_context(self.job_id, |ctx| ctx.add_tokens(total_tokens))
|
||||
.await?
|
||||
{
|
||||
self.mark_failed(&msg).await?;
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
match respond_output.result {
|
||||
RespondResult::Text(response) => {
|
||||
// Check for explicit completion phrases. Use word-boundary
|
||||
@@ -491,17 +522,34 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
}),
|
||||
);
|
||||
|
||||
// Give it one more chance to select a tool
|
||||
if iteration > 3 && iteration % 5 == 0 {
|
||||
reason_ctx.messages.push(ChatMessage::user(
|
||||
"Are you stuck? Do you need help completing this job?",
|
||||
));
|
||||
// Nudge the LLM if it expressed tool intent without calling tools
|
||||
let signals_intent = !reason_ctx.available_tools.is_empty()
|
||||
&& crate::llm::llm_signals_tool_intent(&response);
|
||||
if signals_intent && consecutive_tool_intent_nudges < MAX_TOOL_INTENT_NUDGES
|
||||
{
|
||||
consecutive_tool_intent_nudges += 1;
|
||||
tracing::info!(
|
||||
job_id = %self.job_id,
|
||||
"LLM expressed tool intent without calling a tool, nudging"
|
||||
);
|
||||
reason_ctx
|
||||
.messages
|
||||
.push(ChatMessage::user(crate::llm::TOOL_INTENT_NUDGE));
|
||||
} else if !signals_intent {
|
||||
consecutive_tool_intent_nudges = 0;
|
||||
if iteration > 3 && iteration % 5 == 0 {
|
||||
// Generic fallback nudge
|
||||
reason_ctx.messages.push(ChatMessage::user(
|
||||
"Are you stuck? Do you need help completing this job?",
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
RespondResult::ToolCalls {
|
||||
tool_calls,
|
||||
content,
|
||||
} => {
|
||||
consecutive_tool_intent_nudges = 0;
|
||||
// Model returned tool calls - execute them
|
||||
tracing::debug!(
|
||||
"Job {} respond_with_tools returned {} tool calls",
|
||||
@@ -546,36 +594,54 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
}
|
||||
}
|
||||
}
|
||||
} else if selections.len() == 1 {
|
||||
// Single tool: execute directly
|
||||
let selection = &selections[0];
|
||||
tracing::debug!(
|
||||
"Job {} selecting tool: {} - {}",
|
||||
self.job_id,
|
||||
selection.tool_name,
|
||||
selection.reasoning
|
||||
);
|
||||
|
||||
let result = self
|
||||
.execute_tool(&selection.tool_name, &selection.parameters)
|
||||
.await;
|
||||
|
||||
self.process_tool_result(reason_ctx, selection, result)
|
||||
.await?;
|
||||
} else {
|
||||
// Multiple tools: execute in parallel
|
||||
tracing::debug!(
|
||||
"Job {} executing {} tools in parallel",
|
||||
self.job_id,
|
||||
selections.len()
|
||||
);
|
||||
consecutive_tool_intent_nudges = 0;
|
||||
|
||||
let results = self.execute_tools_parallel(&selections).await;
|
||||
// Record the assistant tool_calls message so that tool_result
|
||||
// messages have a matching parent (prevents orphaned rewrites).
|
||||
let tool_calls: Vec<ToolCall> = selections
|
||||
.iter()
|
||||
.map(|s| ToolCall {
|
||||
id: s.tool_call_id.clone(),
|
||||
name: s.tool_name.clone(),
|
||||
arguments: s.parameters.clone(),
|
||||
})
|
||||
.collect();
|
||||
reason_ctx
|
||||
.messages
|
||||
.push(ChatMessage::assistant_with_tool_calls(None, tool_calls));
|
||||
|
||||
// Process all results
|
||||
for (selection, result) in selections.iter().zip(results) {
|
||||
self.process_tool_result(reason_ctx, selection, result.result)
|
||||
if selections.len() == 1 {
|
||||
// Single tool: execute directly
|
||||
let selection = &selections[0];
|
||||
tracing::debug!(
|
||||
"Job {} selecting tool: {} - {}",
|
||||
self.job_id,
|
||||
selection.tool_name,
|
||||
selection.reasoning
|
||||
);
|
||||
|
||||
let result = self
|
||||
.execute_tool(&selection.tool_name, &selection.parameters)
|
||||
.await;
|
||||
|
||||
self.process_tool_result(reason_ctx, selection, result)
|
||||
.await?;
|
||||
} else {
|
||||
// Multiple tools: execute in parallel
|
||||
tracing::debug!(
|
||||
"Job {} executing {} tools in parallel",
|
||||
self.job_id,
|
||||
selections.len()
|
||||
);
|
||||
|
||||
let results = self.execute_tools_parallel(&selections).await;
|
||||
|
||||
// Process all results
|
||||
for (selection, result) in selections.iter().zip(results) {
|
||||
self.process_tool_result(reason_ctx, selection, result.result)
|
||||
.await?;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -671,8 +737,11 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
name: tool_name.to_string(),
|
||||
})?;
|
||||
|
||||
// Tools requiring approval are blocked in autonomous jobs
|
||||
if tool.requires_approval(params).is_required() {
|
||||
// Check approval: use context-aware check if available, else block all non-Never tools
|
||||
let requirement = tool.requires_approval(params);
|
||||
let blocked =
|
||||
ApprovalContext::is_blocked_or_default(&deps.approval_context, tool_name, requirement);
|
||||
if blocked {
|
||||
return Err(crate::error::ToolError::AuthRequired {
|
||||
name: tool_name.to_string(),
|
||||
}
|
||||
@@ -680,7 +749,11 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
}
|
||||
|
||||
// Fetch job context early so we have the real user_id for hooks and rate limiting
|
||||
let job_ctx = deps.context_manager.get_context(job_id).await?;
|
||||
let mut job_ctx = deps.context_manager.get_context(job_id).await?;
|
||||
// Propagate http_interceptor for trace recording/replay
|
||||
if job_ctx.http_interceptor.is_none() {
|
||||
job_ctx.http_interceptor = deps.http_interceptor.clone();
|
||||
}
|
||||
|
||||
// Check per-tool rate limit before running hooks or executing (cheaper check first)
|
||||
if let Some(config) = tool.rate_limit_config()
|
||||
@@ -1049,11 +1122,6 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
action.reasoning
|
||||
);
|
||||
|
||||
// Execute the planned tool
|
||||
let result = self
|
||||
.execute_tool(&action.tool_name, &action.parameters)
|
||||
.await;
|
||||
|
||||
// Create a synthetic ToolSelection for process_tool_result.
|
||||
// Plan actions don't originate from an LLM tool_call response so
|
||||
// there is no real tool_call_id; generate a unique one.
|
||||
@@ -1065,6 +1133,24 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
tool_call_id: format!("plan_{}_{}", self.job_id, i),
|
||||
};
|
||||
|
||||
// Record the assistant tool_calls message so that the tool_result
|
||||
// has a matching parent (prevents orphaned rewrites).
|
||||
reason_ctx
|
||||
.messages
|
||||
.push(ChatMessage::assistant_with_tool_calls(
|
||||
None,
|
||||
vec![ToolCall {
|
||||
id: selection.tool_call_id.clone(),
|
||||
name: selection.tool_name.clone(),
|
||||
arguments: selection.parameters.clone(),
|
||||
}],
|
||||
));
|
||||
|
||||
// Execute the planned tool
|
||||
let result = self
|
||||
.execute_tool(&action.tool_name, &action.parameters)
|
||||
.await;
|
||||
|
||||
// Process the result
|
||||
let completed = self
|
||||
.process_tool_result(reason_ctx, &selection, result)
|
||||
@@ -1298,6 +1384,8 @@ mod tests {
|
||||
timeout: Duration::from_secs(30),
|
||||
use_planning: false,
|
||||
sse_tx: None,
|
||||
approval_context: None,
|
||||
http_interceptor: None,
|
||||
};
|
||||
|
||||
Worker::new(job_id, deps)
|
||||
@@ -1496,4 +1584,279 @@ mod tests {
|
||||
"Missing tool should produce an error, not a panic"
|
||||
);
|
||||
}
|
||||
|
||||
/// Verify that calling mark_completed on an already-Completed job returns
|
||||
/// an error (Completed → Completed is an invalid state transition).
|
||||
#[tokio::test]
|
||||
async fn test_mark_completed_twice_returns_error() {
|
||||
let worker = make_worker(vec![]).await;
|
||||
|
||||
// Transition to InProgress first (required by state machine)
|
||||
worker
|
||||
.context_manager()
|
||||
.update_context(worker.job_id, |ctx| {
|
||||
ctx.transition_to(JobState::InProgress, None)
|
||||
})
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
|
||||
// First mark_completed should succeed
|
||||
worker.mark_completed().await.unwrap();
|
||||
|
||||
// Verify state is Completed
|
||||
let ctx = worker
|
||||
.context_manager()
|
||||
.get_context(worker.job_id)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(ctx.state, JobState::Completed);
|
||||
|
||||
// Second mark_completed should fail (Completed → Completed is invalid)
|
||||
let result = worker.mark_completed().await;
|
||||
assert!(
|
||||
result.is_err(),
|
||||
"Completed → Completed transition should be rejected by state machine"
|
||||
);
|
||||
}
|
||||
|
||||
/// Build a Worker with the given approval context.
|
||||
async fn make_worker_with_approval(
|
||||
tools: Vec<Arc<dyn Tool>>,
|
||||
approval_context: Option<crate::tools::ApprovalContext>,
|
||||
) -> Worker {
|
||||
let registry = ToolRegistry::new();
|
||||
for t in tools {
|
||||
registry.register(t).await;
|
||||
}
|
||||
|
||||
let cm = Arc::new(crate::context::ContextManager::new(5));
|
||||
let job_id = cm.create_job("test", "test job").await.unwrap();
|
||||
|
||||
let deps = WorkerDeps {
|
||||
context_manager: cm,
|
||||
llm: Arc::new(StubLlm),
|
||||
safety: Arc::new(SafetyLayer::new(&SafetyConfig {
|
||||
max_output_length: 100_000,
|
||||
injection_check_enabled: false,
|
||||
})),
|
||||
tools: Arc::new(registry),
|
||||
store: None,
|
||||
hooks: Arc::new(crate::hooks::HookRegistry::new()),
|
||||
timeout: Duration::from_secs(30),
|
||||
use_planning: false,
|
||||
sse_tx: None,
|
||||
approval_context,
|
||||
http_interceptor: None,
|
||||
};
|
||||
|
||||
Worker::new(job_id, deps)
|
||||
}
|
||||
|
||||
/// A tool that requires approval (UnlessAutoApproved).
|
||||
struct ApprovalTool;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl Tool for ApprovalTool {
|
||||
fn name(&self) -> &str {
|
||||
"needs_approval"
|
||||
}
|
||||
fn description(&self) -> &str {
|
||||
"Tool requiring approval"
|
||||
}
|
||||
fn parameters_schema(&self) -> serde_json::Value {
|
||||
serde_json::json!({"type": "object", "properties": {}})
|
||||
}
|
||||
async fn execute(
|
||||
&self,
|
||||
_params: serde_json::Value,
|
||||
_ctx: &crate::context::JobContext,
|
||||
) -> Result<ToolOutput, crate::tools::ToolError> {
|
||||
Ok(ToolOutput::text(
|
||||
"approved",
|
||||
std::time::Instant::now().elapsed(),
|
||||
))
|
||||
}
|
||||
fn requires_approval(
|
||||
&self,
|
||||
_params: &serde_json::Value,
|
||||
) -> crate::tools::ApprovalRequirement {
|
||||
crate::tools::ApprovalRequirement::UnlessAutoApproved
|
||||
}
|
||||
fn requires_sanitization(&self) -> bool {
|
||||
false
|
||||
}
|
||||
}
|
||||
|
||||
/// A tool that always requires approval.
|
||||
struct AlwaysApprovalTool;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl Tool for AlwaysApprovalTool {
|
||||
fn name(&self) -> &str {
|
||||
"always_approval"
|
||||
}
|
||||
fn description(&self) -> &str {
|
||||
"Tool always requiring approval"
|
||||
}
|
||||
fn parameters_schema(&self) -> serde_json::Value {
|
||||
serde_json::json!({"type": "object", "properties": {}})
|
||||
}
|
||||
async fn execute(
|
||||
&self,
|
||||
_params: serde_json::Value,
|
||||
_ctx: &crate::context::JobContext,
|
||||
) -> Result<ToolOutput, crate::tools::ToolError> {
|
||||
Ok(ToolOutput::text(
|
||||
"always",
|
||||
std::time::Instant::now().elapsed(),
|
||||
))
|
||||
}
|
||||
fn requires_approval(
|
||||
&self,
|
||||
_params: &serde_json::Value,
|
||||
) -> crate::tools::ApprovalRequirement {
|
||||
crate::tools::ApprovalRequirement::Always
|
||||
}
|
||||
fn requires_sanitization(&self) -> bool {
|
||||
false
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_approval_context_unblocks_unless_auto_approved() {
|
||||
// Without approval context, UnlessAutoApproved is blocked
|
||||
let worker_blocked = make_worker_with_approval(vec![Arc::new(ApprovalTool)], None).await;
|
||||
let result = worker_blocked
|
||||
.execute_tool("needs_approval", &serde_json::json!({}))
|
||||
.await;
|
||||
assert!(
|
||||
result.is_err(),
|
||||
"Should be blocked without approval context"
|
||||
);
|
||||
|
||||
// With autonomous approval context, UnlessAutoApproved is allowed
|
||||
let worker_allowed = make_worker_with_approval(
|
||||
vec![Arc::new(ApprovalTool)],
|
||||
Some(crate::tools::ApprovalContext::autonomous()),
|
||||
)
|
||||
.await;
|
||||
let result = worker_allowed
|
||||
.execute_tool("needs_approval", &serde_json::json!({}))
|
||||
.await;
|
||||
assert!(result.is_ok(), "Should be allowed with autonomous context");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_approval_context_blocks_always_unless_permitted() {
|
||||
// Autonomous context without tool_permissions blocks Always tools
|
||||
let worker_blocked = make_worker_with_approval(
|
||||
vec![Arc::new(AlwaysApprovalTool)],
|
||||
Some(crate::tools::ApprovalContext::autonomous()),
|
||||
)
|
||||
.await;
|
||||
let result = worker_blocked
|
||||
.execute_tool("always_approval", &serde_json::json!({}))
|
||||
.await;
|
||||
assert!(
|
||||
result.is_err(),
|
||||
"Always tool should be blocked without permission"
|
||||
);
|
||||
|
||||
// Autonomous context with tool_permissions allows Always tools
|
||||
let worker_allowed = make_worker_with_approval(
|
||||
vec![Arc::new(AlwaysApprovalTool)],
|
||||
Some(crate::tools::ApprovalContext::autonomous_with_tools([
|
||||
"always_approval".to_string(),
|
||||
])),
|
||||
)
|
||||
.await;
|
||||
let result = worker_allowed
|
||||
.execute_tool("always_approval", &serde_json::json!({}))
|
||||
.await;
|
||||
assert!(
|
||||
result.is_ok(),
|
||||
"Always tool should be allowed with permission"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_token_budget_exceeded_fails_job() {
|
||||
let worker = make_worker(vec![]).await;
|
||||
|
||||
// Transition to InProgress (required for mark_failed)
|
||||
worker
|
||||
.context_manager()
|
||||
.update_context(worker.job_id, |ctx| {
|
||||
ctx.transition_to(JobState::InProgress, None)
|
||||
})
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
|
||||
// Set a token budget
|
||||
worker
|
||||
.context_manager()
|
||||
.update_context(worker.job_id, |ctx| {
|
||||
ctx.max_tokens = 100;
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Simulate adding tokens that exceed the budget
|
||||
let budget_result = worker
|
||||
.context_manager()
|
||||
.update_context(worker.job_id, |ctx| ctx.add_tokens(200))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(
|
||||
budget_result.is_err(),
|
||||
"Should return error when token budget exceeded"
|
||||
);
|
||||
|
||||
// Verify that mark_failed transitions job to Failed
|
||||
worker
|
||||
.mark_failed(&budget_result.unwrap_err())
|
||||
.await
|
||||
.unwrap();
|
||||
let ctx = worker
|
||||
.context_manager()
|
||||
.get_context(worker.job_id)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(ctx.state, JobState::Failed);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_iteration_cap_marks_failed_not_stuck() {
|
||||
let worker = make_worker(vec![]).await;
|
||||
|
||||
// Transition to InProgress (required for mark_failed)
|
||||
worker
|
||||
.context_manager()
|
||||
.update_context(worker.job_id, |ctx| {
|
||||
ctx.transition_to(JobState::InProgress, None)
|
||||
})
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
|
||||
// Simulate what the execution loop does when max_iterations is exceeded
|
||||
worker
|
||||
.mark_failed("Maximum iterations exceeded: job hit the iteration cap")
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let ctx = worker
|
||||
.context_manager()
|
||||
.get_context(worker.job_id)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
ctx.state,
|
||||
JobState::Failed,
|
||||
"Iteration cap should transition to Failed, not Stuck"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
+202
-206
@@ -21,7 +21,7 @@ use crate::secrets::SecretsStore;
|
||||
use crate::skills::SkillRegistry;
|
||||
use crate::skills::catalog::SkillCatalog;
|
||||
use crate::tools::ToolRegistry;
|
||||
use crate::tools::mcp::McpSessionManager;
|
||||
use crate::tools::mcp::{McpProcessManager, McpSessionManager};
|
||||
use crate::tools::wasm::SharedCredentialRegistry;
|
||||
use crate::tools::wasm::WasmToolRuntime;
|
||||
use crate::workspace::{EmbeddingProvider, Workspace};
|
||||
@@ -41,6 +41,7 @@ pub struct AppComponents {
|
||||
pub workspace: Option<Arc<Workspace>>,
|
||||
pub extension_manager: Option<Arc<ExtensionManager>>,
|
||||
pub mcp_session_manager: Arc<McpSessionManager>,
|
||||
pub mcp_process_manager: Arc<McpProcessManager>,
|
||||
pub wasm_tool_runtime: Option<Arc<WasmToolRuntime>>,
|
||||
pub log_broadcaster: Arc<LogBroadcaster>,
|
||||
pub context_manager: Arc<ContextManager>,
|
||||
@@ -76,10 +77,7 @@ pub struct AppBuilder {
|
||||
llm_override: Option<Arc<dyn LlmProvider>>,
|
||||
|
||||
// Backend-specific handles needed by secrets store
|
||||
#[cfg(feature = "postgres")]
|
||||
pg_pool: Option<deadpool_postgres::Pool>,
|
||||
#[cfg(feature = "libsql")]
|
||||
libsql_db: Option<Arc<libsql::Database>>,
|
||||
handles: Option<crate::db::DatabaseHandles>,
|
||||
}
|
||||
|
||||
impl AppBuilder {
|
||||
@@ -104,10 +102,7 @@ impl AppBuilder {
|
||||
db: None,
|
||||
secrets_store: None,
|
||||
llm_override: None,
|
||||
#[cfg(feature = "postgres")]
|
||||
pg_pool: None,
|
||||
#[cfg(feature = "libsql")]
|
||||
libsql_db: None,
|
||||
handles: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -136,71 +131,10 @@ impl AppBuilder {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
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
|
||||
.map_err(|e| anyhow::anyhow!("{}", e))?;
|
||||
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."
|
||||
);
|
||||
}
|
||||
};
|
||||
let (db, handles) = crate::db::connect_with_handles(&self.config.database)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("{}", e))?;
|
||||
self.handles = Some(handles);
|
||||
|
||||
// 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 {
|
||||
@@ -211,7 +145,7 @@ impl AppBuilder {
|
||||
match Config::from_db_with_toml(db.as_ref(), "default", toml_path).await {
|
||||
Ok(db_config) => {
|
||||
self.config = db_config;
|
||||
tracing::info!("Configuration reloaded from database");
|
||||
tracing::debug!("Configuration reloaded from database");
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
@@ -244,11 +178,28 @@ impl AppBuilder {
|
||||
let master_key = match self.config.secrets.master_key() {
|
||||
Some(k) => k,
|
||||
None => {
|
||||
// No secrets DB available, but we can still load tokens from
|
||||
// OS credential stores (e.g., Anthropic OAuth via Claude Code's
|
||||
// macOS Keychain / Linux ~/.claude/.credentials.json).
|
||||
crate::config::inject_os_credentials();
|
||||
|
||||
// Consume unused handles
|
||||
#[cfg(feature = "libsql")]
|
||||
self.handles.take();
|
||||
|
||||
// Re-resolve only the LLM config with OS credentials.
|
||||
let store: Option<&(dyn crate::db::SettingsStore + Sync)> =
|
||||
self.db.as_ref().map(|db| db.as_ref() as _);
|
||||
let toml_path = self.toml_path.as_deref();
|
||||
if let Err(e) = self
|
||||
.config
|
||||
.re_resolve_llm(store, "default", toml_path)
|
||||
.await
|
||||
{
|
||||
self.libsql_db.take();
|
||||
tracing::warn!(
|
||||
"Failed to re-resolve LLM config after OS credential injection: {e}"
|
||||
);
|
||||
}
|
||||
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
@@ -257,52 +208,31 @@ impl AppBuilder {
|
||||
Ok(c) => Arc::new(c),
|
||||
Err(e) => {
|
||||
tracing::warn!("Failed to initialize secrets crypto: {}", e);
|
||||
#[cfg(feature = "libsql")]
|
||||
{
|
||||
self.libsql_db.take();
|
||||
}
|
||||
self.handles.take();
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
|
||||
let store: Option<Arc<dyn SecretsStore + Send + Sync>> = None;
|
||||
|
||||
#[cfg(feature = "libsql")]
|
||||
let store = store.or_else(|| {
|
||||
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>
|
||||
})
|
||||
});
|
||||
// Fallback covers the no-database path where `init_database` returned
|
||||
// early before populating `self.handles`.
|
||||
let empty_handles = crate::db::DatabaseHandles::default();
|
||||
let handles = self.handles.as_ref().unwrap_or(&empty_handles);
|
||||
let store = crate::secrets::create_secrets_store(crypto, handles);
|
||||
|
||||
if let Some(ref secrets) = store {
|
||||
// Inject LLM API keys from encrypted storage
|
||||
crate::config::inject_llm_keys_from_secrets(secrets.as_ref(), "default").await;
|
||||
|
||||
// Re-resolve config with newly available keys
|
||||
if let Some(ref db) = self.db {
|
||||
let toml_path = self.toml_path.as_deref();
|
||||
match Config::from_db_with_toml(db.as_ref(), "default", toml_path).await {
|
||||
Ok(refreshed) => {
|
||||
self.config = refreshed;
|
||||
tracing::debug!("LlmConfig re-resolved after secret injection");
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::warn!("Failed to re-resolve config after secret injection: {}", e);
|
||||
}
|
||||
}
|
||||
// Re-resolve only the LLM config with newly available keys.
|
||||
let store: Option<&(dyn crate::db::SettingsStore + Sync)> =
|
||||
self.db.as_ref().map(|db| db.as_ref() as _);
|
||||
let toml_path = self.toml_path.as_deref();
|
||||
if let Err(e) = self
|
||||
.config
|
||||
.re_resolve_llm(store, "default", toml_path)
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Failed to re-resolve LLM config after secret injection: {e}");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -315,7 +245,7 @@ impl AppBuilder {
|
||||
/// Delegates to `build_provider_chain` which applies all decorators
|
||||
/// (retry, smart routing, failover, circuit breaker, response cache).
|
||||
#[allow(clippy::type_complexity)]
|
||||
pub fn init_llm(
|
||||
pub async fn init_llm(
|
||||
&self,
|
||||
) -> Result<
|
||||
(
|
||||
@@ -326,7 +256,7 @@ impl AppBuilder {
|
||||
anyhow::Error,
|
||||
> {
|
||||
let (llm, cheap_llm, recording_handle) =
|
||||
crate::llm::build_provider_chain(&self.config.llm, self.session.clone())?;
|
||||
crate::llm::build_provider_chain(&self.config.llm, self.session.clone()).await?;
|
||||
Ok((llm, cheap_llm, recording_handle))
|
||||
}
|
||||
|
||||
@@ -344,7 +274,7 @@ impl AppBuilder {
|
||||
anyhow::Error,
|
||||
> {
|
||||
let safety = Arc::new(SafetyLayer::new(&self.config.safety));
|
||||
tracing::info!("Safety layer initialized");
|
||||
tracing::debug!("Safety layer initialized");
|
||||
|
||||
// Initialize tool registry with credential injection support
|
||||
let credential_registry = Arc::new(SharedCredentialRegistry::new());
|
||||
@@ -381,18 +311,57 @@ impl AppBuilder {
|
||||
None
|
||||
};
|
||||
|
||||
// Register image/vision tools if we have a workspace and LLM API credentials
|
||||
if workspace.is_some() {
|
||||
let (api_base, api_key_opt) = if let Some(ref provider) = self.config.llm.provider {
|
||||
(
|
||||
provider.base_url.clone(),
|
||||
provider.api_key.as_ref().map(|s| {
|
||||
use secrecy::ExposeSecret;
|
||||
s.expose_secret().to_string()
|
||||
}),
|
||||
)
|
||||
} else {
|
||||
(
|
||||
self.config.llm.nearai.base_url.clone(),
|
||||
self.config.llm.nearai.api_key.as_ref().map(|s| {
|
||||
use secrecy::ExposeSecret;
|
||||
s.expose_secret().to_string()
|
||||
}),
|
||||
)
|
||||
};
|
||||
|
||||
if let Some(api_key) = api_key_opt {
|
||||
// Check for image generation models
|
||||
let model_name = self
|
||||
.config
|
||||
.llm
|
||||
.provider
|
||||
.as_ref()
|
||||
.map(|p| p.model.clone())
|
||||
.unwrap_or_else(|| self.config.llm.nearai.model.clone());
|
||||
let models = vec![model_name.clone()];
|
||||
let gen_model = crate::llm::image_models::suggest_image_model(&models)
|
||||
.unwrap_or("flux-1.1-pro")
|
||||
.to_string();
|
||||
tools.register_image_tools(api_base.clone(), api_key.clone(), gen_model, None);
|
||||
|
||||
// Check for vision models
|
||||
let vision_model = crate::llm::vision_models::suggest_vision_model(&models)
|
||||
.unwrap_or(&model_name)
|
||||
.to_string();
|
||||
tools.register_vision_tools(api_base, api_key, vision_model, None);
|
||||
}
|
||||
}
|
||||
|
||||
// Register builder tool if enabled
|
||||
if self.config.builder.enabled
|
||||
&& (self.config.agent.allow_local_tools || !self.config.sandbox.enabled)
|
||||
{
|
||||
tools
|
||||
.register_builder_tool(
|
||||
llm.clone(),
|
||||
safety.clone(),
|
||||
Some(self.config.builder.to_builder_config()),
|
||||
)
|
||||
.register_builder_tool(llm.clone(), Some(self.config.builder.to_builder_config()))
|
||||
.await;
|
||||
tracing::info!("Builder mode enabled");
|
||||
tracing::debug!("Builder mode enabled");
|
||||
}
|
||||
|
||||
Ok((safety, tools, embeddings, workspace))
|
||||
@@ -406,6 +375,7 @@ impl AppBuilder {
|
||||
) -> Result<
|
||||
(
|
||||
Arc<McpSessionManager>,
|
||||
Arc<McpProcessManager>,
|
||||
Option<Arc<WasmToolRuntime>>,
|
||||
Option<Arc<ExtensionManager>>,
|
||||
Vec<crate::extensions::RegistryEntry>,
|
||||
@@ -413,10 +383,11 @@ impl AppBuilder {
|
||||
),
|
||||
anyhow::Error,
|
||||
> {
|
||||
use crate::tools::mcp::{McpClient, config::load_mcp_servers_from_db, is_authenticated};
|
||||
use crate::tools::mcp::config::load_mcp_servers_from_db;
|
||||
use crate::tools::wasm::{WasmToolLoader, load_dev_tools};
|
||||
|
||||
let mcp_session_manager = Arc::new(McpSessionManager::new());
|
||||
let mcp_process_manager = Arc::new(McpProcessManager::new());
|
||||
|
||||
// Create WASM tool runtime eagerly so extensions installed after startup
|
||||
// (e.g. via the web UI) can still be activated. The tools directory is only
|
||||
@@ -448,7 +419,7 @@ impl AppBuilder {
|
||||
match loader.load_from_dir(&wasm_config.tools_dir).await {
|
||||
Ok(results) => {
|
||||
if !results.loaded.is_empty() {
|
||||
tracing::info!(
|
||||
tracing::debug!(
|
||||
"Loaded {} WASM tools from {}",
|
||||
results.loaded.len(),
|
||||
wasm_config.tools_dir.display()
|
||||
@@ -471,7 +442,7 @@ impl AppBuilder {
|
||||
Ok(results) => {
|
||||
dev_loaded_tool_names.extend(results.loaded.iter().cloned());
|
||||
if !dev_loaded_tool_names.is_empty() {
|
||||
tracing::info!(
|
||||
tracing::debug!(
|
||||
"Loaded {} dev WASM tools from build artifacts",
|
||||
dev_loaded_tool_names.len()
|
||||
);
|
||||
@@ -492,97 +463,107 @@ impl AppBuilder {
|
||||
let db = self.db.clone();
|
||||
let tools = Arc::clone(tools);
|
||||
let mcp_sm = Arc::clone(&mcp_session_manager);
|
||||
let pm = Arc::clone(&mcp_process_manager);
|
||||
async move {
|
||||
if let Some(ref secrets) = secrets_store {
|
||||
let servers_result = if let Some(ref d) = db {
|
||||
load_mcp_servers_from_db(d.as_ref(), "default").await
|
||||
} else {
|
||||
crate::tools::mcp::config::load_mcp_servers().await
|
||||
};
|
||||
match servers_result {
|
||||
Ok(servers) => {
|
||||
let enabled: Vec<_> = servers.enabled_servers().cloned().collect();
|
||||
if !enabled.is_empty() {
|
||||
tracing::info!(
|
||||
"Loading {} configured MCP server(s)...",
|
||||
enabled.len()
|
||||
);
|
||||
}
|
||||
let servers_result = if let Some(ref d) = db {
|
||||
load_mcp_servers_from_db(d.as_ref(), "default").await
|
||||
} else {
|
||||
crate::tools::mcp::config::load_mcp_servers().await
|
||||
};
|
||||
match servers_result {
|
||||
Ok(servers) => {
|
||||
let enabled: Vec<_> = servers.enabled_servers().cloned().collect();
|
||||
if !enabled.is_empty() {
|
||||
tracing::debug!(
|
||||
"Loading {} configured MCP server(s)...",
|
||||
enabled.len()
|
||||
);
|
||||
}
|
||||
|
||||
let mut join_set = tokio::task::JoinSet::new();
|
||||
for server in enabled {
|
||||
let mcp_sm = Arc::clone(&mcp_sm);
|
||||
let secrets = Arc::clone(secrets);
|
||||
let tools = Arc::clone(&tools);
|
||||
let mut join_set = tokio::task::JoinSet::new();
|
||||
for server in enabled {
|
||||
let mcp_sm = Arc::clone(&mcp_sm);
|
||||
let secrets = secrets_store.clone();
|
||||
let tools = Arc::clone(&tools);
|
||||
let pm = Arc::clone(&pm);
|
||||
|
||||
join_set.spawn(async move {
|
||||
let server_name = server.name.clone();
|
||||
let has_tokens =
|
||||
is_authenticated(&server, &secrets, "default").await;
|
||||
join_set.spawn(async move {
|
||||
let server_name = server.name.clone();
|
||||
|
||||
let client = if has_tokens || server.requires_auth() {
|
||||
McpClient::new_authenticated(
|
||||
server, mcp_sm, secrets, "default",
|
||||
)
|
||||
} else {
|
||||
McpClient::new_with_name(&server_name, &server.url)
|
||||
};
|
||||
let client = match crate::tools::mcp::create_client_from_config(
|
||||
server,
|
||||
&mcp_sm,
|
||||
&pm,
|
||||
secrets,
|
||||
"default",
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(c) => c,
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
"Failed to create MCP client for '{}': {}",
|
||||
server_name,
|
||||
e
|
||||
);
|
||||
return;
|
||||
}
|
||||
};
|
||||
|
||||
match client.list_tools().await {
|
||||
Ok(mcp_tools) => {
|
||||
let tool_count = mcp_tools.len();
|
||||
match client.create_tools().await {
|
||||
Ok(tool_impls) => {
|
||||
for tool in tool_impls {
|
||||
tools.register(tool).await;
|
||||
}
|
||||
tracing::info!(
|
||||
"Loaded {} tools from MCP server '{}'",
|
||||
tool_count,
|
||||
server_name
|
||||
);
|
||||
match client.list_tools().await {
|
||||
Ok(mcp_tools) => {
|
||||
let tool_count = mcp_tools.len();
|
||||
match client.create_tools().await {
|
||||
Ok(tool_impls) => {
|
||||
for tool in tool_impls {
|
||||
tools.register(tool).await;
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
"Failed to create tools from MCP server '{}': {}",
|
||||
server_name,
|
||||
e
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
let err_str = e.to_string();
|
||||
if err_str.contains("401")
|
||||
|| err_str.contains("authentication")
|
||||
{
|
||||
tracing::warn!(
|
||||
"MCP server '{}' requires authentication. \
|
||||
Run: ironclaw mcp auth {}",
|
||||
server_name,
|
||||
tracing::debug!(
|
||||
"Loaded {} tools from MCP server '{}'",
|
||||
tool_count,
|
||||
server_name
|
||||
);
|
||||
} else {
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
"Failed to connect to MCP server '{}': {}",
|
||||
"Failed to create tools from MCP server '{}': {}",
|
||||
server_name,
|
||||
e
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
while let Some(result) = join_set.join_next().await {
|
||||
if let Err(e) = result {
|
||||
tracing::warn!("MCP server loading task panicked: {}", e);
|
||||
Err(e) => {
|
||||
let err_str = e.to_string();
|
||||
if err_str.contains("401")
|
||||
|| err_str.contains("authentication")
|
||||
{
|
||||
tracing::warn!(
|
||||
"MCP server '{}' requires authentication. \
|
||||
Run: ironclaw mcp auth {}",
|
||||
server_name,
|
||||
server_name
|
||||
);
|
||||
} else {
|
||||
tracing::warn!(
|
||||
"Failed to connect to MCP server '{}': {}",
|
||||
server_name,
|
||||
e
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
while let Some(result) = join_set.join_next().await {
|
||||
if let Err(e) = result {
|
||||
tracing::warn!("MCP server loading task panicked: {}", e);
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::debug!("No MCP servers configured ({})", e);
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::debug!("No MCP servers configured ({})", e);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -598,7 +579,7 @@ impl AppBuilder {
|
||||
.iter()
|
||||
.map(|m| m.to_registry_entry())
|
||||
.collect();
|
||||
tracing::info!(
|
||||
tracing::debug!(
|
||||
count = entries.len(),
|
||||
"Loaded registry catalog entries for extension discovery"
|
||||
);
|
||||
@@ -627,6 +608,7 @@ impl AppBuilder {
|
||||
let extension_manager = {
|
||||
let manager = Arc::new(ExtensionManager::new(
|
||||
Arc::clone(&mcp_session_manager),
|
||||
Arc::clone(&mcp_process_manager),
|
||||
ext_secrets,
|
||||
Arc::clone(tools),
|
||||
Some(Arc::clone(hooks)),
|
||||
@@ -639,7 +621,7 @@ impl AppBuilder {
|
||||
catalog_entries.clone(),
|
||||
));
|
||||
tools.register_extension_tools(Arc::clone(&manager));
|
||||
tracing::info!("Extension manager initialized with in-chat discovery tools");
|
||||
tracing::debug!("Extension manager initialized with in-chat discovery tools");
|
||||
Some(manager)
|
||||
};
|
||||
|
||||
@@ -653,6 +635,7 @@ impl AppBuilder {
|
||||
|
||||
Ok((
|
||||
mcp_session_manager,
|
||||
mcp_process_manager,
|
||||
wasm_tool_runtime,
|
||||
extension_manager,
|
||||
catalog_entries,
|
||||
@@ -665,10 +648,21 @@ impl AppBuilder {
|
||||
self.init_database().await?;
|
||||
self.init_secrets().await?;
|
||||
|
||||
// Post-init validation: if a non-nearai backend was selected but
|
||||
// credentials were never resolved (deferred resolution found no keys),
|
||||
// fail early with a clear error instead of a confusing runtime failure.
|
||||
if self.config.llm.backend != "nearai" && self.config.llm.provider.is_none() {
|
||||
let backend = &self.config.llm.backend;
|
||||
anyhow::bail!(
|
||||
"LLM_BACKEND={backend} is configured but no credentials were found. \
|
||||
Set the appropriate API key environment variable or run the setup wizard."
|
||||
);
|
||||
}
|
||||
|
||||
let (llm, cheap_llm, recording_handle) = if let Some(llm) = self.llm_override.take() {
|
||||
(llm, None, None)
|
||||
} else {
|
||||
self.init_llm()?
|
||||
self.init_llm().await?
|
||||
};
|
||||
let (safety, tools, embeddings, workspace) = self.init_tools(&llm).await?;
|
||||
|
||||
@@ -677,6 +671,7 @@ impl AppBuilder {
|
||||
|
||||
let (
|
||||
mcp_session_manager,
|
||||
mcp_process_manager,
|
||||
wasm_tool_runtime,
|
||||
extension_manager,
|
||||
catalog_entries,
|
||||
@@ -697,7 +692,7 @@ impl AppBuilder {
|
||||
let import_path = std::path::Path::new(&import_dir);
|
||||
match ws.import_from_directory(import_path).await {
|
||||
Ok(count) if count > 0 => {
|
||||
tracing::info!("Imported {} workspace file(s) from {}", count, import_dir);
|
||||
tracing::debug!("Imported {} workspace file(s) from {}", count, import_dir);
|
||||
}
|
||||
Ok(_) => {}
|
||||
Err(e) => {
|
||||
@@ -722,7 +717,7 @@ impl AppBuilder {
|
||||
tokio::spawn(async move {
|
||||
match ws_bg.backfill_embeddings().await {
|
||||
Ok(count) if count > 0 => {
|
||||
tracing::info!("Backfilled embeddings for {} chunks", count);
|
||||
tracing::debug!("Backfilled embeddings for {} chunks", count);
|
||||
}
|
||||
Ok(_) => {}
|
||||
Err(e) => {
|
||||
@@ -739,7 +734,7 @@ impl AppBuilder {
|
||||
.with_installed_dir(self.config.skills.installed_dir.clone());
|
||||
let loaded = registry.discover_all().await;
|
||||
if !loaded.is_empty() {
|
||||
tracing::info!("Loaded {} skill(s): {}", loaded.len(), loaded.join(", "));
|
||||
tracing::debug!("Loaded {} skill(s): {}", loaded.len(), loaded.join(", "));
|
||||
}
|
||||
let registry = Arc::new(std::sync::RwLock::new(registry));
|
||||
let catalog = crate::skills::catalog::shared_catalog();
|
||||
@@ -757,7 +752,7 @@ impl AppBuilder {
|
||||
},
|
||||
));
|
||||
|
||||
tracing::info!(
|
||||
tracing::debug!(
|
||||
"Tool registry initialized with {} total tools",
|
||||
tools.count()
|
||||
);
|
||||
@@ -774,6 +769,7 @@ impl AppBuilder {
|
||||
workspace,
|
||||
extension_manager,
|
||||
mcp_session_manager,
|
||||
mcp_process_manager,
|
||||
wasm_tool_runtime,
|
||||
log_broadcaster: self.log_broadcaster,
|
||||
context_manager,
|
||||
|
||||
@@ -198,6 +198,58 @@ pub fn save_bootstrap_env_to(path: &std::path::Path, vars: &[(&str, &str)]) -> s
|
||||
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.
|
||||
///
|
||||
/// Unlike `save_bootstrap_env` (which overwrites the entire file), this
|
||||
@@ -414,10 +466,103 @@ pub enum MigrationError {
|
||||
Io(String),
|
||||
}
|
||||
|
||||
// ── PID Lock ──────────────────────────────────────────────────────────────
|
||||
|
||||
/// Path to the PID lock file: `~/.ironclaw/ironclaw.pid`.
|
||||
pub fn pid_lock_path() -> PathBuf {
|
||||
ironclaw_base_dir().join("ironclaw.pid")
|
||||
}
|
||||
|
||||
/// A PID-based lock that prevents multiple IronClaw instances from running
|
||||
/// simultaneously.
|
||||
///
|
||||
/// Uses `fs4::try_lock_exclusive()` for atomic locking (no TOCTOU race),
|
||||
/// then writes the current PID into the locked file for diagnostics.
|
||||
/// The OS-level lock is held for the lifetime of this struct and
|
||||
/// automatically released on drop (along with the PID file cleanup).
|
||||
#[derive(Debug)]
|
||||
pub struct PidLock {
|
||||
path: PathBuf,
|
||||
/// Held open to maintain the OS-level exclusive lock.
|
||||
_file: std::fs::File,
|
||||
}
|
||||
|
||||
/// Errors from PID lock acquisition.
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum PidLockError {
|
||||
#[error("Another IronClaw instance is already running (PID {pid})")]
|
||||
AlreadyRunning { pid: u32 },
|
||||
#[error("Failed to acquire PID lock: {0}")]
|
||||
Io(#[from] std::io::Error),
|
||||
}
|
||||
|
||||
impl PidLock {
|
||||
/// Try to acquire the PID lock.
|
||||
///
|
||||
/// Uses an exclusive file lock (`flock`/`LockFileEx`) so that two
|
||||
/// concurrent processes cannot both acquire the lock — no TOCTOU race.
|
||||
/// If the lock file exists but the holding process is gone (stale),
|
||||
/// the lock is reclaimed automatically by the OS.
|
||||
pub fn acquire() -> Result<Self, PidLockError> {
|
||||
Self::acquire_at(pid_lock_path())
|
||||
}
|
||||
|
||||
/// Acquire at a specific path (for testing).
|
||||
fn acquire_at(path: PathBuf) -> Result<Self, PidLockError> {
|
||||
use fs4::FileExt;
|
||||
use std::fs::OpenOptions;
|
||||
use std::io::Write;
|
||||
|
||||
// Ensure parent directory exists
|
||||
if let Some(parent) = path.parent() {
|
||||
std::fs::create_dir_all(parent)?;
|
||||
}
|
||||
|
||||
// Open (or create) the lock file
|
||||
let mut file = OpenOptions::new()
|
||||
.read(true)
|
||||
.write(true)
|
||||
.create(true)
|
||||
.truncate(false)
|
||||
.open(&path)?;
|
||||
|
||||
// Try non-blocking exclusive lock — if another process holds it,
|
||||
// this fails immediately instead of blocking.
|
||||
if let Err(e) = file.try_lock_exclusive() {
|
||||
if e.kind() == std::io::ErrorKind::WouldBlock {
|
||||
// Lock held by another process — read its PID for the error message
|
||||
let pid = std::fs::read_to_string(&path)
|
||||
.ok()
|
||||
.and_then(|s| s.trim().parse::<u32>().ok())
|
||||
.unwrap_or(0);
|
||||
return Err(PidLockError::AlreadyRunning { pid });
|
||||
}
|
||||
// Other errors (permissions, unsupported filesystem, etc.)
|
||||
return Err(PidLockError::Io(e));
|
||||
}
|
||||
|
||||
// We hold the exclusive lock — write our PID
|
||||
file.set_len(0)?; // truncate
|
||||
write!(file, "{}", std::process::id())?;
|
||||
|
||||
Ok(PidLock { path, _file: file })
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for PidLock {
|
||||
fn drop(&mut self) {
|
||||
// Remove the PID file; the OS-level lock is released when _file is dropped.
|
||||
let _ = std::fs::remove_file(&self.path);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::process::Command;
|
||||
use std::sync::Mutex;
|
||||
use std::thread;
|
||||
use std::time::{Duration, Instant};
|
||||
use tempfile::tempdir;
|
||||
|
||||
static ENV_MUTEX: Mutex<()> = Mutex::new(());
|
||||
@@ -986,4 +1131,266 @@ INJECTED="pwned"#;
|
||||
unsafe { std::env::remove_var("IRONCLAW_BASE_DIR") };
|
||||
}
|
||||
}
|
||||
|
||||
// ── PID Lock tests ───────────────────────────────────────────────
|
||||
|
||||
#[test]
|
||||
fn test_pid_lock_acquire_and_drop() {
|
||||
let dir = tempdir().unwrap();
|
||||
let pid_path = dir.path().join("ironclaw.pid");
|
||||
|
||||
// Acquire lock
|
||||
let lock = PidLock::acquire_at(pid_path.clone()).unwrap();
|
||||
assert!(pid_path.exists());
|
||||
|
||||
// PID file should contain our PID
|
||||
let contents = std::fs::read_to_string(&pid_path).unwrap();
|
||||
assert_eq!(contents.trim().parse::<u32>().unwrap(), std::process::id());
|
||||
|
||||
// Drop should remove the file
|
||||
drop(lock);
|
||||
assert!(!pid_path.exists());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_pid_lock_rejects_second_acquire() {
|
||||
let dir = tempdir().unwrap();
|
||||
let pid_path = dir.path().join("ironclaw.pid");
|
||||
|
||||
// First lock succeeds
|
||||
let _lock1 = PidLock::acquire_at(pid_path.clone()).unwrap();
|
||||
|
||||
// Second acquire on same file must fail (exclusive flock held)
|
||||
let result = PidLock::acquire_at(pid_path.clone());
|
||||
assert!(result.is_err());
|
||||
match result.unwrap_err() {
|
||||
PidLockError::AlreadyRunning { pid } => {
|
||||
assert_eq!(pid, std::process::id());
|
||||
}
|
||||
other => panic!("expected AlreadyRunning, got: {}", other),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_pid_lock_reclaims_after_drop() {
|
||||
let dir = tempdir().unwrap();
|
||||
let pid_path = dir.path().join("ironclaw.pid");
|
||||
|
||||
// Acquire and release
|
||||
let lock = PidLock::acquire_at(pid_path.clone()).unwrap();
|
||||
drop(lock);
|
||||
|
||||
// Should succeed — OS lock was released on drop
|
||||
let lock2 = PidLock::acquire_at(pid_path).unwrap();
|
||||
drop(lock2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_pid_lock_reclaims_stale_file_without_flock() {
|
||||
let dir = tempdir().unwrap();
|
||||
let pid_path = dir.path().join("ironclaw.pid");
|
||||
|
||||
// Write a stale PID file manually (no flock held)
|
||||
std::fs::write(&pid_path, "4294967294").unwrap();
|
||||
|
||||
// Should succeed because no OS lock is held on the file
|
||||
let lock = PidLock::acquire_at(pid_path.clone()).unwrap();
|
||||
let contents = std::fs::read_to_string(&pid_path).unwrap();
|
||||
assert_eq!(contents.trim().parse::<u32>().unwrap(), std::process::id());
|
||||
drop(lock);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_pid_lock_handles_corrupt_pid_file() {
|
||||
let dir = tempdir().unwrap();
|
||||
let pid_path = dir.path().join("ironclaw.pid");
|
||||
|
||||
// Write garbage (no flock held)
|
||||
std::fs::write(&pid_path, "not-a-number").unwrap();
|
||||
|
||||
// Should succeed — no OS lock held, file is reclaimed
|
||||
let lock = PidLock::acquire_at(pid_path).unwrap();
|
||||
drop(lock);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_pid_lock_creates_parent_dirs() {
|
||||
let dir = tempdir().unwrap();
|
||||
let pid_path = dir.path().join("nested").join("deep").join("ironclaw.pid");
|
||||
|
||||
let lock = PidLock::acquire_at(pid_path.clone()).unwrap();
|
||||
assert!(pid_path.exists());
|
||||
drop(lock);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_pid_lock_child_helper_holds_lock() {
|
||||
if std::env::var("IRONCLAW_PID_LOCK_CHILD").ok().as_deref() != Some("1") {
|
||||
return;
|
||||
}
|
||||
|
||||
let pid_path = PathBuf::from(
|
||||
std::env::var("IRONCLAW_PID_LOCK_PATH").expect("IRONCLAW_PID_LOCK_PATH missing"),
|
||||
);
|
||||
let hold_ms = std::env::var("IRONCLAW_PID_LOCK_HOLD_MS")
|
||||
.ok()
|
||||
.and_then(|s| s.parse::<u64>().ok())
|
||||
.unwrap_or(3000);
|
||||
|
||||
let _lock = PidLock::acquire_at(pid_path).expect("child failed to acquire pid lock");
|
||||
thread::sleep(Duration::from_millis(hold_ms));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_pid_lock_rejects_lock_held_by_other_process() {
|
||||
let dir = tempdir().unwrap();
|
||||
let pid_path = dir.path().join("ironclaw.pid");
|
||||
|
||||
let current_exe = std::env::current_exe().unwrap();
|
||||
let mut child = Command::new(current_exe)
|
||||
.args([
|
||||
"--exact",
|
||||
"bootstrap::tests::test_pid_lock_child_helper_holds_lock",
|
||||
"--nocapture",
|
||||
"--test-threads=1",
|
||||
])
|
||||
.env("IRONCLAW_PID_LOCK_CHILD", "1")
|
||||
.env("IRONCLAW_PID_LOCK_PATH", pid_path.display().to_string())
|
||||
.env("IRONCLAW_PID_LOCK_HOLD_MS", "3000")
|
||||
.spawn()
|
||||
.unwrap();
|
||||
|
||||
let started = Instant::now();
|
||||
while started.elapsed() < Duration::from_secs(2) {
|
||||
if pid_path.exists() {
|
||||
break;
|
||||
}
|
||||
if let Some(status) = child.try_wait().unwrap() {
|
||||
panic!("child exited before acquiring lock: {}", status);
|
||||
}
|
||||
thread::sleep(Duration::from_millis(20));
|
||||
}
|
||||
assert!(
|
||||
pid_path.exists(),
|
||||
"child did not create lock file in time: {}",
|
||||
pid_path.display()
|
||||
);
|
||||
|
||||
let result = PidLock::acquire_at(pid_path.clone());
|
||||
match result.unwrap_err() {
|
||||
PidLockError::AlreadyRunning { .. } => {}
|
||||
other => panic!("expected AlreadyRunning, got: {}", other),
|
||||
}
|
||||
|
||||
let status = child.wait().unwrap();
|
||||
assert!(status.success(), "child process failed: {}", status);
|
||||
|
||||
// After the child exits, lock should be released and reacquirable.
|
||||
let lock = PidLock::acquire_at(pid_path).unwrap();
|
||||
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())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -10,6 +10,56 @@ use uuid::Uuid;
|
||||
|
||||
use crate::error::ChannelError;
|
||||
|
||||
/// Kind of attachment carried on an incoming message.
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum AttachmentKind {
|
||||
/// Audio content (voice notes, audio files).
|
||||
Audio,
|
||||
/// Image content (photos, screenshots).
|
||||
Image,
|
||||
/// Document content (PDFs, files).
|
||||
Document,
|
||||
}
|
||||
|
||||
impl AttachmentKind {
|
||||
/// Infer attachment kind from MIME type.
|
||||
pub fn from_mime_type(mime: &str) -> Self {
|
||||
let base = mime.split(';').next().unwrap_or(mime).trim();
|
||||
if base.starts_with("audio/") {
|
||||
Self::Audio
|
||||
} else if base.starts_with("image/") {
|
||||
Self::Image
|
||||
} else {
|
||||
Self::Document
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// A file or media attachment on an incoming message.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct IncomingAttachment {
|
||||
/// Unique identifier within the channel (e.g., Telegram file_id).
|
||||
pub id: String,
|
||||
/// What kind of content this is.
|
||||
pub kind: AttachmentKind,
|
||||
/// MIME type (e.g., "image/jpeg", "audio/ogg", "application/pdf").
|
||||
pub mime_type: String,
|
||||
/// Original filename, if known.
|
||||
pub filename: Option<String>,
|
||||
/// File size in bytes, if known.
|
||||
pub size_bytes: Option<u64>,
|
||||
/// URL to download the file from the channel's API.
|
||||
pub source_url: Option<String>,
|
||||
/// Opaque key for host-side storage (e.g., after download/caching).
|
||||
pub storage_key: Option<String>,
|
||||
/// Extracted text content (e.g., OCR result, PDF text, audio transcript).
|
||||
pub extracted_text: Option<String>,
|
||||
/// Raw file bytes (for small files downloaded by the channel).
|
||||
pub data: Vec<u8>,
|
||||
/// Duration in seconds (for audio/video).
|
||||
pub duration_secs: Option<u32>,
|
||||
}
|
||||
|
||||
/// A message received from an external channel.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct IncomingMessage {
|
||||
@@ -29,6 +79,10 @@ pub struct IncomingMessage {
|
||||
pub received_at: DateTime<Utc>,
|
||||
/// Channel-specific metadata.
|
||||
pub metadata: serde_json::Value,
|
||||
/// IANA timezone string from the client (e.g. "America/New_York").
|
||||
pub timezone: Option<String>,
|
||||
/// File or media attachments on this message.
|
||||
pub attachments: Vec<IncomingAttachment>,
|
||||
}
|
||||
|
||||
impl IncomingMessage {
|
||||
@@ -47,6 +101,8 @@ impl IncomingMessage {
|
||||
thread_id: None,
|
||||
received_at: Utc::now(),
|
||||
metadata: serde_json::Value::Null,
|
||||
timezone: None,
|
||||
attachments: Vec::new(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -67,6 +123,18 @@ impl IncomingMessage {
|
||||
self.user_name = Some(name.into());
|
||||
self
|
||||
}
|
||||
|
||||
/// Set the client timezone.
|
||||
pub fn with_timezone(mut self, tz: impl Into<String>) -> Self {
|
||||
self.timezone = Some(tz.into());
|
||||
self
|
||||
}
|
||||
|
||||
/// Set attachments.
|
||||
pub fn with_attachments(mut self, attachments: Vec<IncomingAttachment>) -> Self {
|
||||
self.attachments = attachments;
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
/// Stream of incoming messages.
|
||||
@@ -163,6 +231,13 @@ pub enum StatusUpdate {
|
||||
success: bool,
|
||||
message: String,
|
||||
},
|
||||
/// An image was generated by a tool.
|
||||
ImageGenerated {
|
||||
/// Base64 data URL of the generated image.
|
||||
data_url: String,
|
||||
/// Optional workspace path where the image was saved.
|
||||
path: Option<String>,
|
||||
},
|
||||
}
|
||||
|
||||
impl StatusUpdate {
|
||||
@@ -395,4 +470,10 @@ mod tests {
|
||||
panic!("expected ToolCompleted variant");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_incoming_message_with_timezone() {
|
||||
let msg = IncomingMessage::new("test", "user1", "hello").with_timezone("America/New_York");
|
||||
assert_eq!(msg.timezone.as_deref(), Some("America/New_York"));
|
||||
}
|
||||
}
|
||||
|
||||
+126
-6
@@ -17,7 +17,9 @@ use tokio::sync::{RwLock, mpsc, oneshot};
|
||||
use tokio_stream::wrappers::ReceiverStream;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::channels::{Channel, IncomingMessage, MessageStream, OutgoingResponse};
|
||||
use crate::channels::{
|
||||
AttachmentKind, Channel, IncomingAttachment, IncomingMessage, MessageStream, OutgoingResponse,
|
||||
};
|
||||
use crate::config::HttpConfig;
|
||||
use crate::error::ChannelError;
|
||||
|
||||
@@ -46,8 +48,9 @@ struct RateLimitState {
|
||||
request_count: u32,
|
||||
}
|
||||
|
||||
/// Maximum JSON body size for webhook requests (64 KB).
|
||||
const MAX_BODY_BYTES: usize = 64 * 1024;
|
||||
/// Maximum JSON body size for webhook requests (15 MB, to support base64 image attachments
|
||||
/// with ~33% overhead from base64 encoding).
|
||||
const MAX_BODY_BYTES: usize = 15 * 1024 * 1024;
|
||||
|
||||
/// Maximum number of pending wait-for-response requests.
|
||||
const MAX_PENDING_RESPONSES: usize = 100;
|
||||
@@ -115,8 +118,34 @@ struct WebhookRequest {
|
||||
/// Whether to wait for a synchronous response.
|
||||
#[serde(default)]
|
||||
wait_for_response: bool,
|
||||
/// Optional file attachments (base64-encoded).
|
||||
#[serde(default)]
|
||||
attachments: Vec<AttachmentData>,
|
||||
}
|
||||
|
||||
/// A file attachment in a webhook request.
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct AttachmentData {
|
||||
/// MIME type (e.g. "image/png", "application/pdf").
|
||||
mime_type: String,
|
||||
/// Optional filename.
|
||||
#[serde(default)]
|
||||
filename: Option<String>,
|
||||
/// Base64-encoded file data.
|
||||
#[serde(default)]
|
||||
data_base64: Option<String>,
|
||||
/// URL to fetch the file from (not downloaded server-side for SSRF prevention).
|
||||
#[serde(default)]
|
||||
url: Option<String>,
|
||||
}
|
||||
|
||||
/// Maximum size per attachment (5 MB decoded).
|
||||
const MAX_ATTACHMENT_BYTES: usize = 5 * 1024 * 1024;
|
||||
/// Maximum total attachment size (10 MB decoded).
|
||||
const MAX_TOTAL_ATTACHMENT_BYTES: usize = 10 * 1024 * 1024;
|
||||
/// Maximum number of attachments per request.
|
||||
const MAX_ATTACHMENTS: usize = 5;
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
struct WebhookResponse {
|
||||
/// Message ID assigned to this request.
|
||||
@@ -211,15 +240,106 @@ async fn webhook_handler(
|
||||
);
|
||||
}
|
||||
|
||||
let msg = IncomingMessage::new("http", &state.user_id, &req.content).with_metadata(
|
||||
// Validate and decode attachments
|
||||
let attachments = if !req.attachments.is_empty() {
|
||||
if req.attachments.len() > MAX_ATTACHMENTS {
|
||||
return (
|
||||
StatusCode::BAD_REQUEST,
|
||||
Json(WebhookResponse {
|
||||
message_id: Uuid::nil(),
|
||||
status: "error".to_string(),
|
||||
response: Some(format!("Too many attachments (max {})", MAX_ATTACHMENTS)),
|
||||
}),
|
||||
);
|
||||
}
|
||||
|
||||
let mut decoded_attachments = Vec::new();
|
||||
let mut total_bytes: usize = 0;
|
||||
for att in &req.attachments {
|
||||
if let Some(ref b64) = att.data_base64 {
|
||||
use base64::Engine;
|
||||
let data = match base64::engine::general_purpose::STANDARD.decode(b64) {
|
||||
Ok(d) => d,
|
||||
Err(_) => {
|
||||
return (
|
||||
StatusCode::BAD_REQUEST,
|
||||
Json(WebhookResponse {
|
||||
message_id: Uuid::nil(),
|
||||
status: "error".to_string(),
|
||||
response: Some("Invalid base64 in attachment".to_string()),
|
||||
}),
|
||||
);
|
||||
}
|
||||
};
|
||||
if data.len() > MAX_ATTACHMENT_BYTES {
|
||||
return (
|
||||
StatusCode::PAYLOAD_TOO_LARGE,
|
||||
Json(WebhookResponse {
|
||||
message_id: Uuid::nil(),
|
||||
status: "error".to_string(),
|
||||
response: Some(format!(
|
||||
"Attachment too large (max {} bytes)",
|
||||
MAX_ATTACHMENT_BYTES
|
||||
)),
|
||||
}),
|
||||
);
|
||||
}
|
||||
total_bytes += data.len();
|
||||
if total_bytes > MAX_TOTAL_ATTACHMENT_BYTES {
|
||||
return (
|
||||
StatusCode::PAYLOAD_TOO_LARGE,
|
||||
Json(WebhookResponse {
|
||||
message_id: Uuid::nil(),
|
||||
status: "error".to_string(),
|
||||
response: Some("Total attachment size exceeds limit".to_string()),
|
||||
}),
|
||||
);
|
||||
}
|
||||
decoded_attachments.push(IncomingAttachment {
|
||||
id: Uuid::new_v4().to_string(),
|
||||
kind: AttachmentKind::from_mime_type(&att.mime_type),
|
||||
mime_type: att.mime_type.clone(),
|
||||
filename: att.filename.clone(),
|
||||
size_bytes: Some(data.len() as u64),
|
||||
source_url: None,
|
||||
storage_key: None,
|
||||
extracted_text: None,
|
||||
data,
|
||||
duration_secs: None,
|
||||
});
|
||||
} else if let Some(ref url) = att.url {
|
||||
// URL-only attachment: set source_url but don't download (SSRF prevention)
|
||||
decoded_attachments.push(IncomingAttachment {
|
||||
id: Uuid::new_v4().to_string(),
|
||||
kind: AttachmentKind::from_mime_type(&att.mime_type),
|
||||
mime_type: att.mime_type.clone(),
|
||||
filename: att.filename.clone(),
|
||||
size_bytes: None,
|
||||
source_url: Some(url.clone()),
|
||||
storage_key: None,
|
||||
extracted_text: None,
|
||||
data: Vec::new(),
|
||||
duration_secs: None,
|
||||
});
|
||||
}
|
||||
}
|
||||
decoded_attachments
|
||||
} else {
|
||||
Vec::new()
|
||||
};
|
||||
|
||||
let mut msg = IncomingMessage::new("http", &state.user_id, &req.content).with_metadata(
|
||||
serde_json::json!({
|
||||
"wait_for_response": req.wait_for_response,
|
||||
}),
|
||||
);
|
||||
|
||||
if !attachments.is_empty() {
|
||||
msg = msg.with_attachments(attachments);
|
||||
}
|
||||
|
||||
if let Some(thread_id) = &req.thread_id {
|
||||
let msg = msg.with_thread(thread_id);
|
||||
return process_message(state, msg, req.wait_for_response).await;
|
||||
msg = msg.with_thread(thread_id);
|
||||
}
|
||||
|
||||
process_message(state, msg, req.wait_for_response).await
|
||||
|
||||
+105
-2
@@ -75,7 +75,7 @@ impl ChannelManager {
|
||||
break;
|
||||
}
|
||||
}
|
||||
tracing::info!(channel = %name, "Hot-added channel stream ended");
|
||||
tracing::debug!(channel = %name, "Hot-added channel stream ended");
|
||||
});
|
||||
|
||||
Ok(())
|
||||
@@ -92,7 +92,7 @@ impl ChannelManager {
|
||||
for (name, channel) in channels.iter() {
|
||||
match channel.start().await {
|
||||
Ok(stream) => {
|
||||
tracing::info!("Started channel: {}", name);
|
||||
tracing::debug!("Started channel: {}", name);
|
||||
streams.push(stream);
|
||||
}
|
||||
Err(e) => {
|
||||
@@ -235,3 +235,106 @@ impl Default for ChannelManager {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::channels::IncomingMessage;
|
||||
use crate::testing::StubChannel;
|
||||
use futures::StreamExt;
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_add_and_start_all() {
|
||||
let manager = ChannelManager::new();
|
||||
let (stub, sender) = StubChannel::new("test");
|
||||
|
||||
manager.add(Box::new(stub)).await;
|
||||
|
||||
let mut stream = manager.start_all().await.expect("start_all failed");
|
||||
|
||||
// Inject a message through the stub
|
||||
sender
|
||||
.send(IncomingMessage::new("test", "user1", "hello"))
|
||||
.await
|
||||
.expect("send failed");
|
||||
|
||||
// Should appear in the merged stream
|
||||
let msg = stream.next().await.expect("stream ended");
|
||||
assert_eq!(msg.content, "hello");
|
||||
assert_eq!(msg.channel, "test");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_respond_routes_to_correct_channel() {
|
||||
let manager = ChannelManager::new();
|
||||
let (stub, _sender) = StubChannel::new("alpha");
|
||||
|
||||
// Keep a reference for response inspection
|
||||
let responses = stub.captured_responses_handle();
|
||||
manager.add(Box::new(stub)).await;
|
||||
|
||||
let msg = IncomingMessage::new("alpha", "user1", "request");
|
||||
manager
|
||||
.respond(&msg, OutgoingResponse::text("reply"))
|
||||
.await
|
||||
.expect("respond failed");
|
||||
|
||||
// Verify the stub captured the response
|
||||
let captured = responses.lock().expect("poisoned");
|
||||
assert_eq!(captured.len(), 1);
|
||||
assert_eq!(captured[0].1.content, "reply");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_respond_unknown_channel_errors() {
|
||||
let manager = ChannelManager::new();
|
||||
let msg = IncomingMessage::new("nonexistent", "user1", "test");
|
||||
let result = manager.respond(&msg, OutgoingResponse::text("hi")).await;
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_health_check_all() {
|
||||
let manager = ChannelManager::new();
|
||||
let (stub1, _) = StubChannel::new("healthy");
|
||||
let (stub2, _) = StubChannel::new("sick");
|
||||
stub2.set_healthy(false);
|
||||
|
||||
manager.add(Box::new(stub1)).await;
|
||||
manager.add(Box::new(stub2)).await;
|
||||
|
||||
let results = manager.health_check_all().await;
|
||||
assert!(results["healthy"].is_ok());
|
||||
assert!(results["sick"].is_err());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_start_all_no_channels_errors() {
|
||||
let manager = ChannelManager::new();
|
||||
let result = manager.start_all().await;
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_injection_channel_merges() {
|
||||
let manager = ChannelManager::new();
|
||||
let (stub, _sender) = StubChannel::new("real");
|
||||
manager.add(Box::new(stub)).await;
|
||||
|
||||
let mut stream = manager.start_all().await.expect("start_all failed");
|
||||
|
||||
// Use the injection channel (simulating background task)
|
||||
let inject_tx = manager.inject_sender();
|
||||
inject_tx
|
||||
.send(IncomingMessage::new(
|
||||
"injected",
|
||||
"system",
|
||||
"background alert",
|
||||
))
|
||||
.await
|
||||
.expect("inject failed");
|
||||
|
||||
let msg = stream.next().await.expect("stream ended");
|
||||
assert_eq!(msg.content, "background alert");
|
||||
}
|
||||
}
|
||||
|
||||
+4
-1
@@ -36,7 +36,10 @@ pub mod wasm;
|
||||
pub mod web;
|
||||
mod webhook_server;
|
||||
|
||||
pub use channel::{Channel, IncomingMessage, MessageStream, OutgoingResponse, StatusUpdate};
|
||||
pub use channel::{
|
||||
AttachmentKind, Channel, IncomingAttachment, IncomingMessage, MessageStream, OutgoingResponse,
|
||||
StatusUpdate,
|
||||
};
|
||||
pub use http::HttpChannel;
|
||||
pub use manager::ChannelManager;
|
||||
pub use repl::ReplChannel;
|
||||
|
||||
+57
-9
@@ -18,7 +18,7 @@
|
||||
//! - `Esc` - Interrupt current operation
|
||||
|
||||
use std::borrow::Cow;
|
||||
use std::io::{self, Write};
|
||||
use std::io::{self, IsTerminal, Write};
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
|
||||
@@ -297,10 +297,15 @@ impl Channel for ReplChannel {
|
||||
let esc_interrupt_triggered_for_thread = Arc::new(AtomicBool::new(false));
|
||||
|
||||
std::thread::spawn(move || {
|
||||
let sys_tz = crate::timezone::detect_system_timezone().name().to_string();
|
||||
|
||||
// Single message mode: send it and return
|
||||
if let Some(msg) = single_message {
|
||||
let incoming = IncomingMessage::new("repl", "default", &msg);
|
||||
let incoming = IncomingMessage::new("repl", "default", &msg).with_timezone(&sys_tz);
|
||||
let _ = tx.blocking_send(incoming);
|
||||
// Ensure the agent exits after handling exactly one turn in -m mode,
|
||||
// even when other channels (gateway/http) are enabled.
|
||||
let _ = tx.blocking_send(IncomingMessage::new("repl", "default", "/quit"));
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -361,7 +366,8 @@ impl Channel for ReplChannel {
|
||||
"/quit" | "/exit" => {
|
||||
// Forward shutdown command so the agent loop exits even
|
||||
// when other channels (e.g. web gateway) are still active.
|
||||
let msg = IncomingMessage::new("repl", "default", "/quit");
|
||||
let msg = IncomingMessage::new("repl", "default", "/quit")
|
||||
.with_timezone(&sys_tz);
|
||||
let _ = tx.blocking_send(msg);
|
||||
break;
|
||||
}
|
||||
@@ -382,7 +388,8 @@ impl Channel for ReplChannel {
|
||||
_ => {}
|
||||
}
|
||||
|
||||
let msg = IncomingMessage::new("repl", "default", line);
|
||||
let msg =
|
||||
IncomingMessage::new("repl", "default", line).with_timezone(&sys_tz);
|
||||
if tx.blocking_send(msg).is_err() {
|
||||
break;
|
||||
}
|
||||
@@ -390,21 +397,29 @@ impl Channel for ReplChannel {
|
||||
Err(ReadlineError::Interrupted) => {
|
||||
if esc_interrupt_triggered_for_thread.swap(false, Ordering::Relaxed) {
|
||||
// Esc: interrupt current operation and keep REPL open.
|
||||
let msg = IncomingMessage::new("repl", "default", "/interrupt");
|
||||
let msg = IncomingMessage::new("repl", "default", "/interrupt")
|
||||
.with_timezone(&sys_tz);
|
||||
if tx.blocking_send(msg).is_err() {
|
||||
break;
|
||||
}
|
||||
} else {
|
||||
// Ctrl+C (VINTR): request graceful shutdown.
|
||||
let msg = IncomingMessage::new("repl", "default", "/quit");
|
||||
let msg = IncomingMessage::new("repl", "default", "/quit")
|
||||
.with_timezone(&sys_tz);
|
||||
let _ = tx.blocking_send(msg);
|
||||
break;
|
||||
}
|
||||
}
|
||||
Err(ReadlineError::Eof) => {
|
||||
// Ctrl+D: send /quit so the agent loop runs graceful shutdown
|
||||
let msg = IncomingMessage::new("repl", "default", "/quit");
|
||||
let _ = tx.blocking_send(msg);
|
||||
// Ctrl+D in interactive mode: graceful shutdown.
|
||||
// In daemon mode (stdin = /dev/null, no TTY), EOF arrives
|
||||
// immediately — just drop the REPL thread silently so other
|
||||
// channels (gateway, telegram, …) keep running.
|
||||
if std::io::stdin().is_terminal() {
|
||||
let msg = IncomingMessage::new("repl", "default", "/quit")
|
||||
.with_timezone(&sys_tz);
|
||||
let _ = tx.blocking_send(msg);
|
||||
}
|
||||
break;
|
||||
}
|
||||
Err(e) => {
|
||||
@@ -585,6 +600,13 @@ impl Channel for ReplChannel {
|
||||
eprintln!("\x1b[31m {extension_name}: {message}\x1b[0m");
|
||||
}
|
||||
}
|
||||
StatusUpdate::ImageGenerated { path, .. } => {
|
||||
if let Some(ref p) = path {
|
||||
eprintln!("\x1b[36m [image] {p}\x1b[0m");
|
||||
} else {
|
||||
eprintln!("\x1b[36m [image generated]\x1b[0m");
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
@@ -614,3 +636,29 @@ impl Channel for ReplChannel {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use futures::StreamExt;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn single_message_mode_sends_message_then_quit() {
|
||||
let repl = ReplChannel::with_message("hi".to_string());
|
||||
let mut stream = repl.start().await.expect("repl start should succeed");
|
||||
|
||||
let first = stream.next().await.expect("first message missing");
|
||||
assert_eq!(first.channel, "repl");
|
||||
assert_eq!(first.content, "hi");
|
||||
|
||||
let second = stream.next().await.expect("quit message missing");
|
||||
assert_eq!(second.channel, "repl");
|
||||
assert_eq!(second.content, "/quit");
|
||||
|
||||
assert!(
|
||||
stream.next().await.is_none(),
|
||||
"stream should end after /quit"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
+329
-1
@@ -5,6 +5,7 @@
|
||||
//! - Workspace write access (scoped to channel namespace)
|
||||
//! - Rate limiting for message emission
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use crate::channels::wasm::capabilities::{ChannelCapabilities, EmitRateLimitConfig};
|
||||
@@ -17,6 +18,52 @@ const MAX_EMITS_PER_EXECUTION: usize = 100;
|
||||
/// Maximum message content size (64 KB).
|
||||
const MAX_MESSAGE_CONTENT_SIZE: usize = 64 * 1024;
|
||||
|
||||
/// A file or media attachment on an incoming message.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct Attachment {
|
||||
/// Unique identifier within the channel (e.g., Telegram file_id).
|
||||
pub id: String,
|
||||
/// MIME type (e.g., "image/jpeg", "audio/ogg", "application/pdf").
|
||||
pub mime_type: String,
|
||||
/// Original filename, if known.
|
||||
pub filename: Option<String>,
|
||||
/// File size in bytes, if known.
|
||||
pub size_bytes: Option<u64>,
|
||||
/// URL to download the file from the channel's API.
|
||||
pub source_url: Option<String>,
|
||||
/// Opaque key for host-side storage (e.g., after download/caching).
|
||||
pub storage_key: Option<String>,
|
||||
/// Extracted text content (e.g., OCR result, PDF text, audio transcript).
|
||||
pub extracted_text: Option<String>,
|
||||
/// Raw file bytes (for small files downloaded by the channel).
|
||||
pub data: Vec<u8>,
|
||||
/// Duration in seconds (for audio/video).
|
||||
pub duration_secs: Option<u32>,
|
||||
}
|
||||
|
||||
/// Maximum total attachment size per message (20 MB).
|
||||
const MAX_ATTACHMENT_TOTAL_SIZE: u64 = 20 * 1024 * 1024;
|
||||
|
||||
/// Maximum number of attachments per message.
|
||||
const MAX_ATTACHMENTS_PER_MESSAGE: usize = 10;
|
||||
|
||||
/// Allowed MIME type prefixes for attachments.
|
||||
const ALLOWED_MIME_PREFIXES: &[&str] = &[
|
||||
"image/",
|
||||
"audio/",
|
||||
"video/",
|
||||
"application/pdf",
|
||||
"application/vnd.",
|
||||
"application/msword",
|
||||
"application/rtf",
|
||||
"text/",
|
||||
"application/json",
|
||||
"application/zip",
|
||||
"application/gzip",
|
||||
"application/x-tar",
|
||||
"application/octet-stream",
|
||||
];
|
||||
|
||||
/// A message emitted by a WASM channel to be sent to the agent.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct EmittedMessage {
|
||||
@@ -35,6 +82,9 @@ pub struct EmittedMessage {
|
||||
/// Channel-specific metadata as JSON string.
|
||||
pub metadata_json: String,
|
||||
|
||||
/// File or media attachments on this message.
|
||||
pub attachments: Vec<Attachment>,
|
||||
|
||||
/// Timestamp when the message was emitted.
|
||||
pub emitted_at_millis: u64,
|
||||
}
|
||||
@@ -48,6 +98,7 @@ impl EmittedMessage {
|
||||
content: content.into(),
|
||||
thread_id: None,
|
||||
metadata_json: "{}".to_string(),
|
||||
attachments: Vec::new(),
|
||||
emitted_at_millis: SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map(|d| d.as_millis() as u64)
|
||||
@@ -72,6 +123,12 @@ impl EmittedMessage {
|
||||
self.metadata_json = metadata_json.into();
|
||||
self
|
||||
}
|
||||
|
||||
/// Set attachments.
|
||||
pub fn with_attachments(mut self, attachments: Vec<Attachment>) -> Self {
|
||||
self.attachments = attachments;
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
/// A pending workspace write operation.
|
||||
@@ -112,6 +169,13 @@ pub struct ChannelHostState {
|
||||
|
||||
/// Count of emits dropped due to rate limiting.
|
||||
emits_dropped: usize,
|
||||
|
||||
/// Binary data stored for attachments via `store-attachment-data`.
|
||||
/// Keyed by attachment ID, cleared after callback completes.
|
||||
attachment_data: HashMap<String, Vec<u8>>,
|
||||
|
||||
/// Total bytes stored in attachment_data (for enforcing limits).
|
||||
attachment_data_total: u64,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for ChannelHostState {
|
||||
@@ -141,6 +205,8 @@ impl ChannelHostState {
|
||||
emit_count: 0,
|
||||
emit_enabled: true,
|
||||
emits_dropped: 0,
|
||||
attachment_data: HashMap::new(),
|
||||
attachment_data_total: 0,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -168,6 +234,7 @@ impl ChannelHostState {
|
||||
///
|
||||
/// Messages are queued and delivered after callback execution completes.
|
||||
/// Rate limiting is enforced per-execution and globally.
|
||||
/// Attachments are validated for count, total size, and MIME type.
|
||||
pub fn emit_message(&mut self, msg: EmittedMessage) -> Result<(), WasmChannelError> {
|
||||
// Check per-execution limit
|
||||
if !self.emit_enabled {
|
||||
@@ -186,6 +253,9 @@ impl ChannelHostState {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
// Validate attachments
|
||||
let msg = self.validate_attachments(msg);
|
||||
|
||||
// Validate message content size
|
||||
if msg.content.len() > MAX_MESSAGE_CONTENT_SIZE {
|
||||
tracing::warn!(
|
||||
@@ -209,6 +279,71 @@ impl ChannelHostState {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Validate and sanitize attachments on an emitted message.
|
||||
///
|
||||
/// Enforces count limits, total size limits, and MIME type allowlist.
|
||||
/// Invalid attachments are dropped with a warning.
|
||||
fn validate_attachments(&self, mut msg: EmittedMessage) -> EmittedMessage {
|
||||
if msg.attachments.is_empty() {
|
||||
return msg;
|
||||
}
|
||||
|
||||
// Enforce attachment count limit
|
||||
if msg.attachments.len() > MAX_ATTACHMENTS_PER_MESSAGE {
|
||||
tracing::warn!(
|
||||
channel = %self.channel_name,
|
||||
count = msg.attachments.len(),
|
||||
max = MAX_ATTACHMENTS_PER_MESSAGE,
|
||||
"Too many attachments, truncating"
|
||||
);
|
||||
msg.attachments.truncate(MAX_ATTACHMENTS_PER_MESSAGE);
|
||||
}
|
||||
|
||||
// Filter by MIME type and enforce total size limit
|
||||
let mut total_size: u64 = 0;
|
||||
msg.attachments.retain(|att| {
|
||||
let mime_ok = ALLOWED_MIME_PREFIXES
|
||||
.iter()
|
||||
.any(|prefix| att.mime_type.starts_with(prefix));
|
||||
if !mime_ok {
|
||||
tracing::warn!(
|
||||
channel = %self.channel_name,
|
||||
mime_type = %att.mime_type,
|
||||
"Attachment MIME type not allowed, dropping"
|
||||
);
|
||||
return false;
|
||||
}
|
||||
|
||||
// Use the larger of reported size_bytes and actual stored data size
|
||||
// to prevent WASM channels from under-reporting to bypass limits.
|
||||
let stored_size = self
|
||||
.attachment_data
|
||||
.get(&att.id)
|
||||
.map(|d| d.len() as u64)
|
||||
.unwrap_or(att.data.len() as u64);
|
||||
let size = att
|
||||
.size_bytes
|
||||
.map(|reported| reported.max(stored_size))
|
||||
.unwrap_or(stored_size);
|
||||
if size > 0 {
|
||||
total_size = total_size.saturating_add(size);
|
||||
if total_size > MAX_ATTACHMENT_TOTAL_SIZE {
|
||||
tracing::warn!(
|
||||
channel = %self.channel_name,
|
||||
total_size,
|
||||
max = MAX_ATTACHMENT_TOTAL_SIZE,
|
||||
"Attachment total size exceeded, dropping"
|
||||
);
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
true
|
||||
});
|
||||
|
||||
msg
|
||||
}
|
||||
|
||||
/// Take all emitted messages (clears the queue).
|
||||
pub fn take_emitted_messages(&mut self) -> Vec<EmittedMessage> {
|
||||
std::mem::take(&mut self.emitted_messages)
|
||||
@@ -224,6 +359,69 @@ impl ChannelHostState {
|
||||
self.emits_dropped
|
||||
}
|
||||
|
||||
/// Store binary data for an attachment.
|
||||
///
|
||||
/// Called by WASM channels to associate downloaded bytes with an attachment ID.
|
||||
/// The data is retrieved after callback completion and merged into `Attachment::data`.
|
||||
pub fn store_attachment_data(
|
||||
&mut self,
|
||||
attachment_id: &str,
|
||||
data: Vec<u8>,
|
||||
) -> Result<(), WasmChannelError> {
|
||||
const MAX_PER_ATTACHMENT: u64 = 20 * 1024 * 1024; // 20 MB
|
||||
const MAX_TOTAL: u64 = 50 * 1024 * 1024; // 50 MB
|
||||
|
||||
let size = data.len() as u64;
|
||||
if size > MAX_PER_ATTACHMENT {
|
||||
return Err(WasmChannelError::CallbackFailed {
|
||||
name: self.channel_name.clone(),
|
||||
reason: format!(
|
||||
"Attachment data too large: {} bytes (max {})",
|
||||
size, MAX_PER_ATTACHMENT
|
||||
),
|
||||
});
|
||||
}
|
||||
|
||||
// Subtract the old entry size (if overwriting) before adding new size
|
||||
let old_size = self
|
||||
.attachment_data
|
||||
.get(attachment_id)
|
||||
.map(|d| d.len() as u64)
|
||||
.unwrap_or(0);
|
||||
let adjusted_total = self.attachment_data_total.saturating_sub(old_size);
|
||||
let new_total = adjusted_total.saturating_add(size);
|
||||
if new_total > MAX_TOTAL {
|
||||
return Err(WasmChannelError::CallbackFailed {
|
||||
name: self.channel_name.clone(),
|
||||
reason: format!(
|
||||
"Total attachment data too large: {} bytes (max {})",
|
||||
new_total, MAX_TOTAL
|
||||
),
|
||||
});
|
||||
}
|
||||
|
||||
self.attachment_data_total = new_total;
|
||||
self.attachment_data.insert(attachment_id.to_string(), data);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Remove stored binary data for a specific attachment ID.
|
||||
pub fn remove_attachment_data(&mut self, id: &str) -> Option<Vec<u8>> {
|
||||
if let Some(data) = self.attachment_data.remove(id) {
|
||||
self.attachment_data_total =
|
||||
self.attachment_data_total.saturating_sub(data.len() as u64);
|
||||
Some(data)
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
/// Take all stored attachment data (clears the store).
|
||||
pub fn take_attachment_data(&mut self) -> HashMap<String, Vec<u8>> {
|
||||
self.attachment_data_total = 0;
|
||||
std::mem::take(&mut self.attachment_data)
|
||||
}
|
||||
|
||||
/// Write to workspace (scoped to channel namespace).
|
||||
///
|
||||
/// Writes are queued and committed after callback execution completes.
|
||||
@@ -431,7 +629,8 @@ impl ChannelEmitRateLimiter {
|
||||
mod tests {
|
||||
use crate::channels::wasm::capabilities::{ChannelCapabilities, EmitRateLimitConfig};
|
||||
use crate::channels::wasm::host::{
|
||||
ChannelEmitRateLimiter, ChannelHostState, EmittedMessage, MAX_EMITS_PER_EXECUTION,
|
||||
Attachment, ChannelEmitRateLimiter, ChannelHostState, EmittedMessage,
|
||||
MAX_ATTACHMENT_TOTAL_SIZE, MAX_ATTACHMENTS_PER_MESSAGE, MAX_EMITS_PER_EXECUTION,
|
||||
};
|
||||
|
||||
#[test]
|
||||
@@ -760,4 +959,133 @@ mod tests {
|
||||
Some("200".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
// === Attachment validation tests ===
|
||||
|
||||
fn make_attachment(id: &str, mime: &str, size: Option<u64>) -> Attachment {
|
||||
Attachment {
|
||||
id: id.to_string(),
|
||||
mime_type: mime.to_string(),
|
||||
filename: None,
|
||||
size_bytes: size,
|
||||
source_url: None,
|
||||
storage_key: None,
|
||||
extracted_text: None,
|
||||
data: Vec::new(),
|
||||
duration_secs: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_emit_message_with_attachments() {
|
||||
let caps = ChannelCapabilities::for_channel("test");
|
||||
let mut state = ChannelHostState::new("test", caps);
|
||||
|
||||
let msg = EmittedMessage::new("user1", "Check this image")
|
||||
.with_attachments(vec![make_attachment("file1", "image/jpeg", Some(1024))]);
|
||||
|
||||
state.emit_message(msg).unwrap();
|
||||
|
||||
let messages = state.take_emitted_messages();
|
||||
assert_eq!(messages.len(), 1);
|
||||
assert_eq!(messages[0].attachments.len(), 1);
|
||||
assert_eq!(messages[0].attachments[0].id, "file1");
|
||||
assert_eq!(messages[0].attachments[0].mime_type, "image/jpeg");
|
||||
assert_eq!(messages[0].attachments[0].size_bytes, Some(1024));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_emit_message_no_attachments_backward_compat() {
|
||||
let caps = ChannelCapabilities::for_channel("test");
|
||||
let mut state = ChannelHostState::new("test", caps);
|
||||
|
||||
let msg = EmittedMessage::new("user1", "Just text");
|
||||
state.emit_message(msg).unwrap();
|
||||
|
||||
let messages = state.take_emitted_messages();
|
||||
assert_eq!(messages.len(), 1);
|
||||
assert!(messages[0].attachments.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_attachment_count_limit() {
|
||||
let caps = ChannelCapabilities::for_channel("test");
|
||||
let mut state = ChannelHostState::new("test", caps);
|
||||
|
||||
let attachments: Vec<Attachment> = (0..MAX_ATTACHMENTS_PER_MESSAGE + 5)
|
||||
.map(|i| make_attachment(&format!("file{}", i), "image/png", Some(100)))
|
||||
.collect();
|
||||
|
||||
let msg = EmittedMessage::new("user1", "Many files").with_attachments(attachments);
|
||||
state.emit_message(msg).unwrap();
|
||||
|
||||
let messages = state.take_emitted_messages();
|
||||
assert_eq!(messages[0].attachments.len(), MAX_ATTACHMENTS_PER_MESSAGE);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_attachment_total_size_limit() {
|
||||
let caps = ChannelCapabilities::for_channel("test");
|
||||
let mut state = ChannelHostState::new("test", caps);
|
||||
|
||||
// Each file is 1/3 of the limit, so 3 fit but 4th does not
|
||||
let chunk_size = MAX_ATTACHMENT_TOTAL_SIZE / 3;
|
||||
let attachments = vec![
|
||||
make_attachment("file1", "image/png", Some(chunk_size)),
|
||||
make_attachment("file2", "image/png", Some(chunk_size)),
|
||||
make_attachment("file3", "image/png", Some(chunk_size)),
|
||||
make_attachment("file4", "image/png", Some(chunk_size)),
|
||||
];
|
||||
|
||||
let msg = EmittedMessage::new("user1", "Big files").with_attachments(attachments);
|
||||
state.emit_message(msg).unwrap();
|
||||
|
||||
let messages = state.take_emitted_messages();
|
||||
// Only first 3 fit within the total size limit
|
||||
assert_eq!(messages[0].attachments.len(), 3);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_attachment_mime_type_filtering() {
|
||||
let caps = ChannelCapabilities::for_channel("test");
|
||||
let mut state = ChannelHostState::new("test", caps);
|
||||
|
||||
let attachments = vec![
|
||||
make_attachment("ok1", "image/jpeg", Some(100)),
|
||||
make_attachment("bad1", "application/x-executable", Some(100)),
|
||||
make_attachment("ok2", "application/pdf", Some(100)),
|
||||
make_attachment("bad2", "application/x-msdos-program", Some(100)),
|
||||
make_attachment("ok3", "text/plain", Some(100)),
|
||||
make_attachment("ok4", "audio/mpeg", Some(100)),
|
||||
make_attachment("ok5", "video/mp4", Some(100)),
|
||||
];
|
||||
|
||||
let msg = EmittedMessage::new("user1", "Mixed files").with_attachments(attachments);
|
||||
state.emit_message(msg).unwrap();
|
||||
|
||||
let messages = state.take_emitted_messages();
|
||||
let ids: Vec<&str> = messages[0]
|
||||
.attachments
|
||||
.iter()
|
||||
.map(|a| a.id.as_str())
|
||||
.collect();
|
||||
assert_eq!(ids, vec!["ok1", "ok2", "ok3", "ok4", "ok5"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_attachment_unknown_size_allowed() {
|
||||
let caps = ChannelCapabilities::for_channel("test");
|
||||
let mut state = ChannelHostState::new("test", caps);
|
||||
|
||||
let attachments = vec![
|
||||
make_attachment("file1", "image/jpeg", None),
|
||||
make_attachment("file2", "image/png", None),
|
||||
];
|
||||
|
||||
let msg = EmittedMessage::new("user1", "No sizes").with_attachments(attachments);
|
||||
state.emit_message(msg).unwrap();
|
||||
|
||||
let messages = state.take_emitted_messages();
|
||||
assert_eq!(messages[0].attachments.len(), 2);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -184,18 +184,32 @@ impl WasmChannelLoader {
|
||||
/// └── telegram.capabilities.json
|
||||
/// ```
|
||||
pub async fn load_from_dir(&self, dir: &Path) -> Result<LoadResults, WasmChannelError> {
|
||||
if !dir.is_dir() {
|
||||
return Err(WasmChannelError::Io(std::io::Error::new(
|
||||
std::io::ErrorKind::NotADirectory,
|
||||
format!("{} is not a directory", dir.display()),
|
||||
)));
|
||||
match fs::metadata(dir).await {
|
||||
Ok(meta) if meta.is_dir() => {}
|
||||
Ok(_) => {
|
||||
return Err(WasmChannelError::Io(std::io::Error::new(
|
||||
std::io::ErrorKind::NotADirectory,
|
||||
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();
|
||||
|
||||
// Collect all .wasm entries first, then load in parallel
|
||||
let mut channel_entries = Vec::new();
|
||||
let mut entries = fs::read_dir(dir).await?;
|
||||
// Handle TOCTOU: if read_dir fails with NotFound, treat as empty
|
||||
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? {
|
||||
let path = entry.path();
|
||||
@@ -486,4 +500,21 @@ mod tests {
|
||||
let result = loader.load_from_files("", &wasm_path, None).await;
|
||||
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());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -86,6 +86,7 @@ mod loader;
|
||||
mod router;
|
||||
mod runtime;
|
||||
mod schema;
|
||||
pub mod setup;
|
||||
pub(crate) mod signature;
|
||||
#[allow(dead_code)]
|
||||
pub(crate) mod storage;
|
||||
@@ -105,4 +106,5 @@ pub use runtime::{PreparedChannelModule, WasmChannelRuntime, WasmChannelRuntimeC
|
||||
pub use schema::{
|
||||
ChannelCapabilitiesFile, ChannelConfig, SecretSetupSchema, SetupSchema, WebhookSchema,
|
||||
};
|
||||
pub use setup::{WasmChannelSetup, inject_channel_credentials, setup_wasm_channels};
|
||||
pub use wrapper::{HttpResponse, SharedWasmChannel, WasmChannel};
|
||||
|
||||
@@ -0,0 +1,324 @@
|
||||
//! 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.
|
||||
if let Some(secrets) = secrets_store {
|
||||
match inject_channel_credentials(&channel_arc, secrets.as_ref(), &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 with the uppercase name if not found
|
||||
/// in the secrets store (e.g., `TELEGRAM_BOT_TOKEN`).
|
||||
pub async fn inject_channel_credentials(
|
||||
channel: &Arc<WasmChannel>,
|
||||
secrets: &dyn SecretsStore,
|
||||
channel_name: &str,
|
||||
) -> anyhow::Result<usize> {
|
||||
let all_secrets = secrets
|
||||
.list("default")
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("Failed to list secrets: {}", e))?;
|
||||
|
||||
let prefix = format!("{}_", channel_name);
|
||||
let mut count = 0;
|
||||
let mut injected_placeholders = HashSet::new();
|
||||
|
||||
for secret_meta in all_secrets {
|
||||
if !secret_meta.name.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;
|
||||
}
|
||||
|
||||
// Fall back to environment variables for required secrets not found in the store.
|
||||
// This allows channels to work when configured via env vars (e.g., TELEGRAM_BOT_TOKEN)
|
||||
// without requiring the setup wizard to have run.
|
||||
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 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)
|
||||
}
|
||||
+449
-15
@@ -532,9 +532,45 @@ impl near::agent::channel_host::Host for ChannelStoreData {
|
||||
user_id = %msg.user_id,
|
||||
user_name = ?msg.user_name,
|
||||
content_len = msg.content.len(),
|
||||
attachment_count = msg.attachments.len(),
|
||||
"WASM emit_message called"
|
||||
);
|
||||
|
||||
let attachments: Vec<crate::channels::wasm::host::Attachment> = msg
|
||||
.attachments
|
||||
.into_iter()
|
||||
.map(|a| {
|
||||
// Parse extras-json for well-known fields
|
||||
let extras: serde_json::Value = if a.extras_json.is_empty() {
|
||||
serde_json::Value::Null
|
||||
} else {
|
||||
serde_json::from_str(&a.extras_json).unwrap_or(serde_json::Value::Null)
|
||||
};
|
||||
let duration_secs = extras
|
||||
.get("duration_secs")
|
||||
.and_then(|v| v.as_u64())
|
||||
.map(|v| v as u32);
|
||||
|
||||
// Merge stored binary data (from store-attachment-data host call)
|
||||
let data = self
|
||||
.host_state
|
||||
.remove_attachment_data(&a.id)
|
||||
.unwrap_or_default();
|
||||
|
||||
crate::channels::wasm::host::Attachment {
|
||||
id: a.id,
|
||||
mime_type: a.mime_type,
|
||||
filename: a.filename,
|
||||
size_bytes: a.size_bytes,
|
||||
source_url: a.source_url,
|
||||
storage_key: a.storage_key,
|
||||
extracted_text: a.extracted_text,
|
||||
data,
|
||||
duration_secs,
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
let mut emitted = EmittedMessage::new(msg.user_id.clone(), msg.content.clone());
|
||||
if let Some(name) = msg.user_name {
|
||||
emitted = emitted.with_user_name(name);
|
||||
@@ -543,6 +579,7 @@ impl near::agent::channel_host::Host for ChannelStoreData {
|
||||
emitted = emitted.with_thread_id(tid);
|
||||
}
|
||||
emitted = emitted.with_metadata(msg.metadata_json);
|
||||
emitted = emitted.with_attachments(attachments);
|
||||
|
||||
match self.host_state.emit_message(emitted) {
|
||||
Ok(()) => {
|
||||
@@ -554,6 +591,21 @@ impl near::agent::channel_host::Host for ChannelStoreData {
|
||||
}
|
||||
}
|
||||
|
||||
fn store_attachment_data(
|
||||
&mut self,
|
||||
attachment_id: String,
|
||||
data: Vec<u8>,
|
||||
) -> Result<(), String> {
|
||||
tracing::debug!(
|
||||
attachment_id = %attachment_id,
|
||||
size = data.len(),
|
||||
"WASM store_attachment_data called"
|
||||
);
|
||||
self.host_state
|
||||
.store_attachment_data(&attachment_id, data)
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
fn pairing_upsert_request(
|
||||
&mut self,
|
||||
channel: String,
|
||||
@@ -1327,12 +1379,14 @@ impl WasmChannel {
|
||||
content: &str,
|
||||
thread_id: Option<&str>,
|
||||
metadata_json: &str,
|
||||
attachments: &[String],
|
||||
) -> Result<(), WasmChannelError> {
|
||||
tracing::info!(
|
||||
channel = %self.name,
|
||||
message_id = %message_id,
|
||||
content_len = content.len(),
|
||||
thread_id = ?thread_id,
|
||||
attachment_count = attachments.len(),
|
||||
"call_on_respond invoked"
|
||||
);
|
||||
|
||||
@@ -1370,12 +1424,21 @@ impl WasmChannel {
|
||||
let content = content.to_string();
|
||||
let thread_id = thread_id.map(|s| s.to_string());
|
||||
let metadata_json = metadata_json.to_string();
|
||||
let attachments = attachments.to_vec();
|
||||
|
||||
// Execute in blocking task with timeout
|
||||
tracing::info!(channel = %channel_name, "Starting on_respond WASM execution");
|
||||
|
||||
let result = tokio::time::timeout(timeout, async move {
|
||||
tokio::task::spawn_blocking(move || {
|
||||
// Read attachment files from disk before entering WASM
|
||||
let wit_attachments = read_attachments(&attachments).map_err(|e| {
|
||||
WasmChannelError::CallbackFailed {
|
||||
name: prepared.name.clone(),
|
||||
reason: e,
|
||||
}
|
||||
})?;
|
||||
|
||||
tracing::info!("Creating WASM store for on_respond");
|
||||
let mut store = Self::create_store(
|
||||
&runtime,
|
||||
@@ -1395,6 +1458,7 @@ impl WasmChannel {
|
||||
content: content.clone(),
|
||||
thread_id,
|
||||
metadata_json,
|
||||
attachments: wit_attachments,
|
||||
};
|
||||
|
||||
// Truncate at char boundary for logging (avoid panic on multi-byte UTF-8)
|
||||
@@ -1458,6 +1522,124 @@ impl WasmChannel {
|
||||
}
|
||||
}
|
||||
|
||||
/// Execute the on_broadcast callback.
|
||||
///
|
||||
/// Called to send a proactive message to a user without a prior incoming message.
|
||||
pub async fn call_on_broadcast(
|
||||
&self,
|
||||
user_id: &str,
|
||||
content: &str,
|
||||
thread_id: Option<&str>,
|
||||
attachments: &[String],
|
||||
) -> Result<(), WasmChannelError> {
|
||||
tracing::info!(
|
||||
channel = %self.name,
|
||||
user_id = %user_id,
|
||||
content_len = content.len(),
|
||||
attachment_count = attachments.len(),
|
||||
"call_on_broadcast invoked"
|
||||
);
|
||||
|
||||
// If no WASM bytes, do nothing (for testing)
|
||||
if self.prepared.component().is_none() {
|
||||
tracing::debug!(
|
||||
channel = %self.name,
|
||||
"WASM channel on_broadcast called (no WASM module)"
|
||||
);
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let runtime = Arc::clone(&self.runtime);
|
||||
let prepared = Arc::clone(&self.prepared);
|
||||
let capabilities = self.capabilities.clone();
|
||||
let timeout = self.runtime.config().callback_timeout;
|
||||
let channel_name = self.name.clone();
|
||||
let credentials = self.get_credentials().await;
|
||||
let host_credentials =
|
||||
resolve_channel_host_credentials(&self.capabilities, self.secrets_store.as_deref())
|
||||
.await;
|
||||
let pairing_store = self.pairing_store.clone();
|
||||
|
||||
let user_id = user_id.to_string();
|
||||
let content = content.to_string();
|
||||
let thread_id = thread_id.map(|s| s.to_string());
|
||||
let attachments = attachments.to_vec();
|
||||
|
||||
let result = tokio::time::timeout(timeout, async move {
|
||||
tokio::task::spawn_blocking(move || {
|
||||
// Read attachment files from disk
|
||||
let wit_attachments = read_attachments(&attachments).map_err(|e| {
|
||||
WasmChannelError::CallbackFailed {
|
||||
name: prepared.name.clone(),
|
||||
reason: e,
|
||||
}
|
||||
})?;
|
||||
|
||||
let mut store = Self::create_store(
|
||||
&runtime,
|
||||
&prepared,
|
||||
&capabilities,
|
||||
credentials,
|
||||
host_credentials,
|
||||
pairing_store,
|
||||
)?;
|
||||
|
||||
let instance = Self::instantiate_component(&runtime, &prepared, &mut store)?;
|
||||
|
||||
let wit_response = wit_channel::AgentResponse {
|
||||
message_id: String::new(),
|
||||
content: content.clone(),
|
||||
thread_id,
|
||||
metadata_json: String::new(),
|
||||
attachments: wit_attachments,
|
||||
};
|
||||
|
||||
let channel_iface = instance.near_agent_channel();
|
||||
let wasm_result = channel_iface
|
||||
.call_on_broadcast(&mut store, &user_id, &wit_response)
|
||||
.map_err(|e| {
|
||||
tracing::error!(error = %e, "WASM on_broadcast call failed");
|
||||
Self::map_wasm_error(e, &prepared.name, prepared.limits.fuel)
|
||||
})?;
|
||||
|
||||
if let Err(ref err_msg) = wasm_result {
|
||||
tracing::error!(error = %err_msg, "WASM on_broadcast returned error");
|
||||
return Err(WasmChannelError::CallbackFailed {
|
||||
name: prepared.name.clone(),
|
||||
reason: err_msg.clone(),
|
||||
});
|
||||
}
|
||||
|
||||
let host_state =
|
||||
Self::extract_host_state(&mut store, &prepared.name, &capabilities);
|
||||
tracing::info!("on_broadcast WASM execution completed successfully");
|
||||
Ok(((), host_state))
|
||||
})
|
||||
.await
|
||||
.map_err(|e| WasmChannelError::ExecutionPanicked {
|
||||
name: channel_name.clone(),
|
||||
reason: e.to_string(),
|
||||
})?
|
||||
})
|
||||
.await;
|
||||
|
||||
let channel_name = self.name.clone();
|
||||
match result {
|
||||
Ok(Ok(((), _host_state))) => {
|
||||
tracing::debug!(
|
||||
channel = %channel_name,
|
||||
"WASM channel on_broadcast completed"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
Ok(Err(e)) => Err(e),
|
||||
Err(_) => Err(WasmChannelError::Timeout {
|
||||
name: channel_name,
|
||||
callback: "on_broadcast".to_string(),
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
/// Execute the on_status callback.
|
||||
///
|
||||
/// Called to notify the WASM channel of agent status changes (e.g., typing).
|
||||
@@ -1745,7 +1927,7 @@ impl WasmChannel {
|
||||
|
||||
let metadata_json = serde_json::to_string(metadata).unwrap_or_default();
|
||||
if let Err(e) = self
|
||||
.call_on_respond(uuid::Uuid::new_v4(), &prompt, None, &metadata_json)
|
||||
.call_on_respond(uuid::Uuid::new_v4(), &prompt, None, &metadata_json, &[])
|
||||
.await
|
||||
{
|
||||
tracing::warn!(
|
||||
@@ -1847,6 +2029,27 @@ impl WasmChannel {
|
||||
msg = msg.with_thread(thread_id);
|
||||
}
|
||||
|
||||
// Convert attachments
|
||||
if !emitted.attachments.is_empty() {
|
||||
let incoming_attachments = emitted
|
||||
.attachments
|
||||
.iter()
|
||||
.map(|a| crate::channels::IncomingAttachment {
|
||||
id: a.id.clone(),
|
||||
kind: crate::channels::AttachmentKind::from_mime_type(&a.mime_type),
|
||||
mime_type: a.mime_type.clone(),
|
||||
filename: a.filename.clone(),
|
||||
size_bytes: a.size_bytes,
|
||||
source_url: a.source_url.clone(),
|
||||
storage_key: a.storage_key.clone(),
|
||||
extracted_text: a.extracted_text.clone(),
|
||||
data: a.data.clone(),
|
||||
duration_secs: a.duration_secs,
|
||||
})
|
||||
.collect();
|
||||
msg = msg.with_attachments(incoming_attachments);
|
||||
}
|
||||
|
||||
// Parse metadata JSON
|
||||
if let Ok(metadata) = serde_json::from_str(&emitted.metadata_json) {
|
||||
msg = msg.with_metadata(metadata);
|
||||
@@ -1859,6 +2062,7 @@ impl WasmChannel {
|
||||
channel = %self.name,
|
||||
user_id = %emitted.user_id,
|
||||
content_len = emitted.content.len(),
|
||||
attachment_count = msg.attachments.len(),
|
||||
"Sending emitted message to agent"
|
||||
);
|
||||
|
||||
@@ -2112,6 +2316,27 @@ impl WasmChannel {
|
||||
msg = msg.with_thread(thread_id);
|
||||
}
|
||||
|
||||
// Convert attachments
|
||||
if !emitted.attachments.is_empty() {
|
||||
let incoming_attachments = emitted
|
||||
.attachments
|
||||
.iter()
|
||||
.map(|a| crate::channels::IncomingAttachment {
|
||||
id: a.id.clone(),
|
||||
kind: crate::channels::AttachmentKind::from_mime_type(&a.mime_type),
|
||||
mime_type: a.mime_type.clone(),
|
||||
filename: a.filename.clone(),
|
||||
size_bytes: a.size_bytes,
|
||||
source_url: a.source_url.clone(),
|
||||
storage_key: a.storage_key.clone(),
|
||||
extracted_text: a.extracted_text.clone(),
|
||||
data: a.data.clone(),
|
||||
duration_secs: a.duration_secs,
|
||||
})
|
||||
.collect();
|
||||
msg = msg.with_attachments(incoming_attachments);
|
||||
}
|
||||
|
||||
// Parse metadata JSON
|
||||
if let Ok(metadata) = serde_json::from_str(&emitted.metadata_json) {
|
||||
msg = msg.with_metadata(metadata);
|
||||
@@ -2130,6 +2355,7 @@ impl WasmChannel {
|
||||
channel = %channel_name,
|
||||
user_id = %emitted.user_id,
|
||||
content_len = emitted.content.len(),
|
||||
attachment_count = msg.attachments.len(),
|
||||
"Sending polled message to agent"
|
||||
);
|
||||
|
||||
@@ -2257,6 +2483,7 @@ impl Channel for WasmChannel {
|
||||
&response.content,
|
||||
response.thread_id.as_deref(),
|
||||
&metadata_json,
|
||||
&response.attachments,
|
||||
)
|
||||
.await
|
||||
.map_err(|e| ChannelError::SendFailed {
|
||||
@@ -2269,24 +2496,15 @@ impl Channel for WasmChannel {
|
||||
|
||||
async fn broadcast(
|
||||
&self,
|
||||
_user_id: &str,
|
||||
user_id: &str,
|
||||
response: OutgoingResponse,
|
||||
) -> Result<(), ChannelError> {
|
||||
let metadata_json = self
|
||||
.last_broadcast_metadata
|
||||
.read()
|
||||
.await
|
||||
.clone()
|
||||
.ok_or_else(|| ChannelError::SendFailed {
|
||||
name: self.name.clone(),
|
||||
reason: "No messages received yet — no chat_id available for broadcast".into(),
|
||||
})?;
|
||||
|
||||
self.call_on_respond(
|
||||
uuid::Uuid::new_v4(),
|
||||
self.cancel_typing_task().await;
|
||||
self.call_on_broadcast(
|
||||
user_id,
|
||||
&response.content,
|
||||
response.thread_id.as_deref(),
|
||||
&metadata_json,
|
||||
&response.attachments,
|
||||
)
|
||||
.await
|
||||
.map_err(|e| ChannelError::SendFailed {
|
||||
@@ -2591,6 +2809,14 @@ fn status_to_wit(status: &StatusUpdate, metadata: &serde_json::Value) -> wit_cha
|
||||
),
|
||||
metadata_json,
|
||||
},
|
||||
StatusUpdate::ImageGenerated { path, .. } => wit_channel::StatusUpdate {
|
||||
status: wit_channel::StatusType::Status,
|
||||
message: match path {
|
||||
Some(p) => format!("[image] {}", p),
|
||||
None => "[image generated]".to_string(),
|
||||
},
|
||||
metadata_json,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2749,6 +2975,79 @@ async fn resolve_channel_host_credentials(
|
||||
resolved
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Attachment Helpers
|
||||
// ============================================================================
|
||||
|
||||
/// Maximum total attachment size (50 MB).
|
||||
const MAX_TOTAL_ATTACHMENT_BYTES: u64 = 50 * 1024 * 1024;
|
||||
|
||||
/// Detect MIME type from file extension using the `mime_guess` crate.
|
||||
fn mime_from_extension(path: &str) -> String {
|
||||
mime_guess::from_path(path)
|
||||
.first_or_octet_stream()
|
||||
.to_string()
|
||||
}
|
||||
|
||||
/// Read attachment files from disk and build WIT attachment records.
|
||||
///
|
||||
/// Validates total size against `MAX_TOTAL_ATTACHMENT_BYTES`.
|
||||
fn read_attachments(paths: &[String]) -> Result<Vec<wit_channel::Attachment>, String> {
|
||||
if paths.is_empty() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let mut attachments = Vec::with_capacity(paths.len());
|
||||
let mut total_bytes: u64 = 0;
|
||||
let tmp_base = std::path::Path::new("/tmp");
|
||||
let home_base = dirs::home_dir()
|
||||
.map(|h| h.join(".ironclaw"))
|
||||
.unwrap_or_default();
|
||||
|
||||
for path in paths {
|
||||
// Validate paths are under /tmp/ or ~/.ironclaw/ to prevent arbitrary file reads
|
||||
let validated = crate::tools::builtin::path_utils::validate_path(path, Some(tmp_base))
|
||||
.or_else(|_| crate::tools::builtin::path_utils::validate_path(path, Some(&home_base)));
|
||||
let validated = validated.map_err(|e| {
|
||||
format!(
|
||||
"Invalid attachment path '{}': must be under /tmp/ or ~/.ironclaw/: {}",
|
||||
path, e
|
||||
)
|
||||
})?;
|
||||
|
||||
// Pre-check file size before reading into memory to avoid OOM
|
||||
let file_size = std::fs::metadata(&validated)
|
||||
.map_err(|e| format!("Failed to stat attachment '{}': {}", validated.display(), e))?
|
||||
.len();
|
||||
total_bytes += file_size;
|
||||
if total_bytes > MAX_TOTAL_ATTACHMENT_BYTES {
|
||||
return Err(format!(
|
||||
"Total attachment size exceeds {} MB limit",
|
||||
MAX_TOTAL_ATTACHMENT_BYTES / (1024 * 1024)
|
||||
));
|
||||
}
|
||||
|
||||
let data = std::fs::read(&validated)
|
||||
.map_err(|e| format!("Failed to read attachment '{}': {}", validated.display(), e))?;
|
||||
|
||||
let filename = validated
|
||||
.file_name()
|
||||
.and_then(|n| n.to_str())
|
||||
.unwrap_or("file")
|
||||
.to_string();
|
||||
|
||||
let mime_type = mime_from_extension(path);
|
||||
|
||||
attachments.push(wit_channel::Attachment {
|
||||
filename,
|
||||
mime_type,
|
||||
data,
|
||||
});
|
||||
}
|
||||
|
||||
Ok(attachments)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
@@ -3871,4 +4170,139 @@ mod tests {
|
||||
// 404 because "000" is not a valid bot token
|
||||
assert_eq!(result, 404);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_dispatch_emitted_messages_preserves_attachments() {
|
||||
use crate::channels::wasm::host::{Attachment, EmittedMessage};
|
||||
|
||||
let (tx, mut rx) = tokio::sync::mpsc::channel(10);
|
||||
let message_tx = Arc::new(tokio::sync::RwLock::new(Some(tx)));
|
||||
|
||||
let rate_limiter = Arc::new(tokio::sync::RwLock::new(
|
||||
crate::channels::wasm::host::ChannelEmitRateLimiter::new(
|
||||
crate::channels::wasm::capabilities::EmitRateLimitConfig::default(),
|
||||
),
|
||||
));
|
||||
|
||||
let attachments = vec![
|
||||
Attachment {
|
||||
id: "photo123".to_string(),
|
||||
mime_type: "image/jpeg".to_string(),
|
||||
filename: Some("cat.jpg".to_string()),
|
||||
size_bytes: Some(50_000),
|
||||
source_url: Some("https://api.telegram.org/file/photo123".to_string()),
|
||||
storage_key: None,
|
||||
extracted_text: None,
|
||||
data: Vec::new(),
|
||||
duration_secs: None,
|
||||
},
|
||||
Attachment {
|
||||
id: "doc456".to_string(),
|
||||
mime_type: "application/pdf".to_string(),
|
||||
filename: Some("report.pdf".to_string()),
|
||||
size_bytes: Some(120_000),
|
||||
source_url: None,
|
||||
storage_key: Some("store/doc456".to_string()),
|
||||
extracted_text: Some("Report contents...".to_string()),
|
||||
data: Vec::new(),
|
||||
duration_secs: None,
|
||||
},
|
||||
];
|
||||
|
||||
let messages =
|
||||
vec![EmittedMessage::new("user1", "Check these files").with_attachments(attachments)];
|
||||
|
||||
let last_broadcast_metadata = Arc::new(tokio::sync::RwLock::new(None));
|
||||
let result = WasmChannel::dispatch_emitted_messages(
|
||||
"test-channel",
|
||||
messages,
|
||||
&message_tx,
|
||||
&rate_limiter,
|
||||
&last_broadcast_metadata,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(result.is_ok());
|
||||
|
||||
let msg = rx.try_recv().expect("Should receive message");
|
||||
assert_eq!(msg.content, "Check these files");
|
||||
assert_eq!(msg.attachments.len(), 2);
|
||||
|
||||
// Verify first attachment
|
||||
assert_eq!(msg.attachments[0].id, "photo123");
|
||||
assert_eq!(msg.attachments[0].mime_type, "image/jpeg");
|
||||
assert_eq!(msg.attachments[0].filename, Some("cat.jpg".to_string()));
|
||||
assert_eq!(msg.attachments[0].size_bytes, Some(50_000));
|
||||
assert_eq!(
|
||||
msg.attachments[0].source_url,
|
||||
Some("https://api.telegram.org/file/photo123".to_string())
|
||||
);
|
||||
|
||||
// Verify second attachment
|
||||
assert_eq!(msg.attachments[1].id, "doc456");
|
||||
assert_eq!(msg.attachments[1].mime_type, "application/pdf");
|
||||
assert_eq!(
|
||||
msg.attachments[1].extracted_text,
|
||||
Some("Report contents...".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
msg.attachments[1].storage_key,
|
||||
Some("store/doc456".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_dispatch_emitted_messages_no_attachments_backward_compat() {
|
||||
use crate::channels::wasm::host::EmittedMessage;
|
||||
|
||||
let (tx, mut rx) = tokio::sync::mpsc::channel(10);
|
||||
let message_tx = Arc::new(tokio::sync::RwLock::new(Some(tx)));
|
||||
|
||||
let rate_limiter = Arc::new(tokio::sync::RwLock::new(
|
||||
crate::channels::wasm::host::ChannelEmitRateLimiter::new(
|
||||
crate::channels::wasm::capabilities::EmitRateLimitConfig::default(),
|
||||
),
|
||||
));
|
||||
|
||||
let messages = vec![EmittedMessage::new("user1", "Just text, no attachments")];
|
||||
|
||||
let last_broadcast_metadata = Arc::new(tokio::sync::RwLock::new(None));
|
||||
let result = WasmChannel::dispatch_emitted_messages(
|
||||
"test-channel",
|
||||
messages,
|
||||
&message_tx,
|
||||
&rate_limiter,
|
||||
&last_broadcast_metadata,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(result.is_ok());
|
||||
|
||||
let msg = rx.try_recv().expect("Should receive message");
|
||||
assert_eq!(msg.content, "Just text, no attachments");
|
||||
assert!(msg.attachments.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_mime_from_extension() {
|
||||
use super::mime_from_extension;
|
||||
assert_eq!(mime_from_extension("screenshot.png"), "image/png");
|
||||
assert_eq!(mime_from_extension("photo.JPG"), "image/jpeg");
|
||||
assert_eq!(mime_from_extension("photo.jpeg"), "image/jpeg");
|
||||
assert_eq!(mime_from_extension("animation.gif"), "image/gif");
|
||||
assert_eq!(mime_from_extension("doc.pdf"), "application/pdf");
|
||||
assert_eq!(mime_from_extension("video.mp4"), "video/mp4");
|
||||
assert_eq!(mime_from_extension("data.csv"), "text/csv");
|
||||
assert_eq!(
|
||||
mime_from_extension("unknown.qqqzzz"),
|
||||
"application/octet-stream"
|
||||
);
|
||||
assert_eq!(mime_from_extension("noext"), "application/octet-stream");
|
||||
assert_eq!(
|
||||
mime_from_extension("/home/user/.ironclaw/screenshot.png"),
|
||||
"image/png"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,212 @@
|
||||
# Web Gateway Module
|
||||
|
||||
Browser-facing HTTP API and SSE/WebSocket real-time streaming. Axum-based, single-user with bearer token auth.
|
||||
|
||||
## File Map
|
||||
|
||||
| File | Role |
|
||||
|------|------|
|
||||
| `mod.rs` | Gateway builder, startup, `WebChannel` implementation, `with_*` builder methods |
|
||||
| `server.rs` | `GatewayState`, `start_server()`, all Axum route registrations, inline handlers |
|
||||
| `types.rs` | Request/response DTOs and `SseEvent` enum (source of truth for SSE contract) |
|
||||
| `sse.rs` | `SseManager` — broadcast channel that fans out `SseEvent` to all connected SSE clients |
|
||||
| `ws.rs` | WebSocket handler (`handle_ws_connection`) + `WsConnectionTracker` |
|
||||
| `auth.rs` | Bearer token middleware (`Authorization: Bearer <GATEWAY_AUTH_TOKEN>`) |
|
||||
| `log_layer.rs` | Tracing layer that tees log lines to the `/api/logs/events` SSE stream |
|
||||
| `handlers/` | Handler functions split by domain: `chat`, `extensions`, `jobs`, `memory`, `routines`, `settings`, `skills`, `static_files` |
|
||||
| `openai_compat.rs` | OpenAI-compatible proxy (`/v1/chat/completions`, `/v1/models`) |
|
||||
| `util.rs` | Shared helpers (`build_turns_from_db_messages`, `truncate_preview`) |
|
||||
| `static/` | Single-page app (HTML/CSS/JS) — embedded at compile time via `include_str!`/`include_bytes!` |
|
||||
|
||||
## API Routes
|
||||
|
||||
### Public (no auth)
|
||||
| Method | Path | Description |
|
||||
|--------|------|-------------|
|
||||
| GET | `/api/health` | Health check |
|
||||
| GET | `/oauth/callback` | OAuth callback for extension auth |
|
||||
|
||||
### Chat
|
||||
| Method | Path | Description |
|
||||
|--------|------|-------------|
|
||||
| POST | `/api/chat/send` | Send message → queues to agent loop |
|
||||
| GET | `/api/chat/events` | SSE stream of agent events |
|
||||
| GET | `/api/chat/ws` | WebSocket alternative to SSE |
|
||||
| GET | `/api/chat/history` | Paginated turn history for a thread |
|
||||
| GET | `/api/chat/threads` | List threads (returns `assistant_thread` + regular threads) |
|
||||
| POST | `/api/chat/thread/new` | Create new thread |
|
||||
| POST | `/api/chat/approval` | Approve/deny/always a pending tool call |
|
||||
| POST | `/api/chat/auth-token` | Submit auth token for an extension |
|
||||
| POST | `/api/chat/auth-cancel` | Cancel pending auth flow |
|
||||
|
||||
### Memory
|
||||
| Method | Path | Description |
|
||||
|--------|------|-------------|
|
||||
| GET | `/api/memory/tree` | Workspace directory tree |
|
||||
| GET | `/api/memory/list` | List files at a path |
|
||||
| GET | `/api/memory/read` | Read a workspace file |
|
||||
| POST | `/api/memory/write` | Write a workspace file |
|
||||
| POST | `/api/memory/search` | Hybrid FTS + vector search |
|
||||
|
||||
### Jobs (sandbox)
|
||||
| Method | Path | Description |
|
||||
|--------|------|-------------|
|
||||
| GET | `/api/jobs` | List sandbox jobs |
|
||||
| GET | `/api/jobs/summary` | Aggregated stats |
|
||||
| GET | `/api/jobs/{id}` | Job detail |
|
||||
| POST | `/api/jobs/{id}/cancel` | Cancel a running job |
|
||||
| POST | `/api/jobs/{id}/restart` | Restart a failed job |
|
||||
| POST | `/api/jobs/{id}/prompt` | Send follow-up prompt to Claude Code bridge |
|
||||
| GET | `/api/jobs/{id}/events` | SSE stream for a specific job |
|
||||
| GET | `/api/jobs/{id}/files/list` | List files in job workspace |
|
||||
| GET | `/api/jobs/{id}/files/read` | Read a file from job workspace |
|
||||
|
||||
### Skills
|
||||
| Method | Path | Description |
|
||||
|--------|------|-------------|
|
||||
| GET | `/api/skills` | List installed skills |
|
||||
| POST | `/api/skills/search` | Search ClawHub registry + local skills |
|
||||
| POST | `/api/skills/install` | Install a skill from ClawHub or by URL/content |
|
||||
| DELETE | `/api/skills/{name}` | Remove an installed skill |
|
||||
|
||||
### Extensions
|
||||
| Method | Path | Description |
|
||||
|--------|------|-------------|
|
||||
| GET | `/api/extensions` | Installed extensions |
|
||||
| GET | `/api/extensions/tools` | All registered tools (from tool registry) |
|
||||
| POST | `/api/extensions/install` | Install extension |
|
||||
| GET | `/api/extensions/registry` | Available extensions from registry manifests |
|
||||
| POST | `/api/extensions/{name}/activate` | Activate installed extension |
|
||||
| POST | `/api/extensions/{name}/remove` | Remove extension |
|
||||
| GET/POST | `/api/extensions/{name}/setup` | Extension setup wizard |
|
||||
|
||||
### Routines
|
||||
| Method | Path | Description |
|
||||
|--------|------|-------------|
|
||||
| GET | `/api/routines` | List routines |
|
||||
| GET | `/api/routines/summary` | Aggregated stats (total/enabled/disabled/failing/runs_today) |
|
||||
| GET | `/api/routines/{id}` | Routine detail with recent run history |
|
||||
| POST | `/api/routines/{id}/trigger` | Manually trigger a routine |
|
||||
| POST | `/api/routines/{id}/toggle` | Enable/disable a routine |
|
||||
| DELETE | `/api/routines/{id}` | Delete a routine |
|
||||
| GET | `/api/routines/{id}/runs` | List runs for a specific routine |
|
||||
|
||||
### Settings
|
||||
| Method | Path | Description |
|
||||
|--------|------|-------------|
|
||||
| GET | `/api/settings` | List all settings |
|
||||
| GET | `/api/settings/export` | Export all settings as a map |
|
||||
| POST | `/api/settings/import` | Bulk-import settings from a map |
|
||||
| GET | `/api/settings/{key}` | Get a single setting |
|
||||
| PUT | `/api/settings/{key}` | Set a single setting |
|
||||
| DELETE | `/api/settings/{key}` | Delete a setting |
|
||||
|
||||
### Other
|
||||
| Method | Path | Description |
|
||||
|--------|------|-------------|
|
||||
| GET | `/api/logs/events` | Live log stream (SSE) |
|
||||
| GET/PUT | `/api/logs/level` | Get/set log level at runtime |
|
||||
| GET | `/api/pairing/{channel}` | List pending pairing requests |
|
||||
| POST | `/api/pairing/{channel}/approve` | Approve a pairing request |
|
||||
| GET | `/api/gateway/status` | Server uptime, connected clients, config |
|
||||
| POST | `/v1/chat/completions` | OpenAI-compatible LLM proxy |
|
||||
| GET | `/v1/models` | OpenAI-compatible model list |
|
||||
|
||||
### Static / Project files
|
||||
| Method | Path | Description |
|
||||
|--------|------|-------------|
|
||||
| GET | `/` | Single-page app HTML |
|
||||
| GET | `/style.css` | App stylesheet |
|
||||
| GET | `/app.js` | App JavaScript |
|
||||
| GET | `/favicon.ico` | Favicon (cached 1 day) |
|
||||
| GET | `/projects/{project_id}/` | Job workspace browser (redirects) |
|
||||
| GET | `/projects/{project_id}/{*path}` | Serve file from job workspace (auth required) |
|
||||
|
||||
## SSE Event Types (`SseEvent` in `types.rs`)
|
||||
|
||||
The SSE contract — every field is `#[serde(tag = "type")]`:
|
||||
|
||||
| Type | When emitted |
|
||||
|------|-------------|
|
||||
| `response` | Final text response from agent |
|
||||
| `stream_chunk` | Streaming token (partial response) |
|
||||
| `thinking` | Agent status update during reasoning |
|
||||
| `tool_started` | Tool call began |
|
||||
| `tool_completed` | Tool call finished (includes success/error) |
|
||||
| `tool_result` | Tool output preview |
|
||||
| `status` | Generic status message |
|
||||
| `job_started` | Sandbox job created |
|
||||
| `job_message` | Message from sandbox worker |
|
||||
| `job_tool_use` | Tool invoked inside sandbox |
|
||||
| `job_tool_result` | Tool result from sandbox |
|
||||
| `job_status` | Sandbox job status update |
|
||||
| `job_result` | Sandbox job final result |
|
||||
| `approval_needed` | Tool requires user approval (pauses agent) |
|
||||
| `auth_required` | Extension needs auth credentials |
|
||||
| `auth_completed` | Extension auth flow finished |
|
||||
| `extension_status` | WASM channel activation status changed |
|
||||
| `error` | Error from agent or gateway |
|
||||
| `heartbeat` | SSE keepalive (empty payload) |
|
||||
|
||||
**SSE serialization:** Events use `#[serde(tag = "type")]` — the wire format is `{"type":"<variant>", ...fields}`. The SSE frame's `event:` field is set to the same string as `type` for easy `addEventListener` use in the browser.
|
||||
|
||||
**WebSocket envelope:** Over WebSocket, SSE events are wrapped as `{"type":"event","event_type":"<variant>","data":{...}}`. Ping/pong uses `{"type":"ping"}` / `{"type":"pong"}`. Client-to-server messages (`message`, `approval`, `auth_token`, `auth_cancel`) are defined in `WsClientMessage` in `types.rs`.
|
||||
|
||||
**To add a new SSE event:** Use the `add-sse-event` skill (`/add-sse-event`). It scaffolds the Rust variant, serialization, broadcast call, and frontend handler. Also add a matching arm to `WsServerMessage::from_sse_event()` in `types.rs`.
|
||||
|
||||
## Auth
|
||||
|
||||
All protected routes require `Authorization: Bearer <GATEWAY_AUTH_TOKEN>`. The token is set via `GATEWAY_AUTH_TOKEN` env var. Missing/wrong token → 401. The `Bearer` prefix is compared case-insensitively (RFC 6750).
|
||||
|
||||
**Query-string token auth (`?token=xxx`):** Because `EventSource` and WebSocket upgrades cannot set custom headers from the browser, three endpoints also accept the token as a URL query parameter: `/api/chat/events`, `/api/logs/events`, and `/api/chat/ws`. All other endpoints reject query-string tokens. If you add a new SSE or WebSocket endpoint, register its path in `allows_query_token_auth()` in `auth.rs`.
|
||||
|
||||
**If no `GATEWAY_AUTH_TOKEN` is configured**, a random 32-character alphanumeric token is generated at startup and printed to the console.
|
||||
|
||||
Rate limiting: chat send endpoints are capped at **30 messages per 60 seconds** (sliding window, not per-IP).
|
||||
|
||||
## GatewayState
|
||||
|
||||
The shared state struct (`server.rs`) holds refs to all subsystems. Fields are `Option<Arc<T>>` so the gateway can start even when optional subsystems (workspace, sandbox, skills) are disabled. Always null-check before use in handlers.
|
||||
|
||||
Key fields:
|
||||
- `msg_tx` — `RwLock<Option<mpsc::Sender<IncomingMessage>>>` — sends messages to the agent loop; set when `start()` is called on the `Channel`.
|
||||
- `sse` — `SseManager` — broadcast hub; call `state.sse.broadcast(event)` from any handler.
|
||||
- `ws_tracker` — `Option<Arc<WsConnectionTracker>>` — tracks WS connection count separately from SSE.
|
||||
- `chat_rate_limiter` — `RateLimiter` — 30 req/60 s sliding window shared across all chat send callers.
|
||||
- `scheduler` — `Option<SchedulerSlot>` — used to inject follow-up messages into running agent jobs.
|
||||
- `cost_guard` — `Option<Arc<CostGuard>>` — exposes token usage / cost totals in the status endpoint.
|
||||
- `startup_time` — `Instant` — used to compute uptime in the gateway status response.
|
||||
- `registry_entries` — `Vec<RegistryEntry>` — loaded once at startup from registry manifests; used by the available extensions API without hitting the network.
|
||||
|
||||
Subsystems are wired via `with_*` builder methods on `GatewayChannel` (`mod.rs`). Each call rebuilds `Arc<GatewayState>` — safe to call before `start()`, not after.
|
||||
|
||||
## SSE / WebSocket Connection Limits
|
||||
|
||||
Both SSE and WebSocket share the same `SseManager` broadcast channel. Key characteristics:
|
||||
|
||||
- **Broadcast buffer:** 256 events. A slow client that falls behind will miss events — the `BroadcastStream` silently drops lagged events. SSE clients are expected to reconnect and re-fetch history.
|
||||
- **Max connections:** 100 total (SSE + WebSocket combined). Connections beyond the limit receive a 503 / are immediately dropped.
|
||||
- **SSE keepalive:** Axum's `KeepAlive` sends an empty event every **30 seconds** to prevent proxy timeouts.
|
||||
- **WebSocket:** Two tasks per connection — a sender task (broadcast → WS frames) and a receiver loop (WS frames → agent). When the client disconnects, the sender is aborted and both the SSE connection counter and WS tracker counter are decremented.
|
||||
|
||||
## CORS and Security Headers
|
||||
|
||||
CORS is restricted to the gateway's own origin (same IP+port and `localhost`+port). Allowed methods: GET, POST, PUT, DELETE. Allowed headers: `Content-Type`, `Authorization`. Credentials are allowed.
|
||||
|
||||
All responses include:
|
||||
- `X-Content-Type-Options: nosniff`
|
||||
- `X-Frame-Options: DENY`
|
||||
|
||||
**Request body limit:** 1 MB (`DefaultBodyLimit::max(1024 * 1024)`). Larger payloads return 413.
|
||||
|
||||
## Pending Approvals
|
||||
|
||||
Tool approval state is **in-memory only** (not persisted to DB). Server restart clears all pending approvals. The `pending_approval` field in `HistoryResponse` is re-populated on thread switch from in-memory state.
|
||||
|
||||
## Adding a New API Endpoint
|
||||
|
||||
1. Define request/response types in `types.rs`.
|
||||
2. Implement the handler in the appropriate `handlers/*.rs` file (or inline in `server.rs` for simple handlers).
|
||||
3. Register the route in `start_server()` in `server.rs` under the correct router (`public`, `protected`, or `statics`).
|
||||
4. If it is an SSE or WebSocket endpoint, add its path to `allows_query_token_auth()` in `auth.rs`.
|
||||
5. If it requires a new `GatewayState` field, add it to the struct and to both the `GatewayChannel::new()` initializer and `rebuild_state()` in `mod.rs`, then add a `with_*` builder method.
|
||||
@@ -426,7 +426,7 @@ pub async fn chat_threads_handler(
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
|
||||
if let Ok(summaries) = store
|
||||
.list_conversations_with_preview(&state.user_id, "gateway", 50)
|
||||
.list_conversations_all_channels(&state.user_id, 50)
|
||||
.await
|
||||
{
|
||||
let mut assistant_thread = None;
|
||||
@@ -441,6 +441,7 @@ pub async fn chat_threads_handler(
|
||||
updated_at: s.last_activity.to_rfc3339(),
|
||||
title: s.title.clone(),
|
||||
thread_type: s.thread_type.clone(),
|
||||
channel: Some(s.channel.clone()),
|
||||
};
|
||||
|
||||
if s.id == assistant_id {
|
||||
@@ -460,6 +461,7 @@ pub async fn chat_threads_handler(
|
||||
updated_at: chrono::Utc::now().to_rfc3339(),
|
||||
title: None,
|
||||
thread_type: Some("assistant".to_string()),
|
||||
channel: Some("gateway".to_string()),
|
||||
});
|
||||
}
|
||||
|
||||
@@ -472,9 +474,10 @@ pub async fn chat_threads_handler(
|
||||
}
|
||||
|
||||
// Fallback: in-memory only (no assistant thread without DB)
|
||||
let threads: Vec<ThreadInfo> = sess
|
||||
.threads
|
||||
.values()
|
||||
let mut sorted_threads: Vec<_> = sess.threads.values().collect();
|
||||
sorted_threads.sort_by(|a, b| b.updated_at.cmp(&a.updated_at));
|
||||
let threads: Vec<ThreadInfo> = sorted_threads
|
||||
.into_iter()
|
||||
.map(|t| ThreadInfo {
|
||||
id: t.id,
|
||||
state: format!("{:?}", t.state),
|
||||
@@ -483,6 +486,7 @@ pub async fn chat_threads_handler(
|
||||
updated_at: t.updated_at.to_rfc3339(),
|
||||
title: None,
|
||||
thread_type: None,
|
||||
channel: Some("gateway".to_string()),
|
||||
})
|
||||
.collect();
|
||||
|
||||
@@ -502,38 +506,39 @@ pub async fn chat_new_thread_handler(
|
||||
))?;
|
||||
|
||||
let session = session_manager.get_or_create_session(&state.user_id).await;
|
||||
let mut sess = session.lock().await;
|
||||
let thread = sess.create_thread();
|
||||
let thread_id = thread.id;
|
||||
let info = ThreadInfo {
|
||||
id: thread.id,
|
||||
state: format!("{:?}", thread.state),
|
||||
turn_count: thread.turns.len(),
|
||||
created_at: thread.created_at.to_rfc3339(),
|
||||
updated_at: thread.updated_at.to_rfc3339(),
|
||||
title: None,
|
||||
thread_type: Some("thread".to_string()),
|
||||
let (thread_id, info) = {
|
||||
let mut sess = session.lock().await;
|
||||
let thread = sess.create_thread();
|
||||
let id = thread.id;
|
||||
let info = ThreadInfo {
|
||||
id: thread.id,
|
||||
state: format!("{:?}", thread.state),
|
||||
turn_count: thread.turns.len(),
|
||||
created_at: thread.created_at.to_rfc3339(),
|
||||
updated_at: thread.updated_at.to_rfc3339(),
|
||||
title: None,
|
||||
thread_type: Some("thread".to_string()),
|
||||
channel: Some("gateway".to_string()),
|
||||
};
|
||||
(id, info)
|
||||
};
|
||||
|
||||
// Persist the empty conversation row with thread_type metadata
|
||||
// Persist the empty conversation row with thread_type metadata synchronously
|
||||
// so that the subsequent loadThreads() call from the frontend sees it.
|
||||
if let Some(ref store) = state.store {
|
||||
let store = Arc::clone(store);
|
||||
let user_id = state.user_id.clone();
|
||||
tokio::spawn(async move {
|
||||
if let Err(e) = store
|
||||
.ensure_conversation(thread_id, "gateway", &user_id, None)
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Failed to persist new thread: {}", e);
|
||||
}
|
||||
let metadata_val = serde_json::json!("thread");
|
||||
if let Err(e) = store
|
||||
.update_conversation_metadata_field(thread_id, "thread_type", &metadata_val)
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Failed to set thread_type metadata: {}", e);
|
||||
}
|
||||
});
|
||||
if let Err(e) = store
|
||||
.ensure_conversation(thread_id, "gateway", &state.user_id, None)
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Failed to persist new thread: {}", e);
|
||||
}
|
||||
let metadata_val = serde_json::json!("thread");
|
||||
if let Err(e) = store
|
||||
.update_conversation_metadata_field(thread_id, "thread_type", &metadata_val)
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Failed to set thread_type metadata: {}", e);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(Json(info))
|
||||
|
||||
@@ -62,6 +62,7 @@ pub async fn extensions_list_handler(
|
||||
has_auth: ext.has_auth,
|
||||
activation_status,
|
||||
activation_error: ext.activation_error,
|
||||
version: ext.version,
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
@@ -276,11 +276,25 @@ pub async fn jobs_cancel_handler(
|
||||
})));
|
||||
}
|
||||
|
||||
// Fall back to agent job cancellation via DB status update.
|
||||
// Fall back to agent job cancellation: stop the worker via the scheduler
|
||||
// (which updates the in-memory ContextManager AND aborts the task handle),
|
||||
// then persist the status to the DB as a fallback.
|
||||
if let Some(ref store) = state.store
|
||||
&& let Ok(Some(job)) = store.get_job(job_id).await
|
||||
{
|
||||
if job.state.is_active() {
|
||||
// Try to stop via scheduler (aborts the worker task + updates
|
||||
// in-memory ContextManager). This is best-effort — the job may
|
||||
// not be in the scheduler map if it already finished.
|
||||
if let Some(ref slot) = state.scheduler
|
||||
&& let Some(ref scheduler) = *slot.read().await
|
||||
{
|
||||
let _ = scheduler.stop(job_id).await;
|
||||
}
|
||||
|
||||
// Always persist cancellation to the DB so the state is
|
||||
// consistent even if the scheduler wasn't available or the
|
||||
// job wasn't in its in-memory map.
|
||||
store
|
||||
.update_job_status(
|
||||
job_id,
|
||||
|
||||
@@ -10,9 +10,9 @@ use axum::{
|
||||
use serde::Deserialize;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::channels::IncomingMessage;
|
||||
use crate::channels::web::server::GatewayState;
|
||||
use crate::channels::web::types::*;
|
||||
use crate::error::RoutineError;
|
||||
|
||||
pub async fn routines_list_handler(
|
||||
State(state): State<Arc<GatewayState>>,
|
||||
@@ -108,6 +108,7 @@ pub async fn routines_detail_handler(
|
||||
status: format!("{:?}", run.status),
|
||||
result_summary: run.result_summary.clone(),
|
||||
tokens_used: run.tokens_used,
|
||||
job_id: run.job_id,
|
||||
})
|
||||
.collect();
|
||||
|
||||
@@ -133,56 +134,27 @@ pub async fn routines_trigger_handler(
|
||||
State(state): State<Arc<GatewayState>>,
|
||||
Path(id): Path<String>,
|
||||
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
||||
let store = state.store.as_ref().ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Database not available".to_string(),
|
||||
))?;
|
||||
// Clone the Arc out of the lock to avoid holding the RwLock across .await.
|
||||
let engine = {
|
||||
let guard = state.routine_engine.read().await;
|
||||
guard.as_ref().cloned().ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Routine engine not available".to_string(),
|
||||
))?
|
||||
};
|
||||
|
||||
let routine_id = Uuid::parse_str(&id)
|
||||
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?;
|
||||
|
||||
let routine = store
|
||||
.get_routine(routine_id)
|
||||
let run_id = engine
|
||||
.fire_manual(routine_id, Some(&state.user_id))
|
||||
.await
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
|
||||
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
|
||||
|
||||
if routine.user_id != state.user_id {
|
||||
return Err((StatusCode::FORBIDDEN, "Access denied".to_string()));
|
||||
}
|
||||
|
||||
// Send the routine prompt through the message pipeline as a manual trigger.
|
||||
let prompt = match &routine.action {
|
||||
crate::agent::routine::RoutineAction::Lightweight { prompt, .. } => prompt.clone(),
|
||||
crate::agent::routine::RoutineAction::FullJob {
|
||||
title, description, ..
|
||||
} => format!("{}: {}", title, description),
|
||||
};
|
||||
|
||||
let content = format!("[routine:{}] {}", routine.name, prompt);
|
||||
let thread_id = format!(
|
||||
"routine-{}-{}",
|
||||
routine_id,
|
||||
chrono::Utc::now().timestamp_millis()
|
||||
);
|
||||
let msg = IncomingMessage::new("gateway", &state.user_id, content).with_thread(thread_id);
|
||||
|
||||
let tx_guard = state.msg_tx.read().await;
|
||||
let tx = tx_guard.as_ref().ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Channel not started".to_string(),
|
||||
))?;
|
||||
|
||||
tx.send(msg).await.map_err(|_| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
"Channel closed".to_string(),
|
||||
)
|
||||
})?;
|
||||
.map_err(|e| (routine_error_status(&e), e.to_string()))?;
|
||||
|
||||
Ok(Json(serde_json::json!({
|
||||
"status": "triggered",
|
||||
"routine_id": routine_id,
|
||||
"run_id": run_id,
|
||||
})))
|
||||
}
|
||||
|
||||
@@ -281,6 +253,7 @@ pub async fn routines_runs_handler(
|
||||
status: format!("{:?}", run.status),
|
||||
result_summary: run.result_summary.clone(),
|
||||
tokens_used: run.tokens_used,
|
||||
job_id: run.job_id,
|
||||
})
|
||||
.collect();
|
||||
|
||||
@@ -293,7 +266,7 @@ 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 } => {
|
||||
crate::agent::routine::Trigger::Cron { schedule, .. } => {
|
||||
("cron".to_string(), format!("cron: {}", schedule))
|
||||
}
|
||||
crate::agent::routine::Trigger::Event {
|
||||
@@ -337,3 +310,13 @@ fn routine_to_info(r: &crate::agent::routine::Routine) -> RoutineInfo {
|
||||
status: status.to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Map `RoutineError` variants to appropriate HTTP status codes.
|
||||
fn routine_error_status(err: &RoutineError) -> StatusCode {
|
||||
match err {
|
||||
RoutineError::NotFound { .. } => StatusCode::NOT_FOUND,
|
||||
RoutineError::NotAuthorized { .. } => StatusCode::FORBIDDEN,
|
||||
RoutineError::Disabled { .. } | RoutineError::MaxConcurrent { .. } => StatusCode::CONFLICT,
|
||||
_ => StatusCode::INTERNAL_SERVER_ERROR,
|
||||
}
|
||||
}
|
||||
|
||||
+33
-2
@@ -24,6 +24,13 @@ pub mod types;
|
||||
pub(crate) mod util;
|
||||
pub mod ws;
|
||||
|
||||
/// Test helpers for gateway integration tests.
|
||||
///
|
||||
/// Always compiled (not behind `#[cfg(test)]`) so that integration tests in
|
||||
/// `tests/` -- which import this crate as a regular dependency -- can use
|
||||
/// [`TestGatewayBuilder`](test_helpers::TestGatewayBuilder).
|
||||
pub mod test_helpers;
|
||||
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::Arc;
|
||||
|
||||
@@ -92,6 +99,7 @@ impl GatewayChannel {
|
||||
chat_rate_limiter: server::RateLimiter::new(30, 60),
|
||||
registry_entries: Vec::new(),
|
||||
cost_guard: None,
|
||||
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
|
||||
startup_time: std::time::Instant::now(),
|
||||
});
|
||||
|
||||
@@ -127,6 +135,7 @@ impl GatewayChannel {
|
||||
chat_rate_limiter: server::RateLimiter::new(30, 60),
|
||||
registry_entries: self.state.registry_entries.clone(),
|
||||
cost_guard: self.state.cost_guard.clone(),
|
||||
routine_engine: Arc::clone(&self.state.routine_engine),
|
||||
startup_time: self.state.startup_time,
|
||||
};
|
||||
mutate(&mut new_state);
|
||||
@@ -274,7 +283,15 @@ impl Channel for GatewayChannel {
|
||||
msg: &IncomingMessage,
|
||||
response: OutgoingResponse,
|
||||
) -> Result<(), ChannelError> {
|
||||
let thread_id = msg.thread_id.clone().unwrap_or_default();
|
||||
let thread_id = match &msg.thread_id {
|
||||
Some(tid) => tid.clone(),
|
||||
None => {
|
||||
tracing::warn!(
|
||||
"Gateway respond with no thread_id — skipping (clients would drop it)"
|
||||
);
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
|
||||
self.state.sse.broadcast(SseEvent::Response {
|
||||
content: response.content,
|
||||
@@ -369,6 +386,11 @@ impl Channel for GatewayChannel {
|
||||
success,
|
||||
message,
|
||||
},
|
||||
StatusUpdate::ImageGenerated { data_url, path } => SseEvent::ImageGenerated {
|
||||
data_url,
|
||||
path,
|
||||
thread_id,
|
||||
},
|
||||
};
|
||||
|
||||
self.state.sse.broadcast(event);
|
||||
@@ -380,9 +402,18 @@ impl Channel for GatewayChannel {
|
||||
_user_id: &str,
|
||||
response: OutgoingResponse,
|
||||
) -> Result<(), ChannelError> {
|
||||
let thread_id = match response.thread_id {
|
||||
Some(tid) => tid,
|
||||
None => {
|
||||
tracing::warn!(
|
||||
"Gateway broadcast with no thread_id — skipping (clients would drop it)"
|
||||
);
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
self.state.sse.broadcast(SseEvent::Response {
|
||||
content: response.content,
|
||||
thread_id: String::new(),
|
||||
thread_id,
|
||||
});
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -244,6 +244,7 @@ pub fn convert_messages(messages: &[OpenAiMessage]) -> Result<Vec<ChatMessage>,
|
||||
_ => Ok(ChatMessage {
|
||||
role,
|
||||
content: m.content.as_deref().unwrap_or("").to_string(),
|
||||
content_parts: Vec::new(),
|
||||
tool_call_id: None,
|
||||
name: m.name.clone(),
|
||||
tool_calls: None,
|
||||
|
||||
+149
-72
@@ -57,6 +57,10 @@ pub type PromptQueue = Arc<
|
||||
>,
|
||||
>;
|
||||
|
||||
/// Slot for the routine engine, filled at runtime after the agent starts.
|
||||
pub type RoutineEngineSlot =
|
||||
Arc<tokio::sync::RwLock<Option<Arc<crate::agent::routine_engine::RoutineEngine>>>>;
|
||||
|
||||
/// Simple sliding-window rate limiter.
|
||||
///
|
||||
/// Tracks the number of requests in the current window. Resets when the window expires.
|
||||
@@ -165,6 +169,8 @@ pub struct GatewayState {
|
||||
pub registry_entries: Vec<crate::extensions::RegistryEntry>,
|
||||
/// Cost guard for token/cost tracking.
|
||||
pub cost_guard: Option<Arc<crate::agent::cost_guard::CostGuard>>,
|
||||
/// Routine engine slot for manual routine triggering (filled at runtime).
|
||||
pub routine_engine: RoutineEngineSlot,
|
||||
/// Server startup time for uptime calculation.
|
||||
pub startup_time: std::time::Instant,
|
||||
}
|
||||
@@ -345,7 +351,7 @@ pub async fn start_server(
|
||||
.merge(statics)
|
||||
.merge(projects)
|
||||
.merge(protected)
|
||||
.layer(DefaultBodyLimit::max(1024 * 1024)) // 1 MB max request body
|
||||
.layer(DefaultBodyLimit::max(10 * 1024 * 1024)) // 10 MB max request body (image uploads)
|
||||
.layer(cors)
|
||||
.layer(SetResponseHeaderLayer::if_not_present(
|
||||
header::X_CONTENT_TYPE_OPTIONS,
|
||||
@@ -364,7 +370,7 @@ pub async fn start_server(
|
||||
if let Err(e) = axum::serve(listener, app)
|
||||
.with_graceful_shutdown(async {
|
||||
let _ = shutdown_rx.await;
|
||||
tracing::info!("Web gateway shutting down");
|
||||
tracing::debug!("Web gateway shutting down");
|
||||
})
|
||||
.await
|
||||
{
|
||||
@@ -602,8 +608,59 @@ async fn oauth_callback_handler(
|
||||
|
||||
// --- Chat handlers ---
|
||||
|
||||
/// Convert web gateway `ImageData` to `IncomingAttachment` objects.
|
||||
pub(crate) fn images_to_attachments(
|
||||
images: &[ImageData],
|
||||
) -> Vec<crate::channels::IncomingAttachment> {
|
||||
use base64::Engine;
|
||||
images
|
||||
.iter()
|
||||
.enumerate()
|
||||
.filter_map(|(i, img)| {
|
||||
if !img.media_type.starts_with("image/") {
|
||||
tracing::warn!(
|
||||
"Skipping image {i}: invalid media type '{}' (must start with 'image/')",
|
||||
img.media_type
|
||||
);
|
||||
return None;
|
||||
}
|
||||
let data = match base64::engine::general_purpose::STANDARD.decode(&img.data) {
|
||||
Ok(d) => d,
|
||||
Err(e) => {
|
||||
tracing::warn!("Skipping image {i}: invalid base64 data: {e}");
|
||||
return None;
|
||||
}
|
||||
};
|
||||
Some(crate::channels::IncomingAttachment {
|
||||
id: format!("web-image-{i}"),
|
||||
kind: crate::channels::AttachmentKind::Image,
|
||||
mime_type: img.media_type.clone(),
|
||||
filename: Some(format!("image-{i}.{}", mime_to_ext(&img.media_type))),
|
||||
size_bytes: Some(data.len() as u64),
|
||||
source_url: None,
|
||||
storage_key: None,
|
||||
extracted_text: None,
|
||||
data,
|
||||
duration_secs: None,
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Map MIME type to file extension.
|
||||
fn mime_to_ext(mime: &str) -> &str {
|
||||
match mime {
|
||||
"image/png" => "png",
|
||||
"image/gif" => "gif",
|
||||
"image/webp" => "webp",
|
||||
"image/svg+xml" => "svg",
|
||||
_ => "jpg",
|
||||
}
|
||||
}
|
||||
|
||||
async fn chat_send_handler(
|
||||
State(state): State<Arc<GatewayState>>,
|
||||
headers: axum::http::HeaderMap,
|
||||
Json(req): Json<SendMessageRequest>,
|
||||
) -> Result<(StatusCode, Json<SendMessageResponse>), (StatusCode, String)> {
|
||||
tracing::debug!(
|
||||
@@ -620,17 +677,32 @@ async fn chat_send_handler(
|
||||
}
|
||||
|
||||
let mut msg = IncomingMessage::new("gateway", &state.user_id, &req.content);
|
||||
// Prefer timezone from JSON body, fall back to X-Timezone header
|
||||
let tz = req
|
||||
.timezone
|
||||
.as_deref()
|
||||
.or_else(|| headers.get("X-Timezone").and_then(|v| v.to_str().ok()));
|
||||
if let Some(tz) = tz {
|
||||
msg = msg.with_timezone(tz);
|
||||
}
|
||||
|
||||
if let Some(ref thread_id) = req.thread_id {
|
||||
msg = msg.with_thread(thread_id);
|
||||
msg = msg.with_metadata(serde_json::json!({"thread_id": thread_id}));
|
||||
}
|
||||
|
||||
// Convert uploaded images to IncomingAttachments
|
||||
if !req.images.is_empty() {
|
||||
let attachments = images_to_attachments(&req.images);
|
||||
msg = msg.with_attachments(attachments);
|
||||
}
|
||||
|
||||
let msg_id = msg.id;
|
||||
tracing::debug!(
|
||||
"[chat_send_handler] Created message id={}, content={:?}",
|
||||
"[chat_send_handler] Created message id={}, content={:?}, images={}",
|
||||
msg_id,
|
||||
req.content
|
||||
req.content,
|
||||
req.images.len()
|
||||
);
|
||||
|
||||
let tx_guard = state.msg_tx.read().await;
|
||||
@@ -1037,7 +1109,7 @@ async fn chat_threads_handler(
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
|
||||
if let Ok(summaries) = store
|
||||
.list_conversations_with_preview(&state.user_id, "gateway", 50)
|
||||
.list_conversations_all_channels(&state.user_id, 50)
|
||||
.await
|
||||
{
|
||||
let mut assistant_thread = None;
|
||||
@@ -1052,6 +1124,7 @@ async fn chat_threads_handler(
|
||||
updated_at: s.last_activity.to_rfc3339(),
|
||||
title: s.title.clone(),
|
||||
thread_type: s.thread_type.clone(),
|
||||
channel: Some(s.channel.clone()),
|
||||
};
|
||||
|
||||
if s.id == assistant_id {
|
||||
@@ -1071,6 +1144,7 @@ async fn chat_threads_handler(
|
||||
updated_at: chrono::Utc::now().to_rfc3339(),
|
||||
title: None,
|
||||
thread_type: Some("assistant".to_string()),
|
||||
channel: Some("gateway".to_string()),
|
||||
});
|
||||
}
|
||||
|
||||
@@ -1083,9 +1157,10 @@ async fn chat_threads_handler(
|
||||
}
|
||||
|
||||
// Fallback: in-memory only (no assistant thread without DB)
|
||||
let threads: Vec<ThreadInfo> = sess
|
||||
.threads
|
||||
.values()
|
||||
let mut sorted_threads: Vec<_> = sess.threads.values().collect();
|
||||
sorted_threads.sort_by(|a, b| b.updated_at.cmp(&a.updated_at));
|
||||
let threads: Vec<ThreadInfo> = sorted_threads
|
||||
.into_iter()
|
||||
.map(|t| ThreadInfo {
|
||||
id: t.id,
|
||||
state: format!("{:?}", t.state),
|
||||
@@ -1094,6 +1169,7 @@ async fn chat_threads_handler(
|
||||
updated_at: t.updated_at.to_rfc3339(),
|
||||
title: None,
|
||||
thread_type: None,
|
||||
channel: Some("gateway".to_string()),
|
||||
})
|
||||
.collect();
|
||||
|
||||
@@ -1113,38 +1189,39 @@ async fn chat_new_thread_handler(
|
||||
))?;
|
||||
|
||||
let session = session_manager.get_or_create_session(&state.user_id).await;
|
||||
let mut sess = session.lock().await;
|
||||
let thread = sess.create_thread();
|
||||
let thread_id = thread.id;
|
||||
let info = ThreadInfo {
|
||||
id: thread.id,
|
||||
state: format!("{:?}", thread.state),
|
||||
turn_count: thread.turns.len(),
|
||||
created_at: thread.created_at.to_rfc3339(),
|
||||
updated_at: thread.updated_at.to_rfc3339(),
|
||||
title: None,
|
||||
thread_type: Some("thread".to_string()),
|
||||
let (thread_id, info) = {
|
||||
let mut sess = session.lock().await;
|
||||
let thread = sess.create_thread();
|
||||
let id = thread.id;
|
||||
let info = ThreadInfo {
|
||||
id: thread.id,
|
||||
state: format!("{:?}", thread.state),
|
||||
turn_count: thread.turns.len(),
|
||||
created_at: thread.created_at.to_rfc3339(),
|
||||
updated_at: thread.updated_at.to_rfc3339(),
|
||||
title: None,
|
||||
thread_type: Some("thread".to_string()),
|
||||
channel: Some("gateway".to_string()),
|
||||
};
|
||||
(id, info)
|
||||
};
|
||||
|
||||
// Persist the empty conversation row with thread_type metadata
|
||||
// Persist the empty conversation row with thread_type metadata synchronously
|
||||
// so that the subsequent loadThreads() call from the frontend sees it.
|
||||
if let Some(ref store) = state.store {
|
||||
let store = Arc::clone(store);
|
||||
let user_id = state.user_id.clone();
|
||||
tokio::spawn(async move {
|
||||
if let Err(e) = store
|
||||
.ensure_conversation(thread_id, "gateway", &user_id, None)
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Failed to persist new thread: {}", e);
|
||||
}
|
||||
let metadata_val = serde_json::json!("thread");
|
||||
if let Err(e) = store
|
||||
.update_conversation_metadata_field(thread_id, "thread_type", &metadata_val)
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Failed to set thread_type metadata: {}", e);
|
||||
}
|
||||
});
|
||||
if let Err(e) = store
|
||||
.ensure_conversation(thread_id, "gateway", &state.user_id, None)
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Failed to persist new thread: {}", e);
|
||||
}
|
||||
let metadata_val = serde_json::json!("thread");
|
||||
if let Err(e) = store
|
||||
.update_conversation_metadata_field(thread_id, "thread_type", &metadata_val)
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Failed to set thread_type metadata: {}", e);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(Json(info))
|
||||
@@ -1438,6 +1515,7 @@ async fn extensions_list_handler(
|
||||
has_auth: ext.has_auth,
|
||||
activation_status,
|
||||
activation_error: ext.activation_error,
|
||||
version: ext.version,
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
@@ -1731,6 +1809,7 @@ async fn extensions_registry_handler(
|
||||
kind: kind_str,
|
||||
description: e.description.clone(),
|
||||
keywords: e.keywords.clone(),
|
||||
version: e.version.clone(),
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
@@ -1938,6 +2017,7 @@ async fn routines_detail_handler(
|
||||
status: format!("{:?}", run.status),
|
||||
result_summary: run.result_summary.clone(),
|
||||
tokens_used: run.tokens_used,
|
||||
job_id: run.job_id,
|
||||
})
|
||||
.collect();
|
||||
|
||||
@@ -1963,47 +2043,35 @@ async fn routines_trigger_handler(
|
||||
State(state): State<Arc<GatewayState>>,
|
||||
Path(id): Path<String>,
|
||||
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
||||
let store = state.store.as_ref().ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Database not available".to_string(),
|
||||
))?;
|
||||
let engine = {
|
||||
let guard = state.routine_engine.read().await;
|
||||
guard.as_ref().cloned().ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Routine engine not available".to_string(),
|
||||
))?
|
||||
};
|
||||
|
||||
let routine_id = Uuid::parse_str(&id)
|
||||
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?;
|
||||
|
||||
let routine = store
|
||||
.get_routine(routine_id)
|
||||
let run_id = engine
|
||||
.fire_manual(routine_id, Some(&state.user_id))
|
||||
.await
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
|
||||
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
|
||||
|
||||
// Send the routine prompt through the message pipeline as a manual trigger.
|
||||
let prompt = match &routine.action {
|
||||
crate::agent::routine::RoutineAction::Lightweight { prompt, .. } => prompt.clone(),
|
||||
crate::agent::routine::RoutineAction::FullJob {
|
||||
title, description, ..
|
||||
} => format!("{}: {}", title, description),
|
||||
};
|
||||
|
||||
let content = format!("[routine:{}] {}", routine.name, prompt);
|
||||
let msg = IncomingMessage::new("gateway", &state.user_id, content);
|
||||
|
||||
let tx_guard = state.msg_tx.read().await;
|
||||
let tx = tx_guard.as_ref().ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Channel not started".to_string(),
|
||||
))?;
|
||||
|
||||
tx.send(msg).await.map_err(|_| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
"Channel closed".to_string(),
|
||||
)
|
||||
})?;
|
||||
.map_err(|e| {
|
||||
let status = match &e {
|
||||
crate::error::RoutineError::NotFound { .. } => StatusCode::NOT_FOUND,
|
||||
crate::error::RoutineError::NotAuthorized { .. } => StatusCode::FORBIDDEN,
|
||||
crate::error::RoutineError::Disabled { .. }
|
||||
| crate::error::RoutineError::MaxConcurrent { .. } => StatusCode::CONFLICT,
|
||||
_ => StatusCode::INTERNAL_SERVER_ERROR,
|
||||
};
|
||||
(status, e.to_string())
|
||||
})?;
|
||||
|
||||
Ok(Json(serde_json::json!({
|
||||
"status": "triggered",
|
||||
"routine_id": routine_id,
|
||||
"run_id": run_id,
|
||||
})))
|
||||
}
|
||||
|
||||
@@ -2102,6 +2170,7 @@ async fn routines_runs_handler(
|
||||
status: format!("{:?}", run.status),
|
||||
result_summary: run.result_summary.clone(),
|
||||
tokens_used: run.tokens_used,
|
||||
job_id: run.job_id,
|
||||
})
|
||||
.collect();
|
||||
|
||||
@@ -2114,7 +2183,7 @@ 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 } => {
|
||||
crate::agent::routine::Trigger::Cron { schedule, .. } => {
|
||||
("cron".to_string(), format!("cron: {}", schedule))
|
||||
}
|
||||
crate::agent::routine::Trigger::Event {
|
||||
@@ -2461,6 +2530,7 @@ mod tests {
|
||||
chat_rate_limiter: RateLimiter::new(30, 60),
|
||||
registry_entries: vec![],
|
||||
cost_guard: None,
|
||||
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
|
||||
startup_time: std::time::Instant::now(),
|
||||
})
|
||||
}
|
||||
@@ -2539,6 +2609,7 @@ mod tests {
|
||||
|
||||
let ext_mgr = Arc::new(ExtensionManager::new(
|
||||
mcp_sm,
|
||||
Arc::new(crate::tools::mcp::process::McpProcessManager::new()),
|
||||
secrets,
|
||||
tool_registry,
|
||||
None,
|
||||
@@ -2588,6 +2659,7 @@ mod tests {
|
||||
|
||||
let ext_mgr = Arc::new(ExtensionManager::new(
|
||||
mcp_sm,
|
||||
Arc::new(crate::tools::mcp::process::McpProcessManager::new()),
|
||||
secrets.clone(),
|
||||
tool_registry,
|
||||
None,
|
||||
@@ -2618,7 +2690,9 @@ mod tests {
|
||||
secrets,
|
||||
sse_sender: None,
|
||||
gateway_token: None,
|
||||
created_at: std::time::Instant::now() - std::time::Duration::from_secs(600),
|
||||
created_at: std::time::Instant::now()
|
||||
.checked_sub(std::time::Duration::from_secs(600))
|
||||
.expect("System uptime is too low to run expired flow test"),
|
||||
};
|
||||
|
||||
ext_mgr
|
||||
@@ -2691,6 +2765,7 @@ mod tests {
|
||||
|
||||
let ext_mgr = Arc::new(ExtensionManager::new(
|
||||
mcp_sm,
|
||||
Arc::new(crate::tools::mcp::process::McpProcessManager::new()),
|
||||
secrets.clone(),
|
||||
tool_registry,
|
||||
None,
|
||||
@@ -2725,7 +2800,9 @@ mod tests {
|
||||
sse_sender: None,
|
||||
gateway_token: None,
|
||||
// Expired — handler will reject after lookup (no network I/O)
|
||||
created_at: std::time::Instant::now() - std::time::Duration::from_secs(600),
|
||||
created_at: std::time::Instant::now()
|
||||
.checked_sub(std::time::Duration::from_secs(600))
|
||||
.expect("System uptime is too low to run expired flow test"),
|
||||
};
|
||||
|
||||
ext_mgr
|
||||
|
||||
@@ -142,6 +142,7 @@ impl SseManager {
|
||||
SseEvent::JobStatus { .. } => "job_status",
|
||||
SseEvent::JobResult { .. } => "job_result",
|
||||
SseEvent::Heartbeat => "heartbeat",
|
||||
SseEvent::ImageGenerated { .. } => "image_generated",
|
||||
SseEvent::ExtensionStatus { .. } => "extension_status",
|
||||
};
|
||||
Ok(Event::default().event(event_type).data(data))
|
||||
|
||||
+284
-49
@@ -5,6 +5,7 @@ let eventSource = null;
|
||||
let logEventSource = null;
|
||||
let currentTab = 'chat';
|
||||
let currentThreadId = null;
|
||||
let currentThreadIsReadOnly = false;
|
||||
let assistantThreadId = null;
|
||||
let hasMore = false;
|
||||
let oldestTimestamp = null;
|
||||
@@ -13,8 +14,11 @@ let sseHasConnectedBefore = false;
|
||||
let jobEvents = new Map(); // job_id -> Array of events
|
||||
let jobListRefreshTimer = null;
|
||||
let pairingPollInterval = null;
|
||||
let unreadThreads = new Map(); // thread_id -> unread count
|
||||
let _loadThreadsTimer = null;
|
||||
const JOB_EVENTS_CAP = 500;
|
||||
const MEMORY_SEARCH_QUERY_MAX_LENGTH = 100;
|
||||
let stagedImages = [];
|
||||
|
||||
// --- Slash Commands ---
|
||||
|
||||
@@ -178,6 +182,7 @@ function confirmRestart() {
|
||||
body: {
|
||||
content: '/restart',
|
||||
thread_id: currentThreadId,
|
||||
timezone: Intl.DateTimeFormat().resolvedOptions().timeZone,
|
||||
},
|
||||
})
|
||||
.then((response) => {
|
||||
@@ -220,23 +225,6 @@ function updateRestartButtonVisibility() {
|
||||
}
|
||||
}
|
||||
|
||||
function startGatewayStatusPolling() {
|
||||
fetchGatewayStatus();
|
||||
// Poll every 5 seconds
|
||||
setInterval(fetchGatewayStatus, 5000);
|
||||
}
|
||||
|
||||
function fetchGatewayStatus() {
|
||||
apiFetch('/api/gateway/status')
|
||||
.then((data) => {
|
||||
restartEnabled = data.restart_enabled || false;
|
||||
updateRestartButtonVisibility();
|
||||
})
|
||||
.catch((err) => {
|
||||
console.warn('[gateway status] Failed to fetch:', err);
|
||||
});
|
||||
}
|
||||
|
||||
// --- SSE ---
|
||||
|
||||
function connectSSE() {
|
||||
@@ -273,7 +261,13 @@ function connectSSE() {
|
||||
|
||||
eventSource.addEventListener('response', (e) => {
|
||||
const data = JSON.parse(e.data);
|
||||
if (!isCurrentThread(data.thread_id)) return;
|
||||
if (!isCurrentThread(data.thread_id)) {
|
||||
if (data.thread_id) {
|
||||
unreadThreads.set(data.thread_id, (unreadThreads.get(data.thread_id) || 0) + 1);
|
||||
debouncedLoadThreads();
|
||||
}
|
||||
return;
|
||||
}
|
||||
finalizeActivityGroup();
|
||||
addMessage('assistant', data.content);
|
||||
enableChatInput();
|
||||
@@ -288,7 +282,10 @@ function connectSSE() {
|
||||
|
||||
eventSource.addEventListener('thinking', (e) => {
|
||||
const data = JSON.parse(e.data);
|
||||
if (!isCurrentThread(data.thread_id)) return;
|
||||
if (!isCurrentThread(data.thread_id)) {
|
||||
if (data.thread_id) debouncedLoadThreads();
|
||||
return;
|
||||
}
|
||||
showActivityThinking(data.message);
|
||||
});
|
||||
|
||||
@@ -324,7 +321,10 @@ function connectSSE() {
|
||||
|
||||
eventSource.addEventListener('status', (e) => {
|
||||
const data = JSON.parse(e.data);
|
||||
if (!isCurrentThread(data.thread_id)) return;
|
||||
if (!isCurrentThread(data.thread_id)) {
|
||||
if (data.thread_id) debouncedLoadThreads();
|
||||
return;
|
||||
}
|
||||
// "Done" and "Awaiting approval" are terminal signals from the agent:
|
||||
// the agentic loop finished, so re-enable input as a safety net in case
|
||||
// the response SSE event is empty or lost.
|
||||
@@ -373,6 +373,12 @@ function connectSSE() {
|
||||
if (currentTab === 'extensions') loadExtensions();
|
||||
});
|
||||
|
||||
eventSource.addEventListener('image_generated', (e) => {
|
||||
const data = JSON.parse(e.data);
|
||||
if (!isCurrentThread(data.thread_id)) return;
|
||||
addGeneratedImage(data.data_url, data.path);
|
||||
});
|
||||
|
||||
eventSource.addEventListener('error', (e) => {
|
||||
if (e.data) {
|
||||
const data = JSON.parse(e.data);
|
||||
@@ -414,9 +420,9 @@ function connectSSE() {
|
||||
}
|
||||
|
||||
// Check if an SSE event belongs to the currently viewed thread.
|
||||
// Events without a thread_id (legacy) are always shown.
|
||||
// Events without a thread_id are dropped (prevents notification leaking).
|
||||
function isCurrentThread(threadId) {
|
||||
if (!threadId) return true;
|
||||
if (!threadId) return false;
|
||||
if (!currentThreadId) return true;
|
||||
return threadId === currentThreadId;
|
||||
}
|
||||
@@ -430,23 +436,135 @@ function sendMessage() {
|
||||
return;
|
||||
}
|
||||
const content = input.value.trim();
|
||||
if (!content) return;
|
||||
if (!content && stagedImages.length === 0) return;
|
||||
|
||||
addMessage('user', content);
|
||||
addMessage('user', content || '(images attached)');
|
||||
input.value = '';
|
||||
autoResizeTextarea(input);
|
||||
input.focus();
|
||||
|
||||
const body = { content, thread_id: currentThreadId || undefined, timezone: Intl.DateTimeFormat().resolvedOptions().timeZone };
|
||||
if (stagedImages.length > 0) {
|
||||
body.images = stagedImages.map(img => ({ media_type: img.media_type, data: img.data }));
|
||||
stagedImages = [];
|
||||
renderImagePreviews();
|
||||
}
|
||||
|
||||
apiFetch('/api/chat/send', {
|
||||
method: 'POST',
|
||||
body: { content, thread_id: currentThreadId || undefined },
|
||||
body: body,
|
||||
}).catch((err) => {
|
||||
addMessage('system', 'Failed to send: ' + err.message);
|
||||
});
|
||||
}
|
||||
|
||||
function enableChatInput() {
|
||||
// no-op: input and send button are always enabled
|
||||
if (currentThreadIsReadOnly) return;
|
||||
const input = document.getElementById('chat-input');
|
||||
const btn = document.getElementById('send-btn');
|
||||
if (input) {
|
||||
input.disabled = false;
|
||||
input.placeholder = 'Message or / for commands...';
|
||||
}
|
||||
if (btn) btn.disabled = false;
|
||||
}
|
||||
|
||||
// --- Image Upload ---
|
||||
|
||||
function renderImagePreviews() {
|
||||
const strip = document.getElementById('image-preview-strip');
|
||||
strip.innerHTML = '';
|
||||
stagedImages.forEach((img, idx) => {
|
||||
const container = document.createElement('div');
|
||||
container.className = 'image-preview-container';
|
||||
|
||||
const preview = document.createElement('img');
|
||||
preview.className = 'image-preview';
|
||||
preview.src = img.dataUrl;
|
||||
preview.alt = 'Attached image';
|
||||
|
||||
const removeBtn = document.createElement('button');
|
||||
removeBtn.className = 'image-preview-remove';
|
||||
removeBtn.textContent = '\u00d7';
|
||||
removeBtn.addEventListener('click', () => {
|
||||
stagedImages.splice(idx, 1);
|
||||
renderImagePreviews();
|
||||
});
|
||||
|
||||
container.appendChild(preview);
|
||||
container.appendChild(removeBtn);
|
||||
strip.appendChild(container);
|
||||
});
|
||||
}
|
||||
|
||||
const MAX_IMAGE_SIZE_BYTES = 5 * 1024 * 1024; // 5 MB per image
|
||||
const MAX_STAGED_IMAGES = 5;
|
||||
|
||||
function handleImageFiles(files) {
|
||||
Array.from(files).forEach(file => {
|
||||
if (!file.type.startsWith('image/')) return;
|
||||
if (file.size > MAX_IMAGE_SIZE_BYTES) {
|
||||
alert(`Image "${file.name}" exceeds 5 MB limit (${(file.size / 1024 / 1024).toFixed(1)} MB)`);
|
||||
return;
|
||||
}
|
||||
if (stagedImages.length >= MAX_STAGED_IMAGES) {
|
||||
alert(`Maximum ${MAX_STAGED_IMAGES} images allowed per message`);
|
||||
return;
|
||||
}
|
||||
const reader = new FileReader();
|
||||
reader.onload = function(e) {
|
||||
const dataUrl = e.target.result;
|
||||
const commaIdx = dataUrl.indexOf(',');
|
||||
const meta = dataUrl.substring(0, commaIdx); // e.g. "data:image/png;base64"
|
||||
const base64 = dataUrl.substring(commaIdx + 1);
|
||||
const mediaType = meta.replace('data:', '').replace(';base64', '');
|
||||
stagedImages.push({ media_type: mediaType, data: base64, dataUrl: dataUrl });
|
||||
renderImagePreviews();
|
||||
};
|
||||
reader.readAsDataURL(file);
|
||||
});
|
||||
}
|
||||
|
||||
document.getElementById('attach-btn').addEventListener('click', () => {
|
||||
document.getElementById('image-file-input').click();
|
||||
});
|
||||
|
||||
document.getElementById('image-file-input').addEventListener('change', (e) => {
|
||||
handleImageFiles(e.target.files);
|
||||
e.target.value = '';
|
||||
});
|
||||
|
||||
document.getElementById('chat-input').addEventListener('paste', (e) => {
|
||||
const items = (e.clipboardData || e.originalEvent.clipboardData).items;
|
||||
for (let i = 0; i < items.length; i++) {
|
||||
if (items[i].kind === 'file' && items[i].type.startsWith('image/')) {
|
||||
const file = items[i].getAsFile();
|
||||
if (file) handleImageFiles([file]);
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
function addGeneratedImage(dataUrl, path) {
|
||||
const container = document.getElementById('chat-messages');
|
||||
const card = document.createElement('div');
|
||||
card.className = 'generated-image-card';
|
||||
|
||||
const img = document.createElement('img');
|
||||
img.className = 'generated-image';
|
||||
img.src = dataUrl;
|
||||
img.alt = 'Generated image';
|
||||
|
||||
card.appendChild(img);
|
||||
|
||||
if (path) {
|
||||
const pathLabel = document.createElement('div');
|
||||
pathLabel.className = 'generated-image-path';
|
||||
pathLabel.textContent = path;
|
||||
card.appendChild(pathLabel);
|
||||
}
|
||||
|
||||
container.appendChild(card);
|
||||
container.scrollTop = container.scrollHeight;
|
||||
}
|
||||
|
||||
// --- Slash Autocomplete ---
|
||||
@@ -541,6 +659,13 @@ function sendApprovalAction(requestId, action) {
|
||||
|
||||
function renderMarkdown(text) {
|
||||
if (typeof marked !== 'undefined') {
|
||||
// Escape raw HTML error pages instead of rendering them as markup.
|
||||
// Only triggers when the text *starts with* a doctype or <html> tag
|
||||
// (after optional whitespace), so normal messages that mention HTML
|
||||
// tags in prose or code fences are not affected. See #263.
|
||||
if (/^\s*<!doctype\s/i.test(text) || /^\s*<html[\s>]/i.test(text)) {
|
||||
return escapeHtml(text);
|
||||
}
|
||||
let html = marked.parse(text);
|
||||
// Sanitize HTML output to prevent XSS from tool output or LLM responses.
|
||||
html = sanitizeRenderedHtml(html);
|
||||
@@ -1134,7 +1259,9 @@ function loadHistory(before) {
|
||||
// Fresh load: clear and render
|
||||
container.innerHTML = '';
|
||||
for (const turn of data.turns) {
|
||||
addMessage('user', turn.user_input);
|
||||
if (turn.user_input) {
|
||||
addMessage('user', turn.user_input);
|
||||
}
|
||||
if (turn.tool_calls && turn.tool_calls.length > 0) {
|
||||
addToolCallsSummary(turn.tool_calls);
|
||||
}
|
||||
@@ -1156,8 +1283,10 @@ function loadHistory(before) {
|
||||
const savedHeight = container.scrollHeight;
|
||||
const fragment = document.createDocumentFragment();
|
||||
for (const turn of data.turns) {
|
||||
const userDiv = createMessageElement('user', turn.user_input);
|
||||
fragment.appendChild(userDiv);
|
||||
if (turn.user_input) {
|
||||
const userDiv = createMessageElement('user', turn.user_input);
|
||||
fragment.appendChild(userDiv);
|
||||
}
|
||||
if (turn.tool_calls && turn.tool_calls.length > 0) {
|
||||
fragment.appendChild(createToolCallsSummaryElement(turn.tool_calls));
|
||||
}
|
||||
@@ -1256,6 +1385,37 @@ function removeScrollSpinner() {
|
||||
|
||||
// --- Threads ---
|
||||
|
||||
function threadTitle(thread) {
|
||||
if (thread.title) return thread.title;
|
||||
const ch = thread.channel || 'gateway';
|
||||
if (thread.thread_type === 'heartbeat') return 'Heartbeat Alerts';
|
||||
if (thread.thread_type === 'routine') return 'Routine';
|
||||
if (ch !== 'gateway') return ch.charAt(0).toUpperCase() + ch.slice(1);
|
||||
if (thread.turn_count === 0) return 'New chat';
|
||||
return thread.id.substring(0, 8);
|
||||
}
|
||||
|
||||
function relativeTime(isoStr) {
|
||||
if (!isoStr) return '';
|
||||
const diff = Date.now() - new Date(isoStr).getTime();
|
||||
const mins = Math.floor(diff / 60000);
|
||||
if (mins < 1) return 'now';
|
||||
if (mins < 60) return mins + 'm ago';
|
||||
const hrs = Math.floor(mins / 60);
|
||||
if (hrs < 24) return hrs + 'h ago';
|
||||
const days = Math.floor(hrs / 24);
|
||||
return days + 'd ago';
|
||||
}
|
||||
|
||||
function isReadOnlyChannel(channel) {
|
||||
return channel && channel !== 'gateway' && channel !== 'routine' && channel !== 'heartbeat';
|
||||
}
|
||||
|
||||
function debouncedLoadThreads() {
|
||||
if (_loadThreadsTimer) clearTimeout(_loadThreadsTimer);
|
||||
_loadThreadsTimer = setTimeout(() => { _loadThreadsTimer = null; loadThreads(); }, 500);
|
||||
}
|
||||
|
||||
function loadThreads() {
|
||||
apiFetch('/api/chat/threads').then((data) => {
|
||||
// Pinned assistant thread
|
||||
@@ -1264,9 +1424,13 @@ function loadThreads() {
|
||||
const el = document.getElementById('assistant-thread');
|
||||
const isActive = currentThreadId === assistantThreadId;
|
||||
el.className = 'assistant-item' + (isActive ? ' active' : '');
|
||||
const labelEl = document.getElementById('assistant-label');
|
||||
if (labelEl) {
|
||||
const at = data.assistant_thread;
|
||||
labelEl.textContent = 'Assistant';
|
||||
}
|
||||
const meta = document.getElementById('assistant-meta');
|
||||
const count = data.assistant_thread.turn_count || 0;
|
||||
meta.textContent = count > 0 ? count + ' turns' : '';
|
||||
meta.textContent = relativeTime(data.assistant_thread.updated_at);
|
||||
}
|
||||
|
||||
// Regular threads
|
||||
@@ -1275,16 +1439,38 @@ function loadThreads() {
|
||||
const threads = data.threads || [];
|
||||
for (const thread of threads) {
|
||||
const item = document.createElement('div');
|
||||
item.className = 'thread-item' + (thread.id === currentThreadId ? ' active' : '');
|
||||
const isActive = thread.id === currentThreadId;
|
||||
item.className = 'thread-item' + (isActive ? ' active' : '');
|
||||
|
||||
// Channel badge for non-gateway threads
|
||||
const ch = thread.channel || 'gateway';
|
||||
if (ch !== 'gateway') {
|
||||
const badge = document.createElement('span');
|
||||
badge.className = 'thread-badge thread-badge-' + ch;
|
||||
badge.textContent = ch;
|
||||
item.appendChild(badge);
|
||||
}
|
||||
|
||||
const label = document.createElement('span');
|
||||
label.className = 'thread-label';
|
||||
label.textContent = thread.title || thread.id.substring(0, 8);
|
||||
label.title = thread.title ? thread.title + ' (' + thread.id + ')' : thread.id;
|
||||
label.textContent = threadTitle(thread);
|
||||
label.title = (thread.title || '') + ' (' + thread.id + ')';
|
||||
item.appendChild(label);
|
||||
|
||||
const meta = document.createElement('span');
|
||||
meta.className = 'thread-meta';
|
||||
meta.textContent = (thread.turn_count || 0) + ' turns';
|
||||
meta.textContent = relativeTime(thread.updated_at);
|
||||
item.appendChild(meta);
|
||||
|
||||
// Unread dot
|
||||
const unread = unreadThreads.get(thread.id) || 0;
|
||||
if (unread > 0 && !isActive) {
|
||||
const dot = document.createElement('span');
|
||||
dot.className = 'thread-unread';
|
||||
dot.textContent = unread > 9 ? '9+' : String(unread);
|
||||
item.appendChild(dot);
|
||||
}
|
||||
|
||||
item.addEventListener('click', () => switchThread(thread.id));
|
||||
list.appendChild(item);
|
||||
}
|
||||
@@ -1294,17 +1480,36 @@ function loadThreads() {
|
||||
switchToAssistant();
|
||||
}
|
||||
|
||||
// Enable chat input once a thread is available
|
||||
// Enable/disable chat input based on channel type
|
||||
if (currentThreadId) {
|
||||
enableChatInput();
|
||||
const currentThread = threads.find(t => t.id === currentThreadId);
|
||||
const ch = currentThread ? currentThread.channel : 'gateway';
|
||||
currentThreadIsReadOnly = isReadOnlyChannel(ch);
|
||||
if (currentThreadIsReadOnly) {
|
||||
disableChatInputReadOnly();
|
||||
} else {
|
||||
enableChatInput();
|
||||
}
|
||||
}
|
||||
}).catch(() => {});
|
||||
}
|
||||
|
||||
function disableChatInputReadOnly() {
|
||||
const input = document.getElementById('chat-input');
|
||||
const btn = document.getElementById('send-btn');
|
||||
if (input) {
|
||||
input.disabled = true;
|
||||
input.placeholder = 'Read-only thread (external channel)';
|
||||
}
|
||||
if (btn) btn.disabled = true;
|
||||
}
|
||||
|
||||
function switchToAssistant() {
|
||||
if (!assistantThreadId) return;
|
||||
finalizeActivityGroup();
|
||||
currentThreadId = assistantThreadId;
|
||||
currentThreadIsReadOnly = false;
|
||||
unreadThreads.delete(assistantThreadId);
|
||||
hasMore = false;
|
||||
oldestTimestamp = null;
|
||||
loadHistory();
|
||||
@@ -1314,6 +1519,7 @@ function switchToAssistant() {
|
||||
function switchThread(threadId) {
|
||||
finalizeActivityGroup();
|
||||
currentThreadId = threadId;
|
||||
unreadThreads.delete(threadId);
|
||||
hasMore = false;
|
||||
oldestTimestamp = null;
|
||||
loadHistory();
|
||||
@@ -1370,7 +1576,7 @@ chatInput.addEventListener('keydown', (e) => {
|
||||
}
|
||||
}
|
||||
|
||||
if (e.key === 'Enter' && !e.shiftKey) {
|
||||
if (e.key === 'Enter' && !e.shiftKey && !e.isComposing) {
|
||||
e.preventDefault();
|
||||
hideSlashAutocomplete();
|
||||
sendMessage();
|
||||
@@ -1889,6 +2095,13 @@ function renderAvailableExtensionCard(entry) {
|
||||
kind.textContent = kindLabels[entry.kind] || entry.kind;
|
||||
header.appendChild(kind);
|
||||
|
||||
if (entry.version) {
|
||||
const ver = document.createElement('span');
|
||||
ver.className = 'ext-version';
|
||||
ver.textContent = 'v' + entry.version;
|
||||
header.appendChild(ver);
|
||||
}
|
||||
|
||||
card.appendChild(header);
|
||||
|
||||
const desc = document.createElement('div');
|
||||
@@ -2049,6 +2262,13 @@ function renderExtensionCard(ext) {
|
||||
kind.textContent = kindLabels[ext.kind] || ext.kind;
|
||||
header.appendChild(kind);
|
||||
|
||||
if (ext.version) {
|
||||
const ver = document.createElement('span');
|
||||
ver.className = 'ext-version';
|
||||
ver.textContent = 'v' + ext.version;
|
||||
header.appendChild(ver);
|
||||
}
|
||||
|
||||
// Auth dot only for non-WASM-channel extensions (channels use the stepper instead)
|
||||
if (ext.kind !== 'wasm_channel') {
|
||||
const authDot = document.createElement('span');
|
||||
@@ -3291,6 +3511,10 @@ function shortModelName(model) {
|
||||
|
||||
function fetchGatewayStatus() {
|
||||
apiFetch('/api/gateway/status').then(function(data) {
|
||||
// Update restart button visibility
|
||||
restartEnabled = data.restart_enabled || false;
|
||||
updateRestartButtonVisibility();
|
||||
|
||||
var popover = document.getElementById('gateway-popover');
|
||||
var html = '';
|
||||
|
||||
@@ -3356,10 +3580,15 @@ let teeReportCache = null;
|
||||
let teeReportLoading = false;
|
||||
|
||||
function teeApiBase() {
|
||||
var parts = window.location.hostname.split('.');
|
||||
if (parts.length < 2) return null;
|
||||
var domain = parts.slice(1).join('.');
|
||||
return window.location.protocol + '//api.' + domain;
|
||||
var hostname = window.location.hostname;
|
||||
// Skip IP addresses (IPv4 and IPv6) and localhost
|
||||
if (hostname === "localhost" || /^(?:(?:25[0-5]|2[0-4][0-9]|[01]?[0-9][0-9]?)\.){3}(?:25[0-5]|2[0-4][0-9]|[01]?[0-9][0-9]?)$/.test(hostname) || hostname.indexOf(":") !== -1) {
|
||||
return null;
|
||||
}
|
||||
var parts = hostname.split(".");
|
||||
if (parts.length < 2) return null;
|
||||
var domain = parts.slice(1).join(".");
|
||||
return window.location.protocol + "//api." + domain;
|
||||
}
|
||||
|
||||
function teeInstanceName() {
|
||||
@@ -3370,13 +3599,19 @@ function checkTeeStatus() {
|
||||
var base = teeApiBase();
|
||||
if (!base) return;
|
||||
var name = teeInstanceName();
|
||||
fetch(base + '/instances/' + encodeURIComponent(name) + '/attestation').then(function(res) {
|
||||
if (!res.ok) throw new Error(res.status);
|
||||
return res.json();
|
||||
}).then(function(data) {
|
||||
teeInfo = data;
|
||||
document.getElementById('tee-shield').style.display = 'flex';
|
||||
}).catch(function() {});
|
||||
try {
|
||||
fetch(base + '/instances/' + encodeURIComponent(name) + '/attestation').then(function(res) {
|
||||
if (!res.ok) throw new Error(res.status);
|
||||
return res.json();
|
||||
}).then(function(data) {
|
||||
teeInfo = data;
|
||||
document.getElementById('tee-shield').style.display = 'flex';
|
||||
}).catch(function(err) {
|
||||
console.warn('Failed to fetch TEE attestation:', err);
|
||||
});
|
||||
} catch (e) {
|
||||
console.warn("Failed to check TEE status:", e);
|
||||
}
|
||||
}
|
||||
|
||||
function fetchTeeReport() {
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user