mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-26 15:40:18 +00:00
Compare commits
22
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c1ca3bb91c | ||
|
|
e499795b8c | ||
|
|
5e1da4827a | ||
|
|
d04af5cd75 | ||
|
|
956037c4d3 | ||
|
|
68a1851c19 | ||
|
|
dfa105539b | ||
|
|
8929baf76a | ||
|
|
e07dfab449 | ||
|
|
6783cba4e4 | ||
|
|
63302ab406 | ||
|
|
7c553b0973 | ||
|
|
5e44185e48 | ||
|
|
72623c9e5b | ||
|
|
6895adbcc9 | ||
|
|
f1480f471b | ||
|
|
9db949746f | ||
|
|
61a123a746 | ||
|
|
0e981429ee | ||
|
|
1b38a64e15 | ||
|
|
2e5f8b60d5 | ||
|
|
f0a0642e7d |
@@ -0,0 +1,97 @@
|
|||||||
|
---
|
||||||
|
description: Fetch a GitHub issue, create a branch, research the codebase, plan the fix, implement with tests, and commit
|
||||||
|
disable-model-invocation: true
|
||||||
|
allowed-tools: Bash(gh issue view:*), Bash(gh repo view:*), Bash(git fetch:*), Bash(git checkout:*), Bash(git status:*), Bash(git branch:*), Bash(git add:*), Bash(git commit:*), Bash(cargo fmt:*), Bash(cargo clippy:*), Bash(cargo test:*), Read, Edit, Write, Grep, Glob
|
||||||
|
argument-hint: "<issue-number or github-issue-url>"
|
||||||
|
---
|
||||||
|
|
||||||
|
# Fix GitHub Issue
|
||||||
|
|
||||||
|
## Step 1: Resolve the issue
|
||||||
|
|
||||||
|
Parse `$ARGUMENTS` to extract the issue number:
|
||||||
|
- If it's a URL like `https://github.com/owner/repo/issues/42`, extract `42`.
|
||||||
|
- If it's a bare number, use it directly.
|
||||||
|
- If empty, stop and ask the user for an issue number.
|
||||||
|
|
||||||
|
Fetch the issue:
|
||||||
|
|
||||||
|
```
|
||||||
|
gh issue view {number} --json title,body,labels,assignees,comments,state
|
||||||
|
```
|
||||||
|
|
||||||
|
If the issue is closed, warn the user and ask if they still want to proceed.
|
||||||
|
|
||||||
|
## Step 2: Create a branch
|
||||||
|
|
||||||
|
Create a fresh branch off the latest main:
|
||||||
|
|
||||||
|
1. Fetch latest: `git fetch origin`
|
||||||
|
2. Detect default branch: `gh repo view --json defaultBranchRef --jq .defaultBranchRef.name`
|
||||||
|
3. Create and switch to a new branch: `git checkout -b fix/{number}-{short-slug} origin/{default-branch}`
|
||||||
|
- `{short-slug}` is 3-5 words from the issue title, lowercase, hyphenated (e.g. `fix/42-idor-workspace-check`)
|
||||||
|
|
||||||
|
If the working tree has uncommitted changes, warn the user and stop. Do not stash or discard their work.
|
||||||
|
|
||||||
|
## Step 3: Understand the issue
|
||||||
|
|
||||||
|
Summarize the issue in 2-3 sentences. Identify:
|
||||||
|
- **What's broken or missing** (the symptom or feature request)
|
||||||
|
- **Acceptance criteria** (what "done" looks like, from the issue body or comments)
|
||||||
|
- **Constraints** (mentioned technologies, backward compatibility, performance requirements)
|
||||||
|
|
||||||
|
If the issue is unclear or ambiguous, list the open questions. These will be addressed during planning.
|
||||||
|
|
||||||
|
## Step 4: Research the codebase
|
||||||
|
|
||||||
|
Before planning, gather context:
|
||||||
|
|
||||||
|
1. **Find relevant code** - Search for files, functions, types, and patterns mentioned in the issue. Read them in full.
|
||||||
|
2. **Trace the flow** - If the issue is about a specific behavior, trace the code path from the entry point (route handler, CLI command, etc.) through to the relevant logic.
|
||||||
|
3. **Check existing tests** - Find tests related to the affected code. Understand what's already covered.
|
||||||
|
4. **Check for prior art** - Look for similar patterns in the codebase that solve analogous problems. Prefer consistency with existing patterns.
|
||||||
|
|
||||||
|
## Step 5: Enter planning mode
|
||||||
|
|
||||||
|
Enter planning mode to design the implementation. The plan MUST cover:
|
||||||
|
|
||||||
|
1. **Root cause** (for bugs) or **design approach** (for features)
|
||||||
|
2. **Files to modify** with specific descriptions of what changes in each
|
||||||
|
3. **New files** (if any) with justification for why they're needed
|
||||||
|
4. **Tests to add** - every code path introduced or changed needs a test:
|
||||||
|
- Happy path (expected input produces expected output)
|
||||||
|
- Error paths (invalid input, missing data, permission denied)
|
||||||
|
- Edge cases (empty collections, boundary values, concurrent access)
|
||||||
|
5. **IronClaw-specific concerns**:
|
||||||
|
- If the change touches persistence, both database backends must be updated (`postgres.rs` and `libsql_backend.rs`)
|
||||||
|
- New `Database` trait methods need implementations in both backends
|
||||||
|
- No `.unwrap()` or `.expect()` in production code
|
||||||
|
- Use `crate::` imports, not `super::`
|
||||||
|
- Error types via `thiserror` in `error.rs`
|
||||||
|
6. **Migration or compatibility concerns** (if any)
|
||||||
|
|
||||||
|
Follow the project's CLAUDE.md guidance for architecture decisions.
|
||||||
|
|
||||||
|
Wait for user approval before implementing.
|
||||||
|
|
||||||
|
## Step 6: Implement
|
||||||
|
|
||||||
|
After the plan is approved:
|
||||||
|
|
||||||
|
1. Implement each change from the plan.
|
||||||
|
2. Write all planned tests.
|
||||||
|
3. Run IronClaw's full quality gate:
|
||||||
|
- `cargo fmt`
|
||||||
|
- `cargo clippy --all --benches --tests --examples --all-features` (zero warnings)
|
||||||
|
- `cargo test --lib` (all tests pass)
|
||||||
|
4. If any check fails, fix it before proceeding.
|
||||||
|
|
||||||
|
Note: Integration tests (`--test workspace_integration`) require PostgreSQL and are expected to fail locally. Only `--lib` test failures are blocking.
|
||||||
|
|
||||||
|
## Step 7: Commit and summarize
|
||||||
|
|
||||||
|
1. Commit with a descriptive message referencing the issue (e.g. `fix: prevent IDOR in function call outputs (#42)`).
|
||||||
|
2. Summarize what was done:
|
||||||
|
- Files changed with line references
|
||||||
|
- Tests added and what they cover
|
||||||
|
- Any follow-up work or open questions
|
||||||
@@ -0,0 +1,81 @@
|
|||||||
|
---
|
||||||
|
description: Respond to PR review comments — triage, plan fixes, implement after confirmation, push, and reply to reviewers
|
||||||
|
disable-model-invocation: true
|
||||||
|
allowed-tools: Bash(gh pr list:*), Bash(gh pr comment:*), Bash(gh api:*), Bash(gh repo view:*), Bash(git branch:*), Bash(git status:*), Bash(git add:*), Bash(git commit:*), Bash(git push:*), Bash(cargo fmt:*), Bash(cargo clippy:*), Bash(cargo test:*), Read, Edit, Write, Grep, Glob
|
||||||
|
argument-hint: "[pr-number (optional, auto-detects from branch)]"
|
||||||
|
---
|
||||||
|
|
||||||
|
# Review and Address PR Comments
|
||||||
|
|
||||||
|
## Step 1: Find the PR
|
||||||
|
|
||||||
|
If `$ARGUMENTS` is provided, use that as the PR number. Otherwise, detect the PR for the current branch:
|
||||||
|
|
||||||
|
```
|
||||||
|
gh pr list --head $(git branch --show-current) --json number,title,url --jq '.[0]'
|
||||||
|
```
|
||||||
|
|
||||||
|
If no PR is found, tell the user and stop.
|
||||||
|
|
||||||
|
## Step 2: Fetch all review comments
|
||||||
|
|
||||||
|
Resolve the repo owner and name:
|
||||||
|
|
||||||
|
```
|
||||||
|
gh repo view --json owner,name --jq '"\(.owner.login)/\(.name)"'
|
||||||
|
```
|
||||||
|
|
||||||
|
Fetch the full set of review comments (not issue-level comments):
|
||||||
|
|
||||||
|
```
|
||||||
|
gh api --paginate repos/{owner}/{repo}/pulls/{number}/comments
|
||||||
|
```
|
||||||
|
|
||||||
|
Also fetch the review summaries:
|
||||||
|
|
||||||
|
```
|
||||||
|
gh api --paginate repos/{owner}/{repo}/pulls/{number}/reviews
|
||||||
|
```
|
||||||
|
|
||||||
|
Deduplicate comments that appear multiple times (bots sometimes post the same finding under different IDs). Group by the actual issue being raised, not by comment ID.
|
||||||
|
|
||||||
|
## Step 3: Triage and plan
|
||||||
|
|
||||||
|
For each unique issue raised in the comments:
|
||||||
|
|
||||||
|
1. **Check if already addressed** - Read the current code at the referenced location. If a prior commit already fixed it, note it as "already resolved".
|
||||||
|
2. **Assess validity** - Determine if the comment identifies a real problem or is a false positive. Be honest about false positives but explain why.
|
||||||
|
3. **Classify severity** - Critical (security/data loss), High (bugs/broken behavior), Medium (correctness/robustness), Low (style/naming/nits).
|
||||||
|
4. **Plan the fix** - For each valid unresolved issue, describe the specific code change needed.
|
||||||
|
|
||||||
|
Present the plan as a table to the user:
|
||||||
|
|
||||||
|
| # | Issue | File:Line | Severity | Status | Planned Fix |
|
||||||
|
|---|-------|-----------|----------|--------|-------------|
|
||||||
|
|
||||||
|
Wait for user confirmation before proceeding to implementation.
|
||||||
|
|
||||||
|
## Step 4: Implement fixes
|
||||||
|
|
||||||
|
After user confirms:
|
||||||
|
|
||||||
|
1. Implement each fix in the plan.
|
||||||
|
2. Run IronClaw's quality gate to verify nothing breaks:
|
||||||
|
- `cargo fmt`
|
||||||
|
- `cargo clippy --all --benches --tests --examples --all-features`
|
||||||
|
- `cargo test --lib`
|
||||||
|
3. Commit with a descriptive message referencing the PR review.
|
||||||
|
4. Push to the branch.
|
||||||
|
|
||||||
|
## Step 5: Reply to comments
|
||||||
|
|
||||||
|
For each comment addressed, reply on the PR with a short message stating what was fixed and the commit SHA. For false positives or already-resolved items, reply explaining why no change was needed.
|
||||||
|
|
||||||
|
## Rules
|
||||||
|
|
||||||
|
- Never guess at code you haven't read. Always read the referenced file and line before assessing a comment.
|
||||||
|
- Group duplicate comments (same issue reported by multiple bots) and reply to all of them.
|
||||||
|
- Do not make changes beyond what the review comments ask for. Stay focused.
|
||||||
|
- If a comment suggests a change you disagree with, present your reasoning to the user during the planning phase rather than silently ignoring it.
|
||||||
|
- Follow IronClaw conventions: no `.unwrap()` in production code, use `crate::` imports, `thiserror` errors.
|
||||||
|
- If changes touch persistence, verify both database backends are updated.
|
||||||
@@ -0,0 +1,245 @@
|
|||||||
|
---
|
||||||
|
description: Deep audit of the IronClaw crate for vulnerabilities, bugs, unfinished work, inconsistencies, and oversights
|
||||||
|
disable-model-invocation: true
|
||||||
|
allowed-tools: Bash(cargo fmt:*), Bash(cargo clippy:*), Bash(cargo test:*), Bash(cargo audit:*), Bash(git diff:*), Bash(git log:*), Bash(git show:*), Bash(wc:*), Read, Grep, Glob, Task
|
||||||
|
argument-hint: "[path/to/crate]"
|
||||||
|
---
|
||||||
|
|
||||||
|
# Rust Crate Audit
|
||||||
|
|
||||||
|
You are performing a thorough audit of a Rust crate. Your goal is to find every vulnerability, bug, unfinished piece of work, inconsistency, and oversight before it ships. Leave no stone unturned.
|
||||||
|
|
||||||
|
## Step 1: Locate the crate
|
||||||
|
|
||||||
|
Parse `$ARGUMENTS`:
|
||||||
|
- If a path is provided, use it as the crate root.
|
||||||
|
- If empty, use the current working directory.
|
||||||
|
|
||||||
|
Verify it's a valid Rust crate by checking for `Cargo.toml`. If not found, stop and ask the user.
|
||||||
|
|
||||||
|
## Step 2: Understand the crate
|
||||||
|
|
||||||
|
Read `Cargo.toml` to understand:
|
||||||
|
- Crate name, version, edition
|
||||||
|
- Dependencies (look for outdated, unmaintained, or suspicious crates)
|
||||||
|
- Feature flags and their implications
|
||||||
|
- Build scripts (`build.rs`) if any
|
||||||
|
|
||||||
|
Read `CLAUDE.md`, `README.md`, or top-level documentation if present to understand intent and architecture.
|
||||||
|
|
||||||
|
Read `src/lib.rs` or `src/main.rs` to get the module tree. Then read each module's `mod.rs` or top-level file to build a mental map of the crate's structure before diving into details.
|
||||||
|
Read all Rust files (`src/*.rs`) to make sure everything is in context when you are reasoning.
|
||||||
|
|
||||||
|
## Step 3: Run the compiler's checks
|
||||||
|
|
||||||
|
Run these commands and capture output. Do NOT fix anything, just collect findings:
|
||||||
|
|
||||||
|
```
|
||||||
|
cargo fmt --check 2>&1
|
||||||
|
```
|
||||||
|
|
||||||
|
```
|
||||||
|
cargo clippy --all --benches --tests --examples --all-features -- -W clippy::all -W clippy::pedantic -W clippy::nursery 2>&1
|
||||||
|
```
|
||||||
|
|
||||||
|
```
|
||||||
|
cargo test --lib 2>&1
|
||||||
|
```
|
||||||
|
|
||||||
|
If any of these fail, record the failures as findings. If `cargo test` has ignored tests, note which ones and why.
|
||||||
|
|
||||||
|
Note: Integration tests (`--test workspace_integration`) require a PostgreSQL database and are expected to fail locally. Only report `--lib` test failures as blocking.
|
||||||
|
|
||||||
|
## Step 4: Scan for unfinished work
|
||||||
|
|
||||||
|
Search the entire `src/` tree for:
|
||||||
|
|
||||||
|
```
|
||||||
|
todo!
|
||||||
|
unimplemented!
|
||||||
|
fixme
|
||||||
|
FIXME
|
||||||
|
TODO
|
||||||
|
HACK
|
||||||
|
XXX
|
||||||
|
SAFETY:
|
||||||
|
stub
|
||||||
|
placeholder
|
||||||
|
temporary
|
||||||
|
```
|
||||||
|
|
||||||
|
For each match:
|
||||||
|
- Is it in production code or test code?
|
||||||
|
- Is it a genuine incomplete feature or a deliberate placeholder?
|
||||||
|
- Is there a tracking issue referenced?
|
||||||
|
- Could this panic at runtime?
|
||||||
|
|
||||||
|
Any `todo!()` or `unimplemented!()` in non-test code is **High severity** (runtime panic).
|
||||||
|
|
||||||
|
## Step 5: Audit for vulnerabilities and unsafe code
|
||||||
|
|
||||||
|
### 5a. Unsafe code
|
||||||
|
|
||||||
|
Search for all `unsafe` blocks. For each one:
|
||||||
|
- Is the safety invariant documented with a `// SAFETY:` comment?
|
||||||
|
- Is the invariant actually upheld by the surrounding code?
|
||||||
|
- Could the unsafe block be replaced with a safe alternative?
|
||||||
|
- Are there any pointer dereferences, transmutes, or FFI calls?
|
||||||
|
|
||||||
|
### 5b. Unwrap and panic paths
|
||||||
|
|
||||||
|
Search for `.unwrap()`, `.expect(`, `panic!`, `unreachable!` in non-test code. For each:
|
||||||
|
- Can this actually panic in production?
|
||||||
|
- Is there a code path that reaches this with None/Err?
|
||||||
|
- Should it be replaced with proper error handling (`?`, `.ok()`, `.unwrap_or_default()`)?
|
||||||
|
|
||||||
|
IronClaw convention: `.unwrap()` and `.expect()` are banned in production code. Any occurrence outside `#[cfg(test)]` blocks is a **High severity** finding.
|
||||||
|
|
||||||
|
### 5c. SQL and injection vectors
|
||||||
|
|
||||||
|
Search for string formatting used in SQL queries, shell commands, or HTML:
|
||||||
|
- `format!` used near `.execute(`, `.query(`, `Command::new(`
|
||||||
|
- String interpolation in query construction vs parameterized queries
|
||||||
|
- User input flowing into file paths (`Path::new`, `std::fs::`)
|
||||||
|
|
||||||
|
IronClaw has two database backends (PostgreSQL and libSQL). Check both for injection vectors.
|
||||||
|
|
||||||
|
### 5d. Cryptographic issues
|
||||||
|
|
||||||
|
If the crate uses crypto:
|
||||||
|
- Are comparisons constant-time? (look for `==` on secrets/hashes vs `subtle::ConstantTimeEq`)
|
||||||
|
- Is randomness from `OsRng` / `thread_rng` and not a fixed seed?
|
||||||
|
- Are keys/secrets zeroized after use? (`secrecy`, `zeroize` crates)
|
||||||
|
- Are deprecated algorithms used? (MD5, SHA1 for security, RC4, DES)
|
||||||
|
|
||||||
|
### 5e. Resource exhaustion
|
||||||
|
|
||||||
|
- Are there unbounded allocations? (`Vec` growing from user input without limits)
|
||||||
|
- Are there unbounded loops? (retry loops without max attempts)
|
||||||
|
- Are file reads bounded? (`std::fs::read_to_string` on user-provided paths)
|
||||||
|
- Are timeouts set on all network operations?
|
||||||
|
- Are there connection/resource leaks? (opened but never closed, missing `Drop`)
|
||||||
|
|
||||||
|
### 5f. Error handling
|
||||||
|
|
||||||
|
- Are errors swallowed silently? (`let _ = ...`, `.ok()` discarding errors that matter)
|
||||||
|
- Do error types carry enough context to debug in production?
|
||||||
|
- Are there error type mismatches? (returning generic `anyhow::Error` where a typed error would prevent confusion)
|
||||||
|
- Is `thiserror` used consistently for error types (IronClaw convention)?
|
||||||
|
|
||||||
|
## Step 6: Check for inconsistencies
|
||||||
|
|
||||||
|
### 6a. Naming conventions
|
||||||
|
|
||||||
|
- Are types, functions, modules named consistently? (e.g., mixing `get_` and `fetch_`, `create_` and `new_`)
|
||||||
|
- Do similar operations follow the same patterns?
|
||||||
|
|
||||||
|
### 6b. Duplicate or near-duplicate code
|
||||||
|
|
||||||
|
Look for:
|
||||||
|
- Functions that do nearly the same thing with minor variations (candidates for generics or shared helpers)
|
||||||
|
- Repeated error mapping patterns that should be extracted
|
||||||
|
- Copy-pasted SQL queries or string templates with slight differences
|
||||||
|
- Identical struct definitions or conversion logic in different modules
|
||||||
|
|
||||||
|
### 6c. API consistency
|
||||||
|
|
||||||
|
- Do similar functions take arguments in the same order?
|
||||||
|
- Are return types consistent? (e.g., some functions return `Option<T>`, similar ones return `Result<T, E>`)
|
||||||
|
- Are visibility modifiers consistent? (`pub` where it should be `pub(crate)`, or vice versa)
|
||||||
|
|
||||||
|
### 6d. Dead code and unused items
|
||||||
|
|
||||||
|
- Are there functions, structs, or modules that nothing references?
|
||||||
|
- Are there `#[allow(dead_code)]` annotations that should be investigated?
|
||||||
|
- Are there feature-gated items where the feature is never enabled?
|
||||||
|
|
||||||
|
### 6e. Import style
|
||||||
|
|
||||||
|
IronClaw convention: use `crate::` imports, not `super::`. Flag any `super::` imports in non-test code.
|
||||||
|
|
||||||
|
## Step 7: Inspect for change oversights
|
||||||
|
|
||||||
|
### 7a. Partial refactors
|
||||||
|
|
||||||
|
- Are there old patterns coexisting with new patterns?
|
||||||
|
- Are there renamed types/functions where some call sites still use the old name via a compatibility alias?
|
||||||
|
- Are there comments referencing behavior that no longer exists?
|
||||||
|
|
||||||
|
### 7b. Trait implementation gaps
|
||||||
|
|
||||||
|
- If a trait is defined, do all intended types implement it?
|
||||||
|
- Are there `impl` blocks that look incomplete?
|
||||||
|
- Are `Default` implementations sensible?
|
||||||
|
|
||||||
|
IronClaw key traits: `Database` (~60 methods), `Channel`, `Tool`, `LlmProvider`, `SuccessEvaluator`, `EmbeddingProvider`. If any new methods were added to `Database`, verify both `postgres.rs` and `libsql_backend.rs` implement them.
|
||||||
|
|
||||||
|
### 7c. Test coverage gaps
|
||||||
|
|
||||||
|
- Are there public functions without any test?
|
||||||
|
- Are there error paths without tests?
|
||||||
|
- Are there recently-changed functions where the tests still assert old behavior?
|
||||||
|
|
||||||
|
### 7d. Documentation drift
|
||||||
|
|
||||||
|
- Do doc comments match actual function behavior?
|
||||||
|
- Are examples in doc comments still valid and compilable?
|
||||||
|
|
||||||
|
## Step 8: Dependency audit
|
||||||
|
|
||||||
|
Review `Cargo.toml` and `Cargo.lock`:
|
||||||
|
- Are there duplicate versions of the same crate in the lock file? (potential version conflicts)
|
||||||
|
- Are there dependencies with known security advisories? Run `cargo audit` to check (install with `cargo install cargo-audit` if not present).
|
||||||
|
- Are there heavy dependencies used for trivial functionality?
|
||||||
|
- Are dependency features minimal?
|
||||||
|
|
||||||
|
## Step 9: Present findings
|
||||||
|
|
||||||
|
Compile all findings into a structured report. Group by severity, then by category.
|
||||||
|
|
||||||
|
### Format
|
||||||
|
|
||||||
|
For each finding:
|
||||||
|
|
||||||
|
```
|
||||||
|
### [Severity] Category: One-line summary
|
||||||
|
|
||||||
|
**Location:** `file_path:line_number`
|
||||||
|
**Category:** Vulnerability | Bug | Unfinished | Inconsistency | Duplicate | Oversight | Style
|
||||||
|
|
||||||
|
**Description:**
|
||||||
|
Detailed explanation of the issue, why it matters, and how it could manifest.
|
||||||
|
|
||||||
|
**Suggested fix:**
|
||||||
|
Concrete suggestion with code if applicable.
|
||||||
|
```
|
||||||
|
|
||||||
|
### Severity levels
|
||||||
|
|
||||||
|
- **Critical**: Security vulnerability, data loss, or crash in production
|
||||||
|
- **High**: Bug that causes incorrect behavior, `todo!()`/`unimplemented!()` in prod code, or missing validation on trust boundaries
|
||||||
|
- **Medium**: Inconsistency, duplicate code, incomplete error handling, missing tests for important paths
|
||||||
|
- **Low**: Naming inconsistency, unnecessary complexity, documentation drift, minor dead code
|
||||||
|
- **Nit**: Style preference, optional improvement
|
||||||
|
|
||||||
|
### Summary table
|
||||||
|
|
||||||
|
End with a summary table:
|
||||||
|
|
||||||
|
| # | Severity | Category | File:Line | Finding |
|
||||||
|
|---|----------|----------|-----------|---------|
|
||||||
|
|
||||||
|
And a final tally: X Critical, Y High, Z Medium, W Low, V Nit.
|
||||||
|
|
||||||
|
## Rules
|
||||||
|
|
||||||
|
- Read every file before reporting on it. Never guess about code you haven't seen.
|
||||||
|
- Be specific. "This might have issues" is worthless. "Line 42 calls `.unwrap()` on a `Result` that returns `Err` when the DB connection is dropped" is useful.
|
||||||
|
- Distinguish certainty levels: "this IS a bug" vs "this COULD be a bug if X".
|
||||||
|
- Don't invent problems to look thorough. If the code is solid, say so.
|
||||||
|
- Focus on substance over style. Don't flag formatting unless it causes real confusion.
|
||||||
|
- Respect existing project conventions (check CLAUDE.md). Don't flag patterns the project explicitly endorses.
|
||||||
|
- When in doubt about severity, round up.
|
||||||
|
- For large crates (>50 files), prioritize: core logic > public API > internal utilities > tests > examples.
|
||||||
|
- Use the Task tool to parallelize file reading across modules when the crate is large.
|
||||||
|
- Do NOT fix anything. This is a read-only audit. Report findings for the user to action.
|
||||||
@@ -0,0 +1,170 @@
|
|||||||
|
---
|
||||||
|
description: Paranoid architect review of a PR — fetches diff, reads changed files, deep review across 6 lenses, posts findings as GitHub comments
|
||||||
|
disable-model-invocation: true
|
||||||
|
allowed-tools: Bash(gh pr view:*), Bash(gh pr diff:*), Bash(gh pr comment:*), Bash(gh api:*), Bash(gh repo view:*), Bash(git diff:*), Bash(git log:*), Read, Grep, Glob
|
||||||
|
argument-hint: "<pr-number or github-pr-url>"
|
||||||
|
---
|
||||||
|
|
||||||
|
# Paranoid Architect Code Review
|
||||||
|
|
||||||
|
You are reviewing this PR as a paranoid architect. Your job is to find every bug, vulnerability, race condition, edge case, and undocumented assumption before it ships. Assume adversarial users, concurrent access, and Murphy's law.
|
||||||
|
|
||||||
|
## Step 1: Resolve the PR
|
||||||
|
|
||||||
|
Parse `$ARGUMENTS` to extract the PR number:
|
||||||
|
- If it's a URL like `https://github.com/owner/repo/pull/123`, extract `123`.
|
||||||
|
- If it's a bare number, use it directly.
|
||||||
|
- If empty, stop and ask the user for a PR number.
|
||||||
|
|
||||||
|
Fetch PR metadata (including head commit SHA for posting line comments later):
|
||||||
|
|
||||||
|
```
|
||||||
|
gh pr view {number} --json title,body,baseRefName,headRefName,headRefOid,files,additions,deletions
|
||||||
|
```
|
||||||
|
|
||||||
|
Save the `headRefOid` value, you'll need it as `commit_id` in Step 6.
|
||||||
|
|
||||||
|
## Step 2: Load the full diff
|
||||||
|
|
||||||
|
```
|
||||||
|
gh pr diff {number}
|
||||||
|
```
|
||||||
|
|
||||||
|
Also get the list of changed files:
|
||||||
|
|
||||||
|
```
|
||||||
|
gh pr diff {number} --name-only
|
||||||
|
```
|
||||||
|
|
||||||
|
## Step 3: Read every changed file in full
|
||||||
|
|
||||||
|
For each changed file, read the ENTIRE current file (not just the diff hunks). You need surrounding context to catch:
|
||||||
|
- Callers of modified functions that now behave differently
|
||||||
|
- Trait/interface contracts that the change may violate
|
||||||
|
- Invariants established elsewhere that the diff breaks
|
||||||
|
|
||||||
|
If the PR touches more than 20 files, still read all of them, but process in this priority order: service logic > routes/handlers > models/types > tests > docs. Batch reads in groups of ~20 if needed.
|
||||||
|
|
||||||
|
## Step 4: Deep review
|
||||||
|
|
||||||
|
Go through the changes with each of these lenses. For every finding, note the file, line range, severity, and a concrete description.
|
||||||
|
|
||||||
|
### IronClaw-specific checks
|
||||||
|
|
||||||
|
In addition to the general lenses below, check IronClaw conventions (see CLAUDE.md):
|
||||||
|
- No `.unwrap()` or `.expect()` in production code (tests are fine)
|
||||||
|
- Use `crate::` imports, not `super::`
|
||||||
|
- Error types use `thiserror` in `error.rs`
|
||||||
|
- If the change touches persistence, verify both database backends are updated (PostgreSQL in `postgres.rs` AND libSQL in `libsql_backend.rs`)
|
||||||
|
- New tools must implement the `Tool` trait correctly and be registered in `registry.rs`
|
||||||
|
- External tool output must pass through the safety layer
|
||||||
|
|
||||||
|
### 4a. Correctness and bugs
|
||||||
|
|
||||||
|
- Off-by-one errors, wrong comparison operators, inverted conditions
|
||||||
|
- Unreachable code, dead branches, impossible match arms
|
||||||
|
- Type confusion (mixing up IDs, using wrong enum variant)
|
||||||
|
- Incorrect error propagation (swallowed errors, wrong error type/status code)
|
||||||
|
- Broken invariants (e.g. uniqueness assumptions violated, ordering assumptions wrong)
|
||||||
|
- Concurrency issues (TOCTOU, missing locks, race conditions between check and use)
|
||||||
|
|
||||||
|
### 4b. Edge cases and failure handling
|
||||||
|
|
||||||
|
- What happens with empty input, None/null, zero-length collections?
|
||||||
|
- What happens when external services fail (DB down, HTTP timeout, malformed response)?
|
||||||
|
- What happens at integer boundaries (overflow, underflow, i64::MAX)?
|
||||||
|
- What happens with malformed or adversarial input (invalid UTF-8, huge payloads, deeply nested JSON)?
|
||||||
|
- Are all error paths tested? Does every `?` propagation make sense?
|
||||||
|
- Are partial failures handled (e.g. wrote to DB but failed to emit event)?
|
||||||
|
|
||||||
|
### 4c. Security (assume a malicious actor)
|
||||||
|
|
||||||
|
- **Authentication/Authorization bypass**: Can an unauthenticated user reach this? Can workspace A's user access workspace B's data? Are there IDOR vulnerabilities?
|
||||||
|
- **Injection**: SQL injection via string interpolation? Command injection? Log injection? Header injection?
|
||||||
|
- **Data leakage**: Are secrets, PII, or conversation content logged? Returned in error messages? Exposed in API responses?
|
||||||
|
- **Resource exhaustion / DoS**: Can an attacker send unbounded input? Trigger expensive operations without rate limits? Cause OOM via large allocations?
|
||||||
|
- **Financial abuse**: Can tokens/credits be consumed without being tracked? Can usage limits be bypassed?
|
||||||
|
- **Replay / race conditions**: Can the same request be replayed for double-spend? Can concurrent requests bypass limits?
|
||||||
|
- **Cryptographic issues**: Timing attacks on comparisons? Weak randomness? Missing HMAC verification?
|
||||||
|
|
||||||
|
### 4d. Test coverage
|
||||||
|
|
||||||
|
- Is every new public function/method tested?
|
||||||
|
- Are error paths tested (not just happy paths)?
|
||||||
|
- Are edge cases covered (empty input, boundary values, concurrent access)?
|
||||||
|
- Do existing tests still make sense with the new changes, or do they assert stale behavior?
|
||||||
|
- Are there integration/e2e tests for the full flow?
|
||||||
|
- If a test is missing, describe exactly what test should be written.
|
||||||
|
|
||||||
|
### 4e. Documentation and assumptions
|
||||||
|
|
||||||
|
- Are new assumptions documented in comments? (e.g. "this field is always non-empty because X")
|
||||||
|
- Are non-obvious algorithms or business rules explained?
|
||||||
|
- Are API contracts (request/response shapes, error codes, status codes) documented?
|
||||||
|
- Are there TODO/FIXME/HACK comments that should be tracked as issues?
|
||||||
|
|
||||||
|
### 4f. Architectural concerns
|
||||||
|
|
||||||
|
- Does this change follow existing patterns in the codebase, or does it introduce a new one without justification?
|
||||||
|
- Are there unnecessary abstractions or premature generalizations?
|
||||||
|
- Is there duplicated logic that should be extracted?
|
||||||
|
- Are dependencies between modules clean, or does this create circular/tight coupling?
|
||||||
|
- Will this change make future work harder?
|
||||||
|
|
||||||
|
## Step 5: Present findings
|
||||||
|
|
||||||
|
Summarize findings to the user as a table:
|
||||||
|
|
||||||
|
| # | Severity | Category | File:Line | Finding | Suggested Fix |
|
||||||
|
|---|----------|----------|-----------|---------|---------------|
|
||||||
|
|
||||||
|
Severity levels:
|
||||||
|
- **Critical**: Security vulnerability, data loss, or financial exploit
|
||||||
|
- **High**: Bug that will cause incorrect behavior in production
|
||||||
|
- **Medium**: Robustness issue, missing validation, or incomplete error handling
|
||||||
|
- **Low**: Style, naming, documentation, or minor improvement
|
||||||
|
- **Nit**: Optional suggestion, take-it-or-leave-it
|
||||||
|
|
||||||
|
Ask the user which findings to post as PR comments. Default: all Critical, High, and Medium.
|
||||||
|
|
||||||
|
## Step 6: Post comments on GitHub
|
||||||
|
|
||||||
|
Resolve the repo owner and name if not already known:
|
||||||
|
|
||||||
|
```
|
||||||
|
gh repo view --json owner,name --jq '"\(.owner.login)/\(.name)"'
|
||||||
|
```
|
||||||
|
|
||||||
|
For each approved finding, post a review comment on the PR at the specific file and line. Use the `headRefOid` from Step 1 as the `commit_id`:
|
||||||
|
|
||||||
|
```
|
||||||
|
gh api repos/{owner}/{repo}/pulls/{number}/comments \
|
||||||
|
-f body="..." \
|
||||||
|
-f path="..." \
|
||||||
|
-f commit_id="{headRefOid}" \
|
||||||
|
-F line=... \
|
||||||
|
-f side="RIGHT"
|
||||||
|
```
|
||||||
|
|
||||||
|
For findings that span multiple locations or are architectural, post as a regular PR comment:
|
||||||
|
|
||||||
|
```
|
||||||
|
gh pr comment {number} --body "..."
|
||||||
|
```
|
||||||
|
|
||||||
|
Format each comment clearly:
|
||||||
|
- Severity tag (e.g. `**High Severity**`)
|
||||||
|
- One-line summary
|
||||||
|
- Detailed explanation of the issue
|
||||||
|
- Concrete suggestion for the fix (with code if possible)
|
||||||
|
|
||||||
|
## Rules
|
||||||
|
|
||||||
|
- Read every changed file in full before writing a single finding. Context matters.
|
||||||
|
- Never post a comment about code you haven't actually read. Verify line numbers against the actual file.
|
||||||
|
- Be specific. "This might have issues" is useless. "Line 42 returns 404 but should return 400 because X" is useful.
|
||||||
|
- Distinguish between "this IS a bug" and "this COULD be a bug if X". Be honest about certainty.
|
||||||
|
- Don't nitpick formatting or style unless it causes actual confusion. Focus on substance.
|
||||||
|
- If the code is good and you find nothing, say so. Don't invent problems to look thorough.
|
||||||
|
- Respect the project's CLAUDE.md privacy rules: never include customer data, secrets, or PII in comments.
|
||||||
|
- When in doubt about severity, round up. It's cheaper to dismiss a false alarm than to miss a real bug.
|
||||||
@@ -39,7 +39,6 @@ permissions:
|
|||||||
# If there's a prerelease-style suffix to the version, then the release(s)
|
# If there's a prerelease-style suffix to the version, then the release(s)
|
||||||
# will be marked as a prerelease.
|
# will be marked as a prerelease.
|
||||||
on:
|
on:
|
||||||
pull_request:
|
|
||||||
push:
|
push:
|
||||||
tags:
|
tags:
|
||||||
- '**[0-9]+.[0-9]+.[0-9]+*'
|
- '**[0-9]+.[0-9]+.[0-9]+*'
|
||||||
|
|||||||
@@ -1,6 +1,15 @@
|
|||||||
|
|
||||||
.env
|
.env
|
||||||
.env.local
|
.env.local
|
||||||
|
.env.*
|
||||||
|
!.env.example
|
||||||
|
|
||||||
|
# Claude Code worktrees
|
||||||
|
.claude/worktrees/
|
||||||
|
|
||||||
|
# Sidecar tool data
|
||||||
|
.sidecar/
|
||||||
|
.todos/
|
||||||
|
|
||||||
target/
|
target/
|
||||||
|
|
||||||
|
|||||||
@@ -7,6 +7,60 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
## [Unreleased]
|
## [Unreleased]
|
||||||
|
|
||||||
|
## [0.5.0](https://github.com/nearai/ironclaw/compare/v0.4.0...v0.5.0) - 2026-02-17
|
||||||
|
|
||||||
|
### Added
|
||||||
|
|
||||||
|
- add cooldown management to FailoverProvider ([#114](https://github.com/nearai/ironclaw/pull/114))
|
||||||
|
|
||||||
|
## [0.4.0](https://github.com/nearai/ironclaw/compare/v0.3.0...v0.4.0) - 2026-02-17
|
||||||
|
|
||||||
|
### Added
|
||||||
|
|
||||||
|
- move per-invocation approval check into Tool trait ([#119](https://github.com/nearai/ironclaw/pull/119))
|
||||||
|
- add polished boot screen on CLI startup ([#118](https://github.com/nearai/ironclaw/pull/118))
|
||||||
|
- Add lifecycle hooks system with 6 interception points ([#18](https://github.com/nearai/ironclaw/pull/18))
|
||||||
|
|
||||||
|
### Other
|
||||||
|
|
||||||
|
- remove accidentally committed .sidecar and .todos directories ([#123](https://github.com/nearai/ironclaw/pull/123))
|
||||||
|
|
||||||
|
## [0.3.0](https://github.com/nearai/ironclaw/compare/v0.2.0...v0.3.0) - 2026-02-17
|
||||||
|
|
||||||
|
### Added
|
||||||
|
|
||||||
|
- direct api key and cheap model ([#116](https://github.com/nearai/ironclaw/pull/116))
|
||||||
|
|
||||||
|
## [0.2.0](https://github.com/nearai/ironclaw/compare/v0.1.3...v0.2.0) - 2026-02-16
|
||||||
|
|
||||||
|
### Added
|
||||||
|
|
||||||
|
- mark Ollama + OpenAI-compatible as implemented ([#102](https://github.com/nearai/ironclaw/pull/102))
|
||||||
|
- multi-provider inference + libSQL onboarding selection ([#92](https://github.com/nearai/ironclaw/pull/92))
|
||||||
|
- add multi-provider LLM failover with retry backoff ([#28](https://github.com/nearai/ironclaw/pull/28))
|
||||||
|
- add libSQL/Turso embedded database backend ([#47](https://github.com/nearai/ironclaw/pull/47))
|
||||||
|
- Move debug log truncation from agent loop to REPL channel ([#65](https://github.com/nearai/ironclaw/pull/65))
|
||||||
|
|
||||||
|
### Fixed
|
||||||
|
|
||||||
|
- shell destructive-command check bypassed by Value::Object arguments ([#72](https://github.com/nearai/ironclaw/pull/72))
|
||||||
|
- propagate real tool_call_id instead of hardcoded placeholder ([#73](https://github.com/nearai/ironclaw/pull/73))
|
||||||
|
- Fix wasm tool schemas and runtime ([#42](https://github.com/nearai/ironclaw/pull/42))
|
||||||
|
- flatten tool messages for NEAR AI cloud-api compatibility ([#41](https://github.com/nearai/ironclaw/pull/41))
|
||||||
|
- security hardening across all layers ([#35](https://github.com/nearai/ironclaw/pull/35))
|
||||||
|
|
||||||
|
### Other
|
||||||
|
|
||||||
|
- Explicitly enable cargo-dist caching for binary artifacts building
|
||||||
|
- Skip building binary artifacts on every PR
|
||||||
|
- add module specification rules to CLAUDE.md
|
||||||
|
- add setup/onboarding specification (src/setup/README.md)
|
||||||
|
- deduplicate tool code and remove dead stubs ([#98](https://github.com/nearai/ironclaw/pull/98))
|
||||||
|
- Reformat architecture diagram in README ([#64](https://github.com/nearai/ironclaw/pull/64))
|
||||||
|
- Add review discipline guidelines to CLAUDE.md ([#68](https://github.com/nearai/ironclaw/pull/68))
|
||||||
|
- Bump MSRV to 1.92, add GCP deployment files ([#40](https://github.com/nearai/ironclaw/pull/40))
|
||||||
|
- Add OpenAI-compatible HTTP API (/v1/chat/completions, /v1/models) ([#31](https://github.com/nearai/ironclaw/pull/31))
|
||||||
|
|
||||||
## [0.1.3](https://github.com/nearai/ironclaw/compare/v0.1.2...v0.1.3) - 2026-02-12
|
## [0.1.3](https://github.com/nearai/ironclaw/compare/v0.1.2...v0.1.3) - 2026-02-12
|
||||||
|
|
||||||
### Other
|
### Other
|
||||||
|
|||||||
@@ -630,6 +630,22 @@ RUST_LOG=ironclaw::agent=debug cargo run
|
|||||||
RUST_LOG=ironclaw=debug,tower_http=debug cargo run
|
RUST_LOG=ironclaw=debug,tower_http=debug cargo run
|
||||||
```
|
```
|
||||||
|
|
||||||
|
## Module Specifications
|
||||||
|
|
||||||
|
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` |
|
||||||
|
|
||||||
## Code Style
|
## Code Style
|
||||||
|
|
||||||
- Use `crate::` imports, not `super::`
|
- Use `crate::` imports, not `super::`
|
||||||
|
|||||||
Generated
+20
-174
@@ -352,23 +352,6 @@ dependencies = [
|
|||||||
"syn 2.0.114",
|
"syn 2.0.114",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
|
||||||
name = "async-tungstenite"
|
|
||||||
version = "0.32.1"
|
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
|
||||||
checksum = "8acc405d38be14342132609f06f02acaf825ddccfe76c4824a69281e0458ebd4"
|
|
||||||
dependencies = [
|
|
||||||
"atomic-waker",
|
|
||||||
"futures-core",
|
|
||||||
"futures-io",
|
|
||||||
"futures-task",
|
|
||||||
"futures-util",
|
|
||||||
"log",
|
|
||||||
"pin-project-lite",
|
|
||||||
"tokio",
|
|
||||||
"tungstenite 0.28.0",
|
|
||||||
]
|
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "atomic-waker"
|
name = "atomic-waker"
|
||||||
version = "1.1.2"
|
version = "1.1.2"
|
||||||
@@ -522,7 +505,7 @@ dependencies = [
|
|||||||
"rustc-hash 1.1.0",
|
"rustc-hash 1.1.0",
|
||||||
"shlex",
|
"shlex",
|
||||||
"syn 2.0.114",
|
"syn 2.0.114",
|
||||||
"which 4.4.2",
|
"which",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -833,72 +816,6 @@ version = "0.2.1"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724"
|
checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724"
|
||||||
|
|
||||||
[[package]]
|
|
||||||
name = "chromiumoxide"
|
|
||||||
version = "0.8.0"
|
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
|
||||||
checksum = "6c18200611490f523adb497ddd4744d6d536e243f6add13e7eeeb1c05904fbb1"
|
|
||||||
dependencies = [
|
|
||||||
"async-tungstenite",
|
|
||||||
"base64 0.22.1",
|
|
||||||
"cfg-if",
|
|
||||||
"chromiumoxide_cdp",
|
|
||||||
"chromiumoxide_types",
|
|
||||||
"dunce",
|
|
||||||
"fnv",
|
|
||||||
"futures",
|
|
||||||
"futures-timer",
|
|
||||||
"pin-project-lite",
|
|
||||||
"reqwest",
|
|
||||||
"serde",
|
|
||||||
"serde_json",
|
|
||||||
"thiserror 1.0.69",
|
|
||||||
"tokio",
|
|
||||||
"tracing",
|
|
||||||
"url",
|
|
||||||
"which 8.0.0",
|
|
||||||
"windows-registry 0.5.3",
|
|
||||||
]
|
|
||||||
|
|
||||||
[[package]]
|
|
||||||
name = "chromiumoxide_cdp"
|
|
||||||
version = "0.8.0"
|
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
|
||||||
checksum = "b8f78027ced540595dcbaf9e2f3413cbe3708b839ff239d2858acaea73915dcb"
|
|
||||||
dependencies = [
|
|
||||||
"chromiumoxide_pdl",
|
|
||||||
"chromiumoxide_types",
|
|
||||||
"serde",
|
|
||||||
"serde_json",
|
|
||||||
]
|
|
||||||
|
|
||||||
[[package]]
|
|
||||||
name = "chromiumoxide_pdl"
|
|
||||||
version = "0.8.0"
|
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
|
||||||
checksum = "0d2c7b7c6b41a0de36d00a284e619017e0f4aec5c9bc8d90614b9e1687984f20"
|
|
||||||
dependencies = [
|
|
||||||
"chromiumoxide_types",
|
|
||||||
"either",
|
|
||||||
"heck 0.4.1",
|
|
||||||
"once_cell",
|
|
||||||
"proc-macro2",
|
|
||||||
"quote",
|
|
||||||
"regex",
|
|
||||||
"serde",
|
|
||||||
"serde_json",
|
|
||||||
]
|
|
||||||
|
|
||||||
[[package]]
|
|
||||||
name = "chromiumoxide_types"
|
|
||||||
version = "0.8.0"
|
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
|
||||||
checksum = "309ba8f378bbc093c93f06beb7bd4c5ceffdf14107ad99cacbbf063709926795"
|
|
||||||
dependencies = [
|
|
||||||
"serde",
|
|
||||||
"serde_json",
|
|
||||||
]
|
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "chrono"
|
name = "chrono"
|
||||||
version = "0.4.43"
|
version = "0.4.43"
|
||||||
@@ -910,7 +827,7 @@ dependencies = [
|
|||||||
"num-traits",
|
"num-traits",
|
||||||
"serde",
|
"serde",
|
||||||
"wasm-bindgen",
|
"wasm-bindgen",
|
||||||
"windows-link 0.2.1",
|
"windows-link",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -962,7 +879,7 @@ version = "4.5.55"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "a92793da1a46a5f2a02a6f4c46c6496b28c43638adea8306fcb0caa1634f24e5"
|
checksum = "a92793da1a46a5f2a02a6f4c46c6496b28c43638adea8306fcb0caa1634f24e5"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"heck 0.5.0",
|
"heck",
|
||||||
"proc-macro2",
|
"proc-macro2",
|
||||||
"quote",
|
"quote",
|
||||||
"syn 2.0.114",
|
"syn 2.0.114",
|
||||||
@@ -1580,12 +1497,6 @@ version = "0.15.7"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "1aaf95b3e5c8f23aa320147307562d361db0ae0d51242340f558153b4eb2439b"
|
checksum = "1aaf95b3e5c8f23aa320147307562d361db0ae0d51242340f558153b4eb2439b"
|
||||||
|
|
||||||
[[package]]
|
|
||||||
name = "dunce"
|
|
||||||
version = "1.0.5"
|
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
|
||||||
checksum = "92773504d58c093f6de2459af4af33faa518c13451eb8f2b5698ed3d36e7c813"
|
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "dyn-clone"
|
name = "dyn-clone"
|
||||||
version = "1.0.20"
|
version = "1.0.20"
|
||||||
@@ -1652,12 +1563,6 @@ dependencies = [
|
|||||||
"syn 2.0.114",
|
"syn 2.0.114",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
|
||||||
name = "env_home"
|
|
||||||
version = "0.1.0"
|
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
|
||||||
checksum = "c7f84e12ccf0a7ddc17a6c41c93326024c42920d7ee630d04950e6926645c0fe"
|
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "equivalent"
|
name = "equivalent"
|
||||||
version = "1.0.2"
|
version = "1.0.2"
|
||||||
@@ -2115,12 +2020,6 @@ dependencies = [
|
|||||||
"hashbrown 0.14.5",
|
"hashbrown 0.14.5",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
|
||||||
name = "heck"
|
|
||||||
version = "0.4.1"
|
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
|
||||||
checksum = "95505c38b4572b2d910cecb0281560f54b440a19336cbbcb27bf6ce6adc6f5a8"
|
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "heck"
|
name = "heck"
|
||||||
version = "0.5.0"
|
version = "0.5.0"
|
||||||
@@ -2368,7 +2267,7 @@ dependencies = [
|
|||||||
"tokio",
|
"tokio",
|
||||||
"tower-service",
|
"tower-service",
|
||||||
"tracing",
|
"tracing",
|
||||||
"windows-registry 0.6.1",
|
"windows-registry",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -2591,7 +2490,7 @@ dependencies = [
|
|||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "ironclaw"
|
name = "ironclaw"
|
||||||
version = "0.1.3"
|
version = "0.5.0"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"aes-gcm",
|
"aes-gcm",
|
||||||
"aho-corasick",
|
"aho-corasick",
|
||||||
@@ -2602,7 +2501,6 @@ dependencies = [
|
|||||||
"blake3",
|
"blake3",
|
||||||
"bollard",
|
"bollard",
|
||||||
"bytes",
|
"bytes",
|
||||||
"chromiumoxide",
|
|
||||||
"chrono",
|
"chrono",
|
||||||
"clap",
|
"clap",
|
||||||
"cron",
|
"cron",
|
||||||
@@ -2799,7 +2697,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||||||
checksum = "d7c4b02199fee7c5d21a5ae7d8cfa79a6ef5bb2fc834d6e9058e89c825efdc55"
|
checksum = "d7c4b02199fee7c5d21a5ae7d8cfa79a6ef5bb2fc834d6e9058e89c825efdc55"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"cfg-if",
|
"cfg-if",
|
||||||
"windows-link 0.2.1",
|
"windows-link",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -3411,7 +3309,7 @@ dependencies = [
|
|||||||
"libc",
|
"libc",
|
||||||
"redox_syscall 0.5.18",
|
"redox_syscall 0.5.18",
|
||||||
"smallvec",
|
"smallvec",
|
||||||
"windows-link 0.2.1",
|
"windows-link",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -4048,7 +3946,7 @@ version = "0.8.16"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "72c225407d8e52ef8cf094393781ecda9a99d6544ec28d90a6915751de259264"
|
checksum = "72c225407d8e52ef8cf094393781ecda9a99d6544ec28d90a6915751de259264"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"heck 0.5.0",
|
"heck",
|
||||||
"proc-macro2",
|
"proc-macro2",
|
||||||
"quote",
|
"quote",
|
||||||
"refinery-core",
|
"refinery-core",
|
||||||
@@ -6284,7 +6182,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||||||
checksum = "5f38f7a5eb2f06f53fe943e7fb8bf4197f7cf279f1bc52c0ce56e9d3ffd750a4"
|
checksum = "5f38f7a5eb2f06f53fe943e7fb8bf4197f7cf279f1bc52c0ce56e9d3ffd750a4"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"anyhow",
|
"anyhow",
|
||||||
"heck 0.5.0",
|
"heck",
|
||||||
"indexmap 2.13.0",
|
"indexmap 2.13.0",
|
||||||
"wit-parser",
|
"wit-parser",
|
||||||
]
|
]
|
||||||
@@ -6352,17 +6250,6 @@ dependencies = [
|
|||||||
"rustix 0.38.44",
|
"rustix 0.38.44",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
|
||||||
name = "which"
|
|
||||||
version = "8.0.0"
|
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
|
||||||
checksum = "d3fabb953106c3c8eea8306e4393700d7657561cb43122571b172bbfb7c7ba1d"
|
|
||||||
dependencies = [
|
|
||||||
"env_home",
|
|
||||||
"rustix 1.1.3",
|
|
||||||
"winsafe",
|
|
||||||
]
|
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "whoami"
|
name = "whoami"
|
||||||
version = "2.1.0"
|
version = "2.1.0"
|
||||||
@@ -6396,7 +6283,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||||||
checksum = "8738c5a7ef3a9de0fae10f8b84091a2aa4e059d8fef23de202ab689812b6bc6e"
|
checksum = "8738c5a7ef3a9de0fae10f8b84091a2aa4e059d8fef23de202ab689812b6bc6e"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"anyhow",
|
"anyhow",
|
||||||
"heck 0.5.0",
|
"heck",
|
||||||
"proc-macro2",
|
"proc-macro2",
|
||||||
"quote",
|
"quote",
|
||||||
"shellexpand",
|
"shellexpand",
|
||||||
@@ -6472,9 +6359,9 @@ checksum = "b8e83a14d34d0623b51dce9581199302a221863196a1dde71a7663a4c2be9deb"
|
|||||||
dependencies = [
|
dependencies = [
|
||||||
"windows-implement",
|
"windows-implement",
|
||||||
"windows-interface",
|
"windows-interface",
|
||||||
"windows-link 0.2.1",
|
"windows-link",
|
||||||
"windows-result 0.4.1",
|
"windows-result",
|
||||||
"windows-strings 0.5.1",
|
"windows-strings",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -6499,47 +6386,21 @@ dependencies = [
|
|||||||
"syn 2.0.114",
|
"syn 2.0.114",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
|
||||||
name = "windows-link"
|
|
||||||
version = "0.1.3"
|
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
|
||||||
checksum = "5e6ad25900d524eaabdbbb96d20b4311e1e7ae1699af4fb28c17ae66c80d798a"
|
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "windows-link"
|
name = "windows-link"
|
||||||
version = "0.2.1"
|
version = "0.2.1"
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5"
|
checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5"
|
||||||
|
|
||||||
[[package]]
|
|
||||||
name = "windows-registry"
|
|
||||||
version = "0.5.3"
|
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
|
||||||
checksum = "5b8a9ed28765efc97bbc954883f4e6796c33a06546ebafacbabee9696967499e"
|
|
||||||
dependencies = [
|
|
||||||
"windows-link 0.1.3",
|
|
||||||
"windows-result 0.3.4",
|
|
||||||
"windows-strings 0.4.2",
|
|
||||||
]
|
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "windows-registry"
|
name = "windows-registry"
|
||||||
version = "0.6.1"
|
version = "0.6.1"
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "02752bf7fbdcce7f2a27a742f798510f3e5ad88dbe84871e5168e2120c3d5720"
|
checksum = "02752bf7fbdcce7f2a27a742f798510f3e5ad88dbe84871e5168e2120c3d5720"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"windows-link 0.2.1",
|
"windows-link",
|
||||||
"windows-result 0.4.1",
|
"windows-result",
|
||||||
"windows-strings 0.5.1",
|
"windows-strings",
|
||||||
]
|
|
||||||
|
|
||||||
[[package]]
|
|
||||||
name = "windows-result"
|
|
||||||
version = "0.3.4"
|
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
|
||||||
checksum = "56f42bd332cc6c8eac5af113fc0c1fd6a8fd2aa08a0119358686e5160d0586c6"
|
|
||||||
dependencies = [
|
|
||||||
"windows-link 0.1.3",
|
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -6548,16 +6409,7 @@ version = "0.4.1"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "7781fa89eaf60850ac3d2da7af8e5242a5ea78d1a11c49bf2910bb5a73853eb5"
|
checksum = "7781fa89eaf60850ac3d2da7af8e5242a5ea78d1a11c49bf2910bb5a73853eb5"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"windows-link 0.2.1",
|
"windows-link",
|
||||||
]
|
|
||||||
|
|
||||||
[[package]]
|
|
||||||
name = "windows-strings"
|
|
||||||
version = "0.4.2"
|
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
|
||||||
checksum = "56e6c93f3a0c3b36176cb1327a4958a0353d5d166c2a35cb268ace15e91d3b57"
|
|
||||||
dependencies = [
|
|
||||||
"windows-link 0.1.3",
|
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -6566,7 +6418,7 @@ version = "0.5.1"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "7837d08f69c77cf6b07689544538e017c1bfcf57e34b4c0ff58e6c2cd3b37091"
|
checksum = "7837d08f69c77cf6b07689544538e017c1bfcf57e34b4c0ff58e6c2cd3b37091"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"windows-link 0.2.1",
|
"windows-link",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -6611,7 +6463,7 @@ version = "0.61.2"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "ae137229bcbd6cdf0f7b80a31df61766145077ddf49416a728b02cb3921ff3fc"
|
checksum = "ae137229bcbd6cdf0f7b80a31df61766145077ddf49416a728b02cb3921ff3fc"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"windows-link 0.2.1",
|
"windows-link",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -6651,7 +6503,7 @@ version = "0.53.5"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "4945f9f551b88e0d65f3db0bc25c33b8acea4d9e41163edf90dcd0b19f9069f3"
|
checksum = "4945f9f551b88e0d65f3db0bc25c33b8acea4d9e41163edf90dcd0b19f9069f3"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"windows-link 0.2.1",
|
"windows-link",
|
||||||
"windows_aarch64_gnullvm 0.53.1",
|
"windows_aarch64_gnullvm 0.53.1",
|
||||||
"windows_aarch64_msvc 0.53.1",
|
"windows_aarch64_msvc 0.53.1",
|
||||||
"windows_i686_gnu 0.53.1",
|
"windows_i686_gnu 0.53.1",
|
||||||
@@ -6809,12 +6661,6 @@ dependencies = [
|
|||||||
"memchr",
|
"memchr",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
|
||||||
name = "winsafe"
|
|
||||||
version = "0.0.19"
|
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
|
||||||
checksum = "d135d17ab770252ad95e9a872d365cf3090e3be864a34ab46f48555993efc904"
|
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "winx"
|
name = "winx"
|
||||||
version = "0.36.4"
|
version = "0.36.4"
|
||||||
|
|||||||
+13
-6
@@ -1,6 +1,14 @@
|
|||||||
|
[workspace]
|
||||||
|
exclude = [
|
||||||
|
"channels-src/telegram",
|
||||||
|
"channels-src/slack",
|
||||||
|
"channels-src/whatsapp",
|
||||||
|
"tools-src/gmail",
|
||||||
|
]
|
||||||
|
|
||||||
[package]
|
[package]
|
||||||
name = "ironclaw"
|
name = "ironclaw"
|
||||||
version = "0.1.3"
|
version = "0.5.0"
|
||||||
edition = "2024"
|
edition = "2024"
|
||||||
rust-version = "1.92"
|
rust-version = "1.92"
|
||||||
description = "Secure personal AI assistant that protects your data and expands its capabilities on the fly"
|
description = "Secure personal AI assistant that protects your data and expands its capabilities on the fly"
|
||||||
@@ -122,9 +130,6 @@ bytes = "1"
|
|||||||
base64 = "0.22.1"
|
base64 = "0.22.1"
|
||||||
mime_guess = "2.0.5"
|
mime_guess = "2.0.5"
|
||||||
|
|
||||||
# Headless browser automation via Chrome DevTools Protocol
|
|
||||||
chromiumoxide = { version = "0.8", default-features = false, features = ["tokio-runtime"] }
|
|
||||||
|
|
||||||
# macOS keychain
|
# macOS keychain
|
||||||
[target.'cfg(target_os = "macos")'.dependencies]
|
[target.'cfg(target_os = "macos")'.dependencies]
|
||||||
security-framework = "3"
|
security-framework = "3"
|
||||||
@@ -142,7 +147,7 @@ pretty_assertions = "1"
|
|||||||
tempfile = "3"
|
tempfile = "3"
|
||||||
|
|
||||||
[features]
|
[features]
|
||||||
default = ["postgres"]
|
default = ["postgres", "libsql"]
|
||||||
postgres = [
|
postgres = [
|
||||||
"dep:deadpool-postgres",
|
"dep:deadpool-postgres",
|
||||||
"dep:tokio-postgres",
|
"dep:tokio-postgres",
|
||||||
@@ -186,11 +191,13 @@ windows-archive = ".tar.gz"
|
|||||||
# The archive format to use for non-windows builds (defaults .tar.xz)
|
# The archive format to use for non-windows builds (defaults .tar.xz)
|
||||||
unix-archive = ".tar.gz"
|
unix-archive = ".tar.gz"
|
||||||
# Which actions to run on pull requests
|
# Which actions to run on pull requests
|
||||||
pr-run-mode = "upload"
|
pr-run-mode = "skip"
|
||||||
# Path that installers should place binaries in
|
# Path that installers should place binaries in
|
||||||
install-path = "CARGO_HOME"
|
install-path = "CARGO_HOME"
|
||||||
# Whether to install an updater program
|
# Whether to install an updater program
|
||||||
install-updater = true
|
install-updater = true
|
||||||
|
# Cache intermediate build artifacts to speed up the release pipelines
|
||||||
|
cache-builds = true
|
||||||
|
|
||||||
[workspace.metadata.dist.github-custom-runners]
|
[workspace.metadata.dist.github-custom-runners]
|
||||||
aarch64-unknown-linux-gnu = "ubuntu-24.04-arm"
|
aarch64-unknown-linux-gnu = "ubuntu-24.04-arm"
|
||||||
|
|||||||
+11
-15
@@ -112,7 +112,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
|||||||
| `pairing` | ✅ | ✅ | - | list/approve for channel DM pairing |
|
| `pairing` | ✅ | ✅ | - | list/approve for channel DM pairing |
|
||||||
| `nodes` | ✅ | ❌ | P3 | Device management |
|
| `nodes` | ✅ | ❌ | P3 | Device management |
|
||||||
| `plugins` | ✅ | ❌ | P3 | Plugin management |
|
| `plugins` | ✅ | ❌ | P3 | Plugin management |
|
||||||
| `hooks` | ✅ | ❌ | P2 | Lifecycle hooks |
|
| `hooks` | ✅ | ✅ | P2 | Lifecycle hooks |
|
||||||
| `cron` | ✅ | ❌ | P2 | Scheduled jobs |
|
| `cron` | ✅ | ❌ | P2 | Scheduled jobs |
|
||||||
| `webhooks` | ✅ | ❌ | P3 | Webhook config |
|
| `webhooks` | ✅ | ❌ | P3 | Webhook config |
|
||||||
| `message send` | ✅ | ❌ | P2 | Send to channels |
|
| `message send` | ✅ | ❌ | P2 | Send to channels |
|
||||||
@@ -164,7 +164,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
|||||||
| AWS Bedrock | ✅ | ❌ | P3 | |
|
| AWS Bedrock | ✅ | ❌ | P3 | |
|
||||||
| Google Gemini | ✅ | ❌ | P3 | |
|
| Google Gemini | ✅ | ❌ | P3 | |
|
||||||
| OpenRouter | ✅ | ❌ | P3 | |
|
| OpenRouter | ✅ | ❌ | P3 | |
|
||||||
| Ollama (local) | ✅ | ❌ | P2 | Local models |
|
| Ollama (local) | ✅ | ✅ | - | via `rig::providers::ollama` (full support) |
|
||||||
| node-llama-cpp | ✅ | ➖ | - | N/A for Rust |
|
| node-llama-cpp | ✅ | ➖ | - | N/A for Rust |
|
||||||
| llama.cpp (native) | ❌ | 🔮 | P3 | Rust bindings |
|
| llama.cpp (native) | ❌ | 🔮 | P3 | Rust bindings |
|
||||||
|
|
||||||
@@ -174,7 +174,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
|||||||
|---------|----------|----------|-------|
|
|---------|----------|----------|-------|
|
||||||
| Auto-discovery | ✅ | ❌ | |
|
| Auto-discovery | ✅ | ❌ | |
|
||||||
| Failover chains | ✅ | ✅ | `FailoverProvider` with configurable `fallback_model` |
|
| Failover chains | ✅ | ✅ | `FailoverProvider` with configurable `fallback_model` |
|
||||||
| Cooldown management | ✅ | ❌ | Skip failed providers |
|
| Cooldown management | ✅ | ✅ | Lock-free per-provider cooldown in `FailoverProvider` |
|
||||||
| Per-session model override | ✅ | ✅ | Model selector in TUI |
|
| Per-session model override | ✅ | ✅ | Model selector in TUI |
|
||||||
| Model selection UI | ✅ | ✅ | TUI keyboard shortcut |
|
| Model selection UI | ✅ | ✅ | TUI keyboard shortcut |
|
||||||
|
|
||||||
@@ -323,14 +323,14 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
|||||||
| Cron jobs | ✅ | ✅ | - | Routines with cron trigger |
|
| Cron jobs | ✅ | ✅ | - | Routines with cron trigger |
|
||||||
| Timezone support | ✅ | ✅ | - | Via cron expressions |
|
| Timezone support | ✅ | ✅ | - | Via cron expressions |
|
||||||
| One-shot/recurring jobs | ✅ | ✅ | - | Manual + cron triggers |
|
| One-shot/recurring jobs | ✅ | ✅ | - | Manual + cron triggers |
|
||||||
| `beforeInbound` hook | ✅ | ❌ | P2 | |
|
| `beforeInbound` hook | ✅ | ✅ | P2 | |
|
||||||
| `beforeOutbound` hook | ✅ | ❌ | P2 | |
|
| `beforeOutbound` hook | ✅ | ✅ | P2 | |
|
||||||
| `beforeToolCall` hook | ✅ | ❌ | P2 | |
|
| `beforeToolCall` hook | ✅ | ✅ | P2 | |
|
||||||
| `onMessage` hook | ✅ | ✅ | - | Routines with event trigger |
|
| `onMessage` hook | ✅ | ✅ | - | Routines with event trigger |
|
||||||
| `onSessionStart` hook | ✅ | ❌ | P2 | |
|
| `onSessionStart` hook | ✅ | ✅ | P2 | |
|
||||||
| `onSessionEnd` hook | ✅ | ❌ | P2 | |
|
| `onSessionEnd` hook | ✅ | ✅ | P2 | |
|
||||||
| `transcribeAudio` hook | ✅ | ❌ | P3 | |
|
| `transcribeAudio` hook | ✅ | ❌ | P3 | |
|
||||||
| `transformResponse` hook | ✅ | ❌ | P2 | |
|
| `transformResponse` hook | ✅ | ✅ | P2 | |
|
||||||
| Bundled hooks | ✅ | ❌ | P2 | |
|
| Bundled hooks | ✅ | ❌ | P2 | |
|
||||||
| Plugin hooks | ✅ | ❌ | P3 | |
|
| Plugin hooks | ✅ | ❌ | P3 | |
|
||||||
| Workspace hooks | ✅ | ❌ | P2 | Inline code |
|
| Workspace hooks | ✅ | ❌ | P2 | Inline code |
|
||||||
@@ -420,14 +420,10 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
|||||||
- ✅ Telegram channel (WASM, DM pairing, caption, /start)
|
- ✅ Telegram channel (WASM, DM pairing, caption, /start)
|
||||||
- ❌ WhatsApp channel
|
- ❌ WhatsApp channel
|
||||||
- ✅ Multi-provider failover (`FailoverProvider` with retryable error classification)
|
- ✅ Multi-provider failover (`FailoverProvider` with retryable error classification)
|
||||||
- ❌ Hooks system (beforeInbound, beforeToolCall, etc.)
|
- ✅ Hooks system (beforeInbound, beforeToolCall, beforeOutbound, onSessionStart, onSessionEnd, transformResponse)
|
||||||
|
|
||||||
### P2 - Medium Priority
|
### P2 - Medium Priority
|
||||||
- ❌ Cron job scheduling
|
- ❌ Media handling (images, PDFs)
|
||||||
- ❌ Web Control UI
|
|
||||||
- ❌ WebChat channel
|
|
||||||
- 🚧 Media handling (caption support; no image/PDF processing)
|
|
||||||
- ❌ CLI subcommands (config, status, memory, doctor)
|
|
||||||
- ❌ Ollama/local model support
|
- ❌ Ollama/local model support
|
||||||
- ❌ Configuration hot-reload
|
- ❌ Configuration hot-reload
|
||||||
- ❌ Webhook trigger endpoint in web gateway
|
- ❌ Webhook trigger endpoint in web gateway
|
||||||
|
|||||||
@@ -0,0 +1,23 @@
|
|||||||
|
[package]
|
||||||
|
name = "discord-channel"
|
||||||
|
version = "0.1.0"
|
||||||
|
edition = "2021"
|
||||||
|
description = "Discord channel for IronClaw"
|
||||||
|
license = "MIT OR Apache-2.0"
|
||||||
|
publish = false
|
||||||
|
|
||||||
|
[dependencies]
|
||||||
|
serde = { version = "1.0", features = ["derive"] }
|
||||||
|
serde_json = "1.0"
|
||||||
|
wit-bindgen = "0.41.0"
|
||||||
|
|
||||||
|
[lib]
|
||||||
|
crate-type = ["cdylib"]
|
||||||
|
|
||||||
|
[profile.release]
|
||||||
|
strip = true
|
||||||
|
opt-level = "s"
|
||||||
|
lto = true
|
||||||
|
codegen-units = 1
|
||||||
|
|
||||||
|
|
||||||
@@ -0,0 +1,121 @@
|
|||||||
|
# Discord Channel for IronClaw
|
||||||
|
|
||||||
|
WASM channel for Discord integration - handle slash commands and button interactions via webhooks.
|
||||||
|
|
||||||
|
## Features
|
||||||
|
|
||||||
|
- **Slash Commands** - Process Discord slash commands
|
||||||
|
- **Button Interactions** - Handle button clicks
|
||||||
|
- **Thread Support** - Respond in threads
|
||||||
|
- **DM Support** - Handle direct messages
|
||||||
|
|
||||||
|
## Setup
|
||||||
|
|
||||||
|
1. Create a Discord Application at <https://discord.com/developers/applications>
|
||||||
|
2. Create a Bot and get the token
|
||||||
|
3. Set up Interactions URL to point to your IronClaw instance
|
||||||
|
4. Copy the Application ID and Public Key
|
||||||
|
5. Store in IronClaw secrets:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
ironclaw secret set discord_bot_token YOUR_BOT_TOKEN
|
||||||
|
```
|
||||||
|
|
||||||
|
**Note:** The `discord_bot_token` secret is the only value read directly by this
|
||||||
|
Discord channel WASM component. The `discord_app_id` and `discord_public_key`
|
||||||
|
secrets are used by the IronClaw host (for example, to verify Discord
|
||||||
|
interaction signatures and manage slash command registration) and are not
|
||||||
|
accessed from the WASM module itself.
|
||||||
|
|
||||||
|
## Discord Configuration
|
||||||
|
|
||||||
|
### Register Slash Commands
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl -X POST \
|
||||||
|
-H "Authorization: Bot YOUR_BOT_TOKEN" \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
https://discord.com/api/v10/applications/YOUR_APP_ID/commands \
|
||||||
|
-d '{
|
||||||
|
"name": "ask",
|
||||||
|
"description": "Ask the AI agent",
|
||||||
|
"options": [{
|
||||||
|
"name": "question",
|
||||||
|
"description": "Your question",
|
||||||
|
"type": 3,
|
||||||
|
"required": true
|
||||||
|
}]
|
||||||
|
}'
|
||||||
|
```
|
||||||
|
|
||||||
|
### Set Interactions Endpoint
|
||||||
|
|
||||||
|
In your Discord app settings, set:
|
||||||
|
|
||||||
|
- Interactions Endpoint URL: `https://your-ironclaw.com/webhook/discord`
|
||||||
|
|
||||||
|
## Usage Examples
|
||||||
|
|
||||||
|
### Slash Command
|
||||||
|
|
||||||
|
User types: `/ask question: What is the weather?`
|
||||||
|
|
||||||
|
The agent receives:
|
||||||
|
|
||||||
|
```text
|
||||||
|
User: @username
|
||||||
|
Content: /ask question: What is the weather?
|
||||||
|
```
|
||||||
|
|
||||||
|
### Button Click
|
||||||
|
|
||||||
|
When a user clicks a button in a message, the agent receives:
|
||||||
|
|
||||||
|
```text
|
||||||
|
User: @username
|
||||||
|
Content: [Button clicked] Original message content
|
||||||
|
```
|
||||||
|
|
||||||
|
## Error Handling
|
||||||
|
|
||||||
|
If an internal error occurs (e.g., metadata serialization failure), the tool attempts to send an ephemeral message to the user:
|
||||||
|
|
||||||
|
```text
|
||||||
|
❌ Internal Error: Failed to process command metadata.
|
||||||
|
```
|
||||||
|
|
||||||
|
Check the host logs for detailed error information.
|
||||||
|
|
||||||
|
## Advanced Usage
|
||||||
|
|
||||||
|
### Embeds
|
||||||
|
|
||||||
|
To send embeds, include an `embeds` array in the `metadata_json` field of the agent's response. The structure should match the Discord API `embed` object.
|
||||||
|
|
||||||
|
## Troubleshooting
|
||||||
|
|
||||||
|
### "Invalid Signature"
|
||||||
|
|
||||||
|
- Check that `discord_public_key` is set correctly in IronClaw secrets.
|
||||||
|
- This validation happens on the host before reaching the WASM.
|
||||||
|
|
||||||
|
### "401 Unauthorized"
|
||||||
|
|
||||||
|
- Check that `discord_bot_token` is set correctly in IronClaw secrets.
|
||||||
|
- Ensure the bot is added to the server.
|
||||||
|
|
||||||
|
### "Interaction Failed"
|
||||||
|
|
||||||
|
- The interaction might have timed out (Discord requires a response within 3 seconds).
|
||||||
|
- The `interactions_endpoint_url` might be unreachable.
|
||||||
|
|
||||||
|
## Building
|
||||||
|
|
||||||
|
```bash
|
||||||
|
cd channels-src/discord
|
||||||
|
cargo build --target wasm32-wasi --release
|
||||||
|
```
|
||||||
|
|
||||||
|
## License
|
||||||
|
|
||||||
|
MIT/Apache-2.0
|
||||||
@@ -0,0 +1,39 @@
|
|||||||
|
{
|
||||||
|
"type": "channel",
|
||||||
|
"name": "discord",
|
||||||
|
"description": "Discord Gateway/Webhook channel for handling slash commands, buttons, and messages",
|
||||||
|
"capabilities": {
|
||||||
|
"http": {
|
||||||
|
"allowlist": [
|
||||||
|
{ "host": "discord.com", "path_prefix": "/api/v10" }
|
||||||
|
],
|
||||||
|
"credentials": {
|
||||||
|
"discord_bot_token": {
|
||||||
|
"secret_name": "discord_bot_token",
|
||||||
|
"location": { "type": "header", "header_name": "Authorization", "prefix": "Bot " },
|
||||||
|
"host_patterns": ["discord.com"]
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"rate_limit": {
|
||||||
|
"requests_per_minute": 60,
|
||||||
|
"requests_per_hour": 3600
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"secrets": {
|
||||||
|
"allowed_names": ["discord_bot_token", "discord_*"]
|
||||||
|
},
|
||||||
|
"channel": {
|
||||||
|
"allowed_paths": ["/webhook/discord"],
|
||||||
|
"allow_polling": false,
|
||||||
|
"callback_timeout_secs": 45,
|
||||||
|
"workspace_prefix": "channels/discord/",
|
||||||
|
"emit_rate_limit": {
|
||||||
|
"messages_per_minute": 100,
|
||||||
|
"messages_per_hour": 5000
|
||||||
|
}
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"config": {
|
||||||
|
"require_signature_verification": true
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,476 @@
|
|||||||
|
//! Discord Gateway/Webhook channel for IronClaw.
|
||||||
|
//!
|
||||||
|
//! This WASM component implements the channel interface for handling Discord
|
||||||
|
//! interactions via webhooks and sending messages back to Discord.
|
||||||
|
//!
|
||||||
|
//! # Features
|
||||||
|
//!
|
||||||
|
//! - URL verification for Discord interactions
|
||||||
|
//! - Slash command handling
|
||||||
|
//! - Message event parsing (@mentions, DMs)
|
||||||
|
//! - Thread support for conversations
|
||||||
|
//! - Response posting via Discord Web API
|
||||||
|
//! - Automatic message truncation (> 2000 chars)
|
||||||
|
//!
|
||||||
|
//! # Security
|
||||||
|
//!
|
||||||
|
//! - Signature validation is handled by the host (webhook secrets)
|
||||||
|
//! - Bot token is injected by host during HTTP requests
|
||||||
|
//! - WASM never sees raw credentials
|
||||||
|
|
||||||
|
wit_bindgen::generate!({
|
||||||
|
world: "sandboxed-channel",
|
||||||
|
path: "../../wit/channel.wit",
|
||||||
|
});
|
||||||
|
|
||||||
|
use serde::{Deserialize, Serialize};
|
||||||
|
|
||||||
|
use exports::near::agent::channel::{
|
||||||
|
AgentResponse, ChannelConfig, Guest, HttpEndpointConfig, IncomingHttpRequest,
|
||||||
|
OutgoingHttpResponse, StatusUpdate,
|
||||||
|
};
|
||||||
|
use near::agent::channel_host::{self, EmittedMessage};
|
||||||
|
|
||||||
|
/// Discord interaction wrapper.
|
||||||
|
#[derive(Debug, Deserialize)]
|
||||||
|
struct DiscordInteraction {
|
||||||
|
/// Interaction type (1=Ping, 2=ApplicationCommand, 3=MessageComponent)
|
||||||
|
#[serde(rename = "type")]
|
||||||
|
interaction_type: u8,
|
||||||
|
|
||||||
|
/// Interaction ID
|
||||||
|
id: String,
|
||||||
|
|
||||||
|
/// Application ID
|
||||||
|
application_id: String,
|
||||||
|
|
||||||
|
/// Guild ID (if in server)
|
||||||
|
#[allow(dead_code)] // Part of API payload, currently unused
|
||||||
|
guild_id: Option<String>,
|
||||||
|
|
||||||
|
/// Channel ID
|
||||||
|
channel_id: Option<String>,
|
||||||
|
|
||||||
|
/// Member info (if in server)
|
||||||
|
member: Option<DiscordMember>,
|
||||||
|
|
||||||
|
/// User info (if DM)
|
||||||
|
user: Option<DiscordUser>,
|
||||||
|
|
||||||
|
/// Command data (for slash commands)
|
||||||
|
data: Option<DiscordCommandData>,
|
||||||
|
|
||||||
|
/// Message (for component interactions)
|
||||||
|
message: Option<DiscordMessage>,
|
||||||
|
|
||||||
|
/// Token for responding
|
||||||
|
token: String,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Deserialize, Clone)]
|
||||||
|
struct DiscordMember {
|
||||||
|
user: DiscordUser,
|
||||||
|
#[allow(dead_code)] // Part of API payload, currently unused
|
||||||
|
nick: Option<String>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Deserialize, Clone)]
|
||||||
|
struct DiscordUser {
|
||||||
|
id: String,
|
||||||
|
username: String,
|
||||||
|
global_name: Option<String>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Deserialize, Clone)]
|
||||||
|
struct DiscordCommandData {
|
||||||
|
#[allow(dead_code)] // Part of API payload, currently unused
|
||||||
|
id: String,
|
||||||
|
name: String,
|
||||||
|
options: Option<Vec<DiscordCommandOption>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Deserialize, Clone)]
|
||||||
|
struct DiscordCommandOption {
|
||||||
|
name: String,
|
||||||
|
value: serde_json::Value,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Deserialize, Clone)]
|
||||||
|
struct DiscordMessage {
|
||||||
|
#[allow(dead_code)] // Part of API payload, currently unused
|
||||||
|
id: String,
|
||||||
|
content: String,
|
||||||
|
channel_id: String,
|
||||||
|
#[allow(dead_code)] // Part of API payload, currently unused
|
||||||
|
author: DiscordUser,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Metadata stored with emitted messages for response routing.
|
||||||
|
#[derive(Debug, Serialize, Deserialize)]
|
||||||
|
struct DiscordMessageMetadata {
|
||||||
|
/// Discord channel ID
|
||||||
|
channel_id: String,
|
||||||
|
|
||||||
|
/// Interaction ID for followups
|
||||||
|
interaction_id: String,
|
||||||
|
|
||||||
|
/// Interaction token for responding
|
||||||
|
token: String,
|
||||||
|
|
||||||
|
/// Application ID
|
||||||
|
application_id: String,
|
||||||
|
|
||||||
|
/// Thread ID (for forum threads)
|
||||||
|
thread_id: Option<String>,
|
||||||
|
}
|
||||||
|
|
||||||
|
struct DiscordChannel;
|
||||||
|
|
||||||
|
impl Guest for DiscordChannel {
|
||||||
|
fn on_start(_config_json: String) -> Result<ChannelConfig, String> {
|
||||||
|
channel_host::log(channel_host::LogLevel::Info, "Discord channel starting");
|
||||||
|
|
||||||
|
Ok(ChannelConfig {
|
||||||
|
display_name: "Discord".to_string(),
|
||||||
|
http_endpoints: vec![HttpEndpointConfig {
|
||||||
|
path: "/webhook/discord".to_string(),
|
||||||
|
methods: vec!["POST".to_string()],
|
||||||
|
require_secret: true,
|
||||||
|
}],
|
||||||
|
poll: None,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
fn on_http_request(req: IncomingHttpRequest) -> OutgoingHttpResponse {
|
||||||
|
let body_str = match std::str::from_utf8(&req.body) {
|
||||||
|
Ok(s) => s,
|
||||||
|
Err(_) => {
|
||||||
|
return json_response(400, serde_json::json!({"error": "Invalid UTF-8 body"}));
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
let interaction: DiscordInteraction = match serde_json::from_str(body_str) {
|
||||||
|
Ok(i) => i,
|
||||||
|
Err(e) => {
|
||||||
|
channel_host::log(
|
||||||
|
channel_host::LogLevel::Error,
|
||||||
|
&format!("Failed to parse Discord interaction: {}", e),
|
||||||
|
);
|
||||||
|
return json_response(400, serde_json::json!({"error": "Invalid interaction"}));
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
match interaction.interaction_type {
|
||||||
|
// Ping - Discord verification
|
||||||
|
1 => {
|
||||||
|
channel_host::log(channel_host::LogLevel::Info, "Responding to Discord ping");
|
||||||
|
json_response(200, serde_json::json!({"type": 1}))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Application Command (slash command)
|
||||||
|
2 => {
|
||||||
|
handle_slash_command(&interaction);
|
||||||
|
json_response(
|
||||||
|
200,
|
||||||
|
serde_json::json!({
|
||||||
|
"type": 5,
|
||||||
|
"data": {
|
||||||
|
"content": "🤔 Thinking..."
|
||||||
|
}
|
||||||
|
}),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Message Component (buttons, selects)
|
||||||
|
3 => {
|
||||||
|
if let Some(ref message) = interaction.message {
|
||||||
|
handle_message_component(&interaction, message);
|
||||||
|
}
|
||||||
|
json_response(200, serde_json::json!({"type": 6}))
|
||||||
|
}
|
||||||
|
|
||||||
|
_ => {
|
||||||
|
channel_host::log(
|
||||||
|
channel_host::LogLevel::Warn,
|
||||||
|
&format!(
|
||||||
|
"Unknown Discord interaction type: {}",
|
||||||
|
interaction.interaction_type
|
||||||
|
),
|
||||||
|
);
|
||||||
|
json_response(200, serde_json::json!({"type": 6}))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn on_poll() {}
|
||||||
|
|
||||||
|
fn on_respond(response: AgentResponse) -> Result<(), String> {
|
||||||
|
let metadata: DiscordMessageMetadata = serde_json::from_str(&response.metadata_json)
|
||||||
|
.map_err(|e| format!("Failed to parse metadata: {}", e))?;
|
||||||
|
|
||||||
|
// Use webhook endpoint for followup
|
||||||
|
let url = format!(
|
||||||
|
"https://discord.com/api/v10/webhooks/{}/{}",
|
||||||
|
metadata.application_id, metadata.token
|
||||||
|
);
|
||||||
|
|
||||||
|
// Truncate content to 2000 characters to comply with Discord limits
|
||||||
|
let content = truncate_message(&response.content);
|
||||||
|
|
||||||
|
let mut payload = serde_json::json!({
|
||||||
|
"content": content,
|
||||||
|
});
|
||||||
|
|
||||||
|
// Check for embeds in metadata
|
||||||
|
if let Ok(meta_json) = serde_json::from_str::<serde_json::Value>(&response.metadata_json) {
|
||||||
|
if let Some(embeds) = meta_json.get("embeds") {
|
||||||
|
payload["embeds"] = embeds.clone();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
let payload_bytes =
|
||||||
|
serde_json::to_vec(&payload).map_err(|e| format!("Failed to serialize: {}", e))?;
|
||||||
|
|
||||||
|
let headers = serde_json::json!({
|
||||||
|
"Content-Type": "application/json"
|
||||||
|
});
|
||||||
|
|
||||||
|
let result = channel_host::http_request(
|
||||||
|
"POST",
|
||||||
|
&url,
|
||||||
|
&headers.to_string(),
|
||||||
|
Some(&payload_bytes),
|
||||||
|
None,
|
||||||
|
);
|
||||||
|
|
||||||
|
match result {
|
||||||
|
Ok(http_response) => {
|
||||||
|
if http_response.status >= 200 && http_response.status < 300 {
|
||||||
|
channel_host::log(channel_host::LogLevel::Debug, "Posted followup to Discord");
|
||||||
|
Ok(())
|
||||||
|
} else {
|
||||||
|
let body_str = String::from_utf8_lossy(&http_response.body);
|
||||||
|
Err(format!(
|
||||||
|
"Discord API error: {} - {}",
|
||||||
|
http_response.status, body_str
|
||||||
|
))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Err(e) => Err(format!("HTTP request failed: {}", e)),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn on_status(_update: StatusUpdate) {}
|
||||||
|
|
||||||
|
fn on_shutdown() {
|
||||||
|
channel_host::log(
|
||||||
|
channel_host::LogLevel::Info,
|
||||||
|
"Discord channel shutting down",
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn handle_slash_command(interaction: &DiscordInteraction) {
|
||||||
|
let user = interaction
|
||||||
|
.member
|
||||||
|
.as_ref()
|
||||||
|
.map(|m| &m.user)
|
||||||
|
.or(interaction.user.as_ref());
|
||||||
|
let user_id = user.map(|u| u.id.clone()).unwrap_or_default();
|
||||||
|
let user_name = user
|
||||||
|
.map(|u| {
|
||||||
|
u.global_name
|
||||||
|
.as_ref()
|
||||||
|
.filter(|s| !s.is_empty())
|
||||||
|
.unwrap_or(&u.username)
|
||||||
|
.clone()
|
||||||
|
})
|
||||||
|
.unwrap_or_default();
|
||||||
|
|
||||||
|
let channel_id = interaction.channel_id.clone().unwrap_or_default();
|
||||||
|
|
||||||
|
let command_name = interaction
|
||||||
|
.data
|
||||||
|
.as_ref()
|
||||||
|
.map(|d| d.name.clone())
|
||||||
|
.unwrap_or_default();
|
||||||
|
let options = interaction.data.as_ref().and_then(|d| d.options.clone());
|
||||||
|
|
||||||
|
let content = if let Some(opts) = options {
|
||||||
|
let opt_str = opts
|
||||||
|
.iter()
|
||||||
|
.map(|o| format!("{}: {}", o.name, o.value))
|
||||||
|
.collect::<Vec<_>>()
|
||||||
|
.join(", ");
|
||||||
|
format!("/{} {}", command_name, opt_str)
|
||||||
|
} else {
|
||||||
|
format!("/{}", command_name)
|
||||||
|
};
|
||||||
|
|
||||||
|
let metadata = DiscordMessageMetadata {
|
||||||
|
channel_id: channel_id.clone(),
|
||||||
|
interaction_id: interaction.id.clone(),
|
||||||
|
token: interaction.token.clone(),
|
||||||
|
application_id: interaction.application_id.clone(),
|
||||||
|
thread_id: None,
|
||||||
|
};
|
||||||
|
|
||||||
|
let metadata_json = match serde_json::to_string(&metadata) {
|
||||||
|
Ok(json) => json,
|
||||||
|
Err(e) => {
|
||||||
|
channel_host::log(
|
||||||
|
channel_host::LogLevel::Error,
|
||||||
|
&format!("Failed to serialize metadata: {}", e),
|
||||||
|
);
|
||||||
|
// Attempt to notify user of internal error
|
||||||
|
let url = format!(
|
||||||
|
"https://discord.com/api/v10/webhooks/{}/{}",
|
||||||
|
interaction.application_id, interaction.token
|
||||||
|
);
|
||||||
|
let payload = serde_json::json!({
|
||||||
|
"content": "❌ Internal Error: Failed to process command metadata.",
|
||||||
|
"flags": 64 // Ephemeral
|
||||||
|
});
|
||||||
|
let _ = channel_host::http_request(
|
||||||
|
"POST",
|
||||||
|
&url,
|
||||||
|
&serde_json::json!({"Content-Type": "application/json"}).to_string(),
|
||||||
|
Some(&serde_json::to_vec(&payload).unwrap_or_default()),
|
||||||
|
None,
|
||||||
|
);
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
channel_host::emit_message(&EmittedMessage {
|
||||||
|
user_id,
|
||||||
|
user_name: Some(user_name),
|
||||||
|
content,
|
||||||
|
thread_id: None,
|
||||||
|
metadata_json,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
fn handle_message_component(interaction: &DiscordInteraction, message: &DiscordMessage) {
|
||||||
|
// Check member first (for server contexts), then user (for DMs)
|
||||||
|
let user = interaction
|
||||||
|
.member
|
||||||
|
.as_ref()
|
||||||
|
.map(|m| &m.user)
|
||||||
|
.or(interaction.user.as_ref());
|
||||||
|
let user_id = user.map(|u| u.id.clone()).unwrap_or_default();
|
||||||
|
let user_name = user
|
||||||
|
.map(|u| {
|
||||||
|
u.global_name
|
||||||
|
.as_ref()
|
||||||
|
.filter(|s| !s.is_empty())
|
||||||
|
.unwrap_or(&u.username)
|
||||||
|
.clone()
|
||||||
|
})
|
||||||
|
.unwrap_or_default();
|
||||||
|
|
||||||
|
let channel_id = message.channel_id.clone();
|
||||||
|
|
||||||
|
let metadata = DiscordMessageMetadata {
|
||||||
|
channel_id: channel_id.clone(),
|
||||||
|
interaction_id: interaction.id.clone(),
|
||||||
|
token: interaction.token.clone(),
|
||||||
|
application_id: interaction.application_id.clone(),
|
||||||
|
thread_id: None,
|
||||||
|
};
|
||||||
|
|
||||||
|
let metadata_json = match serde_json::to_string(&metadata) {
|
||||||
|
Ok(json) => json,
|
||||||
|
Err(e) => {
|
||||||
|
channel_host::log(
|
||||||
|
channel_host::LogLevel::Error,
|
||||||
|
&format!("Failed to serialize metadata: {}", e),
|
||||||
|
);
|
||||||
|
return; // Don't emit message if metadata can't be serialized
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
channel_host::emit_message(&EmittedMessage {
|
||||||
|
user_id,
|
||||||
|
user_name: Some(user_name),
|
||||||
|
content: format!("[Button clicked] {}", message.content),
|
||||||
|
thread_id: None,
|
||||||
|
metadata_json,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
fn json_response(status: u16, value: serde_json::Value) -> OutgoingHttpResponse {
|
||||||
|
let body = serde_json::to_vec(&value).unwrap_or_default();
|
||||||
|
let headers = serde_json::json!({"Content-Type": "application/json"});
|
||||||
|
|
||||||
|
OutgoingHttpResponse {
|
||||||
|
status,
|
||||||
|
headers_json: headers.to_string(),
|
||||||
|
body,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
export!(DiscordChannel);
|
||||||
|
|
||||||
|
fn truncate_message(content: &str) -> String {
|
||||||
|
if content.len() <= 2000 {
|
||||||
|
content.to_string()
|
||||||
|
} else {
|
||||||
|
let max_bytes = 1990;
|
||||||
|
let cutoff = content
|
||||||
|
.char_indices()
|
||||||
|
.map(|(i, c)| i + c.len_utf8())
|
||||||
|
.take_while(|&end| end <= max_bytes)
|
||||||
|
.last()
|
||||||
|
.unwrap_or(0);
|
||||||
|
let mut truncated = content[..cutoff].to_string();
|
||||||
|
truncated.push_str("\n... (truncated)");
|
||||||
|
truncated
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_truncate_message() {
|
||||||
|
let short = "Hello world";
|
||||||
|
assert_eq!(truncate_message(short), short);
|
||||||
|
|
||||||
|
let long = "a".repeat(2005);
|
||||||
|
let truncated = truncate_message(&long);
|
||||||
|
assert_eq!(truncated.len(), 2006); // 1990 + 16 chars suffix
|
||||||
|
assert!(truncated.ends_with("\n... (truncated)"));
|
||||||
|
|
||||||
|
// Test with multibyte characters (Euro sign is 3 bytes)
|
||||||
|
// 1000 chars * 3 bytes = 3000 bytes
|
||||||
|
let multi = "€".repeat(1000);
|
||||||
|
let truncated_multi = truncate_message(&multi);
|
||||||
|
|
||||||
|
// 1990 bytes limit. 1990 / 3 = 663 with remainder 1.
|
||||||
|
// Should truncate at 663 chars (1989 bytes).
|
||||||
|
// Suffix is 16 bytes. Total: 1989 + 16 = 2005 bytes.
|
||||||
|
assert!(truncated_multi.len() <= 2006);
|
||||||
|
assert!(truncated_multi.len() >= 2006 - 4); // Allow for max utf8 char width variance
|
||||||
|
assert!(truncated_multi.ends_with("\n... (truncated)"));
|
||||||
|
|
||||||
|
let content_part = &truncated_multi[..truncated_multi.len() - 16];
|
||||||
|
assert!(content_part.chars().all(|c| c == '€'));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_metadata_serialization() {
|
||||||
|
let metadata = DiscordMessageMetadata {
|
||||||
|
channel_id: "123".into(),
|
||||||
|
interaction_id: "456".into(),
|
||||||
|
token: "abc".into(),
|
||||||
|
application_id: "789".into(),
|
||||||
|
thread_id: None,
|
||||||
|
};
|
||||||
|
let json = serde_json::to_string(&metadata).unwrap();
|
||||||
|
let parsed: DiscordMessageMetadata = serde_json::from_str(&json).unwrap();
|
||||||
|
assert_eq!(parsed.channel_id, "123");
|
||||||
|
assert_eq!(parsed.interaction_id, "456");
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -338,7 +338,13 @@ fn emit_message(
|
|||||||
team_id,
|
team_id,
|
||||||
};
|
};
|
||||||
|
|
||||||
let metadata_json = serde_json::to_string(&metadata).unwrap_or_else(|_| "{}".to_string());
|
let metadata_json = serde_json::to_string(&metadata).unwrap_or_else(|e| {
|
||||||
|
channel_host::log(
|
||||||
|
channel_host::LogLevel::Error,
|
||||||
|
&format!("Failed to serialize Slack metadata: {}", e),
|
||||||
|
);
|
||||||
|
"{}".to_string()
|
||||||
|
});
|
||||||
|
|
||||||
// Strip @ mentions of the bot from the text for cleaner messages
|
// Strip @ mentions of the bot from the text for cleaner messages
|
||||||
let cleaned_text = strip_bot_mention(&text);
|
let cleaned_text = strip_bot_mention(&text);
|
||||||
@@ -366,7 +372,13 @@ fn strip_bot_mention(text: &str) -> String {
|
|||||||
|
|
||||||
/// Create a JSON HTTP response.
|
/// Create a JSON HTTP response.
|
||||||
fn json_response(status: u16, value: serde_json::Value) -> OutgoingHttpResponse {
|
fn json_response(status: u16, value: serde_json::Value) -> OutgoingHttpResponse {
|
||||||
let body = serde_json::to_vec(&value).unwrap_or_default();
|
let body = serde_json::to_vec(&value).unwrap_or_else(|e| {
|
||||||
|
channel_host::log(
|
||||||
|
channel_host::LogLevel::Error,
|
||||||
|
&format!("Failed to serialize JSON response: {}", e),
|
||||||
|
);
|
||||||
|
Vec::new()
|
||||||
|
});
|
||||||
let headers = serde_json::json!({"Content-Type": "application/json"});
|
let headers = serde_json::json!({"Content-Type": "application/json"});
|
||||||
|
|
||||||
OutgoingHttpResponse {
|
OutgoingHttpResponse {
|
||||||
|
|||||||
@@ -285,11 +285,7 @@ impl Guest for TelegramChannel {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Persist dm_policy and allow_from for DM pairing in handle_message
|
// Persist dm_policy and allow_from for DM pairing in handle_message
|
||||||
let dm_policy = config
|
let dm_policy = config.dm_policy.as_deref().unwrap_or("pairing").to_string();
|
||||||
.dm_policy
|
|
||||||
.as_deref()
|
|
||||||
.unwrap_or("pairing")
|
|
||||||
.to_string();
|
|
||||||
let _ = channel_host::workspace_write(DM_POLICY_PATH, &dm_policy);
|
let _ = channel_host::workspace_write(DM_POLICY_PATH, &dm_policy);
|
||||||
|
|
||||||
let allow_from_json = serde_json::to_string(&config.allow_from.unwrap_or_default())
|
let allow_from_json = serde_json::to_string(&config.allow_from.unwrap_or_default())
|
||||||
@@ -844,8 +840,8 @@ fn send_pairing_reply(chat_id: i64, code: &str) -> Result<(), String> {
|
|||||||
"parse_mode": "Markdown",
|
"parse_mode": "Markdown",
|
||||||
});
|
});
|
||||||
|
|
||||||
let payload_bytes = serde_json::to_vec(&payload)
|
let payload_bytes =
|
||||||
.map_err(|e| format!("Failed to serialize payload: {}", e))?;
|
serde_json::to_vec(&payload).map_err(|e| format!("Failed to serialize payload: {}", e))?;
|
||||||
|
|
||||||
let headers = serde_json::json!({
|
let headers = serde_json::json!({
|
||||||
"Content-Type": "application/json"
|
"Content-Type": "application/json"
|
||||||
@@ -915,15 +911,10 @@ fn handle_message(message: TelegramMessage) {
|
|||||||
let is_private = message.chat.chat_type == "private";
|
let is_private = message.chat.chat_type == "private";
|
||||||
|
|
||||||
// Owner validation: when owner_id is set, only that user can message
|
// Owner validation: when owner_id is set, only that user can message
|
||||||
let owner_configured = channel_host::workspace_read(OWNER_ID_PATH)
|
let owner_id_str = channel_host::workspace_read(OWNER_ID_PATH).filter(|s| !s.is_empty());
|
||||||
.map(|s| !s.is_empty())
|
|
||||||
.unwrap_or(false);
|
|
||||||
|
|
||||||
if owner_configured {
|
if let Some(ref id_str) = owner_id_str {
|
||||||
if let Ok(owner_id) = channel_host::workspace_read(OWNER_ID_PATH)
|
if let Ok(owner_id) = id_str.parse::<i64>() {
|
||||||
.unwrap()
|
|
||||||
.parse::<i64>()
|
|
||||||
{
|
|
||||||
if from.id != owner_id {
|
if from.id != owner_id {
|
||||||
channel_host::log(
|
channel_host::log(
|
||||||
channel_host::LogLevel::Debug,
|
channel_host::LogLevel::Debug,
|
||||||
@@ -937,8 +928,8 @@ fn handle_message(message: TelegramMessage) {
|
|||||||
}
|
}
|
||||||
} else if is_private {
|
} else if is_private {
|
||||||
// No owner_id: apply dm_policy for private chats
|
// No owner_id: apply dm_policy for private chats
|
||||||
let dm_policy = channel_host::workspace_read(DM_POLICY_PATH)
|
let dm_policy =
|
||||||
.unwrap_or_else(|| "pairing".to_string());
|
channel_host::workspace_read(DM_POLICY_PATH).unwrap_or_else(|| "pairing".to_string());
|
||||||
|
|
||||||
if dm_policy != "open" {
|
if dm_policy != "open" {
|
||||||
// Build effective allow list: config allow_from + pairing store
|
// Build effective allow list: config allow_from + pairing store
|
||||||
@@ -1001,8 +992,7 @@ fn handle_message(message: TelegramMessage) {
|
|||||||
|
|
||||||
if !respond_to_all {
|
if !respond_to_all {
|
||||||
let has_command = content.starts_with('/');
|
let has_command = content.starts_with('/');
|
||||||
let bot_username = channel_host::workspace_read(BOT_USERNAME_PATH)
|
let bot_username = channel_host::workspace_read(BOT_USERNAME_PATH).unwrap_or_default();
|
||||||
.unwrap_or_default();
|
|
||||||
let has_bot_mention = if bot_username.is_empty() {
|
let has_bot_mention = if bot_username.is_empty() {
|
||||||
content.contains('@')
|
content.contains('@')
|
||||||
} else {
|
} else {
|
||||||
|
|||||||
@@ -254,10 +254,19 @@ struct WhatsAppChannel;
|
|||||||
|
|
||||||
impl Guest for WhatsAppChannel {
|
impl Guest for WhatsAppChannel {
|
||||||
fn on_start(config_json: String) -> Result<ChannelConfig, String> {
|
fn on_start(config_json: String) -> Result<ChannelConfig, String> {
|
||||||
let config: WhatsAppConfig = serde_json::from_str(&config_json).unwrap_or(WhatsAppConfig {
|
let config: WhatsAppConfig = match serde_json::from_str(&config_json) {
|
||||||
|
Ok(c) => c,
|
||||||
|
Err(e) => {
|
||||||
|
channel_host::log(
|
||||||
|
channel_host::LogLevel::Warn,
|
||||||
|
&format!("Failed to parse WhatsApp config, using defaults: {}", e),
|
||||||
|
);
|
||||||
|
WhatsAppConfig {
|
||||||
api_version: default_api_version(),
|
api_version: default_api_version(),
|
||||||
reply_to_message: default_reply_to_message(),
|
reply_to_message: default_reply_to_message(),
|
||||||
});
|
}
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
channel_host::log(
|
channel_host::log(
|
||||||
channel_host::LogLevel::Info,
|
channel_host::LogLevel::Info,
|
||||||
@@ -267,6 +276,9 @@ impl Guest for WhatsAppChannel {
|
|||||||
),
|
),
|
||||||
);
|
);
|
||||||
|
|
||||||
|
// Persist api_version in workspace so on_respond() can read it
|
||||||
|
let _ = channel_host::workspace_write("channels/whatsapp/api_version", &config.api_version);
|
||||||
|
|
||||||
// WhatsApp Cloud API is webhook-only, no polling available
|
// WhatsApp Cloud API is webhook-only, no polling available
|
||||||
Ok(ChannelConfig {
|
Ok(ChannelConfig {
|
||||||
display_name: "WhatsApp".to_string(),
|
display_name: "WhatsApp".to_string(),
|
||||||
@@ -327,11 +339,16 @@ impl Guest for WhatsAppChannel {
|
|||||||
let metadata: WhatsAppMessageMetadata = serde_json::from_str(&response.metadata_json)
|
let metadata: WhatsAppMessageMetadata = serde_json::from_str(&response.metadata_json)
|
||||||
.map_err(|e| format!("Failed to parse metadata: {}", e))?;
|
.map_err(|e| format!("Failed to parse metadata: {}", e))?;
|
||||||
|
|
||||||
|
// Read api_version from workspace (set during on_start), fallback to default
|
||||||
|
let api_version = channel_host::workspace_read("channels/whatsapp/api_version")
|
||||||
|
.filter(|s| !s.is_empty())
|
||||||
|
.unwrap_or_else(|| "v18.0".to_string());
|
||||||
|
|
||||||
// Build WhatsApp API URL with token placeholder
|
// Build WhatsApp API URL with token placeholder
|
||||||
// Host will replace {WHATSAPP_ACCESS_TOKEN} with actual token in Authorization header
|
// Host will replace {WHATSAPP_ACCESS_TOKEN} with actual token in Authorization header
|
||||||
let api_url = format!(
|
let api_url = format!(
|
||||||
"https://graph.facebook.com/v18.0/{}/messages",
|
"https://graph.facebook.com/{}/{}/messages",
|
||||||
metadata.phone_number_id
|
api_version, metadata.phone_number_id
|
||||||
);
|
);
|
||||||
|
|
||||||
// Build sendMessage payload
|
// Build sendMessage payload
|
||||||
|
|||||||
+146
-34
@@ -22,6 +22,7 @@ use crate::context::JobContext;
|
|||||||
use crate::db::Database;
|
use crate::db::Database;
|
||||||
use crate::error::Error;
|
use crate::error::Error;
|
||||||
use crate::extensions::ExtensionManager;
|
use crate::extensions::ExtensionManager;
|
||||||
|
use crate::hooks::HookRegistry;
|
||||||
use crate::llm::{ChatMessage, LlmProvider, Reasoning, ReasoningContext, RespondResult};
|
use crate::llm::{ChatMessage, LlmProvider, Reasoning, ReasoningContext, RespondResult};
|
||||||
use crate::safety::SafetyLayer;
|
use crate::safety::SafetyLayer;
|
||||||
use crate::tools::ToolRegistry;
|
use crate::tools::ToolRegistry;
|
||||||
@@ -67,10 +68,14 @@ enum AgenticLoopResult {
|
|||||||
pub struct AgentDeps {
|
pub struct AgentDeps {
|
||||||
pub store: Option<Arc<dyn Database>>,
|
pub store: Option<Arc<dyn Database>>,
|
||||||
pub llm: Arc<dyn LlmProvider>,
|
pub llm: Arc<dyn LlmProvider>,
|
||||||
|
/// Cheap/fast LLM for lightweight tasks (heartbeat, routing, evaluation).
|
||||||
|
/// Falls back to the main `llm` if None.
|
||||||
|
pub cheap_llm: Option<Arc<dyn LlmProvider>>,
|
||||||
pub safety: Arc<SafetyLayer>,
|
pub safety: Arc<SafetyLayer>,
|
||||||
pub tools: Arc<ToolRegistry>,
|
pub tools: Arc<ToolRegistry>,
|
||||||
pub workspace: Option<Arc<Workspace>>,
|
pub workspace: Option<Arc<Workspace>>,
|
||||||
pub extension_manager: Option<Arc<ExtensionManager>>,
|
pub extension_manager: Option<Arc<ExtensionManager>>,
|
||||||
|
pub hooks: Arc<HookRegistry>,
|
||||||
}
|
}
|
||||||
|
|
||||||
/// The main agent that coordinates all components.
|
/// The main agent that coordinates all components.
|
||||||
@@ -113,6 +118,7 @@ impl Agent {
|
|||||||
deps.safety.clone(),
|
deps.safety.clone(),
|
||||||
deps.tools.clone(),
|
deps.tools.clone(),
|
||||||
deps.store.clone(),
|
deps.store.clone(),
|
||||||
|
deps.hooks.clone(),
|
||||||
));
|
));
|
||||||
|
|
||||||
Self {
|
Self {
|
||||||
@@ -138,6 +144,11 @@ impl Agent {
|
|||||||
&self.deps.llm
|
&self.deps.llm
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Get the cheap/fast LLM provider, falling back to the main one.
|
||||||
|
fn cheap_llm(&self) -> &Arc<dyn LlmProvider> {
|
||||||
|
self.deps.cheap_llm.as_ref().unwrap_or(&self.deps.llm)
|
||||||
|
}
|
||||||
|
|
||||||
fn safety(&self) -> &Arc<SafetyLayer> {
|
fn safety(&self) -> &Arc<SafetyLayer> {
|
||||||
&self.deps.safety
|
&self.deps.safety
|
||||||
}
|
}
|
||||||
@@ -150,6 +161,10 @@ impl Agent {
|
|||||||
self.deps.workspace.as_ref()
|
self.deps.workspace.as_ref()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn hooks(&self) -> &Arc<HookRegistry> {
|
||||||
|
&self.deps.hooks
|
||||||
|
}
|
||||||
|
|
||||||
/// Run the agent main loop.
|
/// Run the agent main loop.
|
||||||
pub async fn run(self) -> Result<(), Error> {
|
pub async fn run(self) -> Result<(), Error> {
|
||||||
// Start channels
|
// Start channels
|
||||||
@@ -301,7 +316,7 @@ impl Agent {
|
|||||||
Some(spawn_heartbeat(
|
Some(spawn_heartbeat(
|
||||||
config,
|
config,
|
||||||
workspace.clone(),
|
workspace.clone(),
|
||||||
self.llm().clone(),
|
self.cheap_llm().clone(),
|
||||||
Some(notify_tx),
|
Some(notify_tx),
|
||||||
))
|
))
|
||||||
} else {
|
} else {
|
||||||
@@ -417,11 +432,33 @@ impl Agent {
|
|||||||
|
|
||||||
match self.handle_message(&message).await {
|
match self.handle_message(&message).await {
|
||||||
Ok(Some(response)) if !response.is_empty() => {
|
Ok(Some(response)) if !response.is_empty() => {
|
||||||
|
// Hook: BeforeOutbound — allow hooks to modify or suppress outbound
|
||||||
|
let event = crate::hooks::HookEvent::Outbound {
|
||||||
|
user_id: message.user_id.clone(),
|
||||||
|
channel: message.channel.clone(),
|
||||||
|
content: response.clone(),
|
||||||
|
thread_id: message.thread_id.clone(),
|
||||||
|
};
|
||||||
|
match self.hooks().run(&event).await {
|
||||||
|
Err(err) => {
|
||||||
|
tracing::warn!("BeforeOutbound hook blocked response: {}", err);
|
||||||
|
}
|
||||||
|
Ok(crate::hooks::HookOutcome::Continue {
|
||||||
|
modified: Some(new_content),
|
||||||
|
}) => {
|
||||||
|
let _ = self
|
||||||
|
.channels
|
||||||
|
.respond(&message, OutgoingResponse::text(new_content))
|
||||||
|
.await;
|
||||||
|
}
|
||||||
|
_ => {
|
||||||
let _ = self
|
let _ = self
|
||||||
.channels
|
.channels
|
||||||
.respond(&message, OutgoingResponse::text(response))
|
.respond(&message, OutgoingResponse::text(response))
|
||||||
.await;
|
.await;
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
Ok(Some(_)) => {
|
Ok(Some(_)) => {
|
||||||
// Empty response, nothing to send (e.g. approval handled via send_status)
|
// Empty response, nothing to send (e.g. approval handled via send_status)
|
||||||
}
|
}
|
||||||
@@ -466,7 +503,33 @@ impl Agent {
|
|||||||
|
|
||||||
async fn handle_message(&self, message: &IncomingMessage) -> Result<Option<String>, Error> {
|
async fn handle_message(&self, message: &IncomingMessage) -> Result<Option<String>, Error> {
|
||||||
// Parse submission type first
|
// Parse submission type first
|
||||||
let submission = SubmissionParser::parse(&message.content);
|
let mut submission = SubmissionParser::parse(&message.content);
|
||||||
|
|
||||||
|
// Hook: BeforeInbound — allow hooks to modify or reject user input
|
||||||
|
if let Submission::UserInput { ref content } = submission {
|
||||||
|
let event = crate::hooks::HookEvent::Inbound {
|
||||||
|
user_id: message.user_id.clone(),
|
||||||
|
channel: message.channel.clone(),
|
||||||
|
content: content.clone(),
|
||||||
|
thread_id: message.thread_id.clone(),
|
||||||
|
};
|
||||||
|
match self.hooks().run(&event).await {
|
||||||
|
Err(crate::hooks::HookError::Rejected { reason }) => {
|
||||||
|
return Ok(Some(format!("[Message rejected: {}]", reason)));
|
||||||
|
}
|
||||||
|
Err(err) => {
|
||||||
|
return Ok(Some(format!("[Message blocked by hook policy: {}]", err)));
|
||||||
|
}
|
||||||
|
Ok(crate::hooks::HookOutcome::Continue {
|
||||||
|
modified: Some(new_content),
|
||||||
|
}) => {
|
||||||
|
submission = Submission::UserInput {
|
||||||
|
content: new_content,
|
||||||
|
};
|
||||||
|
}
|
||||||
|
_ => {} // Continue, fail-open errors already logged in registry
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Hydrate thread from DB if it's a historical thread not in memory
|
// Hydrate thread from DB if it's a historical thread not in memory
|
||||||
if let Some(ref external_thread_id) = message.thread_id {
|
if let Some(ref external_thread_id) = message.thread_id {
|
||||||
@@ -875,6 +938,27 @@ impl Agent {
|
|||||||
// Complete, fail, or request approval
|
// Complete, fail, or request approval
|
||||||
match result {
|
match result {
|
||||||
Ok(AgenticLoopResult::Response(response)) => {
|
Ok(AgenticLoopResult::Response(response)) => {
|
||||||
|
// Hook: TransformResponse — allow hooks to modify or reject the final response
|
||||||
|
let response = {
|
||||||
|
let event = crate::hooks::HookEvent::ResponseTransform {
|
||||||
|
user_id: message.user_id.clone(),
|
||||||
|
thread_id: thread_id.to_string(),
|
||||||
|
response: response.clone(),
|
||||||
|
};
|
||||||
|
match self.hooks().run(&event).await {
|
||||||
|
Err(crate::hooks::HookError::Rejected { reason }) => {
|
||||||
|
format!("[Response filtered: {}]", reason)
|
||||||
|
}
|
||||||
|
Err(err) => {
|
||||||
|
format!("[Response blocked by hook policy: {}]", err)
|
||||||
|
}
|
||||||
|
Ok(crate::hooks::HookOutcome::Continue {
|
||||||
|
modified: Some(new_response),
|
||||||
|
}) => new_response,
|
||||||
|
_ => response, // fail-open: use original
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
thread.complete_turn(&response);
|
thread.complete_turn(&response);
|
||||||
self.persist_response_chain(thread);
|
self.persist_response_chain(thread);
|
||||||
let _ = self
|
let _ = self
|
||||||
@@ -1152,8 +1236,8 @@ impl Agent {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Execute each tool (with approval checking)
|
// Execute each tool (with approval checking and hook interception)
|
||||||
for tc in tool_calls {
|
for mut tc in tool_calls {
|
||||||
// Check if tool requires approval
|
// Check if tool requires approval
|
||||||
if let Some(tool) = self.tools().get(&tc.name).await
|
if let Some(tool) = self.tools().get(&tc.name).await
|
||||||
&& tool.requires_approval()
|
&& tool.requires_approval()
|
||||||
@@ -1164,31 +1248,12 @@ impl Agent {
|
|||||||
sess.is_tool_auto_approved(&tc.name)
|
sess.is_tool_auto_approved(&tc.name)
|
||||||
};
|
};
|
||||||
|
|
||||||
// For shell commands, override auto-approval for
|
// Let the tool inspect the specific parameters and
|
||||||
// destructive patterns that should always require
|
// override auto-approval (e.g. destructive shell commands).
|
||||||
// explicit per-invocation approval.
|
if is_auto_approved && tool.requires_approval_for(&tc.arguments) {
|
||||||
if is_auto_approved
|
|
||||||
&& tc.name == "shell"
|
|
||||||
&& let Some(cmd) = tc
|
|
||||||
.arguments
|
|
||||||
.get("command")
|
|
||||||
.and_then(|c| c.as_str().map(String::from))
|
|
||||||
.or_else(|| {
|
|
||||||
tc.arguments
|
|
||||||
.as_str()
|
|
||||||
.and_then(|s| {
|
|
||||||
serde_json::from_str::<serde_json::Value>(s).ok()
|
|
||||||
})
|
|
||||||
.and_then(|v| {
|
|
||||||
v.get("command")
|
|
||||||
.and_then(|c| c.as_str().map(String::from))
|
|
||||||
})
|
|
||||||
})
|
|
||||||
&& crate::tools::builtin::shell::requires_explicit_approval(&cmd)
|
|
||||||
{
|
|
||||||
tracing::info!(
|
tracing::info!(
|
||||||
"Shell command '{}' requires explicit approval despite auto-approve",
|
tool = %tc.name,
|
||||||
cmd.chars().take(80).collect::<String>()
|
"Tool requires explicit approval for these parameters despite auto-approve"
|
||||||
);
|
);
|
||||||
is_auto_approved = false;
|
is_auto_approved = false;
|
||||||
}
|
}
|
||||||
@@ -1208,6 +1273,47 @@ impl Agent {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Hook: BeforeToolCall — allow hooks to modify or reject tool calls
|
||||||
|
{
|
||||||
|
let event = crate::hooks::HookEvent::ToolCall {
|
||||||
|
tool_name: tc.name.clone(),
|
||||||
|
parameters: tc.arguments.clone(),
|
||||||
|
user_id: message.user_id.clone(),
|
||||||
|
context: "chat".to_string(),
|
||||||
|
};
|
||||||
|
match self.hooks().run(&event).await {
|
||||||
|
Err(crate::hooks::HookError::Rejected { reason }) => {
|
||||||
|
context_messages.push(ChatMessage::tool_result(
|
||||||
|
&tc.id,
|
||||||
|
&tc.name,
|
||||||
|
format!("Tool call rejected by hook: {}", reason),
|
||||||
|
));
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
Err(err) => {
|
||||||
|
context_messages.push(ChatMessage::tool_result(
|
||||||
|
&tc.id,
|
||||||
|
&tc.name,
|
||||||
|
format!("Tool call blocked by hook policy: {}", err),
|
||||||
|
));
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
Ok(crate::hooks::HookOutcome::Continue {
|
||||||
|
modified: Some(new_params),
|
||||||
|
}) => match serde_json::from_str(&new_params) {
|
||||||
|
Ok(parsed) => tc.arguments = parsed,
|
||||||
|
Err(e) => {
|
||||||
|
tracing::warn!(
|
||||||
|
tool = %tc.name,
|
||||||
|
"Hook returned non-JSON modification for ToolCall, ignoring: {}",
|
||||||
|
e
|
||||||
|
);
|
||||||
|
}
|
||||||
|
},
|
||||||
|
_ => {} // Continue, fail-open errors already logged
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
let _ = self
|
let _ = self
|
||||||
.channels
|
.channels
|
||||||
.send_status(
|
.send_status(
|
||||||
@@ -1473,6 +1579,9 @@ impl Agent {
|
|||||||
session: Arc<Mutex<Session>>,
|
session: Arc<Mutex<Session>>,
|
||||||
thread_id: Uuid,
|
thread_id: Uuid,
|
||||||
) -> Result<SubmissionResult, Error> {
|
) -> Result<SubmissionResult, Error> {
|
||||||
|
// Lock session first, then undo manager -- consistent with process_user_input
|
||||||
|
// to avoid potential deadlocks.
|
||||||
|
let mut sess = session.lock().await;
|
||||||
let undo_mgr = self.session_manager.get_undo_manager(thread_id).await;
|
let undo_mgr = self.session_manager.get_undo_manager(thread_id).await;
|
||||||
let mut mgr = undo_mgr.lock().await;
|
let mut mgr = undo_mgr.lock().await;
|
||||||
|
|
||||||
@@ -1480,7 +1589,6 @@ impl Agent {
|
|||||||
return Ok(SubmissionResult::ok_with_message("Nothing to undo."));
|
return Ok(SubmissionResult::ok_with_message("Nothing to undo."));
|
||||||
}
|
}
|
||||||
|
|
||||||
let mut sess = session.lock().await;
|
|
||||||
let thread = sess
|
let thread = sess
|
||||||
.threads
|
.threads
|
||||||
.get_mut(&thread_id)
|
.get_mut(&thread_id)
|
||||||
@@ -1491,12 +1599,10 @@ impl Agent {
|
|||||||
let current_turn = thread.turn_number();
|
let current_turn = thread.turn_number();
|
||||||
|
|
||||||
if let Some(checkpoint) = mgr.undo(current_turn, current_messages) {
|
if let Some(checkpoint) = mgr.undo(current_turn, current_messages) {
|
||||||
// Extract values before consuming the reference
|
|
||||||
let turn_number = checkpoint.turn_number;
|
let turn_number = checkpoint.turn_number;
|
||||||
let messages = checkpoint.messages.clone();
|
|
||||||
let undo_count = mgr.undo_count();
|
let undo_count = mgr.undo_count();
|
||||||
// Restore thread from checkpoint
|
// Restore thread from checkpoint
|
||||||
thread.restore_from_messages(messages);
|
thread.restore_from_messages(checkpoint.messages);
|
||||||
Ok(SubmissionResult::ok_with_message(format!(
|
Ok(SubmissionResult::ok_with_message(format!(
|
||||||
"Undone to turn {}. {} undo(s) remaining.",
|
"Undone to turn {}. {} undo(s) remaining.",
|
||||||
turn_number, undo_count
|
turn_number, undo_count
|
||||||
@@ -1511,6 +1617,9 @@ impl Agent {
|
|||||||
session: Arc<Mutex<Session>>,
|
session: Arc<Mutex<Session>>,
|
||||||
thread_id: Uuid,
|
thread_id: Uuid,
|
||||||
) -> Result<SubmissionResult, Error> {
|
) -> Result<SubmissionResult, Error> {
|
||||||
|
// Lock session first, then undo manager -- consistent with process_user_input
|
||||||
|
// to avoid potential deadlocks.
|
||||||
|
let mut sess = session.lock().await;
|
||||||
let undo_mgr = self.session_manager.get_undo_manager(thread_id).await;
|
let undo_mgr = self.session_manager.get_undo_manager(thread_id).await;
|
||||||
let mut mgr = undo_mgr.lock().await;
|
let mut mgr = undo_mgr.lock().await;
|
||||||
|
|
||||||
@@ -1518,12 +1627,15 @@ impl Agent {
|
|||||||
return Ok(SubmissionResult::ok_with_message("Nothing to redo."));
|
return Ok(SubmissionResult::ok_with_message("Nothing to redo."));
|
||||||
}
|
}
|
||||||
|
|
||||||
if let Some(checkpoint) = mgr.redo() {
|
|
||||||
let mut sess = session.lock().await;
|
|
||||||
let thread = sess
|
let thread = sess
|
||||||
.threads
|
.threads
|
||||||
.get_mut(&thread_id)
|
.get_mut(&thread_id)
|
||||||
.ok_or_else(|| Error::from(crate::error::JobError::NotFound { id: thread_id }))?;
|
.ok_or_else(|| Error::from(crate::error::JobError::NotFound { id: thread_id }))?;
|
||||||
|
|
||||||
|
let current_messages = thread.messages();
|
||||||
|
let current_turn = thread.turn_number();
|
||||||
|
|
||||||
|
if let Some(checkpoint) = mgr.redo(current_turn, current_messages) {
|
||||||
thread.restore_from_messages(checkpoint.messages);
|
thread.restore_from_messages(checkpoint.messages);
|
||||||
Ok(SubmissionResult::ok_with_message(format!(
|
Ok(SubmissionResult::ok_with_message(format!(
|
||||||
"Redone to turn {}.",
|
"Redone to turn {}.",
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ use crate::config::AgentConfig;
|
|||||||
use crate::context::{ContextManager, JobContext, JobState};
|
use crate::context::{ContextManager, JobContext, JobState};
|
||||||
use crate::db::Database;
|
use crate::db::Database;
|
||||||
use crate::error::{Error, JobError};
|
use crate::error::{Error, JobError};
|
||||||
|
use crate::hooks::HookRegistry;
|
||||||
use crate::llm::LlmProvider;
|
use crate::llm::LlmProvider;
|
||||||
use crate::safety::SafetyLayer;
|
use crate::safety::SafetyLayer;
|
||||||
use crate::tools::ToolRegistry;
|
use crate::tools::ToolRegistry;
|
||||||
@@ -49,6 +50,7 @@ pub struct Scheduler {
|
|||||||
safety: Arc<SafetyLayer>,
|
safety: Arc<SafetyLayer>,
|
||||||
tools: Arc<ToolRegistry>,
|
tools: Arc<ToolRegistry>,
|
||||||
store: Option<Arc<dyn Database>>,
|
store: Option<Arc<dyn Database>>,
|
||||||
|
hooks: Arc<HookRegistry>,
|
||||||
/// Running jobs (main LLM-driven jobs).
|
/// Running jobs (main LLM-driven jobs).
|
||||||
jobs: Arc<RwLock<HashMap<Uuid, ScheduledJob>>>,
|
jobs: Arc<RwLock<HashMap<Uuid, ScheduledJob>>>,
|
||||||
/// Running sub-tasks (tool executions, background tasks).
|
/// Running sub-tasks (tool executions, background tasks).
|
||||||
@@ -64,6 +66,7 @@ impl Scheduler {
|
|||||||
safety: Arc<SafetyLayer>,
|
safety: Arc<SafetyLayer>,
|
||||||
tools: Arc<ToolRegistry>,
|
tools: Arc<ToolRegistry>,
|
||||||
store: Option<Arc<dyn Database>>,
|
store: Option<Arc<dyn Database>>,
|
||||||
|
hooks: Arc<HookRegistry>,
|
||||||
) -> Self {
|
) -> Self {
|
||||||
Self {
|
Self {
|
||||||
config,
|
config,
|
||||||
@@ -72,6 +75,7 @@ impl Scheduler {
|
|||||||
safety,
|
safety,
|
||||||
tools,
|
tools,
|
||||||
store,
|
store,
|
||||||
|
hooks,
|
||||||
jobs: Arc::new(RwLock::new(HashMap::new())),
|
jobs: Arc::new(RwLock::new(HashMap::new())),
|
||||||
subtasks: Arc::new(RwLock::new(HashMap::new())),
|
subtasks: Arc::new(RwLock::new(HashMap::new())),
|
||||||
}
|
}
|
||||||
@@ -118,6 +122,7 @@ impl Scheduler {
|
|||||||
safety: self.safety.clone(),
|
safety: self.safety.clone(),
|
||||||
tools: self.tools.clone(),
|
tools: self.tools.clone(),
|
||||||
store: self.store.clone(),
|
store: self.store.clone(),
|
||||||
|
hooks: self.hooks.clone(),
|
||||||
timeout: self.config.job_timeout,
|
timeout: self.config.job_timeout,
|
||||||
use_planning: self.config.use_planning,
|
use_planning: self.config.use_planning,
|
||||||
};
|
};
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ use uuid::Uuid;
|
|||||||
|
|
||||||
use crate::agent::session::Session;
|
use crate::agent::session::Session;
|
||||||
use crate::agent::undo::UndoManager;
|
use crate::agent::undo::UndoManager;
|
||||||
|
use crate::hooks::HookRegistry;
|
||||||
|
|
||||||
/// Key for mapping external thread IDs to internal ones.
|
/// Key for mapping external thread IDs to internal ones.
|
||||||
#[derive(Clone, Hash, Eq, PartialEq)]
|
#[derive(Clone, Hash, Eq, PartialEq)]
|
||||||
@@ -25,6 +26,7 @@ pub struct SessionManager {
|
|||||||
sessions: RwLock<HashMap<String, Arc<Mutex<Session>>>>,
|
sessions: RwLock<HashMap<String, Arc<Mutex<Session>>>>,
|
||||||
thread_map: RwLock<HashMap<ThreadKey, Uuid>>,
|
thread_map: RwLock<HashMap<ThreadKey, Uuid>>,
|
||||||
undo_managers: RwLock<HashMap<Uuid, Arc<Mutex<UndoManager>>>>,
|
undo_managers: RwLock<HashMap<Uuid, Arc<Mutex<UndoManager>>>>,
|
||||||
|
hooks: Option<Arc<HookRegistry>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl SessionManager {
|
impl SessionManager {
|
||||||
@@ -34,9 +36,16 @@ impl SessionManager {
|
|||||||
sessions: RwLock::new(HashMap::new()),
|
sessions: RwLock::new(HashMap::new()),
|
||||||
thread_map: RwLock::new(HashMap::new()),
|
thread_map: RwLock::new(HashMap::new()),
|
||||||
undo_managers: RwLock::new(HashMap::new()),
|
undo_managers: RwLock::new(HashMap::new()),
|
||||||
|
hooks: None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Attach a hook registry for session lifecycle events.
|
||||||
|
pub fn with_hooks(mut self, hooks: Arc<HookRegistry>) -> Self {
|
||||||
|
self.hooks = Some(hooks);
|
||||||
|
self
|
||||||
|
}
|
||||||
|
|
||||||
/// Get or create a session for a user.
|
/// Get or create a session for a user.
|
||||||
pub async fn get_or_create_session(&self, user_id: &str) -> Arc<Mutex<Session>> {
|
pub async fn get_or_create_session(&self, user_id: &str) -> Arc<Mutex<Session>> {
|
||||||
// Fast path: check if session exists
|
// Fast path: check if session exists
|
||||||
@@ -54,8 +63,28 @@ impl SessionManager {
|
|||||||
return Arc::clone(session);
|
return Arc::clone(session);
|
||||||
}
|
}
|
||||||
|
|
||||||
let session = Arc::new(Mutex::new(Session::new(user_id)));
|
let new_session = Session::new(user_id);
|
||||||
|
let session_id = new_session.id.to_string();
|
||||||
|
let session = Arc::new(Mutex::new(new_session));
|
||||||
sessions.insert(user_id.to_string(), Arc::clone(&session));
|
sessions.insert(user_id.to_string(), Arc::clone(&session));
|
||||||
|
|
||||||
|
// Fire OnSessionStart hook (fire-and-forget)
|
||||||
|
if let Some(ref hooks) = self.hooks {
|
||||||
|
let hooks = hooks.clone();
|
||||||
|
let uid = user_id.to_string();
|
||||||
|
let sid = session_id;
|
||||||
|
tokio::spawn(async move {
|
||||||
|
use crate::hooks::HookEvent;
|
||||||
|
let event = HookEvent::SessionStart {
|
||||||
|
user_id: uid,
|
||||||
|
session_id: sid,
|
||||||
|
};
|
||||||
|
if let Err(e) = hooks.run(&event).await {
|
||||||
|
tracing::warn!("OnSessionStart hook error: {}", e);
|
||||||
|
}
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
session
|
session
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -173,8 +202,8 @@ impl SessionManager {
|
|||||||
pub async fn prune_stale_sessions(&self, max_idle: std::time::Duration) -> usize {
|
pub async fn prune_stale_sessions(&self, max_idle: std::time::Duration) -> usize {
|
||||||
let cutoff = chrono::Utc::now() - chrono::TimeDelta::seconds(max_idle.as_secs() as i64);
|
let cutoff = chrono::Utc::now() - chrono::TimeDelta::seconds(max_idle.as_secs() as i64);
|
||||||
|
|
||||||
// Find stale session user_ids
|
// Find stale sessions (user_id + session_id)
|
||||||
let stale_users: Vec<String> = {
|
let stale_sessions: Vec<(String, String)> = {
|
||||||
let sessions = self.sessions.read().await;
|
let sessions = self.sessions.read().await;
|
||||||
sessions
|
sessions
|
||||||
.iter()
|
.iter()
|
||||||
@@ -182,7 +211,7 @@ impl SessionManager {
|
|||||||
// Try to lock; skip if contended (someone is actively using it)
|
// Try to lock; skip if contended (someone is actively using it)
|
||||||
let sess = session.try_lock().ok()?;
|
let sess = session.try_lock().ok()?;
|
||||||
if sess.last_active_at < cutoff {
|
if sess.last_active_at < cutoff {
|
||||||
Some(user_id.clone())
|
Some((user_id.clone(), sess.id.to_string()))
|
||||||
} else {
|
} else {
|
||||||
None
|
None
|
||||||
}
|
}
|
||||||
@@ -190,6 +219,11 @@ impl SessionManager {
|
|||||||
.collect()
|
.collect()
|
||||||
};
|
};
|
||||||
|
|
||||||
|
let stale_users: Vec<String> = stale_sessions
|
||||||
|
.iter()
|
||||||
|
.map(|(user_id, _)| user_id.clone())
|
||||||
|
.collect();
|
||||||
|
|
||||||
if stale_users.is_empty() {
|
if stale_users.is_empty() {
|
||||||
return 0;
|
return 0;
|
||||||
}
|
}
|
||||||
@@ -207,6 +241,25 @@ impl SessionManager {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Fire OnSessionEnd hooks for stale sessions (fire-and-forget)
|
||||||
|
if let Some(ref hooks) = self.hooks {
|
||||||
|
for (user_id, session_id) in &stale_sessions {
|
||||||
|
let hooks = hooks.clone();
|
||||||
|
let uid = user_id.clone();
|
||||||
|
let sid = session_id.clone();
|
||||||
|
tokio::spawn(async move {
|
||||||
|
use crate::hooks::HookEvent;
|
||||||
|
let event = HookEvent::SessionEnd {
|
||||||
|
user_id: uid,
|
||||||
|
session_id: sid,
|
||||||
|
};
|
||||||
|
if let Err(e) = hooks.run(&event).await {
|
||||||
|
tracing::warn!("OnSessionEnd hook error: {}", e);
|
||||||
|
}
|
||||||
|
});
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Remove sessions
|
// Remove sessions
|
||||||
let count = {
|
let count = {
|
||||||
let mut sessions = self.sessions.write().await;
|
let mut sessions = self.sessions.write().await;
|
||||||
|
|||||||
+136
-16
@@ -43,6 +43,10 @@ impl Checkpoint {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Manager for undo/redo functionality.
|
/// Manager for undo/redo functionality.
|
||||||
|
///
|
||||||
|
/// Each undo/redo operation pops from one stack and pushes the current state
|
||||||
|
/// onto the other, so `undo_count() + redo_count()` stays constant across
|
||||||
|
/// undo/redo cycles (only `checkpoint()` and `clear()` change the total).
|
||||||
pub struct UndoManager {
|
pub struct UndoManager {
|
||||||
/// Stack of past checkpoints (for undo).
|
/// Stack of past checkpoints (for undo).
|
||||||
undo_stack: VecDeque<Checkpoint>,
|
undo_stack: VecDeque<Checkpoint>,
|
||||||
@@ -68,6 +72,14 @@ impl UndoManager {
|
|||||||
self
|
self
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Push a checkpoint onto the undo stack, trimming oldest entries if over limit.
|
||||||
|
fn push_undo(&mut self, checkpoint: Checkpoint) {
|
||||||
|
self.undo_stack.push_back(checkpoint);
|
||||||
|
while self.undo_stack.len() > self.max_checkpoints {
|
||||||
|
self.undo_stack.pop_front();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
/// Create a checkpoint at the current state.
|
/// Create a checkpoint at the current state.
|
||||||
///
|
///
|
||||||
/// This clears the redo stack since we're creating a new history branch.
|
/// This clears the redo stack since we're creating a new history branch.
|
||||||
@@ -80,24 +92,23 @@ impl UndoManager {
|
|||||||
// Clear redo stack (new branch of history)
|
// Clear redo stack (new branch of history)
|
||||||
self.redo_stack.clear();
|
self.redo_stack.clear();
|
||||||
|
|
||||||
// Create and push checkpoint
|
|
||||||
let checkpoint = Checkpoint::new(turn_number, messages, description);
|
let checkpoint = Checkpoint::new(turn_number, messages, description);
|
||||||
self.undo_stack.push_back(checkpoint);
|
self.push_undo(checkpoint);
|
||||||
|
|
||||||
// Trim if over limit
|
|
||||||
while self.undo_stack.len() > self.max_checkpoints {
|
|
||||||
self.undo_stack.pop_front();
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Undo: pop the last checkpoint and return it.
|
/// Undo: pop the last checkpoint and return it.
|
||||||
///
|
///
|
||||||
/// The current state should be saved to redo stack before calling this.
|
/// Saves the current state to the redo stack and pops the most recent
|
||||||
|
/// checkpoint from the undo stack so that repeated undos walk backwards
|
||||||
|
/// through history.
|
||||||
|
///
|
||||||
|
/// Takes ownership of `current_messages`; callers must clone first if
|
||||||
|
/// they need to retain a copy.
|
||||||
pub fn undo(
|
pub fn undo(
|
||||||
&mut self,
|
&mut self,
|
||||||
current_turn: usize,
|
current_turn: usize,
|
||||||
current_messages: Vec<ChatMessage>,
|
current_messages: Vec<ChatMessage>,
|
||||||
) -> Option<&Checkpoint> {
|
) -> Option<Checkpoint> {
|
||||||
if self.undo_stack.is_empty() {
|
if self.undo_stack.is_empty() {
|
||||||
return None;
|
return None;
|
||||||
}
|
}
|
||||||
@@ -110,9 +121,8 @@ impl UndoManager {
|
|||||||
);
|
);
|
||||||
self.redo_stack.push(current);
|
self.redo_stack.push(current);
|
||||||
|
|
||||||
// Return the most recent checkpoint without removing it
|
// Pop and return the most recent checkpoint
|
||||||
// (we keep it so multiple undos can work)
|
self.undo_stack.pop_back()
|
||||||
self.undo_stack.back()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Pop the last checkpoint from the undo stack.
|
/// Pop the last checkpoint from the undo stack.
|
||||||
@@ -121,7 +131,29 @@ impl UndoManager {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Redo: restore a previously undone state.
|
/// Redo: restore a previously undone state.
|
||||||
pub fn redo(&mut self) -> Option<Checkpoint> {
|
///
|
||||||
|
/// Saves the current state to the undo stack and pops the most recent
|
||||||
|
/// checkpoint from the redo stack.
|
||||||
|
///
|
||||||
|
/// Takes ownership of `current_messages`; callers must clone first if
|
||||||
|
/// they need to retain a copy.
|
||||||
|
pub fn redo(
|
||||||
|
&mut self,
|
||||||
|
current_turn: usize,
|
||||||
|
current_messages: Vec<ChatMessage>,
|
||||||
|
) -> Option<Checkpoint> {
|
||||||
|
if self.redo_stack.is_empty() {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Save current state to undo stack
|
||||||
|
let current = Checkpoint::new(
|
||||||
|
current_turn,
|
||||||
|
current_messages,
|
||||||
|
format!("Turn {}", current_turn),
|
||||||
|
);
|
||||||
|
self.push_undo(current);
|
||||||
|
|
||||||
self.redo_stack.pop()
|
self.redo_stack.pop()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -214,14 +246,16 @@ mod tests {
|
|||||||
assert!(manager.can_undo());
|
assert!(manager.can_undo());
|
||||||
assert!(!manager.can_redo());
|
assert!(!manager.can_redo());
|
||||||
|
|
||||||
// Undo
|
// Undo - returns owned Checkpoint now
|
||||||
let current = vec![ChatMessage::user("Hello"), ChatMessage::assistant("Hi")];
|
let current = vec![ChatMessage::user("Hello"), ChatMessage::assistant("Hi")];
|
||||||
let checkpoint = manager.undo(2, current);
|
let checkpoint = manager.undo(2, current);
|
||||||
assert!(checkpoint.is_some());
|
assert!(checkpoint.is_some());
|
||||||
|
let checkpoint = checkpoint.unwrap();
|
||||||
|
assert_eq!(checkpoint.turn_number, 1);
|
||||||
assert!(manager.can_redo());
|
assert!(manager.can_redo());
|
||||||
|
|
||||||
// Redo
|
// Redo - now requires current state parameters
|
||||||
let restored = manager.redo();
|
let restored = manager.redo(checkpoint.turn_number, checkpoint.messages);
|
||||||
assert!(restored.is_some());
|
assert!(restored.is_some());
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -249,4 +283,90 @@ mod tests {
|
|||||||
assert!(restored.is_some());
|
assert!(restored.is_some());
|
||||||
assert_eq!(manager.undo_count(), 0);
|
assert_eq!(manager.undo_count(), 0);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_repeated_undo_advances_through_stack() {
|
||||||
|
let mut manager = UndoManager::new();
|
||||||
|
|
||||||
|
// Create 3 checkpoints at turns 0, 1, 2
|
||||||
|
manager.checkpoint(0, vec![], "Turn 0");
|
||||||
|
manager.checkpoint(1, vec![ChatMessage::user("msg1")], "Turn 1");
|
||||||
|
manager.checkpoint(2, vec![ChatMessage::user("msg2")], "Turn 2");
|
||||||
|
assert_eq!(manager.undo_count(), 3);
|
||||||
|
|
||||||
|
// First undo: should return turn 2 checkpoint, stack shrinks to 2
|
||||||
|
let cp1 = manager
|
||||||
|
.undo(3, vec![ChatMessage::user("msg3")])
|
||||||
|
.expect("first undo should succeed");
|
||||||
|
assert_eq!(cp1.turn_number, 2);
|
||||||
|
assert_eq!(manager.undo_count(), 2);
|
||||||
|
|
||||||
|
// Second undo: should return turn 1 checkpoint (different!), stack shrinks to 1
|
||||||
|
let cp2 = manager
|
||||||
|
.undo(cp1.turn_number, cp1.messages)
|
||||||
|
.expect("second undo should succeed");
|
||||||
|
assert_eq!(cp2.turn_number, 1);
|
||||||
|
assert_eq!(manager.undo_count(), 1);
|
||||||
|
|
||||||
|
// Verify we walked backwards through distinct checkpoints
|
||||||
|
assert_ne!(cp1.turn_number, cp2.turn_number);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_undo_redo_cycle_preserves_state() {
|
||||||
|
let mut manager = UndoManager::new();
|
||||||
|
|
||||||
|
let msgs_t0: Vec<ChatMessage> = vec![];
|
||||||
|
let msgs_t1 = vec![ChatMessage::user("hello")];
|
||||||
|
let msgs_t2 = vec![ChatMessage::user("hello"), ChatMessage::assistant("hi")];
|
||||||
|
|
||||||
|
manager.checkpoint(0, msgs_t0, "Turn 0");
|
||||||
|
manager.checkpoint(1, msgs_t1, "Turn 1");
|
||||||
|
|
||||||
|
// Undo from turn 2 -> get turn 1 checkpoint
|
||||||
|
let cp_undo1 = manager
|
||||||
|
.undo(2, msgs_t2.clone())
|
||||||
|
.expect("undo should succeed");
|
||||||
|
assert_eq!(cp_undo1.turn_number, 1);
|
||||||
|
|
||||||
|
// Redo from turn 1 -> get turn 2 state back
|
||||||
|
let cp_redo = manager
|
||||||
|
.redo(cp_undo1.turn_number, cp_undo1.messages)
|
||||||
|
.expect("redo should succeed");
|
||||||
|
assert_eq!(cp_redo.turn_number, 2);
|
||||||
|
assert_eq!(cp_redo.messages.len(), 2);
|
||||||
|
|
||||||
|
// Undo again from turn 2 -> should go back to turn 1 again
|
||||||
|
let cp_undo2 = manager
|
||||||
|
.undo(cp_redo.turn_number, cp_redo.messages)
|
||||||
|
.expect("second undo should succeed");
|
||||||
|
assert_eq!(cp_undo2.turn_number, 1);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_undo_redo_stack_sizes_consistent() {
|
||||||
|
let mut manager = UndoManager::new();
|
||||||
|
|
||||||
|
manager.checkpoint(0, vec![], "Turn 0");
|
||||||
|
manager.checkpoint(1, vec![ChatMessage::user("a")], "Turn 1");
|
||||||
|
manager.checkpoint(2, vec![ChatMessage::user("b")], "Turn 2");
|
||||||
|
|
||||||
|
// Start: undo=3, redo=0, total=3
|
||||||
|
let total = manager.undo_count() + manager.redo_count();
|
||||||
|
assert_eq!(total, 3);
|
||||||
|
|
||||||
|
// After undo: total should still be 3 (one moved from undo to redo,
|
||||||
|
// plus the current state pushed to redo)
|
||||||
|
// Actually: undo pops one (3->2), pushes current to redo (0->1), total=3
|
||||||
|
let cp = manager.undo(3, vec![]).unwrap();
|
||||||
|
assert_eq!(manager.undo_count() + manager.redo_count(), 3);
|
||||||
|
|
||||||
|
// After redo: redo pops one (1->0), pushes current to undo (2->3), total=3
|
||||||
|
let cp2 = manager.redo(cp.turn_number, cp.messages).unwrap();
|
||||||
|
assert_eq!(manager.undo_count() + manager.redo_count(), 3);
|
||||||
|
|
||||||
|
// After another undo: same invariant
|
||||||
|
let _cp3 = manager.undo(cp2.turn_number, cp2.messages).unwrap();
|
||||||
|
assert_eq!(manager.undo_count() + manager.redo_count(), 3);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+56
-37
@@ -12,6 +12,7 @@ use crate::agent::task::TaskOutput;
|
|||||||
use crate::context::{ContextManager, JobState};
|
use crate::context::{ContextManager, JobState};
|
||||||
use crate::db::Database;
|
use crate::db::Database;
|
||||||
use crate::error::Error;
|
use crate::error::Error;
|
||||||
|
use crate::hooks::HookRegistry;
|
||||||
use crate::llm::{
|
use crate::llm::{
|
||||||
ActionPlan, ChatMessage, LlmProvider, Reasoning, ReasoningContext, RespondResult, ToolSelection,
|
ActionPlan, ChatMessage, LlmProvider, Reasoning, ReasoningContext, RespondResult, ToolSelection,
|
||||||
};
|
};
|
||||||
@@ -29,6 +30,7 @@ pub struct WorkerDeps {
|
|||||||
pub safety: Arc<SafetyLayer>,
|
pub safety: Arc<SafetyLayer>,
|
||||||
pub tools: Arc<ToolRegistry>,
|
pub tools: Arc<ToolRegistry>,
|
||||||
pub store: Option<Arc<dyn Database>>,
|
pub store: Option<Arc<dyn Database>>,
|
||||||
|
pub hooks: Arc<HookRegistry>,
|
||||||
pub timeout: Duration,
|
pub timeout: Duration,
|
||||||
pub use_planning: bool,
|
pub use_planning: bool,
|
||||||
}
|
}
|
||||||
@@ -352,23 +354,11 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
|||||||
.map(|selection| {
|
.map(|selection| {
|
||||||
let tool_name = selection.tool_name.clone();
|
let tool_name = selection.tool_name.clone();
|
||||||
let params = selection.parameters.clone();
|
let params = selection.parameters.clone();
|
||||||
let tools = self.tools().clone();
|
let deps = self.deps.clone();
|
||||||
let context_manager = self.context_manager().clone();
|
|
||||||
let safety = self.safety().clone();
|
|
||||||
let job_id = self.job_id;
|
let job_id = self.job_id;
|
||||||
let store = self.deps.store.clone();
|
|
||||||
|
|
||||||
async move {
|
async move {
|
||||||
let result = Self::execute_tool_inner(
|
let result = Self::execute_tool_inner(&deps, job_id, &tool_name, ¶ms).await;
|
||||||
tools,
|
|
||||||
context_manager,
|
|
||||||
safety,
|
|
||||||
store,
|
|
||||||
job_id,
|
|
||||||
&tool_name,
|
|
||||||
¶ms,
|
|
||||||
)
|
|
||||||
.await;
|
|
||||||
ToolExecResult { result }
|
ToolExecResult { result }
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
@@ -379,15 +369,13 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
|||||||
|
|
||||||
/// Inner tool execution logic that can be called from both single and parallel paths.
|
/// Inner tool execution logic that can be called from both single and parallel paths.
|
||||||
async fn execute_tool_inner(
|
async fn execute_tool_inner(
|
||||||
tools: Arc<ToolRegistry>,
|
deps: &WorkerDeps,
|
||||||
context_manager: Arc<ContextManager>,
|
|
||||||
safety: Arc<SafetyLayer>,
|
|
||||||
store: Option<Arc<dyn Database>>,
|
|
||||||
job_id: Uuid,
|
job_id: Uuid,
|
||||||
tool_name: &str,
|
tool_name: &str,
|
||||||
params: &serde_json::Value,
|
params: &serde_json::Value,
|
||||||
) -> Result<String, Error> {
|
) -> Result<String, Error> {
|
||||||
let tool = tools
|
let tool =
|
||||||
|
deps.tools
|
||||||
.get(tool_name)
|
.get(tool_name)
|
||||||
.await
|
.await
|
||||||
.ok_or_else(|| crate::error::ToolError::NotFound {
|
.ok_or_else(|| crate::error::ToolError::NotFound {
|
||||||
@@ -402,8 +390,46 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
|||||||
.into());
|
.into());
|
||||||
}
|
}
|
||||||
|
|
||||||
// Get job context for the tool
|
// Fetch job context early so we have the real user_id for hooks
|
||||||
let job_ctx = context_manager.get_context(job_id).await?;
|
let job_ctx = deps.context_manager.get_context(job_id).await?;
|
||||||
|
|
||||||
|
// Run BeforeToolCall hook
|
||||||
|
let params = {
|
||||||
|
use crate::hooks::{HookError, HookEvent, HookOutcome};
|
||||||
|
let event = HookEvent::ToolCall {
|
||||||
|
tool_name: tool_name.to_string(),
|
||||||
|
parameters: params.clone(),
|
||||||
|
user_id: job_ctx.user_id.clone(),
|
||||||
|
context: format!("job:{}", job_id),
|
||||||
|
};
|
||||||
|
match deps.hooks.run(&event).await {
|
||||||
|
Err(HookError::Rejected { reason }) => {
|
||||||
|
return Err(crate::error::ToolError::ExecutionFailed {
|
||||||
|
name: tool_name.to_string(),
|
||||||
|
reason: format!("Blocked by hook: {}", reason),
|
||||||
|
}
|
||||||
|
.into());
|
||||||
|
}
|
||||||
|
Err(err) => {
|
||||||
|
return Err(crate::error::ToolError::ExecutionFailed {
|
||||||
|
name: tool_name.to_string(),
|
||||||
|
reason: format!("Blocked by hook failure mode: {}", err),
|
||||||
|
}
|
||||||
|
.into());
|
||||||
|
}
|
||||||
|
Ok(HookOutcome::Continue {
|
||||||
|
modified: Some(new_params),
|
||||||
|
}) => serde_json::from_str(&new_params).unwrap_or_else(|e| {
|
||||||
|
tracing::warn!(
|
||||||
|
tool = %tool_name,
|
||||||
|
"Hook returned non-JSON modification for ToolCall, ignoring: {}",
|
||||||
|
e
|
||||||
|
);
|
||||||
|
params.clone()
|
||||||
|
}),
|
||||||
|
_ => params.clone(),
|
||||||
|
}
|
||||||
|
};
|
||||||
if job_ctx.state == JobState::Cancelled {
|
if job_ctx.state == JobState::Cancelled {
|
||||||
return Err(crate::error::ToolError::ExecutionFailed {
|
return Err(crate::error::ToolError::ExecutionFailed {
|
||||||
name: tool_name.to_string(),
|
name: tool_name.to_string(),
|
||||||
@@ -413,7 +439,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Validate tool parameters
|
// Validate tool parameters
|
||||||
let validation = safety.validator().validate_tool_params(params);
|
let validation = deps.safety.validator().validate_tool_params(¶ms);
|
||||||
if !validation.is_valid {
|
if !validation.is_valid {
|
||||||
let details = validation
|
let details = validation
|
||||||
.errors
|
.errors
|
||||||
@@ -478,8 +504,8 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
|||||||
Ok(Ok(output)) => {
|
Ok(Ok(output)) => {
|
||||||
let output_str = serde_json::to_string_pretty(&output.result)
|
let output_str = serde_json::to_string_pretty(&output.result)
|
||||||
.ok()
|
.ok()
|
||||||
.map(|s| safety.sanitize_tool_output(tool_name, &s).content);
|
.map(|s| deps.safety.sanitize_tool_output(tool_name, &s).content);
|
||||||
context_manager
|
deps.context_manager
|
||||||
.update_memory(job_id, |mem| {
|
.update_memory(job_id, |mem| {
|
||||||
let rec = mem.create_action(tool_name, params.clone()).succeed(
|
let rec = mem.create_action(tool_name, params.clone()).succeed(
|
||||||
output_str.clone(),
|
output_str.clone(),
|
||||||
@@ -492,7 +518,8 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
|||||||
.await
|
.await
|
||||||
.ok()
|
.ok()
|
||||||
}
|
}
|
||||||
Ok(Err(e)) => context_manager
|
Ok(Err(e)) => deps
|
||||||
|
.context_manager
|
||||||
.update_memory(job_id, |mem| {
|
.update_memory(job_id, |mem| {
|
||||||
let rec = mem
|
let rec = mem
|
||||||
.create_action(tool_name, params.clone())
|
.create_action(tool_name, params.clone())
|
||||||
@@ -502,7 +529,8 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
|||||||
})
|
})
|
||||||
.await
|
.await
|
||||||
.ok(),
|
.ok(),
|
||||||
Err(_) => context_manager
|
Err(_) => deps
|
||||||
|
.context_manager
|
||||||
.update_memory(job_id, |mem| {
|
.update_memory(job_id, |mem| {
|
||||||
let rec = mem
|
let rec = mem
|
||||||
.create_action(tool_name, params.clone())
|
.create_action(tool_name, params.clone())
|
||||||
@@ -515,7 +543,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
|||||||
};
|
};
|
||||||
|
|
||||||
// Persist action to database (fire-and-forget)
|
// Persist action to database (fire-and-forget)
|
||||||
if let (Some(action), Some(store)) = (action, store) {
|
if let (Some(action), Some(store)) = (action, deps.store.clone()) {
|
||||||
tokio::spawn(async move {
|
tokio::spawn(async move {
|
||||||
if let Err(e) = store.save_action(job_id, &action).await {
|
if let Err(e) = store.save_action(job_id, &action).await {
|
||||||
tracing::warn!("Failed to persist action for job {}: {}", job_id, e);
|
tracing::warn!("Failed to persist action for job {}: {}", job_id, e);
|
||||||
@@ -701,16 +729,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
|||||||
tool_name: &str,
|
tool_name: &str,
|
||||||
params: &serde_json::Value,
|
params: &serde_json::Value,
|
||||||
) -> Result<String, Error> {
|
) -> Result<String, Error> {
|
||||||
Self::execute_tool_inner(
|
Self::execute_tool_inner(&self.deps, self.job_id, tool_name, params).await
|
||||||
self.tools().clone(),
|
|
||||||
self.context_manager().clone(),
|
|
||||||
self.safety().clone(),
|
|
||||||
self.deps.store.clone(),
|
|
||||||
self.job_id,
|
|
||||||
tool_name,
|
|
||||||
params,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn mark_completed(&self) -> Result<(), Error> {
|
async fn mark_completed(&self) -> Result<(), Error> {
|
||||||
|
|||||||
@@ -0,0 +1,208 @@
|
|||||||
|
//! Boot screen displayed after all initialization completes.
|
||||||
|
//!
|
||||||
|
//! Shows a polished ANSI-styled status panel summarizing the agent's runtime
|
||||||
|
//! state: model, database, tool count, enabled features, active channels,
|
||||||
|
//! and the gateway URL.
|
||||||
|
|
||||||
|
/// All displayable fields for the boot screen.
|
||||||
|
pub struct BootInfo {
|
||||||
|
pub version: String,
|
||||||
|
pub agent_name: String,
|
||||||
|
pub llm_backend: String,
|
||||||
|
pub llm_model: String,
|
||||||
|
pub cheap_model: Option<String>,
|
||||||
|
pub db_backend: String,
|
||||||
|
pub db_connected: bool,
|
||||||
|
pub tool_count: usize,
|
||||||
|
pub gateway_url: Option<String>,
|
||||||
|
pub embeddings_enabled: bool,
|
||||||
|
pub embeddings_provider: Option<String>,
|
||||||
|
pub heartbeat_enabled: bool,
|
||||||
|
pub heartbeat_interval_secs: u64,
|
||||||
|
pub sandbox_enabled: bool,
|
||||||
|
pub claude_code_enabled: bool,
|
||||||
|
pub routines_enabled: bool,
|
||||||
|
pub channels: Vec<String>,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Print the boot screen to stdout.
|
||||||
|
pub fn print_boot_screen(info: &BootInfo) {
|
||||||
|
// ANSI codes matching existing REPL palette
|
||||||
|
let bold = "\x1b[1m";
|
||||||
|
let cyan = "\x1b[36m";
|
||||||
|
let dim = "\x1b[90m";
|
||||||
|
let yellow_underline = "\x1b[33;4m";
|
||||||
|
let reset = "\x1b[0m";
|
||||||
|
|
||||||
|
let border = format!(" {dim}{}{reset}", "\u{2576}".repeat(58));
|
||||||
|
|
||||||
|
println!();
|
||||||
|
println!("{border}");
|
||||||
|
println!();
|
||||||
|
println!(" {bold}{}{reset} v{}", info.agent_name, info.version);
|
||||||
|
println!();
|
||||||
|
|
||||||
|
// Model line
|
||||||
|
let model_display = if let Some(ref cheap) = info.cheap_model {
|
||||||
|
format!(
|
||||||
|
"{cyan}{}{reset} {dim}cheap{reset} {cyan}{}{reset}",
|
||||||
|
info.llm_model, cheap
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
format!("{cyan}{}{reset}", info.llm_model)
|
||||||
|
};
|
||||||
|
println!(
|
||||||
|
" {dim}model{reset} {model_display} {dim}via {}{reset}",
|
||||||
|
info.llm_backend
|
||||||
|
);
|
||||||
|
|
||||||
|
// Database line
|
||||||
|
let db_status = if info.db_connected {
|
||||||
|
"connected"
|
||||||
|
} else {
|
||||||
|
"none"
|
||||||
|
};
|
||||||
|
println!(
|
||||||
|
" {dim}database{reset} {cyan}{}{reset} {dim}({db_status}){reset}",
|
||||||
|
info.db_backend
|
||||||
|
);
|
||||||
|
|
||||||
|
// Tools line
|
||||||
|
println!(
|
||||||
|
" {dim}tools{reset} {cyan}{}{reset} {dim}registered{reset}",
|
||||||
|
info.tool_count
|
||||||
|
);
|
||||||
|
|
||||||
|
// Features line
|
||||||
|
let mut features = Vec::new();
|
||||||
|
if info.embeddings_enabled {
|
||||||
|
if let Some(ref provider) = info.embeddings_provider {
|
||||||
|
features.push(format!("embeddings ({provider})"));
|
||||||
|
} else {
|
||||||
|
features.push("embeddings".to_string());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if info.heartbeat_enabled {
|
||||||
|
let mins = info.heartbeat_interval_secs / 60;
|
||||||
|
features.push(format!("heartbeat ({mins}m)"));
|
||||||
|
}
|
||||||
|
if info.sandbox_enabled {
|
||||||
|
features.push("sandbox".to_string());
|
||||||
|
}
|
||||||
|
if info.claude_code_enabled {
|
||||||
|
features.push("claude-code".to_string());
|
||||||
|
}
|
||||||
|
if info.routines_enabled {
|
||||||
|
features.push("routines".to_string());
|
||||||
|
}
|
||||||
|
if !features.is_empty() {
|
||||||
|
println!(
|
||||||
|
" {dim}features{reset} {cyan}{}{reset}",
|
||||||
|
features.join(" ")
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Channels line
|
||||||
|
if !info.channels.is_empty() {
|
||||||
|
println!(
|
||||||
|
" {dim}channels{reset} {cyan}{}{reset}",
|
||||||
|
info.channels.join(" ")
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Gateway URL (highlighted)
|
||||||
|
if let Some(ref url) = info.gateway_url {
|
||||||
|
println!();
|
||||||
|
println!(" {dim}gateway{reset} {yellow_underline}{url}{reset}");
|
||||||
|
}
|
||||||
|
|
||||||
|
println!();
|
||||||
|
println!("{border}");
|
||||||
|
println!();
|
||||||
|
println!(" /help for commands, /quit to exit");
|
||||||
|
println!();
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_print_boot_screen_full() {
|
||||||
|
let info = BootInfo {
|
||||||
|
version: "0.2.0".to_string(),
|
||||||
|
agent_name: "ironclaw".to_string(),
|
||||||
|
llm_backend: "nearai".to_string(),
|
||||||
|
llm_model: "claude-3-5-sonnet-20241022".to_string(),
|
||||||
|
cheap_model: Some("gpt-4o-mini".to_string()),
|
||||||
|
db_backend: "libsql".to_string(),
|
||||||
|
db_connected: true,
|
||||||
|
tool_count: 24,
|
||||||
|
gateway_url: Some("http://127.0.0.1:3001/?token=abc123".to_string()),
|
||||||
|
embeddings_enabled: true,
|
||||||
|
embeddings_provider: Some("openai".to_string()),
|
||||||
|
heartbeat_enabled: true,
|
||||||
|
heartbeat_interval_secs: 1800,
|
||||||
|
sandbox_enabled: true,
|
||||||
|
claude_code_enabled: false,
|
||||||
|
routines_enabled: true,
|
||||||
|
channels: vec![
|
||||||
|
"repl".to_string(),
|
||||||
|
"gateway".to_string(),
|
||||||
|
"telegram".to_string(),
|
||||||
|
],
|
||||||
|
};
|
||||||
|
// Should not panic
|
||||||
|
print_boot_screen(&info);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_print_boot_screen_minimal() {
|
||||||
|
let info = BootInfo {
|
||||||
|
version: "0.2.0".to_string(),
|
||||||
|
agent_name: "ironclaw".to_string(),
|
||||||
|
llm_backend: "nearai".to_string(),
|
||||||
|
llm_model: "gpt-4o".to_string(),
|
||||||
|
cheap_model: None,
|
||||||
|
db_backend: "none".to_string(),
|
||||||
|
db_connected: false,
|
||||||
|
tool_count: 5,
|
||||||
|
gateway_url: None,
|
||||||
|
embeddings_enabled: false,
|
||||||
|
embeddings_provider: None,
|
||||||
|
heartbeat_enabled: false,
|
||||||
|
heartbeat_interval_secs: 0,
|
||||||
|
sandbox_enabled: false,
|
||||||
|
claude_code_enabled: false,
|
||||||
|
routines_enabled: false,
|
||||||
|
channels: vec![],
|
||||||
|
};
|
||||||
|
// Should not panic
|
||||||
|
print_boot_screen(&info);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_print_boot_screen_no_features() {
|
||||||
|
let info = BootInfo {
|
||||||
|
version: "0.1.0".to_string(),
|
||||||
|
agent_name: "test".to_string(),
|
||||||
|
llm_backend: "openai".to_string(),
|
||||||
|
llm_model: "gpt-4o".to_string(),
|
||||||
|
cheap_model: None,
|
||||||
|
db_backend: "postgres".to_string(),
|
||||||
|
db_connected: true,
|
||||||
|
tool_count: 10,
|
||||||
|
gateway_url: None,
|
||||||
|
embeddings_enabled: false,
|
||||||
|
embeddings_provider: None,
|
||||||
|
heartbeat_enabled: false,
|
||||||
|
heartbeat_interval_secs: 0,
|
||||||
|
sandbox_enabled: false,
|
||||||
|
claude_code_enabled: false,
|
||||||
|
routines_enabled: false,
|
||||||
|
channels: vec!["repl".to_string()],
|
||||||
|
};
|
||||||
|
// Should not panic
|
||||||
|
print_boot_screen(&info);
|
||||||
|
}
|
||||||
|
}
|
||||||
+81
-5
@@ -81,17 +81,34 @@ fn migrate_bootstrap_json_to_env(env_path: &std::path::Path) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Write `DATABASE_URL` to `~/.ironclaw/.env`.
|
/// Write database bootstrap vars to `~/.ironclaw/.env`.
|
||||||
|
///
|
||||||
|
/// These settings form the chicken-and-egg layer: they must be available
|
||||||
|
/// from the filesystem (env vars) BEFORE any database connection, because
|
||||||
|
/// they determine which database to connect to. Everything else is stored
|
||||||
|
/// in the database itself.
|
||||||
///
|
///
|
||||||
/// Creates the parent directory if it doesn't exist.
|
/// Creates the parent directory if it doesn't exist.
|
||||||
/// The value is double-quoted so that `#` (common in URL-encoded passwords)
|
/// Values are double-quoted so that `#` (common in URL-encoded passwords)
|
||||||
/// and other shell-special characters are preserved by dotenvy.
|
/// and other shell-special characters are preserved by dotenvy.
|
||||||
pub fn save_database_url(url: &str) -> std::io::Result<()> {
|
pub fn save_bootstrap_env(vars: &[(&str, &str)]) -> std::io::Result<()> {
|
||||||
let path = ironclaw_env_path();
|
let path = ironclaw_env_path();
|
||||||
if let Some(parent) = path.parent() {
|
if let Some(parent) = path.parent() {
|
||||||
std::fs::create_dir_all(parent)?;
|
std::fs::create_dir_all(parent)?;
|
||||||
}
|
}
|
||||||
std::fs::write(&path, format!("DATABASE_URL=\"{}\"\n", url))
|
let mut content = String::new();
|
||||||
|
for (key, value) in vars {
|
||||||
|
content.push_str(&format!("{}=\"{}\"\n", key, value));
|
||||||
|
}
|
||||||
|
std::fs::write(&path, content)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Write `DATABASE_URL` to `~/.ironclaw/.env`.
|
||||||
|
///
|
||||||
|
/// Convenience wrapper around `save_bootstrap_env` for single-value migration
|
||||||
|
/// paths. Prefer `save_bootstrap_env` for new code.
|
||||||
|
pub fn save_database_url(url: &str) -> std::io::Result<()> {
|
||||||
|
save_bootstrap_env(&[("DATABASE_URL", url)])
|
||||||
}
|
}
|
||||||
|
|
||||||
/// One-time migration of legacy `~/.ironclaw/settings.json` into the database.
|
/// One-time migration of legacy `~/.ironclaw/settings.json` into the database.
|
||||||
@@ -184,7 +201,7 @@ pub async fn migrate_disk_to_db(
|
|||||||
Ok(content) => match serde_json::from_str::<serde_json::Value>(&content) {
|
Ok(content) => match serde_json::from_str::<serde_json::Value>(&content) {
|
||||||
Ok(value) => {
|
Ok(value) => {
|
||||||
store
|
store
|
||||||
.set_setting(user_id, "nearai.session", &value)
|
.set_setting(user_id, "nearai.session_token", &value)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| {
|
.map_err(|e| {
|
||||||
MigrationError::Database(format!(
|
MigrationError::Database(format!(
|
||||||
@@ -385,4 +402,63 @@ mod tests {
|
|||||||
// Nothing should happen
|
// Nothing should happen
|
||||||
assert!(!env_path.exists());
|
assert!(!env_path.exists());
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_save_bootstrap_env_multiple_vars() {
|
||||||
|
let dir = tempdir().unwrap();
|
||||||
|
let env_path = dir.path().join("nested").join(".env");
|
||||||
|
|
||||||
|
std::fs::create_dir_all(env_path.parent().unwrap()).unwrap();
|
||||||
|
|
||||||
|
let vars = [
|
||||||
|
("DATABASE_BACKEND", "libsql"),
|
||||||
|
("LIBSQL_PATH", "/home/user/.ironclaw/ironclaw.db"),
|
||||||
|
];
|
||||||
|
|
||||||
|
// Write manually to the temp path (save_bootstrap_env uses the global path)
|
||||||
|
let mut content = String::new();
|
||||||
|
for (key, value) in &vars {
|
||||||
|
content.push_str(&format!("{}=\"{}\"\n", key, value));
|
||||||
|
}
|
||||||
|
std::fs::write(&env_path, &content).unwrap();
|
||||||
|
|
||||||
|
// Verify dotenvy can parse all entries
|
||||||
|
let parsed: Vec<(String, String)> = dotenvy::from_path_iter(&env_path)
|
||||||
|
.unwrap()
|
||||||
|
.filter_map(|r| r.ok())
|
||||||
|
.collect();
|
||||||
|
assert_eq!(parsed.len(), 2);
|
||||||
|
assert_eq!(
|
||||||
|
parsed[0],
|
||||||
|
("DATABASE_BACKEND".to_string(), "libsql".to_string())
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
parsed[1],
|
||||||
|
(
|
||||||
|
"LIBSQL_PATH".to_string(),
|
||||||
|
"/home/user/.ironclaw/ironclaw.db".to_string()
|
||||||
|
)
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_save_bootstrap_env_overwrites_previous() {
|
||||||
|
let dir = tempdir().unwrap();
|
||||||
|
let env_path = dir.path().join(".env");
|
||||||
|
|
||||||
|
// Write initial content
|
||||||
|
std::fs::write(&env_path, "DATABASE_URL=\"postgres://old\"\n").unwrap();
|
||||||
|
|
||||||
|
// Overwrite with new vars (simulating save_bootstrap_env behavior)
|
||||||
|
let content = "DATABASE_BACKEND=\"libsql\"\nLIBSQL_PATH=\"/new/path.db\"\n";
|
||||||
|
std::fs::write(&env_path, content).unwrap();
|
||||||
|
|
||||||
|
let parsed: Vec<(String, String)> = dotenvy::from_path_iter(&env_path)
|
||||||
|
.unwrap()
|
||||||
|
.filter_map(|r| r.ok())
|
||||||
|
.collect();
|
||||||
|
// Old DATABASE_URL should be gone
|
||||||
|
assert_eq!(parsed.len(), 2);
|
||||||
|
assert!(parsed.iter().all(|(k, _)| k != "DATABASE_URL"));
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -184,6 +184,8 @@ pub struct ReplChannel {
|
|||||||
debug_mode: Arc<AtomicBool>,
|
debug_mode: Arc<AtomicBool>,
|
||||||
/// Whether we're currently streaming (chunks have been printed without a trailing newline).
|
/// Whether we're currently streaming (chunks have been printed without a trailing newline).
|
||||||
is_streaming: Arc<AtomicBool>,
|
is_streaming: Arc<AtomicBool>,
|
||||||
|
/// When true, the one-liner startup banner is suppressed (boot screen shown instead).
|
||||||
|
suppress_banner: Arc<AtomicBool>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl ReplChannel {
|
impl ReplChannel {
|
||||||
@@ -193,6 +195,7 @@ impl ReplChannel {
|
|||||||
single_message: None,
|
single_message: None,
|
||||||
debug_mode: Arc::new(AtomicBool::new(false)),
|
debug_mode: Arc::new(AtomicBool::new(false)),
|
||||||
is_streaming: Arc::new(AtomicBool::new(false)),
|
is_streaming: Arc::new(AtomicBool::new(false)),
|
||||||
|
suppress_banner: Arc::new(AtomicBool::new(false)),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -202,9 +205,15 @@ impl ReplChannel {
|
|||||||
single_message: Some(message),
|
single_message: Some(message),
|
||||||
debug_mode: Arc::new(AtomicBool::new(false)),
|
debug_mode: Arc::new(AtomicBool::new(false)),
|
||||||
is_streaming: Arc::new(AtomicBool::new(false)),
|
is_streaming: Arc::new(AtomicBool::new(false)),
|
||||||
|
suppress_banner: Arc::new(AtomicBool::new(false)),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Suppress the one-liner startup banner (boot screen will be shown instead).
|
||||||
|
pub fn suppress_banner(&self) {
|
||||||
|
self.suppress_banner.store(true, Ordering::Relaxed);
|
||||||
|
}
|
||||||
|
|
||||||
fn is_debug(&self) -> bool {
|
fn is_debug(&self) -> bool {
|
||||||
self.debug_mode.load(Ordering::Relaxed)
|
self.debug_mode.load(Ordering::Relaxed)
|
||||||
}
|
}
|
||||||
@@ -264,6 +273,7 @@ impl Channel for ReplChannel {
|
|||||||
let (tx, rx) = mpsc::channel(32);
|
let (tx, rx) = mpsc::channel(32);
|
||||||
let single_message = self.single_message.clone();
|
let single_message = self.single_message.clone();
|
||||||
let debug_mode = Arc::clone(&self.debug_mode);
|
let debug_mode = Arc::clone(&self.debug_mode);
|
||||||
|
let suppress_banner = Arc::clone(&self.suppress_banner);
|
||||||
|
|
||||||
std::thread::spawn(move || {
|
std::thread::spawn(move || {
|
||||||
// Single message mode: send it and return
|
// Single message mode: send it and return
|
||||||
@@ -298,8 +308,10 @@ impl Channel for ReplChannel {
|
|||||||
}
|
}
|
||||||
let _ = rl.load_history(&hist_path);
|
let _ = rl.load_history(&hist_path);
|
||||||
|
|
||||||
|
if !suppress_banner.load(Ordering::Relaxed) {
|
||||||
println!("\x1b[1mIronClaw\x1b[0m /help for commands, /quit to exit");
|
println!("\x1b[1mIronClaw\x1b[0m /help for commands, /quit to exit");
|
||||||
println!();
|
println!();
|
||||||
|
}
|
||||||
|
|
||||||
loop {
|
loop {
|
||||||
let prompt = if debug_mode.load(Ordering::Relaxed) {
|
let prompt = if debug_mode.load(Ordering::Relaxed) {
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ let loadingOlder = false;
|
|||||||
let jobEvents = new Map(); // job_id -> Array of events
|
let jobEvents = new Map(); // job_id -> Array of events
|
||||||
let jobListRefreshTimer = null;
|
let jobListRefreshTimer = null;
|
||||||
const JOB_EVENTS_CAP = 500;
|
const JOB_EVENTS_CAP = 500;
|
||||||
|
const MEMORY_SEARCH_QUERY_MAX_LENGTH = 100;
|
||||||
|
|
||||||
// --- Auth ---
|
// --- Auth ---
|
||||||
|
|
||||||
@@ -1001,9 +1002,12 @@ function buildBreadcrumb(path) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
function searchMemory(query) {
|
function searchMemory(query) {
|
||||||
|
const normalizedQuery = normalizeSearchQuery(query);
|
||||||
|
if (!normalizedQuery) return;
|
||||||
|
|
||||||
apiFetch('/api/memory/search', {
|
apiFetch('/api/memory/search', {
|
||||||
method: 'POST',
|
method: 'POST',
|
||||||
body: { query, limit: 20 },
|
body: { query: normalizedQuery, limit: 20 },
|
||||||
}).then((data) => {
|
}).then((data) => {
|
||||||
const tree = document.getElementById('memory-tree');
|
const tree = document.getElementById('memory-tree');
|
||||||
tree.innerHTML = '';
|
tree.innerHTML = '';
|
||||||
@@ -1014,18 +1018,23 @@ function searchMemory(query) {
|
|||||||
for (const result of data.results) {
|
for (const result of data.results) {
|
||||||
const item = document.createElement('div');
|
const item = document.createElement('div');
|
||||||
item.className = 'search-result';
|
item.className = 'search-result';
|
||||||
const snippet = snippetAround(result.content, query, 120);
|
const snippet = snippetAround(result.content, normalizedQuery, 120);
|
||||||
item.innerHTML = '<div class="path">' + escapeHtml(result.path) + '</div>'
|
item.innerHTML = '<div class="path">' + escapeHtml(result.path) + '</div>'
|
||||||
+ '<div class="snippet">' + highlightQuery(snippet, query) + '</div>';
|
+ '<div class="snippet">' + highlightQuery(snippet, normalizedQuery) + '</div>';
|
||||||
item.addEventListener('click', () => readMemoryFile(result.path));
|
item.addEventListener('click', () => readMemoryFile(result.path));
|
||||||
tree.appendChild(item);
|
tree.appendChild(item);
|
||||||
}
|
}
|
||||||
}).catch(() => {});
|
}).catch(() => {});
|
||||||
}
|
}
|
||||||
|
|
||||||
|
function normalizeSearchQuery(query) {
|
||||||
|
return (typeof query === 'string' ? query : '').slice(0, MEMORY_SEARCH_QUERY_MAX_LENGTH);
|
||||||
|
}
|
||||||
|
|
||||||
function snippetAround(text, query, len) {
|
function snippetAround(text, query, len) {
|
||||||
|
const normalizedQuery = normalizeSearchQuery(query);
|
||||||
const lower = text.toLowerCase();
|
const lower = text.toLowerCase();
|
||||||
const idx = lower.indexOf(query.toLowerCase());
|
const idx = lower.indexOf(normalizedQuery.toLowerCase());
|
||||||
if (idx < 0) return text.substring(0, len);
|
if (idx < 0) return text.substring(0, len);
|
||||||
const start = Math.max(0, idx - Math.floor(len / 2));
|
const start = Math.max(0, idx - Math.floor(len / 2));
|
||||||
const end = Math.min(text.length, start + len);
|
const end = Math.min(text.length, start + len);
|
||||||
@@ -1038,11 +1047,11 @@ function snippetAround(text, query, len) {
|
|||||||
function highlightQuery(text, query) {
|
function highlightQuery(text, query) {
|
||||||
if (!query) return escapeHtml(text);
|
if (!query) return escapeHtml(text);
|
||||||
const escaped = escapeHtml(text);
|
const escaped = escapeHtml(text);
|
||||||
const queryEscaped = query.replace(/[.*+?^${}()|[\]\\]/g, '\\$&');
|
const normalizedQuery = normalizeSearchQuery(query);
|
||||||
|
const queryEscaped = normalizedQuery.replace(/[.*+?^${}()|[\]\\]/g, '\\$&');
|
||||||
const re = new RegExp('(' + queryEscaped + ')', 'gi');
|
const re = new RegExp('(' + queryEscaped + ')', 'gi');
|
||||||
return escaped.replace(re, '<mark>$1</mark>');
|
return escaped.replace(re, '<mark>$1</mark>');
|
||||||
}
|
}
|
||||||
|
|
||||||
// --- Logs ---
|
// --- Logs ---
|
||||||
|
|
||||||
const LOG_MAX_ENTRIES = 2000;
|
const LOG_MAX_ENTRIES = 2000;
|
||||||
|
|||||||
@@ -5,7 +5,11 @@
|
|||||||
<meta name="viewport" content="width=device-width, initial-scale=1.0">
|
<meta name="viewport" content="width=device-width, initial-scale=1.0">
|
||||||
<title>IronClaw</title>
|
<title>IronClaw</title>
|
||||||
<link rel="stylesheet" href="/style.css">
|
<link rel="stylesheet" href="/style.css">
|
||||||
<script src="https://cdn.jsdelivr.net/npm/marked/marked.min.js"></script>
|
<script
|
||||||
|
src="https://cdn.jsdelivr.net/npm/[email protected]/lib/marked.umd.min.js"
|
||||||
|
integrity="sha384-pN9zSKOnTZwXRtYZAu0PBPEgR2B7DOC1aeLxQ33oJ0oy5iN1we6gm57xldM2irDG"
|
||||||
|
crossorigin="anonymous"
|
||||||
|
></script>
|
||||||
</head>
|
</head>
|
||||||
<body>
|
<body>
|
||||||
<!-- Auth Screen -->
|
<!-- Auth Screen -->
|
||||||
|
|||||||
@@ -80,14 +80,13 @@ pub enum OAuthCallbackError {
|
|||||||
|
|
||||||
/// Bind the OAuth callback listener on the fixed port.
|
/// Bind the OAuth callback listener on the fixed port.
|
||||||
///
|
///
|
||||||
/// Tries IPv6 loopback (`[::1]`) first so that `http://localhost:…` redirects
|
/// Binds to IPv4 `127.0.0.1` first because callback URLs use `127.0.0.1`
|
||||||
/// work on systems where `localhost` resolves to `::1`. Falls back to IPv4
|
/// explicitly (e.g., NEAR AI redirects to `http://127.0.0.1:9876/auth/callback`).
|
||||||
/// (`127.0.0.1`) only if IPv6 fails for a reason other than `AddrInUse`
|
/// Falls back to IPv6 `[::1]` only if IPv4 binding fails for a reason other
|
||||||
/// (e.g., IPv6 not supported on the host). If the port is already occupied
|
/// than `AddrInUse`. If the port is already occupied, fails immediately.
|
||||||
/// on IPv6, the port is occupied period, so we fail immediately.
|
|
||||||
pub async fn bind_callback_listener() -> Result<TcpListener, OAuthCallbackError> {
|
pub async fn bind_callback_listener() -> Result<TcpListener, OAuthCallbackError> {
|
||||||
let ipv6_addr = format!("[::1]:{}", OAUTH_CALLBACK_PORT);
|
let ipv4_addr = format!("127.0.0.1:{}", OAUTH_CALLBACK_PORT);
|
||||||
match TcpListener::bind(&ipv6_addr).await {
|
match TcpListener::bind(&ipv4_addr).await {
|
||||||
Ok(listener) => return Ok(listener),
|
Ok(listener) => return Ok(listener),
|
||||||
Err(e) if e.kind() == std::io::ErrorKind::AddrInUse => {
|
Err(e) if e.kind() == std::io::ErrorKind::AddrInUse => {
|
||||||
return Err(OAuthCallbackError::PortInUse(
|
return Err(OAuthCallbackError::PortInUse(
|
||||||
@@ -96,10 +95,10 @@ pub async fn bind_callback_listener() -> Result<TcpListener, OAuthCallbackError>
|
|||||||
));
|
));
|
||||||
}
|
}
|
||||||
Err(_) => {
|
Err(_) => {
|
||||||
// IPv6 not available on this host, fall back to IPv4
|
// IPv4 not available, fall back to IPv6
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
TcpListener::bind(format!("127.0.0.1:{}", OAUTH_CALLBACK_PORT))
|
TcpListener::bind(format!("[::1]:{}", OAUTH_CALLBACK_PORT))
|
||||||
.await
|
.await
|
||||||
.map_err(|e| {
|
.map_err(|e| {
|
||||||
if e.kind() == std::io::ErrorKind::AddrInUse {
|
if e.kind() == std::io::ErrorKind::AddrInUse {
|
||||||
|
|||||||
+32
-10
@@ -22,16 +22,37 @@ pub async fn run_status_command() -> anyhow::Result<()> {
|
|||||||
);
|
);
|
||||||
|
|
||||||
// Database
|
// Database
|
||||||
let db_url_set = std::env::var("DATABASE_URL").is_ok();
|
|
||||||
print!(" Database: ");
|
print!(" Database: ");
|
||||||
if db_url_set {
|
let db_backend = std::env::var("DATABASE_BACKEND")
|
||||||
|
.ok()
|
||||||
|
.unwrap_or_else(|| "postgres".to_string());
|
||||||
|
match db_backend.as_str() {
|
||||||
|
"libsql" | "turso" | "sqlite" => {
|
||||||
|
let path = std::env::var("LIBSQL_PATH")
|
||||||
|
.map(std::path::PathBuf::from)
|
||||||
|
.unwrap_or_else(|_| crate::config::default_libsql_path());
|
||||||
|
if path.exists() {
|
||||||
|
let turso = if std::env::var("LIBSQL_URL").is_ok() {
|
||||||
|
" + Turso sync"
|
||||||
|
} else {
|
||||||
|
""
|
||||||
|
};
|
||||||
|
println!("libSQL ({}{})", path.display(), turso);
|
||||||
|
} else {
|
||||||
|
println!("libSQL (file missing: {})", path.display());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
_ => {
|
||||||
|
if std::env::var("DATABASE_URL").is_ok() {
|
||||||
match check_database().await {
|
match check_database().await {
|
||||||
Ok(()) => println!("connected"),
|
Ok(()) => println!("connected (PostgreSQL)"),
|
||||||
Err(e) => println!("error ({})", e),
|
Err(e) => println!("error ({})", e),
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
println!("not configured");
|
println!("not configured");
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Session / Auth
|
// Session / Auth
|
||||||
print!(" Session: ");
|
print!(" Session: ");
|
||||||
@@ -42,16 +63,17 @@ pub async fn run_status_command() -> anyhow::Result<()> {
|
|||||||
println!("not found (run `ironclaw onboard`)");
|
println!("not found (run `ironclaw onboard`)");
|
||||||
}
|
}
|
||||||
|
|
||||||
// Secrets (auto-detect: env var or keychain)
|
// Secrets (auto-detect from env only; skip keychain probe to avoid
|
||||||
|
// triggering macOS system password dialogs on a simple status check)
|
||||||
print!(" Secrets: ");
|
print!(" Secrets: ");
|
||||||
let has_env_key = std::env::var("SECRETS_MASTER_KEY").is_ok();
|
if std::env::var("SECRETS_MASTER_KEY").is_ok() {
|
||||||
let has_keychain = crate::secrets::keychain::has_master_key().await;
|
|
||||||
if has_env_key {
|
|
||||||
println!("configured (env)");
|
println!("configured (env)");
|
||||||
} else if has_keychain {
|
|
||||||
println!("configured (keychain)");
|
|
||||||
} else {
|
} else {
|
||||||
println!("not configured");
|
// We don't probe the keychain here because get_generic_password()
|
||||||
|
// triggers macOS unlock+authorization dialogs, which is bad UX for
|
||||||
|
// a read-only status command. If onboarding completed with keychain
|
||||||
|
// storage, the key is there; we just can't cheaply verify it.
|
||||||
|
println!("env not set (keychain may be configured)");
|
||||||
}
|
}
|
||||||
|
|
||||||
// Embeddings
|
// Embeddings
|
||||||
|
|||||||
+105
-8
@@ -5,7 +5,9 @@
|
|||||||
//! in startup). Everything else comes from env vars, the DB settings
|
//! in startup). Everything else comes from env vars, the DB settings
|
||||||
//! table, or auto-detection.
|
//! table, or auto-detection.
|
||||||
|
|
||||||
|
use std::collections::HashMap;
|
||||||
use std::path::PathBuf;
|
use std::path::PathBuf;
|
||||||
|
use std::sync::OnceLock;
|
||||||
use std::time::Duration;
|
use std::time::Duration;
|
||||||
|
|
||||||
use secrecy::{ExposeSecret, SecretString};
|
use secrecy::{ExposeSecret, SecretString};
|
||||||
@@ -13,6 +15,13 @@ use secrecy::{ExposeSecret, SecretString};
|
|||||||
use crate::error::ConfigError;
|
use crate::error::ConfigError;
|
||||||
use crate::settings::Settings;
|
use crate::settings::Settings;
|
||||||
|
|
||||||
|
/// Thread-safe overlay for injected env vars (secrets loaded from DB).
|
||||||
|
///
|
||||||
|
/// Used by `inject_llm_keys_from_secrets()` to make API keys available to
|
||||||
|
/// `optional_env()` without unsafe `set_var` calls. `optional_env()` checks
|
||||||
|
/// real env vars first, then falls back to this overlay.
|
||||||
|
static INJECTED_VARS: OnceLock<HashMap<String, String>> = OnceLock::new();
|
||||||
|
|
||||||
/// Main configuration for the agent.
|
/// Main configuration for the agent.
|
||||||
#[derive(Debug, Clone)]
|
#[derive(Debug, Clone)]
|
||||||
pub struct Config {
|
pub struct Config {
|
||||||
@@ -144,6 +153,15 @@ pub enum DatabaseBackend {
|
|||||||
LibSql,
|
LibSql,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
impl std::fmt::Display for DatabaseBackend {
|
||||||
|
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||||
|
match self {
|
||||||
|
Self::Postgres => write!(f, "postgres"),
|
||||||
|
Self::LibSql => write!(f, "libsql"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
impl std::str::FromStr for DatabaseBackend {
|
impl std::str::FromStr for DatabaseBackend {
|
||||||
type Err = String;
|
type Err = String;
|
||||||
|
|
||||||
@@ -379,6 +397,9 @@ impl std::str::FromStr for NearAiApiMode {
|
|||||||
pub struct NearAiConfig {
|
pub struct NearAiConfig {
|
||||||
/// Model to use (e.g., "claude-3-5-sonnet-20241022", "gpt-4o")
|
/// Model to use (e.g., "claude-3-5-sonnet-20241022", "gpt-4o")
|
||||||
pub model: String,
|
pub model: String,
|
||||||
|
/// Cheap/fast model for lightweight tasks (heartbeat, routing, evaluation).
|
||||||
|
/// Falls back to the main model if not set.
|
||||||
|
pub cheap_model: Option<String>,
|
||||||
/// Base URL for the NEAR AI API (default: https://api.near.ai)
|
/// Base URL for the NEAR AI API (default: https://api.near.ai)
|
||||||
pub base_url: String,
|
pub base_url: String,
|
||||||
/// Base URL for auth/refresh endpoints (default: https://private.near.ai)
|
/// Base URL for auth/refresh endpoints (default: https://private.near.ai)
|
||||||
@@ -398,16 +419,35 @@ pub struct NearAiConfig {
|
|||||||
/// With the default of 3, the provider makes up to 4 total attempts
|
/// With the default of 3, the provider makes up to 4 total attempts
|
||||||
/// (1 initial + 3 retries) before giving up.
|
/// (1 initial + 3 retries) before giving up.
|
||||||
pub max_retries: u32,
|
pub max_retries: u32,
|
||||||
|
/// Cooldown duration in seconds for the failover provider (default: 300).
|
||||||
|
/// When a provider accumulates enough consecutive failures it is skipped
|
||||||
|
/// for this many seconds.
|
||||||
|
pub failover_cooldown_secs: u64,
|
||||||
|
/// Number of consecutive retryable failures before a provider enters
|
||||||
|
/// cooldown (default: 3).
|
||||||
|
pub failover_cooldown_threshold: u32,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl LlmConfig {
|
impl LlmConfig {
|
||||||
fn resolve(settings: &Settings) -> Result<Self, ConfigError> {
|
fn resolve(settings: &Settings) -> Result<Self, ConfigError> {
|
||||||
// Determine backend (default: NearAi)
|
// Determine backend: env var > settings > default (NearAi)
|
||||||
let backend: LlmBackend = if let Some(b) = optional_env("LLM_BACKEND")? {
|
let backend: LlmBackend = if let Some(b) = optional_env("LLM_BACKEND")? {
|
||||||
b.parse().map_err(|e| ConfigError::InvalidValue {
|
b.parse().map_err(|e| ConfigError::InvalidValue {
|
||||||
key: "LLM_BACKEND".to_string(),
|
key: "LLM_BACKEND".to_string(),
|
||||||
message: e,
|
message: e,
|
||||||
})?
|
})?
|
||||||
|
} else if let Some(ref b) = settings.llm_backend {
|
||||||
|
match b.parse() {
|
||||||
|
Ok(backend) => backend,
|
||||||
|
Err(e) => {
|
||||||
|
tracing::warn!(
|
||||||
|
"Invalid llm_backend '{}' in settings: {}. Using default NearAi.",
|
||||||
|
b,
|
||||||
|
e
|
||||||
|
);
|
||||||
|
LlmBackend::NearAi
|
||||||
|
}
|
||||||
|
}
|
||||||
} else {
|
} else {
|
||||||
LlmBackend::NearAi
|
LlmBackend::NearAi
|
||||||
};
|
};
|
||||||
@@ -433,6 +473,7 @@ impl LlmConfig {
|
|||||||
"fireworks::accounts/fireworks/models/llama4-maverick-instruct-basic"
|
"fireworks::accounts/fireworks/models/llama4-maverick-instruct-basic"
|
||||||
.to_string()
|
.to_string()
|
||||||
}),
|
}),
|
||||||
|
cheap_model: optional_env("NEARAI_CHEAP_MODEL")?,
|
||||||
base_url: optional_env("NEARAI_BASE_URL")?
|
base_url: optional_env("NEARAI_BASE_URL")?
|
||||||
.unwrap_or_else(|| "https://cloud-api.near.ai".to_string()),
|
.unwrap_or_else(|| "https://cloud-api.near.ai".to_string()),
|
||||||
auth_base_url: optional_env("NEARAI_AUTH_URL")?
|
auth_base_url: optional_env("NEARAI_AUTH_URL")?
|
||||||
@@ -444,6 +485,8 @@ impl LlmConfig {
|
|||||||
api_key: nearai_api_key,
|
api_key: nearai_api_key,
|
||||||
fallback_model: optional_env("NEARAI_FALLBACK_MODEL")?,
|
fallback_model: optional_env("NEARAI_FALLBACK_MODEL")?,
|
||||||
max_retries: parse_optional_env("NEARAI_MAX_RETRIES", 3)?,
|
max_retries: parse_optional_env("NEARAI_MAX_RETRIES", 3)?,
|
||||||
|
failover_cooldown_secs: parse_optional_env("LLM_FAILOVER_COOLDOWN_SECS", 300)?,
|
||||||
|
failover_cooldown_threshold: parse_optional_env("LLM_FAILOVER_THRESHOLD", 3)?,
|
||||||
};
|
};
|
||||||
|
|
||||||
// Resolve provider-specific configs based on backend
|
// Resolve provider-specific configs based on backend
|
||||||
@@ -476,6 +519,7 @@ impl LlmConfig {
|
|||||||
|
|
||||||
let ollama = if backend == LlmBackend::Ollama {
|
let ollama = if backend == LlmBackend::Ollama {
|
||||||
let base_url = optional_env("OLLAMA_BASE_URL")?
|
let base_url = optional_env("OLLAMA_BASE_URL")?
|
||||||
|
.or_else(|| settings.ollama_base_url.clone())
|
||||||
.unwrap_or_else(|| "http://localhost:11434".to_string());
|
.unwrap_or_else(|| "http://localhost:11434".to_string());
|
||||||
let model = optional_env("OLLAMA_MODEL")?.unwrap_or_else(|| "llama3".to_string());
|
let model = optional_env("OLLAMA_MODEL")?.unwrap_or_else(|| "llama3".to_string());
|
||||||
Some(OllamaConfig { base_url, model })
|
Some(OllamaConfig { base_url, model })
|
||||||
@@ -484,8 +528,9 @@ impl LlmConfig {
|
|||||||
};
|
};
|
||||||
|
|
||||||
let openai_compatible = if backend == LlmBackend::OpenAiCompatible {
|
let openai_compatible = if backend == LlmBackend::OpenAiCompatible {
|
||||||
let base_url =
|
let base_url = optional_env("LLM_BASE_URL")?
|
||||||
optional_env("LLM_BASE_URL")?.ok_or_else(|| ConfigError::MissingRequired {
|
.or_else(|| settings.openai_compatible_base_url.clone())
|
||||||
|
.ok_or_else(|| ConfigError::MissingRequired {
|
||||||
key: "LLM_BASE_URL".to_string(),
|
key: "LLM_BASE_URL".to_string(),
|
||||||
hint: "Set LLM_BASE_URL when LLM_BACKEND=openai_compatible".to_string(),
|
hint: "Set LLM_BASE_URL when LLM_BACKEND=openai_compatible".to_string(),
|
||||||
})?;
|
})?;
|
||||||
@@ -855,6 +900,11 @@ impl std::fmt::Debug for SecretsConfig {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Process-wide cache for the keychain master key.
|
||||||
|
///
|
||||||
|
/// Avoids re-prompting the OS keychain on every `SecretsConfig::resolve()` call
|
||||||
|
/// (e.g. `Config::from_env()` then `Config::from_db()`). Thread-safe alternative
|
||||||
|
/// to caching in a process env var.
|
||||||
impl SecretsConfig {
|
impl SecretsConfig {
|
||||||
/// Auto-detect secrets master key from env var, then OS keychain.
|
/// Auto-detect secrets master key from env var, then OS keychain.
|
||||||
///
|
///
|
||||||
@@ -1338,19 +1388,66 @@ impl ClaudeCodeConfig {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Load API keys from the encrypted secrets store into a thread-safe overlay.
|
||||||
|
///
|
||||||
|
/// This bridges the gap between secrets stored during onboarding and the
|
||||||
|
/// env-var-first resolution in `LlmConfig::resolve()`. Keys in the overlay
|
||||||
|
/// are read by `optional_env()` before falling back to `std::env::var()`,
|
||||||
|
/// so explicit env vars always win.
|
||||||
|
pub async fn inject_llm_keys_from_secrets(
|
||||||
|
secrets: &dyn crate::secrets::SecretsStore,
|
||||||
|
user_id: &str,
|
||||||
|
) {
|
||||||
|
let mappings = [
|
||||||
|
("llm_openai_api_key", "OPENAI_API_KEY"),
|
||||||
|
("llm_anthropic_api_key", "ANTHROPIC_API_KEY"),
|
||||||
|
("llm_compatible_api_key", "LLM_API_KEY"),
|
||||||
|
];
|
||||||
|
|
||||||
|
let mut injected = HashMap::new();
|
||||||
|
|
||||||
|
for (secret_name, env_var) in mappings {
|
||||||
|
match std::env::var(env_var) {
|
||||||
|
Ok(val) if !val.is_empty() => continue,
|
||||||
|
_ => {}
|
||||||
|
}
|
||||||
|
match secrets.get_decrypted(user_id, secret_name).await {
|
||||||
|
Ok(decrypted) => {
|
||||||
|
injected.insert(env_var.to_string(), decrypted.expose().to_string());
|
||||||
|
tracing::debug!("Loaded secret '{}' for env var '{}'", secret_name, env_var);
|
||||||
|
}
|
||||||
|
Err(_) => {
|
||||||
|
// Secret doesn't exist, that's fine
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
let _ = INJECTED_VARS.set(injected);
|
||||||
|
}
|
||||||
|
|
||||||
// Helper functions
|
// Helper functions
|
||||||
|
|
||||||
fn optional_env(key: &str) -> Result<Option<String>, ConfigError> {
|
fn optional_env(key: &str) -> Result<Option<String>, ConfigError> {
|
||||||
|
// Check real env vars first (always win over injected secrets)
|
||||||
match std::env::var(key) {
|
match std::env::var(key) {
|
||||||
Ok(val) if val.is_empty() => Ok(None),
|
Ok(val) if val.is_empty() => {}
|
||||||
Ok(val) => Ok(Some(val)),
|
Ok(val) => return Ok(Some(val)),
|
||||||
Err(std::env::VarError::NotPresent) => Ok(None),
|
Err(std::env::VarError::NotPresent) => {}
|
||||||
Err(e) => Err(ConfigError::ParseError(format!(
|
Err(e) => {
|
||||||
|
return Err(ConfigError::ParseError(format!(
|
||||||
"failed to read {key}: {e}"
|
"failed to read {key}: {e}"
|
||||||
))),
|
)));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Fall back to thread-safe overlay (secrets injected from DB)
|
||||||
|
if let Some(val) = INJECTED_VARS.get().and_then(|map| map.get(key)) {
|
||||||
|
return Ok(Some(val.clone()));
|
||||||
|
}
|
||||||
|
|
||||||
|
Ok(None)
|
||||||
|
}
|
||||||
|
|
||||||
fn parse_optional_env<T>(key: &str, default: T) -> Result<T, ConfigError>
|
fn parse_optional_env<T>(key: &str, default: T) -> Result<T, ConfigError>
|
||||||
where
|
where
|
||||||
T: std::str::FromStr,
|
T: std::str::FromStr,
|
||||||
|
|||||||
@@ -40,6 +40,9 @@ pub enum Error {
|
|||||||
#[error("Workspace error: {0}")]
|
#[error("Workspace error: {0}")]
|
||||||
Workspace(#[from] WorkspaceError),
|
Workspace(#[from] WorkspaceError),
|
||||||
|
|
||||||
|
#[error("Hook error: {0}")]
|
||||||
|
Hook(#[from] crate::hooks::HookError),
|
||||||
|
|
||||||
#[error("Orchestrator error: {0}")]
|
#[error("Orchestrator error: {0}")]
|
||||||
Orchestrator(#[from] OrchestratorError),
|
Orchestrator(#[from] OrchestratorError),
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,199 @@
|
|||||||
|
//! Core hook types and traits.
|
||||||
|
|
||||||
|
use std::time::Duration;
|
||||||
|
|
||||||
|
use async_trait::async_trait;
|
||||||
|
|
||||||
|
/// Points in the agent lifecycle where hooks can be attached.
|
||||||
|
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
|
||||||
|
pub enum HookPoint {
|
||||||
|
/// Before processing an inbound user message.
|
||||||
|
BeforeInbound,
|
||||||
|
/// Before executing a tool call.
|
||||||
|
BeforeToolCall,
|
||||||
|
/// Before sending an outbound response.
|
||||||
|
BeforeOutbound,
|
||||||
|
/// When a new session starts.
|
||||||
|
OnSessionStart,
|
||||||
|
/// When a session ends (pruned or expired).
|
||||||
|
OnSessionEnd,
|
||||||
|
/// Transform the final response before completing a turn.
|
||||||
|
TransformResponse,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Contextual data carried with each hook invocation.
|
||||||
|
#[derive(Debug, Clone)]
|
||||||
|
pub enum HookEvent {
|
||||||
|
/// An inbound user message about to be processed.
|
||||||
|
Inbound {
|
||||||
|
user_id: String,
|
||||||
|
channel: String,
|
||||||
|
content: String,
|
||||||
|
thread_id: Option<String>,
|
||||||
|
},
|
||||||
|
/// A tool call about to be executed.
|
||||||
|
ToolCall {
|
||||||
|
tool_name: String,
|
||||||
|
parameters: serde_json::Value,
|
||||||
|
user_id: String,
|
||||||
|
/// "chat" for interactive, or a job ID string for autonomous jobs.
|
||||||
|
context: String,
|
||||||
|
},
|
||||||
|
/// An outbound response about to be sent.
|
||||||
|
Outbound {
|
||||||
|
user_id: String,
|
||||||
|
channel: String,
|
||||||
|
content: String,
|
||||||
|
thread_id: Option<String>,
|
||||||
|
},
|
||||||
|
/// A new session was created.
|
||||||
|
SessionStart { user_id: String, session_id: String },
|
||||||
|
/// A session was ended (pruned).
|
||||||
|
SessionEnd { user_id: String, session_id: String },
|
||||||
|
/// The final response is being transformed before completing a turn.
|
||||||
|
ResponseTransform {
|
||||||
|
user_id: String,
|
||||||
|
thread_id: String,
|
||||||
|
response: String,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
impl HookEvent {
|
||||||
|
/// Returns the [`HookPoint`] this event corresponds to.
|
||||||
|
pub fn hook_point(&self) -> HookPoint {
|
||||||
|
match self {
|
||||||
|
HookEvent::Inbound { .. } => HookPoint::BeforeInbound,
|
||||||
|
HookEvent::ToolCall { .. } => HookPoint::BeforeToolCall,
|
||||||
|
HookEvent::Outbound { .. } => HookPoint::BeforeOutbound,
|
||||||
|
HookEvent::SessionStart { .. } => HookPoint::OnSessionStart,
|
||||||
|
HookEvent::SessionEnd { .. } => HookPoint::OnSessionEnd,
|
||||||
|
HookEvent::ResponseTransform { .. } => HookPoint::TransformResponse,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Apply a modification string to the event's primary content field.
|
||||||
|
pub fn apply_modification(&mut self, modified: &str) {
|
||||||
|
match self {
|
||||||
|
HookEvent::Inbound { content, .. } | HookEvent::Outbound { content, .. } => {
|
||||||
|
*content = modified.to_string();
|
||||||
|
}
|
||||||
|
HookEvent::ToolCall { parameters, .. } => match serde_json::from_str(modified) {
|
||||||
|
Ok(parsed) => *parameters = parsed,
|
||||||
|
Err(e) => {
|
||||||
|
tracing::warn!(
|
||||||
|
"Hook returned non-JSON modification for ToolCall, ignoring: {}",
|
||||||
|
e
|
||||||
|
);
|
||||||
|
}
|
||||||
|
},
|
||||||
|
HookEvent::ResponseTransform { response, .. } => {
|
||||||
|
*response = modified.to_string();
|
||||||
|
}
|
||||||
|
HookEvent::SessionStart { .. } | HookEvent::SessionEnd { .. } => {
|
||||||
|
// Session events don't have modifiable content
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The result of executing a hook.
|
||||||
|
#[derive(Debug, Clone)]
|
||||||
|
pub enum HookOutcome {
|
||||||
|
/// Continue processing, optionally with modified content.
|
||||||
|
Continue {
|
||||||
|
/// If `Some`, replace the event's primary content with this value.
|
||||||
|
modified: Option<String>,
|
||||||
|
},
|
||||||
|
/// Reject the event entirely.
|
||||||
|
Reject {
|
||||||
|
/// Human-readable reason for the rejection.
|
||||||
|
reason: String,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
impl HookOutcome {
|
||||||
|
/// Shorthand for `Continue { modified: None }`.
|
||||||
|
pub fn ok() -> Self {
|
||||||
|
HookOutcome::Continue { modified: None }
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Shorthand for `Continue { modified: Some(value) }`.
|
||||||
|
pub fn modify(value: String) -> Self {
|
||||||
|
HookOutcome::Continue {
|
||||||
|
modified: Some(value),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Shorthand for `Reject { reason }`.
|
||||||
|
pub fn reject(reason: impl Into<String>) -> Self {
|
||||||
|
HookOutcome::Reject {
|
||||||
|
reason: reason.into(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// How to handle hook execution failures.
|
||||||
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||||
|
pub enum HookFailureMode {
|
||||||
|
/// On error/timeout, continue processing as if the hook returned `ok()`.
|
||||||
|
FailOpen,
|
||||||
|
/// On error/timeout, reject the event.
|
||||||
|
FailClosed,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Hook execution errors.
|
||||||
|
#[derive(Debug, thiserror::Error)]
|
||||||
|
pub enum HookError {
|
||||||
|
#[error("Hook execution failed: {reason}")]
|
||||||
|
ExecutionFailed { reason: String },
|
||||||
|
|
||||||
|
#[error("Hook timed out after {timeout:?}")]
|
||||||
|
Timeout { timeout: Duration },
|
||||||
|
|
||||||
|
#[error("Hook rejected: {reason}")]
|
||||||
|
Rejected { reason: String },
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Context passed to hooks alongside the event.
|
||||||
|
pub struct HookContext {
|
||||||
|
/// Arbitrary metadata hooks can use.
|
||||||
|
pub metadata: serde_json::Value,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Default for HookContext {
|
||||||
|
fn default() -> Self {
|
||||||
|
Self {
|
||||||
|
metadata: serde_json::Value::Null,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Trait for implementing lifecycle hooks.
|
||||||
|
///
|
||||||
|
/// Hooks intercept and can modify agent operations at well-defined points.
|
||||||
|
#[async_trait]
|
||||||
|
pub trait Hook: Send + Sync {
|
||||||
|
/// A unique name for this hook.
|
||||||
|
fn name(&self) -> &str;
|
||||||
|
|
||||||
|
/// The lifecycle points this hook should be called at.
|
||||||
|
fn hook_points(&self) -> &[HookPoint];
|
||||||
|
|
||||||
|
/// How to handle failures in this hook.
|
||||||
|
///
|
||||||
|
/// Default: `FailOpen` (continue on error).
|
||||||
|
fn failure_mode(&self) -> HookFailureMode {
|
||||||
|
HookFailureMode::FailOpen
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Maximum time this hook is allowed to run.
|
||||||
|
///
|
||||||
|
/// Default: 5 seconds.
|
||||||
|
fn timeout(&self) -> Duration {
|
||||||
|
Duration::from_secs(5)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Execute the hook.
|
||||||
|
async fn execute(&self, event: &HookEvent, ctx: &HookContext)
|
||||||
|
-> Result<HookOutcome, HookError>;
|
||||||
|
}
|
||||||
@@ -0,0 +1,19 @@
|
|||||||
|
//! Lifecycle hooks for intercepting and transforming agent operations.
|
||||||
|
//!
|
||||||
|
//! The hook system provides 6 well-defined interception points:
|
||||||
|
//!
|
||||||
|
//! - **BeforeInbound** — Before processing an inbound user message
|
||||||
|
//! - **BeforeToolCall** — Before executing a tool call
|
||||||
|
//! - **BeforeOutbound** — Before sending an outbound response
|
||||||
|
//! - **OnSessionStart** — When a new session starts
|
||||||
|
//! - **OnSessionEnd** — When a session ends
|
||||||
|
//! - **TransformResponse** — Transform the final response before completing a turn
|
||||||
|
//!
|
||||||
|
//! Hooks are executed in priority order (lower number = higher priority).
|
||||||
|
//! Each hook can pass through, modify content, or reject the event.
|
||||||
|
|
||||||
|
pub mod hook;
|
||||||
|
pub mod registry;
|
||||||
|
|
||||||
|
pub use hook::{Hook, HookContext, HookError, HookEvent, HookFailureMode, HookOutcome, HookPoint};
|
||||||
|
pub use registry::HookRegistry;
|
||||||
@@ -0,0 +1,555 @@
|
|||||||
|
//! Hook registry for managing and executing lifecycle hooks.
|
||||||
|
|
||||||
|
use std::sync::Arc;
|
||||||
|
|
||||||
|
use tokio::sync::RwLock;
|
||||||
|
|
||||||
|
use crate::hooks::hook::{Hook, HookContext, HookError, HookEvent, HookFailureMode, HookOutcome};
|
||||||
|
|
||||||
|
/// A registered hook with its priority.
|
||||||
|
struct HookEntry {
|
||||||
|
hook: Arc<dyn Hook>,
|
||||||
|
priority: u32,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Registry that manages hooks and executes them at lifecycle points.
|
||||||
|
///
|
||||||
|
/// Hooks are executed in priority order (lower number = higher priority).
|
||||||
|
/// A `Reject` outcome stops the chain immediately.
|
||||||
|
/// A `Modify` outcome chains through subsequent hooks.
|
||||||
|
pub struct HookRegistry {
|
||||||
|
hooks: RwLock<Vec<HookEntry>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl HookRegistry {
|
||||||
|
/// Create an empty registry.
|
||||||
|
pub fn new() -> Self {
|
||||||
|
Self {
|
||||||
|
hooks: RwLock::new(Vec::new()),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Register a hook with default priority (100).
|
||||||
|
pub async fn register(&self, hook: Arc<dyn Hook>) {
|
||||||
|
self.register_with_priority(hook, 100).await;
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Register a hook with a specific priority.
|
||||||
|
///
|
||||||
|
/// Lower priority number = runs first.
|
||||||
|
pub async fn register_with_priority(&self, hook: Arc<dyn Hook>, priority: u32) {
|
||||||
|
let mut hooks = self.hooks.write().await;
|
||||||
|
hooks.push(HookEntry { hook, priority });
|
||||||
|
hooks.sort_by_key(|e| e.priority);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Unregister a hook by name. Returns `true` if it was found and removed.
|
||||||
|
pub async fn unregister(&self, name: &str) -> bool {
|
||||||
|
let mut hooks = self.hooks.write().await;
|
||||||
|
let before = hooks.len();
|
||||||
|
hooks.retain(|e| e.hook.name() != name);
|
||||||
|
hooks.len() < before
|
||||||
|
}
|
||||||
|
|
||||||
|
/// List all registered hook names (in priority order).
|
||||||
|
pub async fn list(&self) -> Vec<String> {
|
||||||
|
let hooks = self.hooks.read().await;
|
||||||
|
hooks.iter().map(|e| e.hook.name().to_string()).collect()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Run all hooks matching the event's hook point.
|
||||||
|
///
|
||||||
|
/// - Hooks run in priority order (lowest first).
|
||||||
|
/// - `Reject` stops the chain immediately.
|
||||||
|
/// - `Modify` chains the modification through subsequent hooks.
|
||||||
|
/// - Timeout/error handling respects each hook's `failure_mode`.
|
||||||
|
pub async fn run(&self, event: &HookEvent) -> Result<HookOutcome, HookError> {
|
||||||
|
let point = event.hook_point();
|
||||||
|
let ctx = HookContext::default();
|
||||||
|
|
||||||
|
// Clone matching hooks and drop the read guard before executing.
|
||||||
|
// Each hook can run up to its timeout, so holding the guard would
|
||||||
|
// block concurrent register/unregister/run calls.
|
||||||
|
let matching: Vec<Arc<dyn Hook>> = {
|
||||||
|
let hooks = self.hooks.read().await;
|
||||||
|
hooks
|
||||||
|
.iter()
|
||||||
|
.filter(|e| e.hook.hook_points().contains(&point))
|
||||||
|
.map(|e| e.hook.clone())
|
||||||
|
.collect()
|
||||||
|
};
|
||||||
|
|
||||||
|
if matching.is_empty() {
|
||||||
|
return Ok(HookOutcome::ok());
|
||||||
|
}
|
||||||
|
|
||||||
|
let mut current_event = event.clone();
|
||||||
|
|
||||||
|
for hook in &matching {
|
||||||
|
let timeout = hook.timeout();
|
||||||
|
|
||||||
|
let result = tokio::time::timeout(timeout, hook.execute(¤t_event, &ctx)).await;
|
||||||
|
|
||||||
|
match result {
|
||||||
|
Ok(Ok(HookOutcome::Reject { reason })) => {
|
||||||
|
tracing::debug!(hook = hook.name(), "Hook rejected: {}", reason);
|
||||||
|
return Err(HookError::Rejected { reason });
|
||||||
|
}
|
||||||
|
Ok(Ok(HookOutcome::Continue {
|
||||||
|
modified: Some(value),
|
||||||
|
})) => {
|
||||||
|
tracing::debug!(hook = hook.name(), "Hook modified content");
|
||||||
|
current_event.apply_modification(&value);
|
||||||
|
}
|
||||||
|
Ok(Ok(HookOutcome::Continue { modified: None })) => {
|
||||||
|
// No-op, continue chain
|
||||||
|
}
|
||||||
|
Ok(Err(err)) => match hook.failure_mode() {
|
||||||
|
HookFailureMode::FailOpen => {
|
||||||
|
tracing::warn!(hook = hook.name(), "Hook failed (fail-open): {}", err);
|
||||||
|
}
|
||||||
|
HookFailureMode::FailClosed => {
|
||||||
|
tracing::warn!(hook = hook.name(), "Hook failed (fail-closed): {}", err);
|
||||||
|
return Err(HookError::ExecutionFailed {
|
||||||
|
reason: format!("Hook '{}' failed: {}", hook.name(), err),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
},
|
||||||
|
Err(_elapsed) => match hook.failure_mode() {
|
||||||
|
HookFailureMode::FailOpen => {
|
||||||
|
tracing::warn!(
|
||||||
|
hook = hook.name(),
|
||||||
|
"Hook timed out (fail-open) after {:?}",
|
||||||
|
timeout
|
||||||
|
);
|
||||||
|
}
|
||||||
|
HookFailureMode::FailClosed => {
|
||||||
|
tracing::warn!(
|
||||||
|
hook = hook.name(),
|
||||||
|
"Hook timed out (fail-closed) after {:?}",
|
||||||
|
timeout
|
||||||
|
);
|
||||||
|
return Err(HookError::Timeout { timeout });
|
||||||
|
}
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Determine final outcome by comparing with original event
|
||||||
|
let modified = extract_content(¤t_event);
|
||||||
|
let original = extract_content(event);
|
||||||
|
|
||||||
|
if modified != original {
|
||||||
|
Ok(HookOutcome::modify(modified))
|
||||||
|
} else {
|
||||||
|
Ok(HookOutcome::ok())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Default for HookRegistry {
|
||||||
|
fn default() -> Self {
|
||||||
|
Self::new()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Extract the primary content string from a hook event.
|
||||||
|
fn extract_content(event: &HookEvent) -> String {
|
||||||
|
match event {
|
||||||
|
HookEvent::Inbound { content, .. } | HookEvent::Outbound { content, .. } => content.clone(),
|
||||||
|
HookEvent::ToolCall { parameters, .. } => {
|
||||||
|
serde_json::to_string(parameters).unwrap_or_default()
|
||||||
|
}
|
||||||
|
HookEvent::ResponseTransform { response, .. } => response.clone(),
|
||||||
|
HookEvent::SessionStart { session_id, .. } | HookEvent::SessionEnd { session_id, .. } => {
|
||||||
|
session_id.clone()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
use crate::hooks::hook::{HookFailureMode, HookPoint};
|
||||||
|
use async_trait::async_trait;
|
||||||
|
use std::time::Duration;
|
||||||
|
|
||||||
|
/// A test hook that always returns ok.
|
||||||
|
struct PassthroughHook {
|
||||||
|
name: String,
|
||||||
|
points: Vec<HookPoint>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[async_trait]
|
||||||
|
impl Hook for PassthroughHook {
|
||||||
|
fn name(&self) -> &str {
|
||||||
|
&self.name
|
||||||
|
}
|
||||||
|
fn hook_points(&self) -> &[HookPoint] {
|
||||||
|
&self.points
|
||||||
|
}
|
||||||
|
async fn execute(
|
||||||
|
&self,
|
||||||
|
_event: &HookEvent,
|
||||||
|
_ctx: &HookContext,
|
||||||
|
) -> Result<HookOutcome, HookError> {
|
||||||
|
Ok(HookOutcome::ok())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A hook that modifies content by appending a suffix.
|
||||||
|
struct ModifyHook {
|
||||||
|
name: String,
|
||||||
|
suffix: String,
|
||||||
|
points: Vec<HookPoint>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[async_trait]
|
||||||
|
impl Hook for ModifyHook {
|
||||||
|
fn name(&self) -> &str {
|
||||||
|
&self.name
|
||||||
|
}
|
||||||
|
fn hook_points(&self) -> &[HookPoint] {
|
||||||
|
&self.points
|
||||||
|
}
|
||||||
|
async fn execute(
|
||||||
|
&self,
|
||||||
|
event: &HookEvent,
|
||||||
|
_ctx: &HookContext,
|
||||||
|
) -> Result<HookOutcome, HookError> {
|
||||||
|
let content = extract_content(event);
|
||||||
|
Ok(HookOutcome::modify(format!("{}{}", content, self.suffix)))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A hook that always rejects.
|
||||||
|
struct RejectHook {
|
||||||
|
name: String,
|
||||||
|
reason: String,
|
||||||
|
points: Vec<HookPoint>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[async_trait]
|
||||||
|
impl Hook for RejectHook {
|
||||||
|
fn name(&self) -> &str {
|
||||||
|
&self.name
|
||||||
|
}
|
||||||
|
fn hook_points(&self) -> &[HookPoint] {
|
||||||
|
&self.points
|
||||||
|
}
|
||||||
|
async fn execute(
|
||||||
|
&self,
|
||||||
|
_event: &HookEvent,
|
||||||
|
_ctx: &HookContext,
|
||||||
|
) -> Result<HookOutcome, HookError> {
|
||||||
|
Ok(HookOutcome::reject(&self.reason))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A hook that always errors.
|
||||||
|
struct ErrorHook {
|
||||||
|
name: String,
|
||||||
|
points: Vec<HookPoint>,
|
||||||
|
failure_mode: HookFailureMode,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[async_trait]
|
||||||
|
impl Hook for ErrorHook {
|
||||||
|
fn name(&self) -> &str {
|
||||||
|
&self.name
|
||||||
|
}
|
||||||
|
fn hook_points(&self) -> &[HookPoint] {
|
||||||
|
&self.points
|
||||||
|
}
|
||||||
|
fn failure_mode(&self) -> HookFailureMode {
|
||||||
|
self.failure_mode
|
||||||
|
}
|
||||||
|
async fn execute(
|
||||||
|
&self,
|
||||||
|
_event: &HookEvent,
|
||||||
|
_ctx: &HookContext,
|
||||||
|
) -> Result<HookOutcome, HookError> {
|
||||||
|
Err(HookError::ExecutionFailed {
|
||||||
|
reason: "test error".into(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A hook that sleeps longer than its timeout.
|
||||||
|
struct SlowHook {
|
||||||
|
name: String,
|
||||||
|
points: Vec<HookPoint>,
|
||||||
|
failure_mode: HookFailureMode,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[async_trait]
|
||||||
|
impl Hook for SlowHook {
|
||||||
|
fn name(&self) -> &str {
|
||||||
|
&self.name
|
||||||
|
}
|
||||||
|
fn hook_points(&self) -> &[HookPoint] {
|
||||||
|
&self.points
|
||||||
|
}
|
||||||
|
fn failure_mode(&self) -> HookFailureMode {
|
||||||
|
self.failure_mode
|
||||||
|
}
|
||||||
|
fn timeout(&self) -> Duration {
|
||||||
|
Duration::from_millis(50)
|
||||||
|
}
|
||||||
|
async fn execute(
|
||||||
|
&self,
|
||||||
|
_event: &HookEvent,
|
||||||
|
_ctx: &HookContext,
|
||||||
|
) -> Result<HookOutcome, HookError> {
|
||||||
|
tokio::time::sleep(Duration::from_millis(200)).await;
|
||||||
|
Ok(HookOutcome::ok())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn test_event() -> HookEvent {
|
||||||
|
HookEvent::Inbound {
|
||||||
|
user_id: "user-1".into(),
|
||||||
|
channel: "test".into(),
|
||||||
|
content: "hello".into(),
|
||||||
|
thread_id: None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_empty_registry_returns_ok() {
|
||||||
|
let registry = HookRegistry::new();
|
||||||
|
let result = registry.run(&test_event()).await;
|
||||||
|
assert!(result.is_ok());
|
||||||
|
assert!(matches!(
|
||||||
|
result.unwrap(),
|
||||||
|
HookOutcome::Continue { modified: None }
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_register_and_list() {
|
||||||
|
let registry = HookRegistry::new();
|
||||||
|
registry
|
||||||
|
.register(Arc::new(PassthroughHook {
|
||||||
|
name: "hook-a".into(),
|
||||||
|
points: vec![HookPoint::BeforeInbound],
|
||||||
|
}))
|
||||||
|
.await;
|
||||||
|
registry
|
||||||
|
.register(Arc::new(PassthroughHook {
|
||||||
|
name: "hook-b".into(),
|
||||||
|
points: vec![HookPoint::BeforeInbound],
|
||||||
|
}))
|
||||||
|
.await;
|
||||||
|
|
||||||
|
let names = registry.list().await;
|
||||||
|
assert_eq!(names, vec!["hook-a", "hook-b"]);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_priority_ordering() {
|
||||||
|
let registry = HookRegistry::new();
|
||||||
|
|
||||||
|
// Register in reverse priority order
|
||||||
|
registry
|
||||||
|
.register_with_priority(
|
||||||
|
Arc::new(ModifyHook {
|
||||||
|
name: "low-prio".into(),
|
||||||
|
suffix: "-LOW".into(),
|
||||||
|
points: vec![HookPoint::BeforeInbound],
|
||||||
|
}),
|
||||||
|
200,
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
registry
|
||||||
|
.register_with_priority(
|
||||||
|
Arc::new(ModifyHook {
|
||||||
|
name: "high-prio".into(),
|
||||||
|
suffix: "-HIGH".into(),
|
||||||
|
points: vec![HookPoint::BeforeInbound],
|
||||||
|
}),
|
||||||
|
10,
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
|
||||||
|
// Should run in priority order: high-prio first, then low-prio
|
||||||
|
let names = registry.list().await;
|
||||||
|
assert_eq!(names[0], "high-prio");
|
||||||
|
assert_eq!(names[1], "low-prio");
|
||||||
|
|
||||||
|
let result = registry.run(&test_event()).await.unwrap();
|
||||||
|
match result {
|
||||||
|
HookOutcome::Continue { modified: Some(m) } => {
|
||||||
|
// "hello" -> "hello-HIGH" -> "hello-HIGH-LOW"
|
||||||
|
assert_eq!(m, "hello-HIGH-LOW");
|
||||||
|
}
|
||||||
|
other => panic!("Expected modification chain, got: {:?}", other),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_reject_stops_chain() {
|
||||||
|
let registry = HookRegistry::new();
|
||||||
|
|
||||||
|
registry
|
||||||
|
.register_with_priority(
|
||||||
|
Arc::new(RejectHook {
|
||||||
|
name: "blocker".into(),
|
||||||
|
reason: "blocked".into(),
|
||||||
|
points: vec![HookPoint::BeforeInbound],
|
||||||
|
}),
|
||||||
|
10,
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
registry
|
||||||
|
.register_with_priority(
|
||||||
|
Arc::new(ModifyHook {
|
||||||
|
name: "modifier".into(),
|
||||||
|
suffix: "-MODIFIED".into(),
|
||||||
|
points: vec![HookPoint::BeforeInbound],
|
||||||
|
}),
|
||||||
|
20,
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
|
||||||
|
let result = registry.run(&test_event()).await;
|
||||||
|
assert!(result.is_err());
|
||||||
|
match result.unwrap_err() {
|
||||||
|
HookError::Rejected { reason } => assert_eq!(reason, "blocked"),
|
||||||
|
other => panic!("Expected Rejected, got: {:?}", other),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_modification_chaining() {
|
||||||
|
let registry = HookRegistry::new();
|
||||||
|
|
||||||
|
registry
|
||||||
|
.register_with_priority(
|
||||||
|
Arc::new(ModifyHook {
|
||||||
|
name: "first".into(),
|
||||||
|
suffix: "-A".into(),
|
||||||
|
points: vec![HookPoint::BeforeInbound],
|
||||||
|
}),
|
||||||
|
10,
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
registry
|
||||||
|
.register_with_priority(
|
||||||
|
Arc::new(ModifyHook {
|
||||||
|
name: "second".into(),
|
||||||
|
suffix: "-B".into(),
|
||||||
|
points: vec![HookPoint::BeforeInbound],
|
||||||
|
}),
|
||||||
|
20,
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
|
||||||
|
let result = registry.run(&test_event()).await.unwrap();
|
||||||
|
match result {
|
||||||
|
HookOutcome::Continue { modified: Some(m) } => {
|
||||||
|
assert_eq!(m, "hello-A-B");
|
||||||
|
}
|
||||||
|
other => panic!("Expected chained modification, got: {:?}", other),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_fail_open_on_error() {
|
||||||
|
let registry = HookRegistry::new();
|
||||||
|
registry
|
||||||
|
.register(Arc::new(ErrorHook {
|
||||||
|
name: "err-open".into(),
|
||||||
|
points: vec![HookPoint::BeforeInbound],
|
||||||
|
failure_mode: HookFailureMode::FailOpen,
|
||||||
|
}))
|
||||||
|
.await;
|
||||||
|
|
||||||
|
let result = registry.run(&test_event()).await;
|
||||||
|
assert!(result.is_ok());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_fail_closed_on_error() {
|
||||||
|
let registry = HookRegistry::new();
|
||||||
|
registry
|
||||||
|
.register(Arc::new(ErrorHook {
|
||||||
|
name: "err-closed".into(),
|
||||||
|
points: vec![HookPoint::BeforeInbound],
|
||||||
|
failure_mode: HookFailureMode::FailClosed,
|
||||||
|
}))
|
||||||
|
.await;
|
||||||
|
|
||||||
|
let result = registry.run(&test_event()).await;
|
||||||
|
assert!(result.is_err());
|
||||||
|
assert!(matches!(
|
||||||
|
result.unwrap_err(),
|
||||||
|
HookError::ExecutionFailed { .. }
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_fail_open_on_timeout() {
|
||||||
|
let registry = HookRegistry::new();
|
||||||
|
registry
|
||||||
|
.register(Arc::new(SlowHook {
|
||||||
|
name: "slow-open".into(),
|
||||||
|
points: vec![HookPoint::BeforeInbound],
|
||||||
|
failure_mode: HookFailureMode::FailOpen,
|
||||||
|
}))
|
||||||
|
.await;
|
||||||
|
|
||||||
|
let result = registry.run(&test_event()).await;
|
||||||
|
assert!(result.is_ok());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_fail_closed_on_timeout() {
|
||||||
|
let registry = HookRegistry::new();
|
||||||
|
registry
|
||||||
|
.register(Arc::new(SlowHook {
|
||||||
|
name: "slow-closed".into(),
|
||||||
|
points: vec![HookPoint::BeforeInbound],
|
||||||
|
failure_mode: HookFailureMode::FailClosed,
|
||||||
|
}))
|
||||||
|
.await;
|
||||||
|
|
||||||
|
let result = registry.run(&test_event()).await;
|
||||||
|
assert!(result.is_err());
|
||||||
|
assert!(matches!(result.unwrap_err(), HookError::Timeout { .. }));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_unregister() {
|
||||||
|
let registry = HookRegistry::new();
|
||||||
|
registry
|
||||||
|
.register(Arc::new(PassthroughHook {
|
||||||
|
name: "removable".into(),
|
||||||
|
points: vec![HookPoint::BeforeInbound],
|
||||||
|
}))
|
||||||
|
.await;
|
||||||
|
|
||||||
|
assert_eq!(registry.list().await.len(), 1);
|
||||||
|
assert!(registry.unregister("removable").await);
|
||||||
|
assert_eq!(registry.list().await.len(), 0);
|
||||||
|
|
||||||
|
// Unregistering non-existent returns false
|
||||||
|
assert!(!registry.unregister("nonexistent").await);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_hooks_only_match_their_points() {
|
||||||
|
let registry = HookRegistry::new();
|
||||||
|
registry
|
||||||
|
.register(Arc::new(RejectHook {
|
||||||
|
name: "outbound-only".into(),
|
||||||
|
reason: "blocked".into(),
|
||||||
|
points: vec![HookPoint::BeforeOutbound],
|
||||||
|
}))
|
||||||
|
.await;
|
||||||
|
|
||||||
|
// Inbound event should not be affected by outbound-only hook
|
||||||
|
let result = registry.run(&test_event()).await;
|
||||||
|
assert!(result.is_ok());
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -39,6 +39,7 @@
|
|||||||
//! - **Continuous learning** - Improve estimates from historical data
|
//! - **Continuous learning** - Improve estimates from historical data
|
||||||
|
|
||||||
pub mod agent;
|
pub mod agent;
|
||||||
|
pub mod boot_screen;
|
||||||
pub mod bootstrap;
|
pub mod bootstrap;
|
||||||
pub mod channels;
|
pub mod channels;
|
||||||
pub mod cli;
|
pub mod cli;
|
||||||
@@ -50,6 +51,7 @@ pub mod estimation;
|
|||||||
pub mod evaluation;
|
pub mod evaluation;
|
||||||
pub mod extensions;
|
pub mod extensions;
|
||||||
pub mod history;
|
pub mod history;
|
||||||
|
pub mod hooks;
|
||||||
pub mod llm;
|
pub mod llm;
|
||||||
pub mod orchestrator;
|
pub mod orchestrator;
|
||||||
pub mod pairing;
|
pub mod pairing;
|
||||||
|
|||||||
+542
-8
@@ -2,10 +2,15 @@
|
|||||||
//!
|
//!
|
||||||
//! Wraps multiple LlmProvider instances and tries each in sequence
|
//! Wraps multiple LlmProvider instances and tries each in sequence
|
||||||
//! until one succeeds. Transparent to callers --- same LlmProvider trait.
|
//! until one succeeds. Transparent to callers --- same LlmProvider trait.
|
||||||
|
//!
|
||||||
|
//! Providers that fail repeatedly are temporarily placed in cooldown
|
||||||
|
//! so subsequent requests skip them, reducing latency when a provider
|
||||||
|
//! is known to be down. Cooldown state is lock-free (atomics only).
|
||||||
|
|
||||||
use std::future::Future;
|
use std::future::Future;
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
use std::sync::atomic::{AtomicU32, AtomicU64, AtomicUsize, Ordering};
|
||||||
|
use std::time::{Duration, Instant};
|
||||||
|
|
||||||
use async_trait::async_trait;
|
use async_trait::async_trait;
|
||||||
use rust_decimal::Decimal;
|
use rust_decimal::Decimal;
|
||||||
@@ -41,61 +46,217 @@ fn is_retryable(err: &LlmError) -> bool {
|
|||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Configuration for per-provider cooldown behavior.
|
||||||
|
///
|
||||||
|
/// When a provider accumulates `failure_threshold` consecutive retryable
|
||||||
|
/// failures, it enters cooldown for `cooldown_duration`. During cooldown
|
||||||
|
/// the provider is skipped (unless *all* providers are in cooldown, in
|
||||||
|
/// which case the oldest-cooled one is tried).
|
||||||
|
#[derive(Debug, Clone)]
|
||||||
|
pub struct CooldownConfig {
|
||||||
|
/// How long a provider stays in cooldown after exceeding the threshold.
|
||||||
|
pub cooldown_duration: Duration,
|
||||||
|
/// Number of consecutive retryable failures before cooldown activates.
|
||||||
|
pub failure_threshold: u32,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Default for CooldownConfig {
|
||||||
|
fn default() -> Self {
|
||||||
|
Self {
|
||||||
|
cooldown_duration: Duration::from_secs(300),
|
||||||
|
failure_threshold: 3,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Per-provider cooldown state, entirely lock-free.
|
||||||
|
///
|
||||||
|
/// All atomic operations use `Relaxed` ordering — consistent with the
|
||||||
|
/// existing `last_used` field. Stale reads are harmless: the worst case
|
||||||
|
/// is one extra attempt against a provider that just entered cooldown.
|
||||||
|
struct ProviderCooldown {
|
||||||
|
/// Consecutive retryable failures. Reset to 0 on success.
|
||||||
|
failure_count: AtomicU32,
|
||||||
|
/// Nanoseconds since `epoch` when cooldown was activated.
|
||||||
|
/// 0 means the provider is NOT in cooldown.
|
||||||
|
cooldown_activated_nanos: AtomicU64,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ProviderCooldown {
|
||||||
|
fn new() -> Self {
|
||||||
|
Self {
|
||||||
|
failure_count: AtomicU32::new(0),
|
||||||
|
cooldown_activated_nanos: AtomicU64::new(0),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Check whether the provider is currently in cooldown.
|
||||||
|
fn is_in_cooldown(&self, now_nanos: u64, cooldown_nanos: u64) -> bool {
|
||||||
|
let activated = self.cooldown_activated_nanos.load(Ordering::Relaxed);
|
||||||
|
activated != 0 && now_nanos.saturating_sub(activated) < cooldown_nanos
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Record a retryable failure. Returns `true` if the threshold was
|
||||||
|
/// just reached (caller should activate cooldown).
|
||||||
|
fn record_failure(&self, threshold: u32) -> bool {
|
||||||
|
let prev = self.failure_count.fetch_add(1, Ordering::Relaxed);
|
||||||
|
prev + 1 >= threshold
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Activate cooldown at the given timestamp.
|
||||||
|
fn activate_cooldown(&self, now_nanos: u64) {
|
||||||
|
self.cooldown_activated_nanos
|
||||||
|
.store(now_nanos, Ordering::Relaxed);
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Reset failure count and clear cooldown (called on success).
|
||||||
|
fn reset(&self) {
|
||||||
|
self.failure_count.store(0, Ordering::Relaxed);
|
||||||
|
self.cooldown_activated_nanos.store(0, Ordering::Relaxed);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
/// An LLM provider that wraps multiple providers and tries each in sequence
|
/// An LLM provider that wraps multiple providers and tries each in sequence
|
||||||
/// on transient failures.
|
/// on transient failures.
|
||||||
///
|
///
|
||||||
/// The first provider in the list is the primary. If it fails with a retryable
|
/// The first provider in the list is the primary. If it fails with a retryable
|
||||||
/// error, the next provider is tried, and so on. Non-retryable errors
|
/// error, the next provider is tried, and so on. Non-retryable errors
|
||||||
/// (e.g. `AuthFailed`, `ContextLengthExceeded`) propagate immediately.
|
/// (e.g. `AuthFailed`, `ContextLengthExceeded`) propagate immediately.
|
||||||
|
///
|
||||||
|
/// Providers that repeatedly fail with retryable errors are temporarily
|
||||||
|
/// placed in cooldown and skipped, reducing latency.
|
||||||
pub struct FailoverProvider {
|
pub struct FailoverProvider {
|
||||||
providers: Vec<Arc<dyn LlmProvider>>,
|
providers: Vec<Arc<dyn LlmProvider>>,
|
||||||
/// Index of the provider that last handled a request successfully.
|
/// Index of the provider that last handled a request successfully.
|
||||||
/// Used by `model_name()` and `cost_per_token()` so downstream cost
|
/// Used by `model_name()` and `cost_per_token()` so downstream cost
|
||||||
/// tracking reflects the provider that actually served the request.
|
/// tracking reflects the provider that actually served the request.
|
||||||
last_used: AtomicUsize,
|
last_used: AtomicUsize,
|
||||||
|
/// Per-provider cooldown tracking (same length as `providers`).
|
||||||
|
cooldowns: Vec<ProviderCooldown>,
|
||||||
|
/// Reference instant for computing elapsed nanos. Shared across all
|
||||||
|
/// cooldown timestamps so they are comparable.
|
||||||
|
epoch: Instant,
|
||||||
|
/// Cooldown configuration.
|
||||||
|
cooldown_config: CooldownConfig,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl FailoverProvider {
|
impl FailoverProvider {
|
||||||
/// Create a new failover provider.
|
/// Create a new failover provider with default cooldown settings.
|
||||||
///
|
///
|
||||||
/// Returns an error if `providers` is empty.
|
/// Returns an error if `providers` is empty.
|
||||||
pub fn new(providers: Vec<Arc<dyn LlmProvider>>) -> Result<Self, LlmError> {
|
pub fn new(providers: Vec<Arc<dyn LlmProvider>>) -> Result<Self, LlmError> {
|
||||||
|
Self::with_cooldown(providers, CooldownConfig::default())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Create a new failover provider with explicit cooldown configuration.
|
||||||
|
///
|
||||||
|
/// Returns an error if `providers` is empty.
|
||||||
|
pub fn with_cooldown(
|
||||||
|
providers: Vec<Arc<dyn LlmProvider>>,
|
||||||
|
cooldown_config: CooldownConfig,
|
||||||
|
) -> Result<Self, LlmError> {
|
||||||
if providers.is_empty() {
|
if providers.is_empty() {
|
||||||
return Err(LlmError::RequestFailed {
|
return Err(LlmError::RequestFailed {
|
||||||
provider: "failover".to_string(),
|
provider: "failover".to_string(),
|
||||||
reason: "FailoverProvider requires at least one provider".to_string(),
|
reason: "FailoverProvider requires at least one provider".to_string(),
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
let cooldowns = (0..providers.len())
|
||||||
|
.map(|_| ProviderCooldown::new())
|
||||||
|
.collect();
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
providers,
|
providers,
|
||||||
last_used: AtomicUsize::new(0),
|
last_used: AtomicUsize::new(0),
|
||||||
|
cooldowns,
|
||||||
|
epoch: Instant::now(),
|
||||||
|
cooldown_config,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Nanoseconds elapsed since `self.epoch`.
|
||||||
|
///
|
||||||
|
/// Truncates `u128` → `u64` (wraps after ~584 years of continuous
|
||||||
|
/// uptime). Acceptable because `epoch` is set at construction time.
|
||||||
|
fn now_nanos(&self) -> u64 {
|
||||||
|
self.epoch.elapsed().as_nanos() as u64
|
||||||
|
}
|
||||||
|
|
||||||
/// Try each provider in sequence until one succeeds or all fail.
|
/// Try each provider in sequence until one succeeds or all fail.
|
||||||
|
///
|
||||||
|
/// Providers in cooldown are skipped unless *all* providers are in
|
||||||
|
/// cooldown, in which case the one with the oldest cooldown timestamp
|
||||||
|
/// (most likely to have recovered) is tried.
|
||||||
async fn try_providers<T, F, Fut>(&self, mut call: F) -> Result<T, LlmError>
|
async fn try_providers<T, F, Fut>(&self, mut call: F) -> Result<T, LlmError>
|
||||||
where
|
where
|
||||||
F: FnMut(Arc<dyn LlmProvider>) -> Fut,
|
F: FnMut(Arc<dyn LlmProvider>) -> Fut,
|
||||||
Fut: Future<Output = Result<T, LlmError>>,
|
Fut: Future<Output = Result<T, LlmError>>,
|
||||||
{
|
{
|
||||||
|
let now_nanos = self.now_nanos();
|
||||||
|
let cooldown_nanos = self.cooldown_config.cooldown_duration.as_nanos() as u64;
|
||||||
|
|
||||||
|
// Partition providers into available and cooled-down.
|
||||||
|
let (mut available, cooled_down): (Vec<usize>, Vec<usize>) = (0..self.providers.len())
|
||||||
|
.partition(|&i| !self.cooldowns[i].is_in_cooldown(now_nanos, cooldown_nanos));
|
||||||
|
|
||||||
|
// Log skipped providers.
|
||||||
|
for &i in &cooled_down {
|
||||||
|
tracing::info!(
|
||||||
|
provider = %self.providers[i].model_name(),
|
||||||
|
"Skipping provider (in cooldown)"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Never skip ALL providers: if every provider is in cooldown, pick
|
||||||
|
// the one with the oldest cooldown activation (most likely recovered).
|
||||||
|
if available.is_empty() {
|
||||||
|
let oldest = (0..self.providers.len())
|
||||||
|
.min_by_key(|&i| {
|
||||||
|
self.cooldowns[i]
|
||||||
|
.cooldown_activated_nanos
|
||||||
|
.load(Ordering::Relaxed)
|
||||||
|
})
|
||||||
|
.expect("providers list is non-empty");
|
||||||
|
tracing::info!(
|
||||||
|
provider = %self.providers[oldest].model_name(),
|
||||||
|
"All providers in cooldown, trying oldest-cooled provider"
|
||||||
|
);
|
||||||
|
available.push(oldest);
|
||||||
|
}
|
||||||
|
|
||||||
let mut last_error: Option<LlmError> = None;
|
let mut last_error: Option<LlmError> = None;
|
||||||
|
|
||||||
for (i, provider) in self.providers.iter().enumerate() {
|
for (pos, &i) in available.iter().enumerate() {
|
||||||
|
let provider = &self.providers[i];
|
||||||
let result = call(Arc::clone(provider)).await;
|
let result = call(Arc::clone(provider)).await;
|
||||||
match result {
|
match result {
|
||||||
Ok(response) => {
|
Ok(response) => {
|
||||||
self.last_used.store(i, Ordering::Relaxed);
|
self.last_used.store(i, Ordering::Relaxed);
|
||||||
|
self.cooldowns[i].reset();
|
||||||
return Ok(response);
|
return Ok(response);
|
||||||
}
|
}
|
||||||
Err(err) => {
|
Err(err) => {
|
||||||
if !is_retryable(&err) {
|
if !is_retryable(&err) {
|
||||||
return Err(err);
|
return Err(err);
|
||||||
}
|
}
|
||||||
if i + 1 < self.providers.len() {
|
|
||||||
|
// Increment failure count; activate cooldown if threshold reached.
|
||||||
|
if self.cooldowns[i].record_failure(self.cooldown_config.failure_threshold) {
|
||||||
|
let nanos = self.now_nanos();
|
||||||
|
self.cooldowns[i].activate_cooldown(nanos);
|
||||||
|
tracing::warn!(
|
||||||
|
provider = %provider.model_name(),
|
||||||
|
threshold = self.cooldown_config.failure_threshold,
|
||||||
|
cooldown_secs = self.cooldown_config.cooldown_duration.as_secs(),
|
||||||
|
"Provider entered cooldown after repeated failures"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
if pos + 1 < available.len() {
|
||||||
|
let next_i = available[pos + 1];
|
||||||
tracing::warn!(
|
tracing::warn!(
|
||||||
provider = %provider.model_name(),
|
provider = %provider.model_name(),
|
||||||
error = %err,
|
error = %err,
|
||||||
next_provider = %self.providers[i + 1].model_name(),
|
next_provider = %self.providers[next_i].model_name(),
|
||||||
"Provider failed with retryable error, trying next provider"
|
"Provider failed with retryable error, trying next provider"
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
@@ -104,9 +265,9 @@ impl FailoverProvider {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// SAFETY: providers is non-empty (checked in `new`), so at least one
|
// SAFETY: `available` is non-empty (guaranteed above), so at least one
|
||||||
// iteration ran and `last_error` is `Some`.
|
// iteration ran and `last_error` is `Some`.
|
||||||
Err(last_error.expect("providers list is non-empty"))
|
Err(last_error.expect("available providers list is non-empty"))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -166,7 +327,6 @@ mod tests {
|
|||||||
use super::*;
|
use super::*;
|
||||||
|
|
||||||
use std::sync::Mutex;
|
use std::sync::Mutex;
|
||||||
use std::time::Duration;
|
|
||||||
|
|
||||||
use crate::llm::provider::{CompletionResponse, FinishReason, ToolCompletionResponse};
|
use crate::llm::provider::{CompletionResponse, FinishReason, ToolCompletionResponse};
|
||||||
|
|
||||||
@@ -432,6 +592,380 @@ mod tests {
|
|||||||
assert!(models.contains(&"model-b".to_string()));
|
assert!(models.contains(&"model-b".to_string()));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// --- MultiCallMockProvider for cooldown tests ---
|
||||||
|
//
|
||||||
|
// Unlike `MockProvider` which uses `.take()` (single-use), this mock
|
||||||
|
// tracks a call counter and returns errors for the first N calls,
|
||||||
|
// then succeeds.
|
||||||
|
|
||||||
|
struct MultiCallMockProvider {
|
||||||
|
name: String,
|
||||||
|
/// How many calls should fail before succeeding. 0 = always succeed.
|
||||||
|
fail_count: u32,
|
||||||
|
/// Atomically tracks how many times `complete` has been called.
|
||||||
|
calls: AtomicU32,
|
||||||
|
/// If true, failures are non-retryable (AuthFailed).
|
||||||
|
non_retryable: bool,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl MultiCallMockProvider {
|
||||||
|
/// Always succeeds.
|
||||||
|
fn always_ok(name: &str) -> Self {
|
||||||
|
Self {
|
||||||
|
name: name.to_string(),
|
||||||
|
fail_count: 0,
|
||||||
|
calls: AtomicU32::new(0),
|
||||||
|
non_retryable: false,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Fails with retryable error for the first `n` calls, then succeeds.
|
||||||
|
fn fail_then_ok(name: &str, n: u32) -> Self {
|
||||||
|
Self {
|
||||||
|
name: name.to_string(),
|
||||||
|
fail_count: n,
|
||||||
|
calls: AtomicU32::new(0),
|
||||||
|
non_retryable: false,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Always fails with retryable error.
|
||||||
|
fn always_fail(name: &str) -> Self {
|
||||||
|
Self {
|
||||||
|
name: name.to_string(),
|
||||||
|
fail_count: u32::MAX,
|
||||||
|
calls: AtomicU32::new(0),
|
||||||
|
non_retryable: false,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Always fails with non-retryable error.
|
||||||
|
fn always_fail_non_retryable(name: &str) -> Self {
|
||||||
|
Self {
|
||||||
|
name: name.to_string(),
|
||||||
|
fail_count: u32::MAX,
|
||||||
|
calls: AtomicU32::new(0),
|
||||||
|
non_retryable: true,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn call_count(&self) -> u32 {
|
||||||
|
self.calls.load(Ordering::Relaxed)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[async_trait]
|
||||||
|
impl LlmProvider for MultiCallMockProvider {
|
||||||
|
fn model_name(&self) -> &str {
|
||||||
|
&self.name
|
||||||
|
}
|
||||||
|
|
||||||
|
fn cost_per_token(&self) -> (Decimal, Decimal) {
|
||||||
|
(Decimal::ZERO, Decimal::ZERO)
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn complete(
|
||||||
|
&self,
|
||||||
|
_request: CompletionRequest,
|
||||||
|
) -> Result<CompletionResponse, LlmError> {
|
||||||
|
let n = self.calls.fetch_add(1, Ordering::Relaxed);
|
||||||
|
if n < self.fail_count {
|
||||||
|
if self.non_retryable {
|
||||||
|
return Err(LlmError::AuthFailed {
|
||||||
|
provider: self.name.clone(),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
return Err(LlmError::RequestFailed {
|
||||||
|
provider: self.name.clone(),
|
||||||
|
reason: format!("call {} failed", n),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
Ok(CompletionResponse {
|
||||||
|
content: format!("{} ok", self.name),
|
||||||
|
input_tokens: 10,
|
||||||
|
output_tokens: 5,
|
||||||
|
finish_reason: FinishReason::Stop,
|
||||||
|
response_id: None,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn complete_with_tools(
|
||||||
|
&self,
|
||||||
|
_request: ToolCompletionRequest,
|
||||||
|
) -> Result<ToolCompletionResponse, LlmError> {
|
||||||
|
let n = self.calls.fetch_add(1, Ordering::Relaxed);
|
||||||
|
if n < self.fail_count {
|
||||||
|
if self.non_retryable {
|
||||||
|
return Err(LlmError::AuthFailed {
|
||||||
|
provider: self.name.clone(),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
return Err(LlmError::RequestFailed {
|
||||||
|
provider: self.name.clone(),
|
||||||
|
reason: format!("call {} failed", n),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
Ok(ToolCompletionResponse {
|
||||||
|
content: Some(format!("{} ok", self.name)),
|
||||||
|
tool_calls: vec![],
|
||||||
|
input_tokens: 10,
|
||||||
|
output_tokens: 5,
|
||||||
|
finish_reason: FinishReason::Stop,
|
||||||
|
response_id: None,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn list_models(&self) -> Result<Vec<String>, LlmError> {
|
||||||
|
Ok(vec![self.name.clone()])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- Cooldown tests ---
|
||||||
|
|
||||||
|
// Cooldown test 1: Provider enters cooldown after `threshold` consecutive failures.
|
||||||
|
#[tokio::test]
|
||||||
|
async fn cooldown_activates_after_threshold() {
|
||||||
|
let config = CooldownConfig {
|
||||||
|
cooldown_duration: Duration::from_secs(300),
|
||||||
|
failure_threshold: 2,
|
||||||
|
};
|
||||||
|
let p1 = Arc::new(MultiCallMockProvider::always_fail("p1"));
|
||||||
|
let p2 = Arc::new(MultiCallMockProvider::always_ok("p2"));
|
||||||
|
|
||||||
|
let failover =
|
||||||
|
FailoverProvider::with_cooldown(vec![p1.clone(), p2.clone()], config).unwrap();
|
||||||
|
|
||||||
|
// Request 1: p1 fails (count=1, below threshold), p2 succeeds.
|
||||||
|
let r = failover.complete(make_request()).await.unwrap();
|
||||||
|
assert_eq!(r.content, "p2 ok");
|
||||||
|
assert_eq!(p1.call_count(), 1);
|
||||||
|
|
||||||
|
// Request 2: p1 fails again (count=2, reaches threshold → cooldown), p2 succeeds.
|
||||||
|
let r = failover.complete(make_request()).await.unwrap();
|
||||||
|
assert_eq!(r.content, "p2 ok");
|
||||||
|
assert_eq!(p1.call_count(), 2);
|
||||||
|
|
||||||
|
// Request 3: p1 should be skipped (in cooldown), only p2 called.
|
||||||
|
let prev_p1_calls = p1.call_count();
|
||||||
|
let r = failover.complete(make_request()).await.unwrap();
|
||||||
|
assert_eq!(r.content, "p2 ok");
|
||||||
|
// p1 was NOT called again.
|
||||||
|
assert_eq!(p1.call_count(), prev_p1_calls);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Cooldown test 2: Cooldown expires after duration, provider is retried.
|
||||||
|
#[tokio::test]
|
||||||
|
async fn cooldown_expires_after_duration() {
|
||||||
|
let config = CooldownConfig {
|
||||||
|
cooldown_duration: Duration::from_millis(1),
|
||||||
|
failure_threshold: 1,
|
||||||
|
};
|
||||||
|
// p1 fails once then succeeds (fail_then_ok with n=1 would work,
|
||||||
|
// but we use always_fail to prove it's skipped, then swap).
|
||||||
|
let p1 = Arc::new(MultiCallMockProvider::fail_then_ok("p1", 2));
|
||||||
|
let p2 = Arc::new(MultiCallMockProvider::always_ok("p2"));
|
||||||
|
|
||||||
|
let failover =
|
||||||
|
FailoverProvider::with_cooldown(vec![p1.clone(), p2.clone()], config).unwrap();
|
||||||
|
|
||||||
|
// Request 1: p1 fails (threshold=1, enters cooldown immediately), p2 succeeds.
|
||||||
|
let r = failover.complete(make_request()).await.unwrap();
|
||||||
|
assert_eq!(r.content, "p2 ok");
|
||||||
|
assert_eq!(p1.call_count(), 1);
|
||||||
|
|
||||||
|
// Request 2: p1 in cooldown, skipped. Only p2 called.
|
||||||
|
// (But cooldown is 1ms, so wait a bit to let it expire.)
|
||||||
|
tokio::time::sleep(Duration::from_millis(5)).await;
|
||||||
|
|
||||||
|
// After sleep, cooldown should have expired. p1 gets tried again.
|
||||||
|
// p1 is set to fail 2 times total, so call #2 (index 1) still fails.
|
||||||
|
// But it proves p1 was attempted again after cooldown expired.
|
||||||
|
let r = failover.complete(make_request()).await.unwrap();
|
||||||
|
assert_eq!(p1.call_count(), 2); // p1 was retried
|
||||||
|
assert_eq!(r.content, "p2 ok"); // p2 handled it
|
||||||
|
|
||||||
|
// Wait again for cooldown to expire, p1 call #3 (index 2) succeeds.
|
||||||
|
tokio::time::sleep(Duration::from_millis(5)).await;
|
||||||
|
let r = failover.complete(make_request()).await.unwrap();
|
||||||
|
assert_eq!(r.content, "p1 ok");
|
||||||
|
assert_eq!(p1.call_count(), 3);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Cooldown test 3: Never skip all providers — oldest-cooled one is tried.
|
||||||
|
#[tokio::test]
|
||||||
|
async fn never_skip_all_providers() {
|
||||||
|
let config = CooldownConfig {
|
||||||
|
cooldown_duration: Duration::from_secs(300),
|
||||||
|
failure_threshold: 1,
|
||||||
|
};
|
||||||
|
// Both providers always fail.
|
||||||
|
let p1 = Arc::new(MultiCallMockProvider::always_fail("p1"));
|
||||||
|
let p2 = Arc::new(MultiCallMockProvider::always_fail("p2"));
|
||||||
|
|
||||||
|
let failover =
|
||||||
|
FailoverProvider::with_cooldown(vec![p1.clone(), p2.clone()], config).unwrap();
|
||||||
|
|
||||||
|
// Request 1: both tried, both fail, both enter cooldown.
|
||||||
|
let _ = failover.complete(make_request()).await;
|
||||||
|
assert_eq!(p1.call_count(), 1);
|
||||||
|
assert_eq!(p2.call_count(), 1);
|
||||||
|
|
||||||
|
// Request 2: all in cooldown, but the oldest-cooled one (p1, activated
|
||||||
|
// first) should be tried.
|
||||||
|
let prev_total = p1.call_count() + p2.call_count();
|
||||||
|
let _ = failover.complete(make_request()).await;
|
||||||
|
let new_total = p1.call_count() + p2.call_count();
|
||||||
|
// Exactly one more call was made (to the oldest-cooled provider).
|
||||||
|
assert_eq!(new_total, prev_total + 1);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Cooldown test 4: Success resets failure count so it never reaches threshold.
|
||||||
|
//
|
||||||
|
// With threshold=3, accumulate 2 failures then succeed. Verify the
|
||||||
|
// atomic counter is back to 0 and no cooldown was activated. Then
|
||||||
|
// use a second provider pair to show that without the reset, 3
|
||||||
|
// consecutive failures DO trigger cooldown (control case).
|
||||||
|
#[tokio::test]
|
||||||
|
async fn reset_on_success() {
|
||||||
|
let config = CooldownConfig {
|
||||||
|
cooldown_duration: Duration::from_secs(300),
|
||||||
|
failure_threshold: 3,
|
||||||
|
};
|
||||||
|
// p1 fails for calls 0,1 then succeeds on call 2+.
|
||||||
|
let p1 = Arc::new(MultiCallMockProvider::fail_then_ok("p1", 2));
|
||||||
|
let p2 = Arc::new(MultiCallMockProvider::always_ok("p2"));
|
||||||
|
|
||||||
|
let failover =
|
||||||
|
FailoverProvider::with_cooldown(vec![p1.clone(), p2.clone()], config.clone()).unwrap();
|
||||||
|
|
||||||
|
// Request 1: p1 fails (failure_count=1), p2 succeeds.
|
||||||
|
let r = failover.complete(make_request()).await.unwrap();
|
||||||
|
assert_eq!(r.content, "p2 ok");
|
||||||
|
|
||||||
|
// Request 2: p1 fails (failure_count=2, still below threshold=3), p2 succeeds.
|
||||||
|
let r = failover.complete(make_request()).await.unwrap();
|
||||||
|
assert_eq!(r.content, "p2 ok");
|
||||||
|
assert_eq!(p1.call_count(), 2);
|
||||||
|
|
||||||
|
// Request 3: p1 succeeds (call index 2) → counter resets to 0.
|
||||||
|
let r = failover.complete(make_request()).await.unwrap();
|
||||||
|
assert_eq!(r.content, "p1 ok");
|
||||||
|
assert_eq!(p1.call_count(), 3);
|
||||||
|
|
||||||
|
// Verify counter was reset to 0 and no cooldown activated.
|
||||||
|
let nanos = failover.now_nanos();
|
||||||
|
let cooldown_nanos = failover.cooldown_config.cooldown_duration.as_nanos() as u64;
|
||||||
|
assert!(!failover.cooldowns[0].is_in_cooldown(nanos, cooldown_nanos));
|
||||||
|
assert_eq!(
|
||||||
|
failover.cooldowns[0].failure_count.load(Ordering::Relaxed),
|
||||||
|
0
|
||||||
|
);
|
||||||
|
|
||||||
|
// Control: without a success in the middle, 3 failures DO trigger cooldown.
|
||||||
|
let p3 = Arc::new(MultiCallMockProvider::always_fail("p3"));
|
||||||
|
let p4 = Arc::new(MultiCallMockProvider::always_ok("p4"));
|
||||||
|
let control =
|
||||||
|
FailoverProvider::with_cooldown(vec![p3.clone(), p4.clone()], config).unwrap();
|
||||||
|
for _ in 0..3 {
|
||||||
|
let _ = control.complete(make_request()).await.unwrap();
|
||||||
|
}
|
||||||
|
let nanos = control.now_nanos();
|
||||||
|
assert!(control.cooldowns[0].is_in_cooldown(nanos, cooldown_nanos));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Cooldown test 5: threshold-1 failures don't trigger cooldown, threshold does.
|
||||||
|
#[tokio::test]
|
||||||
|
async fn threshold_boundary() {
|
||||||
|
let config = CooldownConfig {
|
||||||
|
cooldown_duration: Duration::from_secs(300),
|
||||||
|
failure_threshold: 3,
|
||||||
|
};
|
||||||
|
let p1 = Arc::new(MultiCallMockProvider::always_fail("p1"));
|
||||||
|
let p2 = Arc::new(MultiCallMockProvider::always_ok("p2"));
|
||||||
|
|
||||||
|
let failover =
|
||||||
|
FailoverProvider::with_cooldown(vec![p1.clone(), p2.clone()], config).unwrap();
|
||||||
|
|
||||||
|
// 2 requests: p1 fails twice (below threshold of 3), not in cooldown.
|
||||||
|
for _ in 0..2 {
|
||||||
|
let r = failover.complete(make_request()).await.unwrap();
|
||||||
|
assert_eq!(r.content, "p2 ok");
|
||||||
|
}
|
||||||
|
assert_eq!(p1.call_count(), 2);
|
||||||
|
|
||||||
|
// p1 should still be available (not in cooldown).
|
||||||
|
let nanos = failover.now_nanos();
|
||||||
|
let cooldown_nanos = failover.cooldown_config.cooldown_duration.as_nanos() as u64;
|
||||||
|
assert!(!failover.cooldowns[0].is_in_cooldown(nanos, cooldown_nanos));
|
||||||
|
|
||||||
|
// 3rd request: p1 fails → reaches threshold → enters cooldown.
|
||||||
|
let r = failover.complete(make_request()).await.unwrap();
|
||||||
|
assert_eq!(r.content, "p2 ok");
|
||||||
|
assert_eq!(p1.call_count(), 3);
|
||||||
|
|
||||||
|
let nanos = failover.now_nanos();
|
||||||
|
assert!(failover.cooldowns[0].is_in_cooldown(nanos, cooldown_nanos));
|
||||||
|
|
||||||
|
// 4th request: p1 should be skipped.
|
||||||
|
let prev = p1.call_count();
|
||||||
|
let r = failover.complete(make_request()).await.unwrap();
|
||||||
|
assert_eq!(r.content, "p2 ok");
|
||||||
|
assert_eq!(p1.call_count(), prev); // not called
|
||||||
|
}
|
||||||
|
|
||||||
|
// Cooldown test 6: Non-retryable error returns immediately, no failure bump.
|
||||||
|
#[tokio::test]
|
||||||
|
async fn non_retryable_does_not_increment_cooldown() {
|
||||||
|
let config = CooldownConfig {
|
||||||
|
cooldown_duration: Duration::from_secs(300),
|
||||||
|
failure_threshold: 1,
|
||||||
|
};
|
||||||
|
let p1 = Arc::new(MultiCallMockProvider::always_fail_non_retryable("p1"));
|
||||||
|
let p2 = Arc::new(MultiCallMockProvider::always_ok("p2"));
|
||||||
|
|
||||||
|
let failover =
|
||||||
|
FailoverProvider::with_cooldown(vec![p1.clone(), p2.clone()], config).unwrap();
|
||||||
|
|
||||||
|
// Non-retryable error should return immediately.
|
||||||
|
let err = failover.complete(make_request()).await.unwrap_err();
|
||||||
|
assert!(matches!(err, LlmError::AuthFailed { .. }));
|
||||||
|
assert_eq!(p1.call_count(), 1);
|
||||||
|
// p2 should NOT have been called (non-retryable = no failover).
|
||||||
|
assert_eq!(p2.call_count(), 0);
|
||||||
|
|
||||||
|
// p1 should NOT be in cooldown (non-retryable doesn't bump count).
|
||||||
|
let nanos = failover.now_nanos();
|
||||||
|
let cooldown_nanos = failover.cooldown_config.cooldown_duration.as_nanos() as u64;
|
||||||
|
assert!(!failover.cooldowns[0].is_in_cooldown(nanos, cooldown_nanos));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Cooldown test 7: Three providers, first in cooldown, second/third available.
|
||||||
|
#[tokio::test]
|
||||||
|
async fn three_providers_mixed_cooldown() {
|
||||||
|
let config = CooldownConfig {
|
||||||
|
cooldown_duration: Duration::from_secs(300),
|
||||||
|
failure_threshold: 1,
|
||||||
|
};
|
||||||
|
let p1 = Arc::new(MultiCallMockProvider::always_fail("p1"));
|
||||||
|
let p2 = Arc::new(MultiCallMockProvider::always_ok("p2"));
|
||||||
|
let p3 = Arc::new(MultiCallMockProvider::always_ok("p3"));
|
||||||
|
|
||||||
|
let failover =
|
||||||
|
FailoverProvider::with_cooldown(vec![p1.clone(), p2.clone(), p3.clone()], config)
|
||||||
|
.unwrap();
|
||||||
|
|
||||||
|
// Request 1: p1 fails → enters cooldown (threshold=1), p2 succeeds.
|
||||||
|
let r = failover.complete(make_request()).await.unwrap();
|
||||||
|
assert_eq!(r.content, "p2 ok");
|
||||||
|
assert_eq!(p1.call_count(), 1);
|
||||||
|
|
||||||
|
// Request 2: p1 skipped (cooldown), p2 and p3 available.
|
||||||
|
let prev = p1.call_count();
|
||||||
|
let r = failover.complete(make_request()).await.unwrap();
|
||||||
|
assert_eq!(r.content, "p2 ok");
|
||||||
|
assert_eq!(p1.call_count(), prev); // p1 skipped
|
||||||
|
}
|
||||||
|
|
||||||
// Test: is_retryable correctly classifies errors.
|
// Test: is_retryable correctly classifies errors.
|
||||||
#[test]
|
#[test]
|
||||||
fn retryable_classification() {
|
fn retryable_classification() {
|
||||||
|
|||||||
+106
-1
@@ -17,7 +17,7 @@ mod retry;
|
|||||||
mod rig_adapter;
|
mod rig_adapter;
|
||||||
pub mod session;
|
pub mod session;
|
||||||
|
|
||||||
pub use failover::FailoverProvider;
|
pub use failover::{CooldownConfig, FailoverProvider};
|
||||||
pub use nearai::{ModelInfo, NearAiProvider};
|
pub use nearai::{ModelInfo, NearAiProvider};
|
||||||
pub use nearai_chat::NearAiChatProvider;
|
pub use nearai_chat::NearAiChatProvider;
|
||||||
pub use provider::{
|
pub use provider::{
|
||||||
@@ -183,3 +183,108 @@ fn create_openai_compatible_provider(config: &LlmConfig) -> Result<Arc<dyn LlmPr
|
|||||||
);
|
);
|
||||||
Ok(Arc::new(RigAdapter::new(model, &compat.model)))
|
Ok(Arc::new(RigAdapter::new(model, &compat.model)))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Create a cheap/fast LLM provider for lightweight tasks (heartbeat, routing, evaluation).
|
||||||
|
///
|
||||||
|
/// Uses `NEARAI_CHEAP_MODEL` if set, otherwise falls back to the main provider.
|
||||||
|
/// Currently only supports NEAR AI backends (Responses and ChatCompletions modes).
|
||||||
|
pub fn create_cheap_llm_provider(
|
||||||
|
config: &LlmConfig,
|
||||||
|
session: Arc<SessionManager>,
|
||||||
|
) -> Result<Option<Arc<dyn LlmProvider>>, LlmError> {
|
||||||
|
let Some(ref cheap_model) = config.nearai.cheap_model else {
|
||||||
|
return Ok(None);
|
||||||
|
};
|
||||||
|
|
||||||
|
if config.backend != LlmBackend::NearAi {
|
||||||
|
tracing::warn!(
|
||||||
|
"NEARAI_CHEAP_MODEL is set but LLM_BACKEND is {:?}, not NearAi. \
|
||||||
|
Cheap model setting will be ignored.",
|
||||||
|
config.backend
|
||||||
|
);
|
||||||
|
return Ok(None);
|
||||||
|
}
|
||||||
|
|
||||||
|
let mut cheap_config = config.nearai.clone();
|
||||||
|
cheap_config.model = cheap_model.clone();
|
||||||
|
|
||||||
|
tracing::info!("Cheap LLM provider: {}", cheap_model);
|
||||||
|
|
||||||
|
match cheap_config.api_mode {
|
||||||
|
NearAiApiMode::Responses => Ok(Some(Arc::new(NearAiProvider::new(cheap_config, session)))),
|
||||||
|
NearAiApiMode::ChatCompletions => {
|
||||||
|
Ok(Some(Arc::new(NearAiChatProvider::new(cheap_config)?)))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
use crate::config::{LlmBackend, NearAiApiMode, NearAiConfig};
|
||||||
|
use std::path::PathBuf;
|
||||||
|
|
||||||
|
fn test_nearai_config() -> NearAiConfig {
|
||||||
|
NearAiConfig {
|
||||||
|
model: "test-model".to_string(),
|
||||||
|
cheap_model: None,
|
||||||
|
base_url: "https://api.near.ai".to_string(),
|
||||||
|
auth_base_url: "https://private.near.ai".to_string(),
|
||||||
|
session_path: PathBuf::from("/tmp/test-session.json"),
|
||||||
|
api_mode: NearAiApiMode::Responses,
|
||||||
|
api_key: None,
|
||||||
|
fallback_model: None,
|
||||||
|
max_retries: 3,
|
||||||
|
failover_cooldown_secs: 300,
|
||||||
|
failover_cooldown_threshold: 3,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn test_llm_config() -> LlmConfig {
|
||||||
|
LlmConfig {
|
||||||
|
backend: LlmBackend::NearAi,
|
||||||
|
nearai: test_nearai_config(),
|
||||||
|
openai: None,
|
||||||
|
anthropic: None,
|
||||||
|
ollama: None,
|
||||||
|
openai_compatible: None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_create_cheap_llm_provider_returns_none_when_not_configured() {
|
||||||
|
let config = test_llm_config();
|
||||||
|
let session = Arc::new(SessionManager::new(SessionConfig::default()));
|
||||||
|
|
||||||
|
let result = create_cheap_llm_provider(&config, session);
|
||||||
|
assert!(result.is_ok());
|
||||||
|
assert!(result.unwrap().is_none());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_create_cheap_llm_provider_creates_provider_when_configured() {
|
||||||
|
let mut config = test_llm_config();
|
||||||
|
config.nearai.cheap_model = Some("cheap-test-model".to_string());
|
||||||
|
|
||||||
|
let session = Arc::new(SessionManager::new(SessionConfig::default()));
|
||||||
|
let result = create_cheap_llm_provider(&config, session);
|
||||||
|
|
||||||
|
assert!(result.is_ok());
|
||||||
|
let provider = result.unwrap();
|
||||||
|
assert!(provider.is_some());
|
||||||
|
assert_eq!(provider.unwrap().model_name(), "cheap-test-model");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_create_cheap_llm_provider_ignored_for_non_nearai_backend() {
|
||||||
|
let mut config = test_llm_config();
|
||||||
|
config.backend = LlmBackend::OpenAi;
|
||||||
|
config.nearai.cheap_model = Some("cheap-test-model".to_string());
|
||||||
|
|
||||||
|
let session = Arc::new(SessionManager::new(SessionConfig::default()));
|
||||||
|
let result = create_cheap_llm_provider(&config, session);
|
||||||
|
|
||||||
|
assert!(result.is_ok());
|
||||||
|
assert!(result.unwrap().is_none());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+16
-3
@@ -428,17 +428,30 @@ impl SessionManager {
|
|||||||
})?;
|
})?;
|
||||||
|
|
||||||
let user_id = self.user_id.read().await.clone();
|
let user_id = self.user_id.read().await.clone();
|
||||||
let value = store
|
let value = if let Some(value) = store
|
||||||
.get_setting(&user_id, "nearai.session_token")
|
.get_setting(&user_id, "nearai.session_token")
|
||||||
.await
|
.await
|
||||||
|
.map_err(|e| LlmError::SessionRenewalFailed {
|
||||||
|
provider: "nearai".to_string(),
|
||||||
|
reason: format!("DB query failed: {}", e),
|
||||||
|
})? {
|
||||||
|
value
|
||||||
|
} else {
|
||||||
|
tracing::warn!(
|
||||||
|
"nearai.session_token missing; falling back to legacy nearai.session for backwards compatibility"
|
||||||
|
);
|
||||||
|
store
|
||||||
|
.get_setting(&user_id, "nearai.session")
|
||||||
|
.await
|
||||||
.map_err(|e| LlmError::SessionRenewalFailed {
|
.map_err(|e| LlmError::SessionRenewalFailed {
|
||||||
provider: "nearai".to_string(),
|
provider: "nearai".to_string(),
|
||||||
reason: format!("DB query failed: {}", e),
|
reason: format!("DB query failed: {}", e),
|
||||||
})?
|
})?
|
||||||
.ok_or_else(|| LlmError::SessionRenewalFailed {
|
.ok_or(LlmError::SessionRenewalFailed {
|
||||||
provider: "nearai".to_string(),
|
provider: "nearai".to_string(),
|
||||||
reason: "No session in DB".to_string(),
|
reason: "No session in DB".to_string(),
|
||||||
})?;
|
})?
|
||||||
|
};
|
||||||
|
|
||||||
let session: SessionData =
|
let session: SessionData =
|
||||||
serde_json::from_value(value).map_err(|e| LlmError::SessionRenewalFailed {
|
serde_json::from_value(value).map_err(|e| LlmError::SessionRenewalFailed {
|
||||||
|
|||||||
+163
-57
@@ -22,9 +22,10 @@ use ironclaw::{
|
|||||||
config::Config,
|
config::Config,
|
||||||
context::ContextManager,
|
context::ContextManager,
|
||||||
extensions::ExtensionManager,
|
extensions::ExtensionManager,
|
||||||
|
hooks::HookRegistry,
|
||||||
llm::{
|
llm::{
|
||||||
FailoverProvider, LlmProvider, SessionConfig, create_llm_provider,
|
CooldownConfig, FailoverProvider, LlmProvider, SessionConfig, create_cheap_llm_provider,
|
||||||
create_llm_provider_with_config, create_session_manager,
|
create_llm_provider, create_llm_provider_with_config, create_session_manager,
|
||||||
},
|
},
|
||||||
orchestrator::{
|
orchestrator::{
|
||||||
ContainerJobConfig, ContainerJobManager, OrchestratorApi, TokenStore,
|
ContainerJobConfig, ContainerJobManager, OrchestratorApi, TokenStore,
|
||||||
@@ -48,7 +49,6 @@ use ironclaw::secrets::PostgresSecretsStore;
|
|||||||
use ironclaw::secrets::SecretsCrypto;
|
use ironclaw::secrets::SecretsCrypto;
|
||||||
#[cfg(any(feature = "postgres", feature = "libsql"))]
|
#[cfg(any(feature = "postgres", feature = "libsql"))]
|
||||||
use ironclaw::setup::{SetupConfig, SetupWizard};
|
use ironclaw::setup::{SetupConfig, SetupWizard};
|
||||||
|
|
||||||
#[tokio::main]
|
#[tokio::main]
|
||||||
async fn main() -> anyhow::Result<()> {
|
async fn main() -> anyhow::Result<()> {
|
||||||
let cli = Cli::parse();
|
let cli = Cli::parse();
|
||||||
@@ -308,8 +308,11 @@ async fn main() -> anyhow::Result<()> {
|
|||||||
};
|
};
|
||||||
let session = create_session_manager(session_config).await;
|
let session = create_session_manager(session_config).await;
|
||||||
|
|
||||||
// Ensure we're authenticated before proceeding (only needed for NEAR AI backend)
|
// Session-based auth is only needed for NEAR AI backend without an API key.
|
||||||
if config.llm.backend == ironclaw::config::LlmBackend::NearAi {
|
// ChatCompletions mode with an API key skips session auth entirely.
|
||||||
|
if config.llm.backend == ironclaw::config::LlmBackend::NearAi
|
||||||
|
&& config.llm.nearai.api_key.is_none()
|
||||||
|
{
|
||||||
session.ensure_authenticated().await?;
|
session.ensure_authenticated().await?;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -335,7 +338,10 @@ async fn main() -> anyhow::Result<()> {
|
|||||||
let repl_channel = if let Some(ref msg) = cli.message {
|
let repl_channel = if let Some(ref msg) = cli.message {
|
||||||
Some(ReplChannel::with_message(msg.clone()))
|
Some(ReplChannel::with_message(msg.clone()))
|
||||||
} else if config.channels.cli.enabled {
|
} else if config.channels.cli.enabled {
|
||||||
Some(ReplChannel::new())
|
let repl = ReplChannel::new();
|
||||||
|
// Suppress the one-liner banner; boot screen will be shown instead.
|
||||||
|
repl.suppress_banner();
|
||||||
|
Some(repl)
|
||||||
} else {
|
} else {
|
||||||
None
|
None
|
||||||
};
|
};
|
||||||
@@ -444,6 +450,72 @@ async fn main() -> anyhow::Result<()> {
|
|||||||
tracing::warn!("Failed to cleanup stale sandbox jobs: {}", e);
|
tracing::warn!("Failed to cleanup stale sandbox jobs: {}", e);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Create secrets store early: needed for injecting LLM API keys from encrypted
|
||||||
|
// storage before creating the LLM provider, and later for MCP auth + WASM channels.
|
||||||
|
//
|
||||||
|
// When both `postgres` and `libsql` features are compiled, the runtime-selected
|
||||||
|
// backend determines which store is created: whichever DB init branch ran will
|
||||||
|
// have set its handle (pg_pool or libsql_db), and the or_else chain picks it up.
|
||||||
|
let secrets_store: Option<Arc<dyn SecretsStore + Send + Sync>> =
|
||||||
|
if let Some(master_key) = config.secrets.master_key() {
|
||||||
|
match SecretsCrypto::new(master_key.clone()) {
|
||||||
|
Ok(crypto) => {
|
||||||
|
let crypto = Arc::new(crypto);
|
||||||
|
let store: Option<Arc<dyn SecretsStore + Send + Sync>> = None;
|
||||||
|
|
||||||
|
#[cfg(feature = "libsql")]
|
||||||
|
let store = store.or_else(|| {
|
||||||
|
libsql_db.take().map(|db| {
|
||||||
|
Arc::new(LibSqlSecretsStore::new(db, Arc::clone(&crypto)))
|
||||||
|
as Arc<dyn SecretsStore + Send + Sync>
|
||||||
|
})
|
||||||
|
});
|
||||||
|
|
||||||
|
#[cfg(feature = "postgres")]
|
||||||
|
let store = store.or_else(|| {
|
||||||
|
pg_pool.as_ref().map(|pool| {
|
||||||
|
Arc::new(PostgresSecretsStore::new(pool.clone(), Arc::clone(&crypto)))
|
||||||
|
as Arc<dyn SecretsStore + Send + Sync>
|
||||||
|
})
|
||||||
|
});
|
||||||
|
|
||||||
|
store
|
||||||
|
}
|
||||||
|
Err(e) => {
|
||||||
|
tracing::warn!("Failed to initialize secrets crypto: {}", e);
|
||||||
|
#[cfg(feature = "libsql")]
|
||||||
|
let _ = libsql_db.take();
|
||||||
|
None
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
#[cfg(feature = "libsql")]
|
||||||
|
let _ = libsql_db.take();
|
||||||
|
None
|
||||||
|
};
|
||||||
|
|
||||||
|
// Inject LLM API keys from the encrypted secrets store into a thread-safe
|
||||||
|
// overlay so that optional_env() (used by LlmConfig::resolve()) picks them
|
||||||
|
// up. Then re-resolve LlmConfig with the newly available keys (backend may
|
||||||
|
// have been set during onboarding but the API key is in the secrets store).
|
||||||
|
if let Some(ref secrets) = secrets_store {
|
||||||
|
ironclaw::config::inject_llm_keys_from_secrets(secrets.as_ref(), "default").await;
|
||||||
|
|
||||||
|
// Re-resolve LlmConfig now that secrets overlay has been populated
|
||||||
|
if let Some(ref db_ref) = db {
|
||||||
|
match Config::from_db(db_ref.as_ref(), "default").await {
|
||||||
|
Ok(refreshed) => {
|
||||||
|
config = refreshed;
|
||||||
|
tracing::debug!("LlmConfig re-resolved after secret injection");
|
||||||
|
}
|
||||||
|
Err(e) => {
|
||||||
|
tracing::warn!("Failed to re-resolve config after secret injection: {}", e);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Initialize LLM provider (clone session so we can reuse it for embeddings)
|
// Initialize LLM provider (clone session so we can reuse it for embeddings)
|
||||||
let llm = create_llm_provider(&config.llm, session.clone())?;
|
let llm = create_llm_provider(&config.llm, session.clone())?;
|
||||||
tracing::info!("LLM provider initialized: {}", llm.model_name());
|
tracing::info!("LLM provider initialized: {}", llm.model_name());
|
||||||
@@ -464,11 +536,26 @@ async fn main() -> anyhow::Result<()> {
|
|||||||
fallback = %fallback.model_name(),
|
fallback = %fallback.model_name(),
|
||||||
"LLM failover enabled"
|
"LLM failover enabled"
|
||||||
);
|
);
|
||||||
Arc::new(FailoverProvider::new(vec![llm, fallback])?)
|
let cooldown_config = CooldownConfig {
|
||||||
|
cooldown_duration: std::time::Duration::from_secs(
|
||||||
|
config.llm.nearai.failover_cooldown_secs,
|
||||||
|
),
|
||||||
|
failure_threshold: config.llm.nearai.failover_cooldown_threshold,
|
||||||
|
};
|
||||||
|
Arc::new(FailoverProvider::with_cooldown(
|
||||||
|
vec![llm, fallback],
|
||||||
|
cooldown_config,
|
||||||
|
)?)
|
||||||
} else {
|
} else {
|
||||||
llm
|
llm
|
||||||
};
|
};
|
||||||
|
|
||||||
|
// Initialize cheap LLM provider for lightweight tasks (heartbeat, evaluation)
|
||||||
|
let cheap_llm = create_cheap_llm_provider(&config.llm, session.clone())?;
|
||||||
|
if let Some(ref cheap) = cheap_llm {
|
||||||
|
tracing::info!("Cheap LLM provider initialized: {}", cheap.model_name());
|
||||||
|
}
|
||||||
|
|
||||||
// Initialize safety layer
|
// Initialize safety layer
|
||||||
let safety = Arc::new(SafetyLayer::new(&config.safety));
|
let safety = Arc::new(SafetyLayer::new(&config.safety));
|
||||||
tracing::info!("Safety layer initialized");
|
tracing::info!("Safety layer initialized");
|
||||||
@@ -542,49 +629,6 @@ async fn main() -> anyhow::Result<()> {
|
|||||||
tracing::info!("Builder mode enabled");
|
tracing::info!("Builder mode enabled");
|
||||||
}
|
}
|
||||||
|
|
||||||
// Create secrets store if master key is configured (needed for MCP auth and WASM channels).
|
|
||||||
//
|
|
||||||
// When both `postgres` and `libsql` features are compiled, the runtime-selected
|
|
||||||
// backend determines which store is created: whichever DB init branch ran will
|
|
||||||
// have set its handle (pg_pool or libsql_db), and the or_else chain picks it up.
|
|
||||||
let secrets_store: Option<Arc<dyn SecretsStore + Send + Sync>> =
|
|
||||||
if let Some(master_key) = config.secrets.master_key() {
|
|
||||||
match SecretsCrypto::new(master_key.clone()) {
|
|
||||||
Ok(crypto) => {
|
|
||||||
let crypto = Arc::new(crypto);
|
|
||||||
let store: Option<Arc<dyn SecretsStore + Send + Sync>> = None;
|
|
||||||
|
|
||||||
#[cfg(feature = "libsql")]
|
|
||||||
let store = store.or_else(|| {
|
|
||||||
libsql_db.take().map(|db| {
|
|
||||||
Arc::new(LibSqlSecretsStore::new(db, Arc::clone(&crypto)))
|
|
||||||
as Arc<dyn SecretsStore + Send + Sync>
|
|
||||||
})
|
|
||||||
});
|
|
||||||
|
|
||||||
#[cfg(feature = "postgres")]
|
|
||||||
let store = store.or_else(|| {
|
|
||||||
pg_pool.as_ref().map(|pool| {
|
|
||||||
Arc::new(PostgresSecretsStore::new(pool.clone(), Arc::clone(&crypto)))
|
|
||||||
as Arc<dyn SecretsStore + Send + Sync>
|
|
||||||
})
|
|
||||||
});
|
|
||||||
|
|
||||||
store
|
|
||||||
}
|
|
||||||
Err(e) => {
|
|
||||||
tracing::warn!("Failed to initialize secrets crypto: {}", e);
|
|
||||||
#[cfg(feature = "libsql")]
|
|
||||||
let _ = libsql_db.take();
|
|
||||||
None
|
|
||||||
}
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
#[cfg(feature = "libsql")]
|
|
||||||
let _ = libsql_db.take();
|
|
||||||
None
|
|
||||||
};
|
|
||||||
|
|
||||||
let mcp_session_manager = Arc::new(McpSessionManager::new());
|
let mcp_session_manager = Arc::new(McpSessionManager::new());
|
||||||
|
|
||||||
// Create WASM tool runtime (sync, just builds the wasmtime engine)
|
// Create WASM tool runtime (sync, just builds the wasmtime engine)
|
||||||
@@ -854,12 +898,14 @@ async fn main() -> anyhow::Result<()> {
|
|||||||
|
|
||||||
// Initialize channel manager
|
// Initialize channel manager
|
||||||
let mut channels = ChannelManager::new();
|
let mut channels = ChannelManager::new();
|
||||||
|
let mut channel_names: Vec<String> = Vec::new();
|
||||||
|
|
||||||
if let Some(repl) = repl_channel {
|
if let Some(repl) = repl_channel {
|
||||||
channels.add(Box::new(repl));
|
channels.add(Box::new(repl));
|
||||||
if cli.message.is_some() {
|
if cli.message.is_some() {
|
||||||
tracing::info!("Single message mode");
|
tracing::info!("Single message mode");
|
||||||
} else {
|
} else {
|
||||||
|
channel_names.push("repl".to_string());
|
||||||
tracing::info!("REPL mode enabled");
|
tracing::info!("REPL mode enabled");
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -995,6 +1041,7 @@ async fn main() -> anyhow::Result<()> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
channel_names.push(channel_name.clone());
|
||||||
channels.add(Box::new(SharedWasmChannel::new(channel_arc)));
|
channels.add(Box::new(SharedWasmChannel::new(channel_arc)));
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1039,6 +1086,7 @@ async fn main() -> anyhow::Result<()> {
|
|||||||
.parse()
|
.parse()
|
||||||
.expect("HttpConfig host:port must be a valid SocketAddr"),
|
.expect("HttpConfig host:port must be a valid SocketAddr"),
|
||||||
);
|
);
|
||||||
|
channel_names.push("http".to_string());
|
||||||
channels.add(Box::new(http_channel));
|
channels.add(Box::new(http_channel));
|
||||||
tracing::info!(
|
tracing::info!(
|
||||||
"HTTP channel enabled on {}:{}",
|
"HTTP channel enabled on {}:{}",
|
||||||
@@ -1101,8 +1149,11 @@ async fn main() -> anyhow::Result<()> {
|
|||||||
// Create context manager (shared between job tools and agent)
|
// Create context manager (shared between job tools and agent)
|
||||||
let context_manager = Arc::new(ContextManager::new(config.agent.max_parallel_jobs));
|
let context_manager = Arc::new(ContextManager::new(config.agent.max_parallel_jobs));
|
||||||
|
|
||||||
|
// Create hook registry
|
||||||
|
let hooks = Arc::new(HookRegistry::new());
|
||||||
|
|
||||||
// Create session manager (shared between agent and web gateway)
|
// Create session manager (shared between agent and web gateway)
|
||||||
let session_manager = Arc::new(SessionManager::new());
|
let session_manager = Arc::new(SessionManager::new().with_hooks(hooks.clone()));
|
||||||
|
|
||||||
// Register job tools (sandbox deps auto-injected when container_job_manager is available)
|
// Register job tools (sandbox deps auto-injected when container_job_manager is available)
|
||||||
tools.register_job_tools(
|
tools.register_job_tools(
|
||||||
@@ -1112,6 +1163,7 @@ async fn main() -> anyhow::Result<()> {
|
|||||||
);
|
);
|
||||||
|
|
||||||
// Add web gateway channel if configured
|
// Add web gateway channel if configured
|
||||||
|
let mut gateway_url: Option<String> = None;
|
||||||
if let Some(ref gw_config) = config.channels.gateway {
|
if let Some(ref gw_config) = config.channels.gateway {
|
||||||
let mut gw = GatewayChannel::new(gw_config.clone());
|
let mut gw = GatewayChannel::new(gw_config.clone());
|
||||||
if let Some(ref ws) = workspace {
|
if let Some(ref ws) = workspace {
|
||||||
@@ -1144,29 +1196,39 @@ async fn main() -> anyhow::Result<()> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
gateway_url = Some(format!(
|
||||||
|
"http://{}:{}/?token={}",
|
||||||
|
gw_config.host,
|
||||||
|
gw_config.port,
|
||||||
|
gw.auth_token()
|
||||||
|
));
|
||||||
|
|
||||||
tracing::info!(
|
tracing::info!(
|
||||||
"Web gateway enabled on {}:{}",
|
"Web gateway enabled on {}:{}",
|
||||||
gw_config.host,
|
gw_config.host,
|
||||||
gw_config.port
|
gw_config.port
|
||||||
);
|
);
|
||||||
tracing::info!(
|
tracing::info!("Web UI: http://{}:{}/", gw_config.host, gw_config.port);
|
||||||
"Web UI: http://{}:{}/?token={}",
|
|
||||||
gw_config.host,
|
|
||||||
gw_config.port,
|
|
||||||
gw.auth_token()
|
|
||||||
);
|
|
||||||
|
|
||||||
|
channel_names.push("gateway".to_string());
|
||||||
channels.add(Box::new(gw));
|
channels.add(Box::new(gw));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Capture boot screen info before moving Arcs into AgentDeps.
|
||||||
|
let boot_tool_count = tools.count();
|
||||||
|
let boot_llm_model = llm.model_name().to_string();
|
||||||
|
let boot_cheap_model = cheap_llm.as_ref().map(|c| c.model_name().to_string());
|
||||||
|
|
||||||
// Create and run the agent
|
// Create and run the agent
|
||||||
let deps = AgentDeps {
|
let deps = AgentDeps {
|
||||||
store: db,
|
store: db,
|
||||||
llm,
|
llm,
|
||||||
|
cheap_llm,
|
||||||
safety,
|
safety,
|
||||||
tools,
|
tools,
|
||||||
workspace,
|
workspace,
|
||||||
extension_manager,
|
extension_manager,
|
||||||
|
hooks,
|
||||||
};
|
};
|
||||||
let agent = Agent::new(
|
let agent = Agent::new(
|
||||||
config.agent.clone(),
|
config.agent.clone(),
|
||||||
@@ -1180,6 +1242,38 @@ async fn main() -> anyhow::Result<()> {
|
|||||||
|
|
||||||
tracing::info!("Agent initialized, starting main loop...");
|
tracing::info!("Agent initialized, starting main loop...");
|
||||||
|
|
||||||
|
// Print boot screen for interactive CLI mode (not single-message mode).
|
||||||
|
if config.channels.cli.enabled && cli.message.is_none() {
|
||||||
|
let boot_info = ironclaw::boot_screen::BootInfo {
|
||||||
|
version: env!("CARGO_PKG_VERSION").to_string(),
|
||||||
|
agent_name: config.agent.name.clone(),
|
||||||
|
llm_backend: config.llm.backend.to_string(),
|
||||||
|
llm_model: boot_llm_model,
|
||||||
|
cheap_model: boot_cheap_model,
|
||||||
|
db_backend: if cli.no_db {
|
||||||
|
"none".to_string()
|
||||||
|
} else {
|
||||||
|
config.database.backend.to_string()
|
||||||
|
},
|
||||||
|
db_connected: !cli.no_db,
|
||||||
|
tool_count: boot_tool_count,
|
||||||
|
gateway_url,
|
||||||
|
embeddings_enabled: config.embeddings.enabled,
|
||||||
|
embeddings_provider: if config.embeddings.enabled {
|
||||||
|
Some(config.embeddings.provider.clone())
|
||||||
|
} else {
|
||||||
|
None
|
||||||
|
},
|
||||||
|
heartbeat_enabled: config.heartbeat.enabled,
|
||||||
|
heartbeat_interval_secs: config.heartbeat.interval_secs,
|
||||||
|
sandbox_enabled: config.sandbox.enabled,
|
||||||
|
claude_code_enabled: config.claude_code.enabled,
|
||||||
|
routines_enabled: config.routines.enabled,
|
||||||
|
channels: channel_names,
|
||||||
|
};
|
||||||
|
ironclaw::boot_screen::print_boot_screen(&boot_info);
|
||||||
|
}
|
||||||
|
|
||||||
// Run the agent (blocks until shutdown)
|
// Run the agent (blocks until shutdown)
|
||||||
agent.run().await?;
|
agent.run().await?;
|
||||||
|
|
||||||
@@ -1207,6 +1301,18 @@ fn check_onboard_needed() -> Option<&'static str> {
|
|||||||
return Some("Database not configured");
|
return Some("Database not configured");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// First run (onboarding never completed and no session).
|
||||||
|
// Reads NEARAI_API_KEY env var directly because this function runs
|
||||||
|
// before Config is loaded -- Config::from_env() may fail without a
|
||||||
|
// database URL, which is what triggers onboarding in the first place.
|
||||||
|
if std::env::var("NEARAI_API_KEY").is_err() {
|
||||||
|
let settings = ironclaw::settings::Settings::load();
|
||||||
|
let session_path = ironclaw::llm::session::default_session_path();
|
||||||
|
if !settings.onboard_completed && !session_path.exists() {
|
||||||
|
return Some("First run");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
None
|
None
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+51
-5
@@ -40,8 +40,18 @@ pub struct Settings {
|
|||||||
#[serde(default)]
|
#[serde(default)]
|
||||||
pub secrets_master_key_source: KeySource,
|
pub secrets_master_key_source: KeySource,
|
||||||
|
|
||||||
// === Step 3: NEAR AI Auth ===
|
// === Step 3: Inference Provider ===
|
||||||
// Session stored separately in session.json
|
/// LLM backend: "nearai", "anthropic", "openai", "ollama", "openai_compatible".
|
||||||
|
#[serde(default)]
|
||||||
|
pub llm_backend: Option<String>,
|
||||||
|
|
||||||
|
/// Ollama base URL (when llm_backend = "ollama").
|
||||||
|
#[serde(default)]
|
||||||
|
pub ollama_base_url: Option<String>,
|
||||||
|
|
||||||
|
/// OpenAI-compatible endpoint base URL (when llm_backend = "openai_compatible").
|
||||||
|
#[serde(default)]
|
||||||
|
pub openai_compatible_base_url: Option<String>,
|
||||||
|
|
||||||
// === Step 4: Model Selection ===
|
// === Step 4: Model Selection ===
|
||||||
/// Currently selected model.
|
/// Currently selected model.
|
||||||
@@ -504,7 +514,11 @@ impl Settings {
|
|||||||
/// Each key is a dotted path (e.g., "agent.name"), value is a JSONB value.
|
/// Each key is a dotted path (e.g., "agent.name"), value is a JSONB value.
|
||||||
/// Missing keys get their default value.
|
/// Missing keys get their default value.
|
||||||
pub fn from_db_map(map: &std::collections::HashMap<String, serde_json::Value>) -> Self {
|
pub fn from_db_map(map: &std::collections::HashMap<String, serde_json::Value>) -> Self {
|
||||||
// Start with defaults, then overlay each DB setting
|
// Start with defaults, then overlay each DB setting.
|
||||||
|
//
|
||||||
|
// The settings table stores both Settings struct fields and app-specific
|
||||||
|
// data (e.g. nearai.session_token). Skip keys that don't correspond to
|
||||||
|
// a known Settings path.
|
||||||
let mut settings = Self::default();
|
let mut settings = Self::default();
|
||||||
|
|
||||||
for (key, value) in map {
|
for (key, value) in map {
|
||||||
@@ -513,11 +527,16 @@ impl Settings {
|
|||||||
serde_json::Value::String(s) => s.clone(),
|
serde_json::Value::String(s) => s.clone(),
|
||||||
serde_json::Value::Bool(b) => b.to_string(),
|
serde_json::Value::Bool(b) => b.to_string(),
|
||||||
serde_json::Value::Number(n) => n.to_string(),
|
serde_json::Value::Number(n) => n.to_string(),
|
||||||
serde_json::Value::Null => "null".to_string(),
|
serde_json::Value::Null => continue, // null means default, skip
|
||||||
other => other.to_string(),
|
other => other.to_string(),
|
||||||
};
|
};
|
||||||
|
|
||||||
if let Err(e) = settings.set(key, &value_str) {
|
match settings.set(key, &value_str) {
|
||||||
|
Ok(()) => {}
|
||||||
|
// The settings table stores both Settings fields and app-specific
|
||||||
|
// data (e.g. nearai.session_token). Silently skip unknown paths.
|
||||||
|
Err(e) if e.starts_with("Path not found") => {}
|
||||||
|
Err(e) => {
|
||||||
tracing::warn!(
|
tracing::warn!(
|
||||||
"Failed to apply DB setting '{}' = '{}': {}",
|
"Failed to apply DB setting '{}' = '{}': {}",
|
||||||
key,
|
key,
|
||||||
@@ -526,6 +545,7 @@ impl Settings {
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
settings
|
settings
|
||||||
}
|
}
|
||||||
@@ -858,4 +878,30 @@ mod tests {
|
|||||||
.unwrap();
|
.unwrap();
|
||||||
assert_eq!(settings.channels.telegram_owner_id, Some(987654321));
|
assert_eq!(settings.channels.telegram_owner_id, Some(987654321));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_llm_backend_round_trip() {
|
||||||
|
let dir = tempfile::tempdir().unwrap();
|
||||||
|
let path = dir.path().join("settings.json");
|
||||||
|
|
||||||
|
let settings = Settings {
|
||||||
|
llm_backend: Some("anthropic".to_string()),
|
||||||
|
ollama_base_url: Some("http://localhost:11434".to_string()),
|
||||||
|
openai_compatible_base_url: Some("http://my-vllm:8000/v1".to_string()),
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
let json = serde_json::to_string_pretty(&settings).unwrap();
|
||||||
|
std::fs::write(&path, json).unwrap();
|
||||||
|
|
||||||
|
let loaded = Settings::load_from(&path);
|
||||||
|
assert_eq!(loaded.llm_backend, Some("anthropic".to_string()));
|
||||||
|
assert_eq!(
|
||||||
|
loaded.ollama_base_url,
|
||||||
|
Some("http://localhost:11434".to_string())
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
loaded.openai_compatible_base_url,
|
||||||
|
Some("http://my-vllm:8000/v1".to_string())
|
||||||
|
);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,539 @@
|
|||||||
|
# Setup / Onboarding Specification
|
||||||
|
|
||||||
|
This document is the authoritative specification for IronClaw's onboarding
|
||||||
|
wizard. Any code change to `src/setup/` **must** keep this document in sync.
|
||||||
|
If a future contributor or coding agent modifies setup behavior, update this
|
||||||
|
file first, then adjust the code to match.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Entry Points
|
||||||
|
|
||||||
|
```
|
||||||
|
ironclaw onboard [--skip-auth] [--channels-only]
|
||||||
|
```
|
||||||
|
|
||||||
|
Explicit invocation. Loads `.env` files, runs the wizard, exits.
|
||||||
|
|
||||||
|
```
|
||||||
|
ironclaw (first run, no database configured)
|
||||||
|
```
|
||||||
|
|
||||||
|
Auto-detection via `check_onboard_needed()` in `main.rs`. Triggers when
|
||||||
|
none of these are true:
|
||||||
|
- `DATABASE_URL` env var is set
|
||||||
|
- `LIBSQL_PATH` env var is set
|
||||||
|
- `~/.ironclaw/ironclaw.db` exists on disk
|
||||||
|
|
||||||
|
The `--no-onboard` CLI flag suppresses auto-detection.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Startup Sequence (main.rs)
|
||||||
|
|
||||||
|
```
|
||||||
|
1. Parse CLI args
|
||||||
|
2. If Command::Onboard → load .env, run wizard, exit
|
||||||
|
3. If Command::Run or no command:
|
||||||
|
a. Load .env files (dotenvy::dotenv() then load_ironclaw_env())
|
||||||
|
b. check_onboard_needed() → run wizard if needed
|
||||||
|
c. Config::from_env() → build config from env vars
|
||||||
|
d. Create SessionManager → load session token
|
||||||
|
e. ensure_authenticated() → validate session (NEAR AI only)
|
||||||
|
f. ... rest of agent startup
|
||||||
|
```
|
||||||
|
|
||||||
|
**Critical ordering:** `.env` files must be loaded (step 3a) before
|
||||||
|
`Config::from_env()` (step 3c) because bootstrap vars like
|
||||||
|
`DATABASE_BACKEND` live in `~/.ironclaw/.env`.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## The 7-Step Wizard
|
||||||
|
|
||||||
|
### Overview
|
||||||
|
|
||||||
|
```
|
||||||
|
Step 1: Database Connection
|
||||||
|
Step 2: Security (master key)
|
||||||
|
Step 3: Inference Provider ← skipped if --skip-auth
|
||||||
|
Step 4: Model Selection
|
||||||
|
Step 5: Embeddings
|
||||||
|
Step 6: Channel Configuration
|
||||||
|
Step 7: Background Tasks (heartbeat)
|
||||||
|
↓
|
||||||
|
save_and_summarize()
|
||||||
|
```
|
||||||
|
|
||||||
|
`--channels-only` mode runs only Step 6, skipping everything else.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### Step 1: Database Connection
|
||||||
|
|
||||||
|
**Module:** `wizard.rs` → `step_database()`
|
||||||
|
|
||||||
|
**Goal:** Select backend, establish connection, run migrations.
|
||||||
|
|
||||||
|
**Decision tree:**
|
||||||
|
|
||||||
|
```
|
||||||
|
Both features compiled?
|
||||||
|
├─ Yes → DATABASE_BACKEND env var set?
|
||||||
|
│ ├─ Yes → use that backend
|
||||||
|
│ └─ No → interactive selection (PostgreSQL vs libSQL)
|
||||||
|
├─ Only postgres feature → step_database_postgres()
|
||||||
|
└─ Only libsql feature → step_database_libsql()
|
||||||
|
```
|
||||||
|
|
||||||
|
**PostgreSQL path** (`step_database_postgres`):
|
||||||
|
1. Check `DATABASE_URL` from env or settings
|
||||||
|
2. Test connection (creates `deadpool_postgres::Pool`)
|
||||||
|
3. Optionally run refinery migrations
|
||||||
|
4. Store pool in `self.db_pool`
|
||||||
|
|
||||||
|
**libSQL path** (`step_database_libsql`):
|
||||||
|
1. Offer local path (default: `~/.ironclaw/ironclaw.db`)
|
||||||
|
2. Optional Turso cloud sync (URL + auth token)
|
||||||
|
3. Test connection (creates `LibSqlBackend`)
|
||||||
|
4. Always run migrations (idempotent CREATE IF NOT EXISTS)
|
||||||
|
5. Store backend in `self.db_backend`
|
||||||
|
|
||||||
|
**Invariant:** After Step 1, exactly one of `self.db_pool` or
|
||||||
|
`self.db_backend` is `Some`. This is required for settings persistence
|
||||||
|
in `save_and_summarize()`.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### Step 2: Security (Master Key)
|
||||||
|
|
||||||
|
**Module:** `wizard.rs` → `step_security()`
|
||||||
|
|
||||||
|
**Goal:** Configure encryption for API tokens and secrets.
|
||||||
|
|
||||||
|
**Decision tree:**
|
||||||
|
|
||||||
|
```
|
||||||
|
SECRETS_MASTER_KEY env var set?
|
||||||
|
├─ Yes → use env var, done
|
||||||
|
└─ No → try get_master_key() from OS keychain
|
||||||
|
├─ Ok(bytes) → cache in self.secrets_crypto, ask "use existing?"
|
||||||
|
│ ├─ Yes → done (keychain)
|
||||||
|
│ └─ No → clear cache, fall through to options
|
||||||
|
└─ Err → fall through to options
|
||||||
|
├─ OS Keychain: generate + store + build SecretsCrypto
|
||||||
|
├─ Env variable: generate + print export command
|
||||||
|
└─ Skip: disable secrets features
|
||||||
|
```
|
||||||
|
|
||||||
|
**CRITICAL CAVEAT: macOS Keychain Dialogs**
|
||||||
|
|
||||||
|
On macOS, `security_framework::get_generic_password()` can trigger TWO
|
||||||
|
system dialogs:
|
||||||
|
1. "Enter your password to unlock the keychain" (keychain locked)
|
||||||
|
2. "Allow ironclaw to access this keychain item" (per-app authorization)
|
||||||
|
|
||||||
|
This is OS-level behavior we cannot prevent. To minimize pain:
|
||||||
|
|
||||||
|
- **Use `get_master_key()` not `has_master_key()`** in step 2. Both call
|
||||||
|
the same underlying API, but `get_master_key()` returns the key bytes
|
||||||
|
so we can cache them. `has_master_key()` throws them away, forcing a
|
||||||
|
second keychain access later.
|
||||||
|
|
||||||
|
- **Build `SecretsCrypto` eagerly.** When the keychain key is retrieved,
|
||||||
|
immediately construct `SecretsCrypto` and store in `self.secrets_crypto`.
|
||||||
|
Later calls to `init_secrets_context()` check this field first, avoiding
|
||||||
|
redundant keychain probes.
|
||||||
|
|
||||||
|
- **Never probe the keychain in read-only commands** (e.g., `ironclaw status`).
|
||||||
|
The status command reports "env not set (keychain may be configured)"
|
||||||
|
rather than triggering system dialogs.
|
||||||
|
|
||||||
|
**Invariant:** After Step 2, `self.secrets_crypto` is `Some` if the user
|
||||||
|
chose Keychain or generated a new key. It may be `None` if the user chose
|
||||||
|
env-var mode or skipped secrets.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### Step 3: Inference Provider
|
||||||
|
|
||||||
|
**Module:** `wizard.rs` → `step_inference_provider()`
|
||||||
|
|
||||||
|
**Goal:** Choose LLM backend and authenticate.
|
||||||
|
|
||||||
|
**Providers:**
|
||||||
|
|
||||||
|
| Provider | Auth Method | Secret Name | Env Var |
|
||||||
|
|----------|-------------|-------------|---------|
|
||||||
|
| NEAR AI | Browser OAuth | (session token) | `NEARAI_SESSION_TOKEN` |
|
||||||
|
| Anthropic | API key | `anthropic_api_key` | `ANTHROPIC_API_KEY` |
|
||||||
|
| OpenAI | API key | `openai_api_key` | `OPENAI_API_KEY` |
|
||||||
|
| Ollama | None | - | - |
|
||||||
|
| OpenAI-compatible | Optional API key | `llm_compatible_api_key` | `LLM_API_KEY` |
|
||||||
|
|
||||||
|
**API-key providers** (`setup_api_key_provider`):
|
||||||
|
1. Check env var → if set, ask to reuse, persist to secrets store
|
||||||
|
2. Otherwise prompt for key entry via `secret_input()`
|
||||||
|
3. Store encrypted in secrets via `init_secrets_context()`
|
||||||
|
4. **Cache key in `self.llm_api_key`** for model fetching in Step 4
|
||||||
|
|
||||||
|
**NEAR AI** (`setup_nearai`):
|
||||||
|
- Calls `session_manager.ensure_authenticated()` which opens browser
|
||||||
|
- Session token saved to `~/.ironclaw/session.json`
|
||||||
|
|
||||||
|
**`self.llm_api_key` caching:** The wizard caches the API key as
|
||||||
|
`Option<SecretString>` so that Step 4 (model fetching) and Step 5
|
||||||
|
(embeddings) can use it without re-reading from the secrets store or
|
||||||
|
mutating environment variables.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### Step 4: Model Selection
|
||||||
|
|
||||||
|
**Module:** `wizard.rs` → `step_model_selection()`
|
||||||
|
|
||||||
|
**Goal:** Choose which model to use.
|
||||||
|
|
||||||
|
**Flow:**
|
||||||
|
1. If model already set → offer to keep it
|
||||||
|
2. Fetch models from provider API (5-second timeout)
|
||||||
|
3. On timeout or error → use static fallback list
|
||||||
|
4. Present list + "Custom model ID" escape hatch
|
||||||
|
5. Store in `self.settings.selected_model`
|
||||||
|
|
||||||
|
**Model fetchers pass the cached API key explicitly:**
|
||||||
|
```rust
|
||||||
|
let cached = self.llm_api_key.as_ref().map(|k| k.expose_secret().to_string());
|
||||||
|
let models = fetch_anthropic_models(cached.as_deref()).await;
|
||||||
|
```
|
||||||
|
|
||||||
|
This avoids mutating environment variables. The fetcher checks the explicit
|
||||||
|
key first, then falls back to the standard env var.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### Step 5: Embeddings
|
||||||
|
|
||||||
|
**Module:** `wizard.rs` → `step_embeddings()`
|
||||||
|
|
||||||
|
**Goal:** Configure semantic search for workspace memory.
|
||||||
|
|
||||||
|
**Flow:**
|
||||||
|
1. Ask "Enable semantic search?" (default: yes)
|
||||||
|
2. Detect available providers:
|
||||||
|
- NEAR AI: if backend is `nearai` OR valid session exists
|
||||||
|
- OpenAI: if `OPENAI_API_KEY` in env OR (backend is `openai` AND cached key)
|
||||||
|
3. If both available → let user choose
|
||||||
|
4. If only one → use it
|
||||||
|
5. If neither → disable embeddings
|
||||||
|
|
||||||
|
**Default model:** `text-embedding-3-small` (for both providers)
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### Step 6: Channel Configuration
|
||||||
|
|
||||||
|
**Module:** `wizard.rs` → `step_channels()`, delegating to `channels.rs`
|
||||||
|
|
||||||
|
**Goal:** Enable input channels (TUI, HTTP, Telegram, etc.).
|
||||||
|
|
||||||
|
**Sub-steps:**
|
||||||
|
|
||||||
|
```
|
||||||
|
6a. Tunnel setup (if webhook channels needed)
|
||||||
|
6b. Discover WASM channels from ~/.ironclaw/channels/
|
||||||
|
6c. Multi-select: CLI/TUI, HTTP, discovered channels, bundled channels
|
||||||
|
6d. Install missing bundled channels (copy WASM binaries)
|
||||||
|
6e. Initialize SecretsContext (for token storage)
|
||||||
|
6f. Setup HTTP webhook (if selected)
|
||||||
|
6g. Setup each WASM channel (secrets, owner binding)
|
||||||
|
```
|
||||||
|
|
||||||
|
**Tunnel setup** (`setup_tunnel`):
|
||||||
|
- Options: ngrok, Cloudflare Tunnel, localtunnel, custom URL
|
||||||
|
- Validates HTTPS requirement
|
||||||
|
- Stored in `self.settings.tunnel.public_url`
|
||||||
|
|
||||||
|
**WASM channel setup** (`setup_wasm_channel`):
|
||||||
|
- Reads `capabilities.json` for `setup.required_secrets`
|
||||||
|
- For each secret: check existing, prompt or auto-generate, validate regex
|
||||||
|
- Save each secret via `SecretsContext`
|
||||||
|
|
||||||
|
**Telegram special case** (`setup_telegram`):
|
||||||
|
- Validates bot token via Telegram `getMe` API
|
||||||
|
- Owner binding: polls `getUpdates` for 120s to capture sender's user ID
|
||||||
|
- Optional webhook secret generation
|
||||||
|
|
||||||
|
**SecretsContext creation** (`init_secrets_context`):
|
||||||
|
1. Check `self.secrets_crypto` (set in Step 2) → use if available
|
||||||
|
2. Else try `SECRETS_MASTER_KEY` env var
|
||||||
|
3. Else try `get_master_key()` from keychain (only in `channels_only` mode)
|
||||||
|
4. Create backend-appropriate secrets store (respects selected database backend)
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### Step 7: Heartbeat
|
||||||
|
|
||||||
|
**Module:** `wizard.rs` → `step_heartbeat()`
|
||||||
|
|
||||||
|
**Goal:** Configure periodic background execution.
|
||||||
|
|
||||||
|
**Flow:**
|
||||||
|
1. Ask "Enable heartbeat?" (default: no)
|
||||||
|
2. If yes: interval in minutes (default: 30), notification channel
|
||||||
|
3. Store in `self.settings.heartbeat`
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Settings Persistence
|
||||||
|
|
||||||
|
### Two-Layer Architecture
|
||||||
|
|
||||||
|
Settings are persisted in two places:
|
||||||
|
|
||||||
|
**Layer 1: `~/.ironclaw/.env`** (bootstrap vars)
|
||||||
|
|
||||||
|
Contains only the settings needed BEFORE database connection. Written by
|
||||||
|
`save_bootstrap_env()` in `bootstrap.rs`.
|
||||||
|
|
||||||
|
```env
|
||||||
|
DATABASE_BACKEND="libsql"
|
||||||
|
LIBSQL_PATH="/Users/name/.ironclaw/ironclaw.db"
|
||||||
|
```
|
||||||
|
|
||||||
|
Or for PostgreSQL:
|
||||||
|
```env
|
||||||
|
DATABASE_BACKEND="postgres"
|
||||||
|
DATABASE_URL="postgres://user:pass@localhost/ironclaw"
|
||||||
|
```
|
||||||
|
|
||||||
|
**Why separate?** Chicken-and-egg: you need `DATABASE_BACKEND` to know
|
||||||
|
which database to connect to, so it can't be stored in the database.
|
||||||
|
|
||||||
|
**Layer 2: Database settings table** (everything else)
|
||||||
|
|
||||||
|
All other settings are stored as key-value pairs in the `settings` table,
|
||||||
|
keyed by `(user_id, key)`. Written by `set_all_settings()`.
|
||||||
|
|
||||||
|
Settings are serialized via `Settings::to_db_map()` as dotted paths:
|
||||||
|
```
|
||||||
|
database_backend = "libsql"
|
||||||
|
llm_backend = "nearai"
|
||||||
|
selected_model = "anthropic/claude-sonnet-4-5"
|
||||||
|
embeddings.enabled = "true"
|
||||||
|
embeddings.provider = "nearai"
|
||||||
|
channels.http_enabled = "true"
|
||||||
|
heartbeat.enabled = "true"
|
||||||
|
heartbeat.interval_secs = "300"
|
||||||
|
```
|
||||||
|
|
||||||
|
### save_and_summarize()
|
||||||
|
|
||||||
|
Final step of the wizard:
|
||||||
|
|
||||||
|
```
|
||||||
|
1. Mark onboard_completed = true
|
||||||
|
2. Write ALL settings to database (try postgres pool, then libSQL backend)
|
||||||
|
3. Write bootstrap vars to ~/.ironclaw/.env:
|
||||||
|
- DATABASE_BACKEND (always)
|
||||||
|
- DATABASE_URL (if postgres)
|
||||||
|
- LIBSQL_PATH (if libsql)
|
||||||
|
- LIBSQL_URL (if turso sync)
|
||||||
|
4. Print configuration summary
|
||||||
|
```
|
||||||
|
|
||||||
|
**Invariant:** Both Layer 1 and Layer 2 must be written. If the database
|
||||||
|
write fails, the wizard returns an error and the `.env` file is not written.
|
||||||
|
|
||||||
|
### Legacy Migration
|
||||||
|
|
||||||
|
`bootstrap.rs` handles one-time upgrades from older config formats:
|
||||||
|
- `bootstrap.json` → extracts `DATABASE_URL`, writes `.env`, renames to `.migrated`
|
||||||
|
- `settings.json` → migrated to database via `migrate_disk_to_db()`
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Settings Struct
|
||||||
|
|
||||||
|
**Module:** `settings.rs`
|
||||||
|
|
||||||
|
```rust
|
||||||
|
pub struct Settings {
|
||||||
|
// Meta
|
||||||
|
pub onboard_completed: bool,
|
||||||
|
|
||||||
|
// Step 1: Database
|
||||||
|
pub database_backend: Option<String>, // "postgres" | "libsql"
|
||||||
|
pub database_url: Option<String>,
|
||||||
|
pub libsql_path: Option<String>,
|
||||||
|
pub libsql_url: Option<String>,
|
||||||
|
|
||||||
|
// Step 2: Security
|
||||||
|
pub secrets_master_key_source: KeySource, // Keychain | Env | None
|
||||||
|
|
||||||
|
// Step 3: Inference
|
||||||
|
pub llm_backend: Option<String>, // "nearai" | "anthropic" | "openai" | "ollama" | "openai_compatible"
|
||||||
|
pub ollama_base_url: Option<String>,
|
||||||
|
pub openai_compatible_base_url: Option<String>,
|
||||||
|
|
||||||
|
// Step 4: Model
|
||||||
|
pub selected_model: Option<String>,
|
||||||
|
|
||||||
|
// Step 5: Embeddings
|
||||||
|
pub embeddings: EmbeddingsSettings, // enabled, provider, model
|
||||||
|
|
||||||
|
// Step 6: Channels
|
||||||
|
pub tunnel: TunnelSettings, // provider, public_url
|
||||||
|
pub channels: ChannelSettings, // http config, telegram owner, etc.
|
||||||
|
|
||||||
|
// Step 7: Heartbeat
|
||||||
|
pub heartbeat: HeartbeatSettings, // enabled, interval, notify
|
||||||
|
|
||||||
|
// Advanced (not in wizard, set via `ironclaw config set`)
|
||||||
|
pub agent: AgentSettings,
|
||||||
|
pub wasm: WasmSettings,
|
||||||
|
pub sandbox: SandboxSettings,
|
||||||
|
pub safety: SafetySettings,
|
||||||
|
pub builder: BuilderSettings,
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
**KeySource enum:** `Keychain | Env | None`
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Secrets Flow
|
||||||
|
|
||||||
|
### SecretsContext
|
||||||
|
|
||||||
|
Thin wrapper for setup-time secret operations:
|
||||||
|
|
||||||
|
```rust
|
||||||
|
pub struct SecretsContext {
|
||||||
|
store: Arc<dyn SecretsStore>,
|
||||||
|
user_id: String,
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Created by `init_secrets_context()` which:
|
||||||
|
1. Gets `SecretsCrypto` from `self.secrets_crypto` or loads from keychain/env
|
||||||
|
2. Creates the appropriate backend store:
|
||||||
|
- If both features compiled: respects `self.settings.database_backend`
|
||||||
|
- Tries selected backend first, falls back to the other
|
||||||
|
3. Returns `SecretsContext` wrapping the store
|
||||||
|
|
||||||
|
### Secret Storage
|
||||||
|
|
||||||
|
Secrets are encrypted with AES-256-GCM using the master key, then stored
|
||||||
|
in the database `secrets` table. The wizard writes secrets like:
|
||||||
|
|
||||||
|
```
|
||||||
|
telegram_bot_token → encrypted bot token
|
||||||
|
telegram_webhook_secret → encrypted webhook HMAC secret
|
||||||
|
anthropic_api_key → encrypted API key
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Prompt Utilities
|
||||||
|
|
||||||
|
**Module:** `prompts.rs`
|
||||||
|
|
||||||
|
| Function | Description |
|
||||||
|
|----------|-------------|
|
||||||
|
| `select_one(label, options)` | Numbered single-choice menu |
|
||||||
|
| `select_many(label, options, defaults)` | Checkbox multi-select (raw terminal mode) |
|
||||||
|
| `input(label)` | Single line text input |
|
||||||
|
| `optional_input(label, hint)` | Text input that can be empty |
|
||||||
|
| `secret_input(label)` | Hidden input (shows `*` per char), returns `SecretString` |
|
||||||
|
| `confirm(label, default)` | `[Y/n]` or `[y/N]` prompt |
|
||||||
|
| `print_header(text)` | Bold section header with underline |
|
||||||
|
| `print_step(n, total, text)` | `[1/7] Step Name` |
|
||||||
|
| `print_success(text)` | Green checkmark prefix |
|
||||||
|
| `print_error(text)` | Red X prefix |
|
||||||
|
| `print_info(text)` | Blue info prefix |
|
||||||
|
|
||||||
|
`select_many` uses `crossterm` raw mode for arrow key navigation.
|
||||||
|
Must properly restore terminal state on all exit paths.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Platform Caveats
|
||||||
|
|
||||||
|
### macOS Keychain
|
||||||
|
|
||||||
|
- `get_generic_password()` triggers system dialogs (unlock + authorize)
|
||||||
|
- Two dialogs per call is normal, not a bug
|
||||||
|
- Cache the result after first access to avoid repeat prompts
|
||||||
|
- Never probe keychain in read-only commands (`status`, `--help`)
|
||||||
|
- Service name: `"ironclaw"`, account: `"master_key"`
|
||||||
|
|
||||||
|
### Linux Secret Service
|
||||||
|
|
||||||
|
- Uses GNOME Keyring or KWallet via `secret-service` crate
|
||||||
|
- May need `gnome-keyring` daemon running
|
||||||
|
- Collection unlock may prompt for password
|
||||||
|
|
||||||
|
### URL Passwords
|
||||||
|
|
||||||
|
- `#` is common in URL-encoded passwords (`%23` decoded)
|
||||||
|
- `.env` values must be double-quoted to preserve `#`
|
||||||
|
- Display masked: `postgres://user:****@host/db`
|
||||||
|
|
||||||
|
### Telegram API
|
||||||
|
|
||||||
|
- Bot token format: `123456:ABC-DEF...`
|
||||||
|
- Token goes in URL path: `https://api.telegram.org/bot{TOKEN}/method`
|
||||||
|
- Webhook secret header: `X-Telegram-Bot-Api-Secret-Token`
|
||||||
|
- Owner binding polls `getUpdates` (must delete webhook first)
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Testing
|
||||||
|
|
||||||
|
Tests live in `mod tests {}` at the bottom of each file.
|
||||||
|
|
||||||
|
**What to test when modifying setup:**
|
||||||
|
|
||||||
|
- Settings round-trip: `to_db_map()` then `from_db_map()` preserves values
|
||||||
|
- Bootstrap `.env`: dotenvy can parse what `save_bootstrap_env()` writes
|
||||||
|
- Model fetchers: static fallback works when API is unreachable
|
||||||
|
- Channel discovery: handles missing dir, invalid JSON, deduplication
|
||||||
|
- Prompt functions: not tested (interactive I/O), but ensure error paths
|
||||||
|
don't panic
|
||||||
|
|
||||||
|
**Run setup tests:**
|
||||||
|
```bash
|
||||||
|
cargo test --lib -- setup
|
||||||
|
cargo test --lib -- bootstrap
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Modification Checklist
|
||||||
|
|
||||||
|
When changing the onboarding flow:
|
||||||
|
|
||||||
|
1. Update this README first with the intended behavior change
|
||||||
|
2. If adding a new wizard step:
|
||||||
|
- Add to the step enum in `run()`, adjust `total_steps`
|
||||||
|
- Add corresponding settings fields to `Settings`
|
||||||
|
- Add `to_db_map` / `from_db_map` serialization
|
||||||
|
- If the setting is needed before DB connection, add to `save_bootstrap_env()`
|
||||||
|
3. If adding a new provider or channel:
|
||||||
|
- Add to the selection menu in the appropriate step
|
||||||
|
- Add authentication flow (API key or OAuth)
|
||||||
|
- Add model fetcher with static fallback + 5s timeout
|
||||||
|
4. If touching keychain:
|
||||||
|
- Cache the result, never call `get_master_key()` twice
|
||||||
|
- Test on macOS (dialog behavior differs from Linux)
|
||||||
|
5. If touching secrets:
|
||||||
|
- Ensure `init_secrets_context()` respects the selected database backend
|
||||||
|
- Test with both postgres and libsql features
|
||||||
|
6. Run the full shipping checklist:
|
||||||
|
```bash
|
||||||
|
cargo fmt
|
||||||
|
cargo clippy --all --benches --tests --examples --all-features -- -D warnings
|
||||||
|
cargo test --lib -- setup bootstrap
|
||||||
|
```
|
||||||
|
7. Test a fresh onboarding: `rm -rf ~/.ironclaw && cargo run`
|
||||||
+124
-81
@@ -20,6 +20,22 @@ use crate::setup::prompts::{
|
|||||||
confirm, input, optional_input, print_error, print_info, print_success, secret_input,
|
confirm, input, optional_input, print_error, print_info, print_success, secret_input,
|
||||||
};
|
};
|
||||||
|
|
||||||
|
/// Typed errors for channel setup flows.
|
||||||
|
#[derive(Debug, thiserror::Error)]
|
||||||
|
pub enum ChannelSetupError {
|
||||||
|
#[error("I/O error: {0}")]
|
||||||
|
Io(#[from] std::io::Error),
|
||||||
|
|
||||||
|
#[error("{0}")]
|
||||||
|
Network(String),
|
||||||
|
|
||||||
|
#[error("{0}")]
|
||||||
|
Secrets(String),
|
||||||
|
|
||||||
|
#[error("{0}")]
|
||||||
|
Validation(String),
|
||||||
|
}
|
||||||
|
|
||||||
/// Context for saving secrets during setup.
|
/// Context for saving secrets during setup.
|
||||||
pub struct SecretsContext {
|
pub struct SecretsContext {
|
||||||
store: Arc<dyn SecretsStore>,
|
store: Arc<dyn SecretsStore>,
|
||||||
@@ -45,32 +61,39 @@ impl SecretsContext {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Save a secret to the database.
|
/// Save a secret to the database.
|
||||||
pub async fn save_secret(&self, name: &str, value: &SecretString) -> Result<(), String> {
|
pub async fn save_secret(
|
||||||
|
&self,
|
||||||
|
name: &str,
|
||||||
|
value: &SecretString,
|
||||||
|
) -> Result<(), ChannelSetupError> {
|
||||||
let params = CreateSecretParams::new(name, value.expose_secret());
|
let params = CreateSecretParams::new(name, value.expose_secret());
|
||||||
|
|
||||||
self.store
|
self.store
|
||||||
.create(&self.user_id, params)
|
.create(&self.user_id, params)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| format!("Failed to save secret: {}", e))?;
|
.map_err(|e| ChannelSetupError::Secrets(format!("Failed to save secret: {}", e)))?;
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Check if a secret exists.
|
/// Check if a secret exists.
|
||||||
pub async fn secret_exists(&self, name: &str) -> bool {
|
pub async fn secret_exists(&self, name: &str) -> bool {
|
||||||
self.store
|
match self.store.exists(&self.user_id, name).await {
|
||||||
.exists(&self.user_id, name)
|
Ok(exists) => exists,
|
||||||
.await
|
Err(e) => {
|
||||||
.unwrap_or(false)
|
tracing::warn!(secret = name, error = %e, "Failed to check if secret exists, assuming absent");
|
||||||
|
false
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Read a secret from the database (decrypted).
|
/// Read a secret from the database (decrypted).
|
||||||
pub async fn get_secret(&self, name: &str) -> Result<SecretString, String> {
|
pub async fn get_secret(&self, name: &str) -> Result<SecretString, ChannelSetupError> {
|
||||||
let decrypted = self
|
let decrypted = self
|
||||||
.store
|
.store
|
||||||
.get_decrypted(&self.user_id, name)
|
.get_decrypted(&self.user_id, name)
|
||||||
.await
|
.await
|
||||||
.map_err(|e| format!("Failed to read secret: {}", e))?;
|
.map_err(|e| ChannelSetupError::Secrets(format!("Failed to read secret: {}", e)))?;
|
||||||
Ok(SecretString::from(decrypted.expose().to_string()))
|
Ok(SecretString::from(decrypted.expose().to_string()))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -107,7 +130,6 @@ struct TelegramGetUpdatesResponse {
|
|||||||
|
|
||||||
#[derive(Debug, Deserialize)]
|
#[derive(Debug, Deserialize)]
|
||||||
struct TelegramUpdate {
|
struct TelegramUpdate {
|
||||||
#[allow(dead_code)]
|
|
||||||
update_id: i64,
|
update_id: i64,
|
||||||
message: Option<TelegramUpdateMessage>,
|
message: Option<TelegramUpdateMessage>,
|
||||||
}
|
}
|
||||||
@@ -134,7 +156,7 @@ struct TelegramUpdateUser {
|
|||||||
pub async fn setup_telegram(
|
pub async fn setup_telegram(
|
||||||
secrets: &SecretsContext,
|
secrets: &SecretsContext,
|
||||||
settings: &Settings,
|
settings: &Settings,
|
||||||
) -> Result<TelegramSetupResult, String> {
|
) -> Result<TelegramSetupResult, ChannelSetupError> {
|
||||||
println!("Telegram Setup:");
|
println!("Telegram Setup:");
|
||||||
println!();
|
println!();
|
||||||
print_info("To create a Telegram bot:");
|
print_info("To create a Telegram bot:");
|
||||||
@@ -146,7 +168,7 @@ pub async fn setup_telegram(
|
|||||||
// Check if token already exists
|
// Check if token already exists
|
||||||
if secrets.secret_exists("telegram_bot_token").await {
|
if secrets.secret_exists("telegram_bot_token").await {
|
||||||
print_info("Existing Telegram token found in database.");
|
print_info("Existing Telegram token found in database.");
|
||||||
if !confirm("Replace existing token?", false).map_err(|e| e.to_string())? {
|
if !confirm("Replace existing token?", false)? {
|
||||||
// Still offer to configure webhook secret and owner binding
|
// Still offer to configure webhook secret and owner binding
|
||||||
let webhook_secret = setup_telegram_webhook_secret(secrets, &settings.tunnel).await?;
|
let webhook_secret = setup_telegram_webhook_secret(secrets, &settings.tunnel).await?;
|
||||||
let owner_id = bind_telegram_owner_flow(secrets, settings).await?;
|
let owner_id = bind_telegram_owner_flow(secrets, settings).await?;
|
||||||
@@ -159,7 +181,8 @@ pub async fn setup_telegram(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
let token = secret_input("Bot token (from @BotFather)").map_err(|e| e.to_string())?;
|
loop {
|
||||||
|
let token = secret_input("Bot token (from @BotFather)")?;
|
||||||
|
|
||||||
// Validate the token
|
// Validate the token
|
||||||
print_info("Validating bot token...");
|
print_info("Validating bot token...");
|
||||||
@@ -179,27 +202,27 @@ pub async fn setup_telegram(
|
|||||||
let owner_id = bind_telegram_owner(&token).await?;
|
let owner_id = bind_telegram_owner(&token).await?;
|
||||||
|
|
||||||
// Offer webhook secret configuration
|
// Offer webhook secret configuration
|
||||||
let webhook_secret = setup_telegram_webhook_secret(secrets, &settings.tunnel).await?;
|
let webhook_secret =
|
||||||
|
setup_telegram_webhook_secret(secrets, &settings.tunnel).await?;
|
||||||
|
|
||||||
Ok(TelegramSetupResult {
|
return Ok(TelegramSetupResult {
|
||||||
enabled: true,
|
enabled: true,
|
||||||
bot_username: username,
|
bot_username: username,
|
||||||
webhook_secret,
|
webhook_secret,
|
||||||
owner_id,
|
owner_id,
|
||||||
})
|
});
|
||||||
}
|
}
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
print_error(&format!("Token validation failed: {}", e));
|
print_error(&format!("Token validation failed: {}", e));
|
||||||
|
|
||||||
if confirm("Try again?", true).map_err(|e| e.to_string())? {
|
if !confirm("Try again?", true)? {
|
||||||
Box::pin(setup_telegram(secrets, settings)).await
|
return Ok(TelegramSetupResult {
|
||||||
} else {
|
|
||||||
Ok(TelegramSetupResult {
|
|
||||||
enabled: false,
|
enabled: false,
|
||||||
bot_username: None,
|
bot_username: None,
|
||||||
webhook_secret: None,
|
webhook_secret: None,
|
||||||
owner_id: None,
|
owner_id: None,
|
||||||
})
|
});
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -209,14 +232,14 @@ pub async fn setup_telegram(
|
|||||||
///
|
///
|
||||||
/// Polls `getUpdates` until a message arrives, then captures the sender's user ID.
|
/// Polls `getUpdates` until a message arrives, then captures the sender's user ID.
|
||||||
/// Returns `None` if the user declines or the flow times out.
|
/// Returns `None` if the user declines or the flow times out.
|
||||||
async fn bind_telegram_owner(token: &SecretString) -> Result<Option<i64>, String> {
|
async fn bind_telegram_owner(token: &SecretString) -> Result<Option<i64>, ChannelSetupError> {
|
||||||
println!();
|
println!();
|
||||||
print_info("Account Binding (recommended):");
|
print_info("Account Binding (recommended):");
|
||||||
print_info("Binding restricts the bot so only YOU can use it.");
|
print_info("Binding restricts the bot so only YOU can use it.");
|
||||||
print_info("Without this, anyone who finds your bot can send it messages.");
|
print_info("Without this, anyone who finds your bot can send it messages.");
|
||||||
println!();
|
println!();
|
||||||
|
|
||||||
if !confirm("Bind bot to your Telegram account?", true).map_err(|e| e.to_string())? {
|
if !confirm("Bind bot to your Telegram account?", true)? {
|
||||||
print_info("Skipping account binding. Bot will accept messages from all users.");
|
print_info("Skipping account binding. Bot will accept messages from all users.");
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
}
|
}
|
||||||
@@ -227,14 +250,16 @@ async fn bind_telegram_owner(token: &SecretString) -> Result<Option<i64>, String
|
|||||||
let client = Client::builder()
|
let client = Client::builder()
|
||||||
.timeout(std::time::Duration::from_secs(35))
|
.timeout(std::time::Duration::from_secs(35))
|
||||||
.build()
|
.build()
|
||||||
.map_err(|e| format!("Failed to create HTTP client: {}", e))?;
|
.map_err(|e| ChannelSetupError::Network(format!("Failed to create HTTP client: {}", e)))?;
|
||||||
|
|
||||||
// Clear any existing webhook so getUpdates works
|
// Clear any existing webhook so getUpdates works
|
||||||
let delete_url = format!(
|
let delete_url = format!(
|
||||||
"https://api.telegram.org/bot{}/deleteWebhook",
|
"https://api.telegram.org/bot{}/deleteWebhook",
|
||||||
token.expose_secret()
|
token.expose_secret()
|
||||||
);
|
);
|
||||||
let _ = client.post(&delete_url).send().await;
|
if let Err(e) = client.post(&delete_url).send().await {
|
||||||
|
tracing::warn!("Failed to delete webhook (getUpdates may not work): {e}");
|
||||||
|
}
|
||||||
|
|
||||||
let updates_url = format!(
|
let updates_url = format!(
|
||||||
"https://api.telegram.org/bot{}/getUpdates",
|
"https://api.telegram.org/bot{}/getUpdates",
|
||||||
@@ -249,19 +274,23 @@ async fn bind_telegram_owner(token: &SecretString) -> Result<Option<i64>, String
|
|||||||
.query(&[("timeout", "30"), ("allowed_updates", "[\"message\"]")])
|
.query(&[("timeout", "30"), ("allowed_updates", "[\"message\"]")])
|
||||||
.send()
|
.send()
|
||||||
.await
|
.await
|
||||||
.map_err(|e| format!("getUpdates request failed: {}", e))?;
|
.map_err(|e| ChannelSetupError::Network(format!("getUpdates request failed: {}", e)))?;
|
||||||
|
|
||||||
if !response.status().is_success() {
|
if !response.status().is_success() {
|
||||||
return Err(format!("getUpdates returned status {}", response.status()));
|
return Err(ChannelSetupError::Network(format!(
|
||||||
|
"getUpdates returned status {}",
|
||||||
|
response.status()
|
||||||
|
)));
|
||||||
}
|
}
|
||||||
|
|
||||||
let body: TelegramGetUpdatesResponse = response
|
let body: TelegramGetUpdatesResponse = response.json().await.map_err(|e| {
|
||||||
.json()
|
ChannelSetupError::Network(format!("Failed to parse getUpdates response: {}", e))
|
||||||
.await
|
})?;
|
||||||
.map_err(|e| format!("Failed to parse getUpdates response: {}", e))?;
|
|
||||||
|
|
||||||
if !body.ok {
|
if !body.ok {
|
||||||
return Err("Telegram API returned error for getUpdates".to_string());
|
return Err(ChannelSetupError::Network(
|
||||||
|
"Telegram API returned error for getUpdates".to_string(),
|
||||||
|
));
|
||||||
}
|
}
|
||||||
|
|
||||||
// Find the first message with a sender
|
// Find the first message with a sender
|
||||||
@@ -285,11 +314,14 @@ async fn bind_telegram_owner(token: &SecretString) -> Result<Option<i64>, String
|
|||||||
"https://api.telegram.org/bot{}/getUpdates",
|
"https://api.telegram.org/bot{}/getUpdates",
|
||||||
token.expose_secret()
|
token.expose_secret()
|
||||||
);
|
);
|
||||||
let _ = client
|
if let Err(e) = client
|
||||||
.get(&ack_url)
|
.get(&ack_url)
|
||||||
.query(&[("offset", &(update.update_id + 1).to_string())])
|
.query(&[("offset", &(update.update_id + 1).to_string())])
|
||||||
.send()
|
.send()
|
||||||
.await;
|
.await
|
||||||
|
{
|
||||||
|
tracing::warn!("Failed to acknowledge Telegram update: {e}");
|
||||||
|
}
|
||||||
|
|
||||||
return Ok(Some(from.id));
|
return Ok(Some(from.id));
|
||||||
}
|
}
|
||||||
@@ -307,10 +339,10 @@ async fn bind_telegram_owner(token: &SecretString) -> Result<Option<i64>, String
|
|||||||
async fn bind_telegram_owner_flow(
|
async fn bind_telegram_owner_flow(
|
||||||
secrets: &SecretsContext,
|
secrets: &SecretsContext,
|
||||||
settings: &Settings,
|
settings: &Settings,
|
||||||
) -> Result<Option<i64>, String> {
|
) -> Result<Option<i64>, ChannelSetupError> {
|
||||||
if settings.channels.telegram_owner_id.is_some() {
|
if settings.channels.telegram_owner_id.is_some() {
|
||||||
print_info("Bot is already bound to a Telegram account.");
|
print_info("Bot is already bound to a Telegram account.");
|
||||||
if !confirm("Re-bind to a different account?", false).map_err(|e| e.to_string())? {
|
if !confirm("Re-bind to a different account?", false)? {
|
||||||
return Ok(settings.channels.telegram_owner_id);
|
return Ok(settings.channels.telegram_owner_id);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -325,10 +357,10 @@ async fn bind_telegram_owner_flow(
|
|||||||
///
|
///
|
||||||
/// This is shared across all channels that need webhook endpoints.
|
/// This is shared across all channels that need webhook endpoints.
|
||||||
/// Returns the tunnel URL if configured.
|
/// Returns the tunnel URL if configured.
|
||||||
pub fn setup_tunnel(settings: &Settings) -> Result<Option<String>, String> {
|
pub fn setup_tunnel(settings: &Settings) -> Result<Option<String>, ChannelSetupError> {
|
||||||
if let Some(ref url) = settings.tunnel.public_url {
|
if let Some(ref url) = settings.tunnel.public_url {
|
||||||
print_info(&format!("Existing tunnel configured: {}", url));
|
print_info(&format!("Existing tunnel configured: {}", url));
|
||||||
if !confirm("Change tunnel configuration?", false).map_err(|e| e.to_string())? {
|
if !confirm("Change tunnel configuration?", false)? {
|
||||||
return Ok(Some(url.clone()));
|
return Ok(Some(url.clone()));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -348,17 +380,18 @@ pub fn setup_tunnel(settings: &Settings) -> Result<Option<String>, String> {
|
|||||||
print_info("Security comes from provider-specific secrets (e.g., Telegram webhook secret).");
|
print_info("Security comes from provider-specific secrets (e.g., Telegram webhook secret).");
|
||||||
println!();
|
println!();
|
||||||
|
|
||||||
if !confirm("Configure a tunnel?", false).map_err(|e| e.to_string())? {
|
if !confirm("Configure a tunnel?", false)? {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
}
|
}
|
||||||
|
|
||||||
let tunnel_url =
|
let tunnel_url = input("Tunnel URL (e.g., https://abc123.ngrok.io)")?;
|
||||||
input("Tunnel URL (e.g., https://abc123.ngrok.io)").map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
// Validate URL format
|
// Validate URL format
|
||||||
if !tunnel_url.starts_with("https://") {
|
if !tunnel_url.starts_with("https://") {
|
||||||
print_error("URL must start with https:// (webhooks require HTTPS)");
|
print_error("URL must start with https:// (webhooks require HTTPS)");
|
||||||
return Err("Invalid tunnel URL: must use HTTPS".to_string());
|
return Err(ChannelSetupError::Validation(
|
||||||
|
"Invalid tunnel URL: must use HTTPS".to_string(),
|
||||||
|
));
|
||||||
}
|
}
|
||||||
|
|
||||||
// Remove trailing slash if present
|
// Remove trailing slash if present
|
||||||
@@ -378,7 +411,7 @@ pub fn setup_tunnel(settings: &Settings) -> Result<Option<String>, String> {
|
|||||||
async fn setup_telegram_webhook_secret(
|
async fn setup_telegram_webhook_secret(
|
||||||
secrets: &SecretsContext,
|
secrets: &SecretsContext,
|
||||||
tunnel: &TunnelSettings,
|
tunnel: &TunnelSettings,
|
||||||
) -> Result<Option<String>, String> {
|
) -> Result<Option<String>, ChannelSetupError> {
|
||||||
if tunnel.public_url.is_none() {
|
if tunnel.public_url.is_none() {
|
||||||
print_info("");
|
print_info("");
|
||||||
print_info("No tunnel configured. Telegram will use polling mode (30s+ delay).");
|
print_info("No tunnel configured. Telegram will use polling mode (30s+ delay).");
|
||||||
@@ -391,7 +424,7 @@ async fn setup_telegram_webhook_secret(
|
|||||||
print_info("A webhook secret adds an extra layer of security by validating");
|
print_info("A webhook secret adds an extra layer of security by validating");
|
||||||
print_info("that requests actually come from Telegram's servers.");
|
print_info("that requests actually come from Telegram's servers.");
|
||||||
|
|
||||||
if !confirm("Generate a webhook secret?", true).map_err(|e| e.to_string())? {
|
if !confirm("Generate a webhook secret?", true)? {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -410,11 +443,13 @@ async fn setup_telegram_webhook_secret(
|
|||||||
/// Validate a Telegram bot token by calling the getMe API.
|
/// Validate a Telegram bot token by calling the getMe API.
|
||||||
///
|
///
|
||||||
/// Returns the bot's username if valid.
|
/// Returns the bot's username if valid.
|
||||||
pub async fn validate_telegram_token(token: &SecretString) -> Result<Option<String>, String> {
|
pub async fn validate_telegram_token(
|
||||||
|
token: &SecretString,
|
||||||
|
) -> Result<Option<String>, ChannelSetupError> {
|
||||||
let client = Client::builder()
|
let client = Client::builder()
|
||||||
.timeout(std::time::Duration::from_secs(10))
|
.timeout(std::time::Duration::from_secs(10))
|
||||||
.build()
|
.build()
|
||||||
.map_err(|e| format!("Failed to create HTTP client: {}", e))?;
|
.map_err(|e| ChannelSetupError::Network(format!("Failed to create HTTP client: {}", e)))?;
|
||||||
|
|
||||||
let url = format!(
|
let url = format!(
|
||||||
"https://api.telegram.org/bot{}/getMe",
|
"https://api.telegram.org/bot{}/getMe",
|
||||||
@@ -425,21 +460,26 @@ pub async fn validate_telegram_token(token: &SecretString) -> Result<Option<Stri
|
|||||||
.get(&url)
|
.get(&url)
|
||||||
.send()
|
.send()
|
||||||
.await
|
.await
|
||||||
.map_err(|e| format!("Request failed: {}", e))?;
|
.map_err(|e| ChannelSetupError::Network(format!("Request failed: {}", e)))?;
|
||||||
|
|
||||||
if !response.status().is_success() {
|
if !response.status().is_success() {
|
||||||
return Err(format!("API returned status {}", response.status()));
|
return Err(ChannelSetupError::Network(format!(
|
||||||
|
"API returned status {}",
|
||||||
|
response.status()
|
||||||
|
)));
|
||||||
}
|
}
|
||||||
|
|
||||||
let body: TelegramGetMeResponse = response
|
let body: TelegramGetMeResponse = response
|
||||||
.json()
|
.json()
|
||||||
.await
|
.await
|
||||||
.map_err(|e| format!("Failed to parse response: {}", e))?;
|
.map_err(|e| ChannelSetupError::Network(format!("Failed to parse response: {}", e)))?;
|
||||||
|
|
||||||
if body.ok {
|
if body.ok {
|
||||||
Ok(body.result.and_then(|u| u.username))
|
Ok(body.result.and_then(|u| u.username))
|
||||||
} else {
|
} else {
|
||||||
Err("Telegram API returned error".to_string())
|
Err(ChannelSetupError::Network(
|
||||||
|
"Telegram API returned error".to_string(),
|
||||||
|
))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -452,38 +492,34 @@ pub struct HttpSetupResult {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Set up HTTP webhook channel.
|
/// Set up HTTP webhook channel.
|
||||||
pub async fn setup_http(secrets: &SecretsContext) -> Result<HttpSetupResult, String> {
|
pub async fn setup_http(secrets: &SecretsContext) -> Result<HttpSetupResult, ChannelSetupError> {
|
||||||
println!("HTTP Webhook Setup:");
|
println!("HTTP Webhook Setup:");
|
||||||
println!();
|
println!();
|
||||||
print_info("The HTTP webhook allows external services to send messages to the agent.");
|
print_info("The HTTP webhook allows external services to send messages to the agent.");
|
||||||
println!();
|
println!();
|
||||||
|
|
||||||
let port_str = optional_input("Port", Some("default: 8080")).map_err(|e| e.to_string())?;
|
let port_str = optional_input("Port", Some("default: 8080"))?;
|
||||||
let port: u16 = port_str
|
let port: u16 = port_str
|
||||||
.as_deref()
|
.as_deref()
|
||||||
.unwrap_or("8080")
|
.unwrap_or("8080")
|
||||||
.parse()
|
.parse()
|
||||||
.map_err(|e| format!("Invalid port: {}", e))?;
|
.map_err(|e| ChannelSetupError::Validation(format!("Invalid port: {}", e)))?;
|
||||||
|
|
||||||
if port < 1024 {
|
if port < 1024 {
|
||||||
print_info("Note: Ports below 1024 may require root privileges");
|
print_info("Note: Ports below 1024 may require root privileges");
|
||||||
}
|
}
|
||||||
|
|
||||||
let host = optional_input("Host", Some("default: 0.0.0.0"))
|
let host =
|
||||||
.map_err(|e| e.to_string())?
|
optional_input("Host", Some("default: 0.0.0.0"))?.unwrap_or_else(|| "0.0.0.0".to_string());
|
||||||
.unwrap_or_else(|| "0.0.0.0".to_string());
|
|
||||||
|
|
||||||
// Generate a webhook secret
|
// Generate a webhook secret
|
||||||
if confirm("Generate a webhook secret for authentication?", true).map_err(|e| e.to_string())? {
|
if confirm("Generate a webhook secret for authentication?", true)? {
|
||||||
let secret = generate_webhook_secret();
|
let secret = generate_webhook_secret();
|
||||||
secrets
|
secrets
|
||||||
.save_secret("http_webhook_secret", &SecretString::from(secret.clone()))
|
.save_secret("http_webhook_secret", &SecretString::from(secret))
|
||||||
.await?;
|
.await?;
|
||||||
print_success("Webhook secret generated and saved to database");
|
print_success("Webhook secret generated and saved to database");
|
||||||
print_info(&format!(
|
print_info("Retrieve it later with: ironclaw secret get http_webhook_secret");
|
||||||
"Secret: {} (store this for your webhook clients)",
|
|
||||||
secret
|
|
||||||
));
|
|
||||||
}
|
}
|
||||||
|
|
||||||
print_success(&format!("HTTP webhook will listen on {}:{}", host, port));
|
print_success(&format!("HTTP webhook will listen on {}:{}", host, port));
|
||||||
@@ -497,11 +533,7 @@ pub async fn setup_http(secrets: &SecretsContext) -> Result<HttpSetupResult, Str
|
|||||||
|
|
||||||
/// Generate a random webhook secret.
|
/// Generate a random webhook secret.
|
||||||
pub fn generate_webhook_secret() -> String {
|
pub fn generate_webhook_secret() -> String {
|
||||||
use rand::RngCore;
|
generate_secret_with_length(32)
|
||||||
let mut rng = rand::thread_rng();
|
|
||||||
let mut bytes = [0u8; 32];
|
|
||||||
rng.fill_bytes(&mut bytes);
|
|
||||||
bytes.iter().map(|b| format!("{:02x}", b)).collect()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Result of WASM channel setup.
|
/// Result of WASM channel setup.
|
||||||
@@ -519,7 +551,7 @@ pub async fn setup_wasm_channel(
|
|||||||
secrets: &SecretsContext,
|
secrets: &SecretsContext,
|
||||||
channel_name: &str,
|
channel_name: &str,
|
||||||
setup: &crate::channels::wasm::SetupSchema,
|
setup: &crate::channels::wasm::SetupSchema,
|
||||||
) -> Result<WasmChannelSetupResult, String> {
|
) -> Result<WasmChannelSetupResult, ChannelSetupError> {
|
||||||
println!("{} Setup:", channel_name);
|
println!("{} Setup:", channel_name);
|
||||||
println!();
|
println!();
|
||||||
|
|
||||||
@@ -530,7 +562,7 @@ pub async fn setup_wasm_channel(
|
|||||||
"Existing {} found in database.",
|
"Existing {} found in database.",
|
||||||
secret_config.name
|
secret_config.name
|
||||||
));
|
));
|
||||||
if !confirm("Replace existing value?", false).map_err(|e| e.to_string())? {
|
if !confirm("Replace existing value?", false)? {
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -538,8 +570,7 @@ pub async fn setup_wasm_channel(
|
|||||||
// Get the value from user or auto-generate
|
// Get the value from user or auto-generate
|
||||||
let value = if secret_config.optional {
|
let value = if secret_config.optional {
|
||||||
let input_value =
|
let input_value =
|
||||||
optional_input(&secret_config.prompt, Some("leave empty to auto-generate"))
|
optional_input(&secret_config.prompt, Some("leave empty to auto-generate"))?;
|
||||||
.map_err(|e| e.to_string())?;
|
|
||||||
|
|
||||||
if let Some(v) = input_value {
|
if let Some(v) = input_value {
|
||||||
if !v.is_empty() {
|
if !v.is_empty() {
|
||||||
@@ -566,18 +597,21 @@ pub async fn setup_wasm_channel(
|
|||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
// Required secret
|
// Required secret
|
||||||
let input_value = secret_input(&secret_config.prompt).map_err(|e| e.to_string())?;
|
let input_value = secret_input(&secret_config.prompt)?;
|
||||||
|
|
||||||
// Validate if pattern is provided
|
// Validate if pattern is provided
|
||||||
if let Some(ref pattern) = secret_config.validation {
|
if let Some(ref pattern) = secret_config.validation {
|
||||||
let re = regex::Regex::new(pattern)
|
let re = regex::Regex::new(pattern).map_err(|e| {
|
||||||
.map_err(|e| format!("Invalid validation pattern: {}", e))?;
|
ChannelSetupError::Validation(format!("Invalid validation pattern: {}", e))
|
||||||
|
})?;
|
||||||
if !re.is_match(input_value.expose_secret()) {
|
if !re.is_match(input_value.expose_secret()) {
|
||||||
print_error(&format!(
|
print_error(&format!(
|
||||||
"Value does not match expected format: {}",
|
"Value does not match expected format: {}",
|
||||||
pattern
|
pattern
|
||||||
));
|
));
|
||||||
return Err("Validation failed".to_string());
|
return Err(ChannelSetupError::Validation(
|
||||||
|
"Validation failed".to_string(),
|
||||||
|
));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -589,14 +623,11 @@ pub async fn setup_wasm_channel(
|
|||||||
print_success(&format!("{} saved to database", secret_config.name));
|
print_success(&format!("{} saved to database", secret_config.name));
|
||||||
}
|
}
|
||||||
|
|
||||||
// Optionally validate the configuration
|
// TODO: Substitute secrets into the validation URL and make a
|
||||||
|
// GET request to verify the configured credentials actually work.
|
||||||
if let Some(ref validation_endpoint) = setup.validation_endpoint {
|
if let Some(ref validation_endpoint) = setup.validation_endpoint {
|
||||||
print_info("Validating configuration...");
|
|
||||||
// The validation endpoint may contain placeholders like {telegram_bot_token}
|
|
||||||
// For now, we skip validation since we'd need to substitute secrets
|
|
||||||
// A full implementation would fetch secrets and substitute them
|
|
||||||
print_info(&format!(
|
print_info(&format!(
|
||||||
"Validation endpoint configured: {} (validation skipped)",
|
"Validation endpoint configured: {} (validation not yet implemented)",
|
||||||
validation_endpoint
|
validation_endpoint
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
@@ -620,11 +651,23 @@ fn generate_secret_with_length(length: usize) -> String {
|
|||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use crate::setup::channels::generate_webhook_secret;
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_generate_webhook_secret() {
|
fn test_generate_webhook_secret() {
|
||||||
let secret = generate_webhook_secret();
|
let secret = generate_webhook_secret();
|
||||||
assert_eq!(secret.len(), 64); // 32 bytes = 64 hex chars
|
assert_eq!(secret.len(), 64); // 32 bytes = 64 hex chars
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_generate_secret_with_length() {
|
||||||
|
use super::generate_secret_with_length;
|
||||||
|
|
||||||
|
let s = generate_secret_with_length(16);
|
||||||
|
assert_eq!(s.len(), 32); // 16 bytes = 32 hex chars
|
||||||
|
assert!(s.chars().all(|c| c.is_ascii_hexdigit()));
|
||||||
|
|
||||||
|
let s2 = generate_secret_with_length(1);
|
||||||
|
assert_eq!(s2.len(), 2);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+3
-2
@@ -3,7 +3,7 @@
|
|||||||
//! Provides a guided setup experience for:
|
//! Provides a guided setup experience for:
|
||||||
//! 1. Database connection
|
//! 1. Database connection
|
||||||
//! 2. Security (secrets master key)
|
//! 2. Security (secrets master key)
|
||||||
//! 3. NEAR AI authentication
|
//! 3. Inference provider selection
|
||||||
//! 4. Model selection
|
//! 4. Model selection
|
||||||
//! 5. Embeddings
|
//! 5. Embeddings
|
||||||
//! 6. Channel configuration (HTTP, Telegram, etc.)
|
//! 6. Channel configuration (HTTP, Telegram, etc.)
|
||||||
@@ -24,7 +24,8 @@ mod prompts;
|
|||||||
mod wizard;
|
mod wizard;
|
||||||
|
|
||||||
pub use channels::{
|
pub use channels::{
|
||||||
SecretsContext, setup_http, setup_telegram, setup_tunnel, validate_telegram_token,
|
ChannelSetupError, SecretsContext, setup_http, setup_telegram, setup_tunnel,
|
||||||
|
validate_telegram_token,
|
||||||
};
|
};
|
||||||
pub use prompts::{
|
pub use prompts::{
|
||||||
confirm, input, optional_input, print_error, print_header, print_info, print_step,
|
confirm, input, optional_input, print_error, print_header, print_info, print_step,
|
||||||
|
|||||||
@@ -21,6 +21,7 @@ use secrecy::SecretString;
|
|||||||
/// Display a numbered menu and get user selection.
|
/// Display a numbered menu and get user selection.
|
||||||
///
|
///
|
||||||
/// Returns the index (0-based) of the selected option.
|
/// Returns the index (0-based) of the selected option.
|
||||||
|
/// Pressing Enter without input selects the first option (index 0).
|
||||||
///
|
///
|
||||||
/// # Example
|
/// # Example
|
||||||
///
|
///
|
||||||
@@ -84,6 +85,10 @@ pub fn select_one(prompt: &str, options: &[&str]) -> io::Result<usize> {
|
|||||||
/// ])?;
|
/// ])?;
|
||||||
/// ```
|
/// ```
|
||||||
pub fn select_many(prompt: &str, options: &[(&str, bool)]) -> io::Result<Vec<usize>> {
|
pub fn select_many(prompt: &str, options: &[(&str, bool)]) -> io::Result<Vec<usize>> {
|
||||||
|
if options.is_empty() {
|
||||||
|
return Ok(vec![]);
|
||||||
|
}
|
||||||
|
|
||||||
let mut stdout = io::stdout();
|
let mut stdout = io::stdout();
|
||||||
let mut selected: Vec<bool> = options.iter().map(|(_, s)| *s).collect();
|
let mut selected: Vec<bool> = options.iter().map(|(_, s)| *s).collect();
|
||||||
let mut cursor_pos = 0;
|
let mut cursor_pos = 0;
|
||||||
|
|||||||
+851
-99
File diff suppressed because it is too large
Load Diff
@@ -1,451 +0,0 @@
|
|||||||
//! Accessibility tree parsing and element reference generation.
|
|
||||||
//!
|
|
||||||
//! Converts Chrome's CDP accessibility tree into a compact, LLM-friendly
|
|
||||||
//! representation with stable element references (`@e1`, `@e2`, ...).
|
|
||||||
//!
|
|
||||||
//! The key insight: sending the full accessibility tree every turn is wasteful.
|
|
||||||
//! Instead, we assign short IDs to interactive elements and let the LLM
|
|
||||||
//! reference them by ID for clicks/typing. This is ~93% cheaper in tokens
|
|
||||||
//! compared to re-sending the full tree each time.
|
|
||||||
//!
|
|
||||||
//! ```text
|
|
||||||
//! Page: https://example.com/login
|
|
||||||
//! @e1: textbox "Email" [focused]
|
|
||||||
//! @e2: textbox "Password" [type=password]
|
|
||||||
//! @e3: button "Sign In"
|
|
||||||
//! @e4: link "Forgot password?"
|
|
||||||
//! ```
|
|
||||||
|
|
||||||
use std::collections::HashMap;
|
|
||||||
use std::fmt;
|
|
||||||
|
|
||||||
use chromiumoxide::cdp::browser_protocol::accessibility::{AxNode, AxPropertyName};
|
|
||||||
use chromiumoxide::cdp::browser_protocol::dom::BackendNodeId;
|
|
||||||
|
|
||||||
/// A resolved element reference that maps `@eN` back to a DOM target.
|
|
||||||
#[derive(Debug, Clone)]
|
|
||||||
pub struct ElementRef {
|
|
||||||
/// The display label shown to the LLM (e.g., `textbox "Email"`).
|
|
||||||
#[allow(dead_code)]
|
|
||||||
pub label: String,
|
|
||||||
/// CDP backend node ID for targeting this element.
|
|
||||||
pub backend_node_id: BackendNodeId,
|
|
||||||
/// CSS selector hint (best-effort, may not be unique).
|
|
||||||
#[allow(dead_code)]
|
|
||||||
pub selector_hint: Option<String>,
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Stores the current set of element references for a page snapshot.
|
|
||||||
#[derive(Debug, Clone, Default)]
|
|
||||||
pub struct ElementRefMap {
|
|
||||||
refs: HashMap<String, ElementRef>,
|
|
||||||
counter: usize,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl ElementRefMap {
|
|
||||||
pub fn new() -> Self {
|
|
||||||
Self::default()
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Look up a reference like `@e1` or just `e1`.
|
|
||||||
pub fn get(&self, ref_id: &str) -> Option<&ElementRef> {
|
|
||||||
let normalized = ref_id.strip_prefix('@').unwrap_or(ref_id);
|
|
||||||
self.refs.get(normalized)
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Number of tracked elements.
|
|
||||||
#[allow(dead_code)]
|
|
||||||
pub fn len(&self) -> usize {
|
|
||||||
self.refs.len()
|
|
||||||
}
|
|
||||||
|
|
||||||
pub fn is_empty(&self) -> bool {
|
|
||||||
self.refs.is_empty()
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Reset all refs. Called before each new `read_page` and when switching tabs.
|
|
||||||
pub fn reset(&mut self) {
|
|
||||||
self.refs.clear();
|
|
||||||
self.counter = 0;
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Allocate the next reference ID and store the element.
|
|
||||||
fn insert(&mut self, elem: ElementRef) -> String {
|
|
||||||
self.counter += 1;
|
|
||||||
let id = format!("e{}", self.counter);
|
|
||||||
self.refs.insert(id.clone(), elem);
|
|
||||||
id
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Which elements to include when building the tree representation.
|
|
||||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
|
||||||
pub enum ElementFilter {
|
|
||||||
/// Only interactive elements (buttons, links, inputs, selects, textareas).
|
|
||||||
Interactive,
|
|
||||||
/// All elements with meaningful content.
|
|
||||||
All,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl ElementFilter {
|
|
||||||
pub fn from_str_opt(s: Option<&str>) -> Self {
|
|
||||||
match s {
|
|
||||||
Some("all") => Self::All,
|
|
||||||
_ => Self::Interactive,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Roles that are considered "interactive" for filtering purposes.
|
|
||||||
const INTERACTIVE_ROLES: &[&str] = &[
|
|
||||||
"button",
|
|
||||||
"link",
|
|
||||||
"textbox",
|
|
||||||
"searchbox",
|
|
||||||
"combobox",
|
|
||||||
"listbox",
|
|
||||||
"option",
|
|
||||||
"menuitem",
|
|
||||||
"menuitemcheckbox",
|
|
||||||
"menuitemradio",
|
|
||||||
"radio",
|
|
||||||
"checkbox",
|
|
||||||
"switch",
|
|
||||||
"slider",
|
|
||||||
"spinbutton",
|
|
||||||
"tab",
|
|
||||||
"treeitem",
|
|
||||||
];
|
|
||||||
|
|
||||||
/// Roles to skip entirely (structural noise).
|
|
||||||
const SKIP_ROLES: &[&str] = &[
|
|
||||||
"none",
|
|
||||||
"presentation",
|
|
||||||
"generic",
|
|
||||||
"InlineTextBox",
|
|
||||||
"LineBreak",
|
|
||||||
];
|
|
||||||
|
|
||||||
/// Build a compact page representation from the CDP accessibility tree.
|
|
||||||
///
|
|
||||||
/// Returns the text representation and populates `ref_map` with element
|
|
||||||
/// references the LLM can use for subsequent actions.
|
|
||||||
pub fn build_page_repr(
|
|
||||||
url: &str,
|
|
||||||
title: &str,
|
|
||||||
nodes: &[AxNode],
|
|
||||||
filter: ElementFilter,
|
|
||||||
ref_map: &mut ElementRefMap,
|
|
||||||
) -> String {
|
|
||||||
ref_map.reset();
|
|
||||||
|
|
||||||
let mut lines = Vec::new();
|
|
||||||
|
|
||||||
// Header
|
|
||||||
lines.push(format!("Page: {}", url));
|
|
||||||
if !title.is_empty() {
|
|
||||||
lines.push(format!("Title: {}", title));
|
|
||||||
}
|
|
||||||
lines.push(String::new());
|
|
||||||
|
|
||||||
// Walk nodes, collecting elements that pass the filter.
|
|
||||||
for node in nodes {
|
|
||||||
let role = node_role(node);
|
|
||||||
|
|
||||||
if SKIP_ROLES.contains(&role.as_str()) {
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
|
|
||||||
// For "interactive" filter, only include interactive roles.
|
|
||||||
if filter == ElementFilter::Interactive && !INTERACTIVE_ROLES.contains(&role.as_str()) {
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
|
|
||||||
// Skip nodes without a name (usually decorative).
|
|
||||||
let name = node_name(node);
|
|
||||||
if name.is_empty() && filter == ElementFilter::Interactive {
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
|
|
||||||
let backend_id = match node.backend_dom_node_id {
|
|
||||||
Some(id) => id,
|
|
||||||
None => continue,
|
|
||||||
};
|
|
||||||
|
|
||||||
// Build display label
|
|
||||||
let mut label = NodeLabel {
|
|
||||||
role: role.clone(),
|
|
||||||
name: truncate_name(&name, 80),
|
|
||||||
properties: Vec::new(),
|
|
||||||
};
|
|
||||||
|
|
||||||
// Add useful properties
|
|
||||||
if node_has_property(node, "focused") {
|
|
||||||
label.properties.push("focused".to_string());
|
|
||||||
}
|
|
||||||
if node_has_property(node, "checked") {
|
|
||||||
label.properties.push("checked".to_string());
|
|
||||||
}
|
|
||||||
if node_has_property(node, "disabled") {
|
|
||||||
label.properties.push("disabled".to_string());
|
|
||||||
}
|
|
||||||
if node_has_property(node, "expanded") {
|
|
||||||
label.properties.push("expanded".to_string());
|
|
||||||
}
|
|
||||||
if node_has_property(node, "required") {
|
|
||||||
label.properties.push("required".to_string());
|
|
||||||
}
|
|
||||||
if let Some(val) = node_value(node) {
|
|
||||||
if !val.is_empty() && val != name {
|
|
||||||
label
|
|
||||||
.properties
|
|
||||||
.push(format!("value=\"{}\"", truncate_name(&val, 40)));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
let display = label.to_string();
|
|
||||||
|
|
||||||
let elem_ref = ElementRef {
|
|
||||||
label: display.clone(),
|
|
||||||
backend_node_id: backend_id,
|
|
||||||
selector_hint: guess_selector(node),
|
|
||||||
};
|
|
||||||
|
|
||||||
let ref_id = ref_map.insert(elem_ref);
|
|
||||||
lines.push(format!("@{}: {}", ref_id, display));
|
|
||||||
}
|
|
||||||
|
|
||||||
if ref_map.is_empty() {
|
|
||||||
lines.push("(no interactive elements found)".to_string());
|
|
||||||
}
|
|
||||||
|
|
||||||
lines.join("\n")
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Extract the role string from an AX node.
|
|
||||||
fn node_role(node: &AxNode) -> String {
|
|
||||||
node.role
|
|
||||||
.as_ref()
|
|
||||||
.and_then(|v| v.value.as_ref())
|
|
||||||
.and_then(|v| v.as_str())
|
|
||||||
.unwrap_or("unknown")
|
|
||||||
.to_string()
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Extract the name (accessible label) from an AX node.
|
|
||||||
fn node_name(node: &AxNode) -> String {
|
|
||||||
node.name
|
|
||||||
.as_ref()
|
|
||||||
.and_then(|v| v.value.as_ref())
|
|
||||||
.and_then(|v| v.as_str())
|
|
||||||
.unwrap_or("")
|
|
||||||
.to_string()
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Extract the value from an AX node (for inputs, etc.).
|
|
||||||
fn node_value(node: &AxNode) -> Option<String> {
|
|
||||||
node.value
|
|
||||||
.as_ref()
|
|
||||||
.and_then(|v| v.value.as_ref())
|
|
||||||
.and_then(|v| v.as_str())
|
|
||||||
.map(|s| s.to_string())
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Map a property name string to the corresponding `AxPropertyName` variant.
|
|
||||||
fn property_by_name(name: &str) -> Option<AxPropertyName> {
|
|
||||||
match name {
|
|
||||||
"focused" => Some(AxPropertyName::Focused),
|
|
||||||
"checked" => Some(AxPropertyName::Checked),
|
|
||||||
"disabled" => Some(AxPropertyName::Disabled),
|
|
||||||
"expanded" => Some(AxPropertyName::Expanded),
|
|
||||||
"required" => Some(AxPropertyName::Required),
|
|
||||||
"selected" => Some(AxPropertyName::Selected),
|
|
||||||
"pressed" => Some(AxPropertyName::Pressed),
|
|
||||||
"readonly" => Some(AxPropertyName::Readonly),
|
|
||||||
"hidden" => Some(AxPropertyName::Hidden),
|
|
||||||
"modal" => Some(AxPropertyName::Modal),
|
|
||||||
_ => None,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Check if a node has a boolean property set to true.
|
|
||||||
fn node_has_property(node: &AxNode, prop_name: &str) -> bool {
|
|
||||||
let Some(props) = &node.properties else {
|
|
||||||
return false;
|
|
||||||
};
|
|
||||||
let Some(target) = property_by_name(prop_name) else {
|
|
||||||
return false;
|
|
||||||
};
|
|
||||||
props.iter().any(|p| {
|
|
||||||
p.name == target
|
|
||||||
&& p.value
|
|
||||||
.value
|
|
||||||
.as_ref()
|
|
||||||
.and_then(|v| v.as_bool())
|
|
||||||
.unwrap_or(false)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Best-effort CSS selector guess from node attributes.
|
|
||||||
fn guess_selector(node: &AxNode) -> Option<String> {
|
|
||||||
// We don't have DOM attributes directly from the AX tree,
|
|
||||||
// so we can only offer role-based hints. The actual targeting
|
|
||||||
// uses backend_node_id which is precise.
|
|
||||||
let role = node_role(node);
|
|
||||||
let name = node_name(node);
|
|
||||||
|
|
||||||
if name.is_empty() {
|
|
||||||
return None;
|
|
||||||
}
|
|
||||||
|
|
||||||
// Build an ARIA selector hint (not used for actual targeting,
|
|
||||||
// just a human-readable hint in debug output).
|
|
||||||
Some(format!(
|
|
||||||
"[role=\"{}\"][name=\"{}\"]",
|
|
||||||
role,
|
|
||||||
truncate_name(&name, 30)
|
|
||||||
))
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Truncate a display name to max chars, adding ellipsis if needed.
|
|
||||||
fn truncate_name(s: &str, max: usize) -> String {
|
|
||||||
if s.chars().count() <= max {
|
|
||||||
s.to_string()
|
|
||||||
} else {
|
|
||||||
format!(
|
|
||||||
"{}...",
|
|
||||||
s.chars().take(max.saturating_sub(3)).collect::<String>()
|
|
||||||
)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Helper for formatting a node's display label.
|
|
||||||
struct NodeLabel {
|
|
||||||
role: String,
|
|
||||||
name: String,
|
|
||||||
properties: Vec<String>,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl fmt::Display for NodeLabel {
|
|
||||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
|
||||||
write!(f, "{}", self.role)?;
|
|
||||||
if !self.name.is_empty() {
|
|
||||||
write!(f, " \"{}\"", self.name)?;
|
|
||||||
}
|
|
||||||
if !self.properties.is_empty() {
|
|
||||||
write!(f, " [{}]", self.properties.join(", "))?;
|
|
||||||
}
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use crate::tools::builtin::browser::accessibility::{
|
|
||||||
ElementFilter, ElementRefMap, build_page_repr, truncate_name,
|
|
||||||
};
|
|
||||||
use chromiumoxide::cdp::browser_protocol::accessibility::{
|
|
||||||
AxNode, AxNodeId, AxValue, AxValueType,
|
|
||||||
};
|
|
||||||
use chromiumoxide::cdp::browser_protocol::dom::BackendNodeId;
|
|
||||||
|
|
||||||
fn make_ax_value(s: &str) -> AxValue {
|
|
||||||
let mut v = AxValue::new(AxValueType::String);
|
|
||||||
v.value = Some(serde_json::Value::String(s.to_string()));
|
|
||||||
v
|
|
||||||
}
|
|
||||||
|
|
||||||
fn make_ax_node(role: &str, name: &str, backend_id: i64) -> AxNode {
|
|
||||||
let mut node = AxNode::new(AxNodeId::from(format!("node_{}", backend_id)), false);
|
|
||||||
node.role = Some(make_ax_value(role));
|
|
||||||
node.name = Some(make_ax_value(name));
|
|
||||||
node.backend_dom_node_id = Some(BackendNodeId::new(backend_id));
|
|
||||||
node
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_build_page_repr_interactive_filter() {
|
|
||||||
let nodes = vec![
|
|
||||||
make_ax_node("button", "Submit", 1),
|
|
||||||
make_ax_node("link", "Home", 2),
|
|
||||||
make_ax_node("textbox", "Email", 3),
|
|
||||||
make_ax_node("heading", "Welcome", 4), // not interactive
|
|
||||||
make_ax_node("generic", "", 5), // skip role
|
|
||||||
];
|
|
||||||
|
|
||||||
let mut ref_map = ElementRefMap::new();
|
|
||||||
let repr = build_page_repr(
|
|
||||||
"https://example.com",
|
|
||||||
"Test Page",
|
|
||||||
&nodes,
|
|
||||||
ElementFilter::Interactive,
|
|
||||||
&mut ref_map,
|
|
||||||
);
|
|
||||||
|
|
||||||
assert!(repr.contains("@e1: button \"Submit\""));
|
|
||||||
assert!(repr.contains("@e2: link \"Home\""));
|
|
||||||
assert!(repr.contains("@e3: textbox \"Email\""));
|
|
||||||
assert!(!repr.contains("heading"));
|
|
||||||
assert!(!repr.contains("generic"));
|
|
||||||
assert_eq!(ref_map.len(), 3);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_build_page_repr_all_filter() {
|
|
||||||
let nodes = vec![
|
|
||||||
make_ax_node("button", "Submit", 1),
|
|
||||||
make_ax_node("heading", "Welcome", 2),
|
|
||||||
];
|
|
||||||
|
|
||||||
let mut ref_map = ElementRefMap::new();
|
|
||||||
let repr = build_page_repr(
|
|
||||||
"https://example.com",
|
|
||||||
"",
|
|
||||||
&nodes,
|
|
||||||
ElementFilter::All,
|
|
||||||
&mut ref_map,
|
|
||||||
);
|
|
||||||
|
|
||||||
assert!(repr.contains("button"));
|
|
||||||
assert!(repr.contains("heading"));
|
|
||||||
assert_eq!(ref_map.len(), 2);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_element_ref_lookup() {
|
|
||||||
let mut ref_map = ElementRefMap::new();
|
|
||||||
let nodes = vec![make_ax_node("button", "Click me", 1)];
|
|
||||||
build_page_repr(
|
|
||||||
"https://x.com",
|
|
||||||
"",
|
|
||||||
&nodes,
|
|
||||||
ElementFilter::Interactive,
|
|
||||||
&mut ref_map,
|
|
||||||
);
|
|
||||||
|
|
||||||
assert!(ref_map.get("e1").is_some());
|
|
||||||
assert!(ref_map.get("@e1").is_some()); // with @ prefix
|
|
||||||
assert!(ref_map.get("e99").is_none());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_empty_page() {
|
|
||||||
let mut ref_map = ElementRefMap::new();
|
|
||||||
let repr = build_page_repr(
|
|
||||||
"https://empty.com",
|
|
||||||
"",
|
|
||||||
&[],
|
|
||||||
ElementFilter::Interactive,
|
|
||||||
&mut ref_map,
|
|
||||||
);
|
|
||||||
|
|
||||||
assert!(repr.contains("no interactive elements"));
|
|
||||||
assert!(ref_map.is_empty());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_truncate_name() {
|
|
||||||
assert_eq!(truncate_name("short", 10), "short");
|
|
||||||
assert_eq!(truncate_name("this is a very long name", 10), "this is...");
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,517 +0,0 @@
|
|||||||
//! Headless browser tool for web interaction.
|
|
||||||
//!
|
|
||||||
//! A single `BrowserTool` that dispatches actions via a tagged enum,
|
|
||||||
//! keeping the tool registry clean (one tool, not ten). The LLM sends
|
|
||||||
//! an `action` field to pick the operation:
|
|
||||||
//!
|
|
||||||
//! ```json
|
|
||||||
//! { "action": "navigate", "url": "https://example.com" }
|
|
||||||
//! { "action": "click", "ref": "@e3" }
|
|
||||||
//! { "action": "type", "ref": "@e1", "text": "hello" }
|
|
||||||
//! { "action": "read_page" }
|
|
||||||
//! { "action": "screenshot" }
|
|
||||||
//! ```
|
|
||||||
//!
|
|
||||||
//! Element references (`@e1`, `@e2`, ...) are assigned by `read_page`
|
|
||||||
//! and remain valid until the next `read_page` call.
|
|
||||||
|
|
||||||
pub mod accessibility;
|
|
||||||
pub mod session;
|
|
||||||
pub mod stealth;
|
|
||||||
|
|
||||||
use std::time::Duration;
|
|
||||||
|
|
||||||
use async_trait::async_trait;
|
|
||||||
use serde::Deserialize;
|
|
||||||
use tokio::sync::RwLock;
|
|
||||||
|
|
||||||
use crate::context::JobContext;
|
|
||||||
use crate::tools::builtin::browser::accessibility::ElementFilter;
|
|
||||||
use crate::tools::builtin::browser::session::BrowserSession;
|
|
||||||
use crate::tools::tool::{Tool, ToolError, ToolOutput};
|
|
||||||
|
|
||||||
/// Actions the LLM can request from the browser tool.
|
|
||||||
///
|
|
||||||
/// Uses serde tagged enum: the JSON `"action"` field selects the variant,
|
|
||||||
/// remaining fields are variant-specific parameters.
|
|
||||||
#[derive(Debug, Deserialize)]
|
|
||||||
#[serde(tag = "action", rename_all = "snake_case")]
|
|
||||||
enum BrowserAction {
|
|
||||||
/// Navigate to a URL.
|
|
||||||
Navigate { url: String },
|
|
||||||
/// Go back in browser history.
|
|
||||||
Back,
|
|
||||||
/// Go forward in browser history.
|
|
||||||
Forward,
|
|
||||||
/// Read the page's accessibility tree (assigns element refs).
|
|
||||||
ReadPage {
|
|
||||||
/// "interactive" (default) or "all"
|
|
||||||
filter: Option<String>,
|
|
||||||
},
|
|
||||||
/// Click an element by reference ID.
|
|
||||||
Click {
|
|
||||||
/// Element reference like "@e1" or "e1".
|
|
||||||
#[serde(alias = "ref")]
|
|
||||||
ref_id: String,
|
|
||||||
},
|
|
||||||
/// Type text into an element by reference ID.
|
|
||||||
Type {
|
|
||||||
/// Element reference like "@e1" or "e1".
|
|
||||||
#[serde(alias = "ref")]
|
|
||||||
ref_id: String,
|
|
||||||
text: String,
|
|
||||||
},
|
|
||||||
/// Scroll the page.
|
|
||||||
Scroll {
|
|
||||||
/// "up", "down", "left", "right"
|
|
||||||
direction: String,
|
|
||||||
/// Number of scroll steps (default 3).
|
|
||||||
amount: Option<u32>,
|
|
||||||
},
|
|
||||||
/// Capture a screenshot (returns base64 PNG).
|
|
||||||
Screenshot {
|
|
||||||
/// Capture full scrollable page (default false).
|
|
||||||
full_page: Option<bool>,
|
|
||||||
},
|
|
||||||
/// Extract text content from the page or a CSS selector.
|
|
||||||
Extract {
|
|
||||||
/// Optional CSS selector. If omitted, extracts all body text.
|
|
||||||
selector: Option<String>,
|
|
||||||
},
|
|
||||||
/// Wait for a CSS selector to appear or a fixed delay.
|
|
||||||
Wait {
|
|
||||||
/// CSS selector to wait for. If omitted, just sleeps.
|
|
||||||
selector: Option<String>,
|
|
||||||
/// Timeout in milliseconds (default 5000).
|
|
||||||
timeout_ms: Option<u64>,
|
|
||||||
},
|
|
||||||
/// Execute JavaScript (requires user approval).
|
|
||||||
EvalJs { expression: String },
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Headless browser tool for navigating web pages, interacting with
|
|
||||||
/// elements, and extracting content.
|
|
||||||
///
|
|
||||||
/// Uses Chrome/Chromium via the DevTools Protocol. The browser is launched
|
|
||||||
/// lazily on first use and includes basic anti-detection patches.
|
|
||||||
///
|
|
||||||
/// ## Workflow
|
|
||||||
///
|
|
||||||
/// 1. `navigate` to a URL
|
|
||||||
/// 2. `read_page` to get the accessibility tree with element refs
|
|
||||||
/// 3. `click` / `type` using the refs
|
|
||||||
/// 4. `extract` or `screenshot` to get results
|
|
||||||
///
|
|
||||||
/// Element refs (`@e1`, `@e2`) are valid until the next `read_page`.
|
|
||||||
pub struct BrowserTool {
|
|
||||||
/// Lazily initialized browser session. RwLock because `execute` takes `&self`.
|
|
||||||
session: RwLock<Option<BrowserSession>>,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl BrowserTool {
|
|
||||||
pub fn new() -> Self {
|
|
||||||
Self {
|
|
||||||
session: RwLock::new(None),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Ensure the browser session is initialized, launching Chrome if needed.
|
|
||||||
async fn ensure_session(&self) -> Result<(), ToolError> {
|
|
||||||
let needs_launch = self.session.read().await.is_none();
|
|
||||||
if needs_launch {
|
|
||||||
let new_session = BrowserSession::launch().await?;
|
|
||||||
let mut guard = self.session.write().await;
|
|
||||||
if guard.is_none() {
|
|
||||||
*guard = Some(new_session);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
impl Default for BrowserTool {
|
|
||||||
fn default() -> Self {
|
|
||||||
Self::new()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[async_trait]
|
|
||||||
impl Tool for BrowserTool {
|
|
||||||
fn name(&self) -> &str {
|
|
||||||
"browser"
|
|
||||||
}
|
|
||||||
|
|
||||||
fn description(&self) -> &str {
|
|
||||||
"Control a headless web browser. Navigate pages, read content, click elements, type text, \
|
|
||||||
take screenshots. Use 'read_page' to get an accessibility tree with element references \
|
|
||||||
(@e1, @e2...), then use those refs for 'click' and 'type' actions.\n\n\
|
|
||||||
Actions: navigate, back, forward, read_page, click, type, scroll, screenshot, extract, \
|
|
||||||
wait, eval_js"
|
|
||||||
}
|
|
||||||
|
|
||||||
fn parameters_schema(&self) -> serde_json::Value {
|
|
||||||
serde_json::json!({
|
|
||||||
"type": "object",
|
|
||||||
"properties": {
|
|
||||||
"action": {
|
|
||||||
"type": "string",
|
|
||||||
"enum": [
|
|
||||||
"navigate", "back", "forward", "read_page", "click",
|
|
||||||
"type", "scroll", "screenshot", "extract", "wait", "eval_js"
|
|
||||||
],
|
|
||||||
"description": "The browser action to perform"
|
|
||||||
},
|
|
||||||
"url": {
|
|
||||||
"type": "string",
|
|
||||||
"description": "URL to navigate to (for 'navigate' action)"
|
|
||||||
},
|
|
||||||
"ref_id": {
|
|
||||||
"type": "string",
|
|
||||||
"description": "Element reference like '@e1' (for 'click' and 'type' actions)"
|
|
||||||
},
|
|
||||||
"text": {
|
|
||||||
"type": "string",
|
|
||||||
"description": "Text to type (for 'type' action)"
|
|
||||||
},
|
|
||||||
"direction": {
|
|
||||||
"type": "string",
|
|
||||||
"enum": ["up", "down", "left", "right"],
|
|
||||||
"description": "Scroll direction (for 'scroll' action)"
|
|
||||||
},
|
|
||||||
"amount": {
|
|
||||||
"type": "integer",
|
|
||||||
"description": "Scroll steps, default 3 (for 'scroll' action)"
|
|
||||||
},
|
|
||||||
"full_page": {
|
|
||||||
"type": "boolean",
|
|
||||||
"description": "Capture full scrollable page (for 'screenshot' action)"
|
|
||||||
},
|
|
||||||
"selector": {
|
|
||||||
"type": "string",
|
|
||||||
"description": "CSS selector (for 'extract' and 'wait' actions)"
|
|
||||||
},
|
|
||||||
"timeout_ms": {
|
|
||||||
"type": "integer",
|
|
||||||
"description": "Timeout in milliseconds (for 'wait' action, default 5000)"
|
|
||||||
},
|
|
||||||
"filter": {
|
|
||||||
"type": "string",
|
|
||||||
"enum": ["interactive", "all"],
|
|
||||||
"description": "Element filter for 'read_page' (default: interactive)"
|
|
||||||
},
|
|
||||||
"expression": {
|
|
||||||
"type": "string",
|
|
||||||
"description": "JavaScript expression (for 'eval_js' action)"
|
|
||||||
}
|
|
||||||
},
|
|
||||||
"required": ["action"]
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
async fn execute(
|
|
||||||
&self,
|
|
||||||
params: serde_json::Value,
|
|
||||||
_ctx: &JobContext,
|
|
||||||
) -> Result<ToolOutput, ToolError> {
|
|
||||||
let start = std::time::Instant::now();
|
|
||||||
|
|
||||||
let action: BrowserAction = serde_json::from_value(params)
|
|
||||||
.map_err(|e| ToolError::InvalidParameters(format!("Invalid browser action: {}", e)))?;
|
|
||||||
|
|
||||||
// Launch browser on first use.
|
|
||||||
self.ensure_session().await?;
|
|
||||||
|
|
||||||
match action {
|
|
||||||
BrowserAction::Navigate { url } => {
|
|
||||||
let session = self.session.read().await;
|
|
||||||
let session = session.as_ref().ok_or_else(|| {
|
|
||||||
ToolError::ExecutionFailed("Browser session not initialized".to_string())
|
|
||||||
})?;
|
|
||||||
|
|
||||||
let title = session.navigate(&url).await?;
|
|
||||||
let current_url = session.current_url().await?;
|
|
||||||
|
|
||||||
Ok(ToolOutput::success(
|
|
||||||
serde_json::json!({
|
|
||||||
"url": current_url,
|
|
||||||
"title": title,
|
|
||||||
"status": "navigated"
|
|
||||||
}),
|
|
||||||
start.elapsed(),
|
|
||||||
))
|
|
||||||
}
|
|
||||||
|
|
||||||
BrowserAction::Back => {
|
|
||||||
let session = self.session.read().await;
|
|
||||||
let session = session.as_ref().ok_or_else(|| {
|
|
||||||
ToolError::ExecutionFailed("Browser session not initialized".to_string())
|
|
||||||
})?;
|
|
||||||
|
|
||||||
session.go_back().await?;
|
|
||||||
let url = session.current_url().await?;
|
|
||||||
|
|
||||||
Ok(ToolOutput::success(
|
|
||||||
serde_json::json!({ "url": url, "status": "navigated_back" }),
|
|
||||||
start.elapsed(),
|
|
||||||
))
|
|
||||||
}
|
|
||||||
|
|
||||||
BrowserAction::Forward => {
|
|
||||||
let session = self.session.read().await;
|
|
||||||
let session = session.as_ref().ok_or_else(|| {
|
|
||||||
ToolError::ExecutionFailed("Browser session not initialized".to_string())
|
|
||||||
})?;
|
|
||||||
|
|
||||||
session.go_forward().await?;
|
|
||||||
let url = session.current_url().await?;
|
|
||||||
|
|
||||||
Ok(ToolOutput::success(
|
|
||||||
serde_json::json!({ "url": url, "status": "navigated_forward" }),
|
|
||||||
start.elapsed(),
|
|
||||||
))
|
|
||||||
}
|
|
||||||
|
|
||||||
BrowserAction::ReadPage { filter } => {
|
|
||||||
let session = self.session.read().await;
|
|
||||||
let session = session.as_ref().ok_or_else(|| {
|
|
||||||
ToolError::ExecutionFailed("Browser session not initialized".to_string())
|
|
||||||
})?;
|
|
||||||
|
|
||||||
let element_filter = ElementFilter::from_str_opt(filter.as_deref());
|
|
||||||
let repr = session.read_page(element_filter).await?;
|
|
||||||
|
|
||||||
Ok(ToolOutput::text(repr, start.elapsed()))
|
|
||||||
}
|
|
||||||
|
|
||||||
BrowserAction::Click { ref_id } => {
|
|
||||||
let session = self.session.read().await;
|
|
||||||
let session = session.as_ref().ok_or_else(|| {
|
|
||||||
ToolError::ExecutionFailed("Browser session not initialized".to_string())
|
|
||||||
})?;
|
|
||||||
|
|
||||||
session.click_element(&ref_id).await?;
|
|
||||||
|
|
||||||
Ok(ToolOutput::success(
|
|
||||||
serde_json::json!({ "status": "clicked", "ref": ref_id }),
|
|
||||||
start.elapsed(),
|
|
||||||
))
|
|
||||||
}
|
|
||||||
|
|
||||||
BrowserAction::Type { ref_id, text } => {
|
|
||||||
let session = self.session.read().await;
|
|
||||||
let session = session.as_ref().ok_or_else(|| {
|
|
||||||
ToolError::ExecutionFailed("Browser session not initialized".to_string())
|
|
||||||
})?;
|
|
||||||
|
|
||||||
session.type_text(&ref_id, &text).await?;
|
|
||||||
|
|
||||||
Ok(ToolOutput::success(
|
|
||||||
serde_json::json!({
|
|
||||||
"status": "typed",
|
|
||||||
"ref": ref_id,
|
|
||||||
"length": text.len()
|
|
||||||
}),
|
|
||||||
start.elapsed(),
|
|
||||||
))
|
|
||||||
}
|
|
||||||
|
|
||||||
BrowserAction::Scroll { direction, amount } => {
|
|
||||||
let session = self.session.read().await;
|
|
||||||
let session = session.as_ref().ok_or_else(|| {
|
|
||||||
ToolError::ExecutionFailed("Browser session not initialized".to_string())
|
|
||||||
})?;
|
|
||||||
|
|
||||||
let steps = amount.unwrap_or(3);
|
|
||||||
session.scroll(&direction, steps).await?;
|
|
||||||
|
|
||||||
Ok(ToolOutput::success(
|
|
||||||
serde_json::json!({
|
|
||||||
"status": "scrolled",
|
|
||||||
"direction": direction,
|
|
||||||
"amount": steps
|
|
||||||
}),
|
|
||||||
start.elapsed(),
|
|
||||||
))
|
|
||||||
}
|
|
||||||
|
|
||||||
BrowserAction::Screenshot { full_page } => {
|
|
||||||
let session = self.session.read().await;
|
|
||||||
let session = session.as_ref().ok_or_else(|| {
|
|
||||||
ToolError::ExecutionFailed("Browser session not initialized".to_string())
|
|
||||||
})?;
|
|
||||||
|
|
||||||
let b64 = session.screenshot(full_page.unwrap_or(false)).await?;
|
|
||||||
|
|
||||||
Ok(ToolOutput::success(
|
|
||||||
serde_json::json!({
|
|
||||||
"format": "png",
|
|
||||||
"encoding": "base64",
|
|
||||||
"data": b64,
|
|
||||||
"full_page": full_page.unwrap_or(false)
|
|
||||||
}),
|
|
||||||
start.elapsed(),
|
|
||||||
))
|
|
||||||
}
|
|
||||||
|
|
||||||
BrowserAction::Extract { selector } => {
|
|
||||||
let session = self.session.read().await;
|
|
||||||
let session = session.as_ref().ok_or_else(|| {
|
|
||||||
ToolError::ExecutionFailed("Browser session not initialized".to_string())
|
|
||||||
})?;
|
|
||||||
|
|
||||||
let text = session.extract_text(selector.as_deref()).await?;
|
|
||||||
|
|
||||||
// Truncate very long text to avoid blowing up context.
|
|
||||||
let truncated = if text.len() > 32_000 {
|
|
||||||
format!(
|
|
||||||
"{}...\n\n[truncated, {} total chars]",
|
|
||||||
&text[..32_000],
|
|
||||||
text.len()
|
|
||||||
)
|
|
||||||
} else {
|
|
||||||
text.clone()
|
|
||||||
};
|
|
||||||
|
|
||||||
Ok(ToolOutput::text(&truncated, start.elapsed()).with_raw(text))
|
|
||||||
}
|
|
||||||
|
|
||||||
BrowserAction::Wait {
|
|
||||||
selector,
|
|
||||||
timeout_ms,
|
|
||||||
} => {
|
|
||||||
let session = self.session.read().await;
|
|
||||||
let session = session.as_ref().ok_or_else(|| {
|
|
||||||
ToolError::ExecutionFailed("Browser session not initialized".to_string())
|
|
||||||
})?;
|
|
||||||
|
|
||||||
let timeout = timeout_ms.unwrap_or(5000);
|
|
||||||
let found = session.wait(selector.as_deref(), timeout).await?;
|
|
||||||
|
|
||||||
Ok(ToolOutput::success(
|
|
||||||
serde_json::json!({
|
|
||||||
"found": found,
|
|
||||||
"selector": selector,
|
|
||||||
"timeout_ms": timeout
|
|
||||||
}),
|
|
||||||
start.elapsed(),
|
|
||||||
))
|
|
||||||
}
|
|
||||||
|
|
||||||
BrowserAction::EvalJs { expression } => {
|
|
||||||
let session = self.session.read().await;
|
|
||||||
let session = session.as_ref().ok_or_else(|| {
|
|
||||||
ToolError::ExecutionFailed("Browser session not initialized".to_string())
|
|
||||||
})?;
|
|
||||||
|
|
||||||
let result = session.eval_js(&expression).await?;
|
|
||||||
|
|
||||||
Ok(ToolOutput::success(
|
|
||||||
serde_json::json!({ "result": result }),
|
|
||||||
start.elapsed(),
|
|
||||||
))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
fn estimated_duration(&self, _params: &serde_json::Value) -> Option<Duration> {
|
|
||||||
Some(Duration::from_secs(10))
|
|
||||||
}
|
|
||||||
|
|
||||||
fn requires_sanitization(&self) -> bool {
|
|
||||||
true // Page content is untrusted external data
|
|
||||||
}
|
|
||||||
|
|
||||||
fn requires_approval(&self) -> bool {
|
|
||||||
true // Browser navigates to external sites, executes JS
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use crate::tools::builtin::browser::BrowserTool;
|
|
||||||
use crate::tools::tool::Tool;
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_browser_tool_metadata() {
|
|
||||||
let tool = BrowserTool::new();
|
|
||||||
assert_eq!(tool.name(), "browser");
|
|
||||||
assert!(tool.requires_approval());
|
|
||||||
assert!(tool.requires_sanitization());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_schema_has_action_enum() {
|
|
||||||
let tool = BrowserTool::new();
|
|
||||||
let schema = tool.parameters_schema();
|
|
||||||
|
|
||||||
let action_prop = schema.get("properties").and_then(|p| p.get("action"));
|
|
||||||
assert!(action_prop.is_some());
|
|
||||||
|
|
||||||
let action_enum = action_prop.and_then(|a| a.get("enum"));
|
|
||||||
assert!(action_enum.is_some());
|
|
||||||
|
|
||||||
let actions: Vec<&str> = action_enum
|
|
||||||
.and_then(|e| e.as_array())
|
|
||||||
.map(|arr| arr.iter().filter_map(|v| v.as_str()).collect())
|
|
||||||
.unwrap_or_default();
|
|
||||||
|
|
||||||
assert!(actions.contains(&"navigate"));
|
|
||||||
assert!(actions.contains(&"click"));
|
|
||||||
assert!(actions.contains(&"type"));
|
|
||||||
assert!(actions.contains(&"read_page"));
|
|
||||||
assert!(actions.contains(&"screenshot"));
|
|
||||||
assert!(actions.contains(&"eval_js"));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_action_deserialization() {
|
|
||||||
use super::BrowserAction;
|
|
||||||
|
|
||||||
// Navigate
|
|
||||||
let action: BrowserAction = serde_json::from_value(
|
|
||||||
serde_json::json!({"action": "navigate", "url": "https://x.com"}),
|
|
||||||
)
|
|
||||||
.unwrap();
|
|
||||||
assert!(matches!(action, BrowserAction::Navigate { url } if url == "https://x.com"));
|
|
||||||
|
|
||||||
// Click with "ref" alias
|
|
||||||
let action: BrowserAction =
|
|
||||||
serde_json::from_value(serde_json::json!({"action": "click", "ref": "@e1"})).unwrap();
|
|
||||||
assert!(matches!(action, BrowserAction::Click { ref_id } if ref_id == "@e1"));
|
|
||||||
|
|
||||||
// Click with "ref_id"
|
|
||||||
let action: BrowserAction =
|
|
||||||
serde_json::from_value(serde_json::json!({"action": "click", "ref_id": "e2"})).unwrap();
|
|
||||||
assert!(matches!(action, BrowserAction::Click { ref_id } if ref_id == "e2"));
|
|
||||||
|
|
||||||
// Type
|
|
||||||
let action: BrowserAction = serde_json::from_value(
|
|
||||||
serde_json::json!({"action": "type", "ref": "@e1", "text": "hello"}),
|
|
||||||
)
|
|
||||||
.unwrap();
|
|
||||||
assert!(
|
|
||||||
matches!(action, BrowserAction::Type { ref_id, text } if ref_id == "@e1" && text == "hello")
|
|
||||||
);
|
|
||||||
|
|
||||||
// ReadPage with default filter
|
|
||||||
let action: BrowserAction =
|
|
||||||
serde_json::from_value(serde_json::json!({"action": "read_page"})).unwrap();
|
|
||||||
assert!(matches!(action, BrowserAction::ReadPage { filter: None }));
|
|
||||||
|
|
||||||
// Screenshot
|
|
||||||
let action: BrowserAction =
|
|
||||||
serde_json::from_value(serde_json::json!({"action": "screenshot", "full_page": true}))
|
|
||||||
.unwrap();
|
|
||||||
assert!(matches!(
|
|
||||||
action,
|
|
||||||
BrowserAction::Screenshot {
|
|
||||||
full_page: Some(true)
|
|
||||||
}
|
|
||||||
));
|
|
||||||
|
|
||||||
// Invalid action
|
|
||||||
let result: Result<BrowserAction, _> =
|
|
||||||
serde_json::from_value(serde_json::json!({"action": "fly_to_moon"}));
|
|
||||||
assert!(result.is_err());
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,587 +0,0 @@
|
|||||||
//! Browser session management.
|
|
||||||
//!
|
|
||||||
//! Owns the Chrome process lifecycle and per-tab state. Sessions are spawned
|
|
||||||
//! lazily on first browser action and torn down when dropped.
|
|
||||||
//!
|
|
||||||
//! ```text
|
|
||||||
//! BrowserSession
|
|
||||||
//! ├── Browser (chromiumoxide, owns Chrome child process)
|
|
||||||
//! ├── handler_task (JoinHandle polling CDP WebSocket)
|
|
||||||
//! ├── tabs: HashMap<tab_id, Page>
|
|
||||||
//! ├── active_tab: current tab id
|
|
||||||
//! └── element_refs: ElementRefMap (valid until next read_page)
|
|
||||||
//! ```
|
|
||||||
|
|
||||||
use std::collections::HashMap;
|
|
||||||
use std::path::PathBuf;
|
|
||||||
use std::sync::Arc;
|
|
||||||
|
|
||||||
use chromiumoxide::Page;
|
|
||||||
use chromiumoxide::browser::{Browser, BrowserConfig};
|
|
||||||
use chromiumoxide::cdp::browser_protocol::accessibility::GetFullAxTreeParams;
|
|
||||||
use chromiumoxide::cdp::browser_protocol::dom::{GetBoxModelParams, ScrollIntoViewIfNeededParams};
|
|
||||||
use chromiumoxide::cdp::browser_protocol::input::{
|
|
||||||
DispatchMouseEventParams, DispatchMouseEventType, InsertTextParams, MouseButton,
|
|
||||||
};
|
|
||||||
use chromiumoxide::cdp::browser_protocol::page::CaptureScreenshotFormat;
|
|
||||||
use chromiumoxide::page::ScreenshotParams;
|
|
||||||
use futures::StreamExt;
|
|
||||||
use tokio::sync::RwLock;
|
|
||||||
use tokio::task::JoinHandle;
|
|
||||||
|
|
||||||
use crate::tools::builtin::browser::accessibility::{
|
|
||||||
ElementFilter, ElementRefMap, build_page_repr,
|
|
||||||
};
|
|
||||||
use crate::tools::builtin::browser::stealth;
|
|
||||||
use crate::tools::tool::ToolError;
|
|
||||||
|
|
||||||
/// Manages a Chrome browser instance and its tabs.
|
|
||||||
pub struct BrowserSession {
|
|
||||||
#[allow(dead_code)] // Used by new_tab() which is reserved for tab management actions
|
|
||||||
browser: Browser,
|
|
||||||
_handler_task: JoinHandle<()>,
|
|
||||||
tabs: HashMap<String, Page>,
|
|
||||||
active_tab: String,
|
|
||||||
element_refs: Arc<RwLock<ElementRefMap>>,
|
|
||||||
#[allow(dead_code)] // Used by new_tab() which is reserved for tab management actions
|
|
||||||
stealth_js: String,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl BrowserSession {
|
|
||||||
/// Launch a new Chrome browser session.
|
|
||||||
///
|
|
||||||
/// Locates Chrome on the system, applies stealth patches, and opens
|
|
||||||
/// an initial blank tab.
|
|
||||||
pub async fn launch() -> Result<Self, ToolError> {
|
|
||||||
let chrome_path = find_chrome().ok_or_else(|| {
|
|
||||||
ToolError::ExecutionFailed(
|
|
||||||
"Chrome/Chromium not found. Install Chrome or set CHROME_PATH.".to_string(),
|
|
||||||
)
|
|
||||||
})?;
|
|
||||||
|
|
||||||
// Shared profile so the agent accumulates useful state across sessions
|
|
||||||
// (logged-in sessions, dismissed cookie banners, local storage).
|
|
||||||
// Delete ~/.ironclaw/browser/profile/ to reset.
|
|
||||||
let profile_dir = browser_profile_dir();
|
|
||||||
std::fs::create_dir_all(&profile_dir).map_err(|e| {
|
|
||||||
ToolError::ExecutionFailed(format!("Failed to create browser profile dir: {}", e))
|
|
||||||
})?;
|
|
||||||
|
|
||||||
let mut config_builder = BrowserConfig::builder()
|
|
||||||
.chrome_executable(&chrome_path)
|
|
||||||
.user_data_dir(&profile_dir)
|
|
||||||
.window_size(1920, 1080)
|
|
||||||
.no_sandbox();
|
|
||||||
|
|
||||||
for arg in stealth::stealth_args() {
|
|
||||||
config_builder = config_builder.arg(arg);
|
|
||||||
}
|
|
||||||
|
|
||||||
let config = config_builder.build().map_err(|e| {
|
|
||||||
ToolError::ExecutionFailed(format!("Failed to build browser config: {}", e))
|
|
||||||
})?;
|
|
||||||
|
|
||||||
let (browser, mut handler) = Browser::launch(config)
|
|
||||||
.await
|
|
||||||
.map_err(|e| ToolError::ExecutionFailed(format!("Failed to launch Chrome: {}", e)))?;
|
|
||||||
|
|
||||||
// The handler must be polled continuously or the CDP connection dies.
|
|
||||||
let handler_task = tokio::spawn(async move {
|
|
||||||
while let Some(event) = handler.next().await {
|
|
||||||
if event.is_err() {
|
|
||||||
tracing::warn!("Browser handler error: {:?}", event);
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
});
|
|
||||||
|
|
||||||
// Open initial tab.
|
|
||||||
let page = browser.new_page("about:blank").await.map_err(|e| {
|
|
||||||
ToolError::ExecutionFailed(format!("Failed to open initial tab: {}", e))
|
|
||||||
})?;
|
|
||||||
|
|
||||||
// Inject stealth JS on every new document load for this page.
|
|
||||||
let stealth_js = stealth::stealth_js().to_string();
|
|
||||||
page.evaluate_on_new_document(stealth_js.clone())
|
|
||||||
.await
|
|
||||||
.map_err(|e| {
|
|
||||||
ToolError::ExecutionFailed(format!("Failed to inject stealth JS: {}", e))
|
|
||||||
})?;
|
|
||||||
|
|
||||||
let tab_id = "tab0".to_string();
|
|
||||||
let mut tabs = HashMap::new();
|
|
||||||
tabs.insert(tab_id.clone(), page);
|
|
||||||
|
|
||||||
Ok(Self {
|
|
||||||
browser,
|
|
||||||
_handler_task: handler_task,
|
|
||||||
tabs,
|
|
||||||
active_tab: tab_id,
|
|
||||||
element_refs: Arc::new(RwLock::new(ElementRefMap::new())),
|
|
||||||
stealth_js,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Get the active page, or error if session is broken.
|
|
||||||
fn active_page(&self) -> Result<&Page, ToolError> {
|
|
||||||
self.tabs.get(&self.active_tab).ok_or_else(|| {
|
|
||||||
ToolError::ExecutionFailed(format!("No active tab: {}", self.active_tab))
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// --- Navigation ---
|
|
||||||
|
|
||||||
pub async fn navigate(&self, url: &str) -> Result<String, ToolError> {
|
|
||||||
let page = self.active_page()?;
|
|
||||||
page.goto(url)
|
|
||||||
.await
|
|
||||||
.map_err(|e| ToolError::ExternalService(format!("Navigation failed: {}", e)))?;
|
|
||||||
|
|
||||||
let title = page
|
|
||||||
.get_title()
|
|
||||||
.await
|
|
||||||
.map_err(|e| ToolError::ExecutionFailed(format!("Failed to get page title: {}", e)))?
|
|
||||||
.unwrap_or_default();
|
|
||||||
|
|
||||||
Ok(title)
|
|
||||||
}
|
|
||||||
|
|
||||||
pub async fn go_back(&self) -> Result<(), ToolError> {
|
|
||||||
let page = self.active_page()?;
|
|
||||||
page.evaluate("window.history.back()")
|
|
||||||
.await
|
|
||||||
.map_err(|e| ToolError::ExecutionFailed(format!("Failed to go back: {}", e)))?;
|
|
||||||
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
|
|
||||||
pub async fn go_forward(&self) -> Result<(), ToolError> {
|
|
||||||
let page = self.active_page()?;
|
|
||||||
page.evaluate("window.history.forward()")
|
|
||||||
.await
|
|
||||||
.map_err(|e| ToolError::ExecutionFailed(format!("Failed to go forward: {}", e)))?;
|
|
||||||
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
|
|
||||||
// --- Page reading ---
|
|
||||||
|
|
||||||
/// Build accessibility tree representation and update element refs.
|
|
||||||
pub async fn read_page(&self, filter: ElementFilter) -> Result<String, ToolError> {
|
|
||||||
let page = self.active_page()?;
|
|
||||||
|
|
||||||
let url = page
|
|
||||||
.url()
|
|
||||||
.await
|
|
||||||
.map_err(|e| ToolError::ExecutionFailed(format!("Failed to get URL: {}", e)))?
|
|
||||||
.unwrap_or_else(|| "about:blank".to_string());
|
|
||||||
|
|
||||||
let title = page
|
|
||||||
.get_title()
|
|
||||||
.await
|
|
||||||
.map_err(|e| ToolError::ExecutionFailed(format!("Failed to get title: {}", e)))?
|
|
||||||
.unwrap_or_default();
|
|
||||||
|
|
||||||
// Fetch full accessibility tree via CDP.
|
|
||||||
let ax_result = page
|
|
||||||
.execute(GetFullAxTreeParams::default())
|
|
||||||
.await
|
|
||||||
.map_err(|e| {
|
|
||||||
ToolError::ExecutionFailed(format!("Failed to get accessibility tree: {}", e))
|
|
||||||
})?;
|
|
||||||
|
|
||||||
let nodes = ax_result.result.nodes;
|
|
||||||
|
|
||||||
let mut ref_map = self.element_refs.write().await;
|
|
||||||
let repr = build_page_repr(&url, &title, &nodes, filter, &mut ref_map);
|
|
||||||
|
|
||||||
Ok(repr)
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Extract text content from the page or a CSS selector.
|
|
||||||
pub async fn extract_text(&self, selector: Option<&str>) -> Result<String, ToolError> {
|
|
||||||
let page = self.active_page()?;
|
|
||||||
|
|
||||||
let js = match selector {
|
|
||||||
Some(sel) => {
|
|
||||||
let escaped = serde_json::to_string(sel).map_err(|e| {
|
|
||||||
ToolError::InvalidParameters(format!("Invalid selector: {}", e))
|
|
||||||
})?;
|
|
||||||
format!(
|
|
||||||
"(() => {{ const el = document.querySelector({}); return el ? el.innerText : null; }})()",
|
|
||||||
escaped
|
|
||||||
)
|
|
||||||
}
|
|
||||||
None => "document.body.innerText".to_string(),
|
|
||||||
};
|
|
||||||
|
|
||||||
let result: Option<String> = page
|
|
||||||
.evaluate(js.as_str())
|
|
||||||
.await
|
|
||||||
.map_err(|e| ToolError::ExecutionFailed(format!("Failed to extract text: {}", e)))?
|
|
||||||
.into_value()
|
|
||||||
.map_err(|e| {
|
|
||||||
ToolError::ExecutionFailed(format!("Failed to deserialize text: {}", e))
|
|
||||||
})?;
|
|
||||||
|
|
||||||
Ok(result.unwrap_or_default())
|
|
||||||
}
|
|
||||||
|
|
||||||
// --- Interaction ---
|
|
||||||
|
|
||||||
/// Click an element by reference ID (e.g., "e1" or "@e1").
|
|
||||||
///
|
|
||||||
/// Uses DOM.scrollIntoViewIfNeeded + DOM.getBoxModel to find the element's
|
|
||||||
/// center coordinates, then dispatches mouse press + release at that point.
|
|
||||||
pub async fn click_element(&self, ref_id: &str) -> Result<(), ToolError> {
|
|
||||||
let page = self.active_page()?;
|
|
||||||
let refs = self.element_refs.read().await;
|
|
||||||
|
|
||||||
let elem_ref = refs.get(ref_id).ok_or_else(|| {
|
|
||||||
ToolError::InvalidParameters(format!(
|
|
||||||
"Unknown element reference '{}'. Call browser with action 'read_page' first.",
|
|
||||||
ref_id
|
|
||||||
))
|
|
||||||
})?;
|
|
||||||
|
|
||||||
let backend_node_id = elem_ref.backend_node_id;
|
|
||||||
drop(refs);
|
|
||||||
|
|
||||||
// Scroll the element into the viewport.
|
|
||||||
page.execute(
|
|
||||||
ScrollIntoViewIfNeededParams::builder()
|
|
||||||
.backend_node_id(backend_node_id)
|
|
||||||
.build(),
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(|e| {
|
|
||||||
ToolError::ExecutionFailed(format!("Failed to scroll element into view: {}", e))
|
|
||||||
})?;
|
|
||||||
|
|
||||||
// Get element's bounding box via DOM.getBoxModel.
|
|
||||||
let box_result = page
|
|
||||||
.execute(
|
|
||||||
GetBoxModelParams::builder()
|
|
||||||
.backend_node_id(backend_node_id)
|
|
||||||
.build(),
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(|e| {
|
|
||||||
ToolError::ExecutionFailed(format!("Failed to get element box model: {}", e))
|
|
||||||
})?;
|
|
||||||
|
|
||||||
// Content quad is [x1,y1, x2,y2, x3,y3, x4,y4]. Center = average of 4 corners.
|
|
||||||
let content = box_result.result.model.content.inner();
|
|
||||||
if content.len() < 8 {
|
|
||||||
return Err(ToolError::ExecutionFailed(
|
|
||||||
"Element has no valid bounding box".to_string(),
|
|
||||||
));
|
|
||||||
}
|
|
||||||
|
|
||||||
let x = (content[0] + content[2] + content[4] + content[6]) / 4.0;
|
|
||||||
let y = (content[1] + content[3] + content[5] + content[7]) / 4.0;
|
|
||||||
|
|
||||||
// Dispatch mouse press + release at center of element.
|
|
||||||
page.execute(
|
|
||||||
DispatchMouseEventParams::builder()
|
|
||||||
.r#type(DispatchMouseEventType::MousePressed)
|
|
||||||
.x(x)
|
|
||||||
.y(y)
|
|
||||||
.button(MouseButton::Left)
|
|
||||||
.click_count(1)
|
|
||||||
.build()
|
|
||||||
.map_err(|e| {
|
|
||||||
ToolError::ExecutionFailed(format!("Failed to build mouse event: {}", e))
|
|
||||||
})?,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(|e| ToolError::ExecutionFailed(format!("Mouse press failed: {}", e)))?;
|
|
||||||
|
|
||||||
page.execute(
|
|
||||||
DispatchMouseEventParams::builder()
|
|
||||||
.r#type(DispatchMouseEventType::MouseReleased)
|
|
||||||
.x(x)
|
|
||||||
.y(y)
|
|
||||||
.button(MouseButton::Left)
|
|
||||||
.click_count(1)
|
|
||||||
.build()
|
|
||||||
.map_err(|e| {
|
|
||||||
ToolError::ExecutionFailed(format!("Failed to build mouse event: {}", e))
|
|
||||||
})?,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
.map_err(|e| ToolError::ExecutionFailed(format!("Mouse release failed: {}", e)))?;
|
|
||||||
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Type text into an element by reference ID.
|
|
||||||
pub async fn type_text(&self, ref_id: &str, text: &str) -> Result<(), ToolError> {
|
|
||||||
// First click to focus the element.
|
|
||||||
self.click_element(ref_id).await?;
|
|
||||||
|
|
||||||
// Brief delay to let focus settle.
|
|
||||||
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
|
|
||||||
|
|
||||||
let page = self.active_page()?;
|
|
||||||
|
|
||||||
// Use CDP insertText for reliable IME-style text entry.
|
|
||||||
page.execute(InsertTextParams::new(text))
|
|
||||||
.await
|
|
||||||
.map_err(|e| ToolError::ExecutionFailed(format!("Failed to type text: {}", e)))?;
|
|
||||||
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Scroll the page.
|
|
||||||
pub async fn scroll(&self, direction: &str, amount: u32) -> Result<(), ToolError> {
|
|
||||||
let page = self.active_page()?;
|
|
||||||
|
|
||||||
let (dx, dy) = match direction {
|
|
||||||
"up" => (0, -(amount as i32 * 100)),
|
|
||||||
"down" => (0, amount as i32 * 100),
|
|
||||||
"left" => (-(amount as i32 * 100), 0),
|
|
||||||
"right" => (amount as i32 * 100, 0),
|
|
||||||
_ => {
|
|
||||||
return Err(ToolError::InvalidParameters(format!(
|
|
||||||
"Invalid scroll direction '{}'. Use: up, down, left, right",
|
|
||||||
direction
|
|
||||||
)));
|
|
||||||
}
|
|
||||||
};
|
|
||||||
|
|
||||||
let js = format!("window.scrollBy({}, {})", dx, dy);
|
|
||||||
page.evaluate(js.as_str())
|
|
||||||
.await
|
|
||||||
.map_err(|e| ToolError::ExecutionFailed(format!("Scroll failed: {}", e)))?;
|
|
||||||
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Wait for a CSS selector to appear, or a fixed timeout.
|
|
||||||
pub async fn wait(&self, selector: Option<&str>, timeout_ms: u64) -> Result<bool, ToolError> {
|
|
||||||
let page = self.active_page()?;
|
|
||||||
let timeout = std::time::Duration::from_millis(timeout_ms);
|
|
||||||
|
|
||||||
match selector {
|
|
||||||
Some(sel) => {
|
|
||||||
let poll_interval = std::time::Duration::from_millis(100);
|
|
||||||
let start = std::time::Instant::now();
|
|
||||||
let escaped = serde_json::to_string(sel).map_err(|e| {
|
|
||||||
ToolError::InvalidParameters(format!("Invalid selector: {}", e))
|
|
||||||
})?;
|
|
||||||
|
|
||||||
loop {
|
|
||||||
let js = format!("!!document.querySelector({})", escaped);
|
|
||||||
let found: bool = page
|
|
||||||
.evaluate(js.as_str())
|
|
||||||
.await
|
|
||||||
.map_err(|e| {
|
|
||||||
ToolError::ExecutionFailed(format!("Wait poll failed: {}", e))
|
|
||||||
})?
|
|
||||||
.into_value()
|
|
||||||
.unwrap_or(false);
|
|
||||||
|
|
||||||
if found {
|
|
||||||
return Ok(true);
|
|
||||||
}
|
|
||||||
|
|
||||||
if start.elapsed() >= timeout {
|
|
||||||
return Ok(false);
|
|
||||||
}
|
|
||||||
|
|
||||||
tokio::time::sleep(poll_interval).await;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
None => {
|
|
||||||
tokio::time::sleep(timeout).await;
|
|
||||||
Ok(true)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// --- Screenshots ---
|
|
||||||
|
|
||||||
/// Capture a screenshot as base64-encoded PNG.
|
|
||||||
pub async fn screenshot(&self, full_page: bool) -> Result<String, ToolError> {
|
|
||||||
let page = self.active_page()?;
|
|
||||||
|
|
||||||
let params = ScreenshotParams::builder()
|
|
||||||
.format(CaptureScreenshotFormat::Png)
|
|
||||||
.full_page(full_page)
|
|
||||||
.build();
|
|
||||||
|
|
||||||
let bytes = page
|
|
||||||
.screenshot(params)
|
|
||||||
.await
|
|
||||||
.map_err(|e| ToolError::ExecutionFailed(format!("Screenshot failed: {}", e)))?;
|
|
||||||
|
|
||||||
use base64::Engine;
|
|
||||||
Ok(base64::engine::general_purpose::STANDARD.encode(&bytes))
|
|
||||||
}
|
|
||||||
|
|
||||||
// --- JavaScript ---
|
|
||||||
|
|
||||||
/// Execute arbitrary JavaScript and return the result.
|
|
||||||
pub async fn eval_js(&self, expression: &str) -> Result<serde_json::Value, ToolError> {
|
|
||||||
let page = self.active_page()?;
|
|
||||||
|
|
||||||
let result = page
|
|
||||||
.evaluate(expression)
|
|
||||||
.await
|
|
||||||
.map_err(|e| ToolError::ExecutionFailed(format!("JS evaluation failed: {}", e)))?;
|
|
||||||
|
|
||||||
let value: serde_json::Value = result.into_value().unwrap_or(serde_json::Value::Null);
|
|
||||||
|
|
||||||
Ok(value)
|
|
||||||
}
|
|
||||||
|
|
||||||
// --- Tab management ---
|
|
||||||
|
|
||||||
/// Open a new tab and make it active.
|
|
||||||
#[allow(dead_code)] // Reserved for tab management actions
|
|
||||||
pub async fn new_tab(&mut self, url: &str) -> Result<String, ToolError> {
|
|
||||||
let page =
|
|
||||||
self.browser.new_page(url).await.map_err(|e| {
|
|
||||||
ToolError::ExecutionFailed(format!("Failed to open new tab: {}", e))
|
|
||||||
})?;
|
|
||||||
|
|
||||||
// Inject stealth JS on the new page too.
|
|
||||||
page.evaluate_on_new_document(self.stealth_js.clone())
|
|
||||||
.await
|
|
||||||
.map_err(|e| {
|
|
||||||
ToolError::ExecutionFailed(format!("Failed to inject stealth JS on new tab: {}", e))
|
|
||||||
})?;
|
|
||||||
|
|
||||||
let tab_id = format!("tab{}", self.tabs.len());
|
|
||||||
self.tabs.insert(tab_id.clone(), page);
|
|
||||||
self.active_tab = tab_id.clone();
|
|
||||||
|
|
||||||
// Clear element refs since we're on a new page.
|
|
||||||
self.element_refs.write().await.reset();
|
|
||||||
|
|
||||||
Ok(tab_id)
|
|
||||||
}
|
|
||||||
|
|
||||||
/// List open tabs.
|
|
||||||
#[allow(dead_code)] // Reserved for tab management actions
|
|
||||||
pub fn list_tabs(&self) -> Vec<String> {
|
|
||||||
self.tabs.keys().cloned().collect()
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Switch to a different tab.
|
|
||||||
#[allow(dead_code)] // Reserved for tab management actions
|
|
||||||
pub async fn switch_tab(&mut self, tab_id: &str) -> Result<(), ToolError> {
|
|
||||||
if !self.tabs.contains_key(tab_id) {
|
|
||||||
return Err(ToolError::InvalidParameters(format!(
|
|
||||||
"Unknown tab '{}'. Open tabs: {:?}",
|
|
||||||
tab_id,
|
|
||||||
self.list_tabs()
|
|
||||||
)));
|
|
||||||
}
|
|
||||||
|
|
||||||
self.active_tab = tab_id.to_string();
|
|
||||||
// Clear element refs when switching tabs.
|
|
||||||
self.element_refs.write().await.reset();
|
|
||||||
Ok(())
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Get current page URL.
|
|
||||||
pub async fn current_url(&self) -> Result<String, ToolError> {
|
|
||||||
let page = self.active_page()?;
|
|
||||||
page.url()
|
|
||||||
.await
|
|
||||||
.map_err(|e| ToolError::ExecutionFailed(format!("Failed to get URL: {}", e)))
|
|
||||||
.map(|u| u.unwrap_or_else(|| "about:blank".to_string()))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
impl Drop for BrowserSession {
|
|
||||||
fn drop(&mut self) {
|
|
||||||
tracing::debug!("Browser session dropping, Chrome process will be cleaned up");
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Returns `~/.ironclaw/browser/profile/`.
|
|
||||||
fn browser_profile_dir() -> PathBuf {
|
|
||||||
dirs::home_dir()
|
|
||||||
.unwrap_or_else(|| PathBuf::from("."))
|
|
||||||
.join(".ironclaw")
|
|
||||||
.join("browser")
|
|
||||||
.join("profile")
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Search common locations for a Chrome/Chromium binary.
|
|
||||||
pub fn find_chrome() -> Option<PathBuf> {
|
|
||||||
// Environment variable override.
|
|
||||||
if let Ok(path) = std::env::var("CHROME_PATH") {
|
|
||||||
let p = PathBuf::from(&path);
|
|
||||||
if p.exists() {
|
|
||||||
return Some(p);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
let candidates = if cfg!(target_os = "macos") {
|
|
||||||
vec![
|
|
||||||
"/Applications/Google Chrome.app/Contents/MacOS/Google Chrome",
|
|
||||||
"/Applications/Chromium.app/Contents/MacOS/Chromium",
|
|
||||||
"/Applications/Google Chrome Canary.app/Contents/MacOS/Google Chrome Canary",
|
|
||||||
"/Applications/Brave Browser.app/Contents/MacOS/Brave Browser",
|
|
||||||
]
|
|
||||||
} else if cfg!(target_os = "linux") {
|
|
||||||
vec![
|
|
||||||
"/usr/bin/google-chrome",
|
|
||||||
"/usr/bin/google-chrome-stable",
|
|
||||||
"/usr/bin/chromium",
|
|
||||||
"/usr/bin/chromium-browser",
|
|
||||||
"/snap/bin/chromium",
|
|
||||||
]
|
|
||||||
} else {
|
|
||||||
// Windows paths.
|
|
||||||
vec![
|
|
||||||
r"C:\Program Files\Google\Chrome\Application\chrome.exe",
|
|
||||||
r"C:\Program Files (x86)\Google\Chrome\Application\chrome.exe",
|
|
||||||
]
|
|
||||||
};
|
|
||||||
|
|
||||||
for candidate in candidates {
|
|
||||||
let p = PathBuf::from(candidate);
|
|
||||||
if p.exists() {
|
|
||||||
return Some(p);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
which_chrome_in_path()
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Check if chrome/chromium is available in PATH.
|
|
||||||
fn which_chrome_in_path() -> Option<PathBuf> {
|
|
||||||
let path_var = std::env::var("PATH").ok()?;
|
|
||||||
let separator = if cfg!(windows) { ';' } else { ':' };
|
|
||||||
for name in &["google-chrome", "chromium", "chromium-browser", "chrome"] {
|
|
||||||
for dir in path_var.split(separator) {
|
|
||||||
let candidate = PathBuf::from(dir).join(name);
|
|
||||||
if candidate.exists() {
|
|
||||||
return Some(candidate);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
None
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use crate::tools::builtin::browser::session::find_chrome;
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn test_find_chrome_returns_path_or_none() {
|
|
||||||
let result = find_chrome();
|
|
||||||
if let Some(path) = &result {
|
|
||||||
assert!(
|
|
||||||
path.exists(),
|
|
||||||
"find_chrome returned non-existent path: {:?}",
|
|
||||||
path
|
|
||||||
);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,158 +0,0 @@
|
|||||||
//! Anti-detection JavaScript patches for headless Chrome.
|
|
||||||
//!
|
|
||||||
//! Injects scripts via `Page.addScriptToEvaluateOnNewDocument` to suppress
|
|
||||||
//! common bot-detection signals. Handles ~80% of detection for legitimate
|
|
||||||
//! browsing (not adversarial scraping against Cloudflare Enterprise).
|
|
||||||
//!
|
|
||||||
//! What we patch:
|
|
||||||
//! - `navigator.webdriver` (trivial but still checked)
|
|
||||||
//! - `navigator.plugins` (headless has empty plugin list)
|
|
||||||
//! - `navigator.languages` (match system locale)
|
|
||||||
//! - `chrome.runtime` (looks like a real extension API)
|
|
||||||
//! - `HeadlessChrome` user-agent substring (suppressed via launch flags)
|
|
||||||
|
|
||||||
/// Chrome launch arguments that reduce detection surface.
|
|
||||||
pub fn stealth_args() -> Vec<&'static str> {
|
|
||||||
vec![
|
|
||||||
"--disable-blink-features=AutomationControlled",
|
|
||||||
"--no-first-run",
|
|
||||||
"--no-default-browser-check",
|
|
||||||
"--disable-infobars",
|
|
||||||
"--disable-background-networking",
|
|
||||||
"--disable-prompt-on-repost",
|
|
||||||
"--disable-hang-monitor",
|
|
||||||
"--disable-sync",
|
|
||||||
"--metrics-recording-only",
|
|
||||||
"--no-service-autorun",
|
|
||||||
]
|
|
||||||
}
|
|
||||||
|
|
||||||
/// JavaScript injected before any page scripts run.
|
|
||||||
///
|
|
||||||
/// This covers the most common fingerprinting checks. Each patch is
|
|
||||||
/// a self-contained IIFE so failures in one don't break the others.
|
|
||||||
pub fn stealth_js() -> &'static str {
|
|
||||||
r#"
|
|
||||||
// --- navigator.webdriver ---
|
|
||||||
// CDP sets this to true; real browsers have it undefined or false.
|
|
||||||
(() => {
|
|
||||||
Object.defineProperty(navigator, 'webdriver', {
|
|
||||||
get: () => undefined,
|
|
||||||
configurable: true,
|
|
||||||
});
|
|
||||||
})();
|
|
||||||
|
|
||||||
// --- navigator.plugins ---
|
|
||||||
// Headless Chrome reports an empty plugin array. Real Chrome on desktop
|
|
||||||
// always has at least these two. We fake the array shape.
|
|
||||||
(() => {
|
|
||||||
const pluginData = [
|
|
||||||
{ name: 'Chrome PDF Plugin', filename: 'internal-pdf-viewer',
|
|
||||||
description: 'Portable Document Format' },
|
|
||||||
{ name: 'Chrome PDF Viewer', filename: 'mhjfbmdgcfjbbpaeojofohoefgiehjai',
|
|
||||||
description: '' },
|
|
||||||
];
|
|
||||||
|
|
||||||
const makeMimeType = (type_, suffixes, desc, plugin) => {
|
|
||||||
const mt = Object.create(MimeType.prototype);
|
|
||||||
Object.defineProperties(mt, {
|
|
||||||
type: { get: () => type_ },
|
|
||||||
suffixes: { get: () => suffixes },
|
|
||||||
description: { get: () => desc },
|
|
||||||
enabledPlugin: { get: () => plugin },
|
|
||||||
});
|
|
||||||
return mt;
|
|
||||||
};
|
|
||||||
|
|
||||||
const makePlugin = (data) => {
|
|
||||||
const p = Object.create(Plugin.prototype);
|
|
||||||
const mimes = [makeMimeType('application/pdf', 'pdf', 'Portable Document Format', p)];
|
|
||||||
Object.defineProperties(p, {
|
|
||||||
name: { get: () => data.name },
|
|
||||||
filename: { get: () => data.filename },
|
|
||||||
description: { get: () => data.description },
|
|
||||||
length: { get: () => mimes.length },
|
|
||||||
0: { get: () => mimes[0] },
|
|
||||||
});
|
|
||||||
p.item = (i) => mimes[i] || null;
|
|
||||||
p.namedItem = (name) => mimes.find(m => m.type === name) || null;
|
|
||||||
return p;
|
|
||||||
};
|
|
||||||
|
|
||||||
const plugins = pluginData.map(makePlugin);
|
|
||||||
const pluginArray = Object.create(PluginArray.prototype);
|
|
||||||
Object.defineProperties(pluginArray, {
|
|
||||||
length: { get: () => plugins.length },
|
|
||||||
0: { get: () => plugins[0] },
|
|
||||||
1: { get: () => plugins[1] },
|
|
||||||
});
|
|
||||||
pluginArray.item = (i) => plugins[i] || null;
|
|
||||||
pluginArray.namedItem = (name) => plugins.find(p => p.name === name) || null;
|
|
||||||
pluginArray.refresh = () => {};
|
|
||||||
pluginArray[Symbol.iterator] = function* () { yield* plugins; };
|
|
||||||
|
|
||||||
Object.defineProperty(navigator, 'plugins', {
|
|
||||||
get: () => pluginArray,
|
|
||||||
configurable: true,
|
|
||||||
});
|
|
||||||
})();
|
|
||||||
|
|
||||||
// --- navigator.languages ---
|
|
||||||
// Headless sometimes reports just ['en'] instead of a realistic list.
|
|
||||||
(() => {
|
|
||||||
Object.defineProperty(navigator, 'languages', {
|
|
||||||
get: () => ['en-US', 'en'],
|
|
||||||
configurable: true,
|
|
||||||
});
|
|
||||||
})();
|
|
||||||
|
|
||||||
// --- chrome.runtime ---
|
|
||||||
// Bot detectors check for chrome.runtime to see if it's a real Chrome
|
|
||||||
// extension environment. CDP-controlled Chrome has a broken stub.
|
|
||||||
(() => {
|
|
||||||
if (!window.chrome) window.chrome = {};
|
|
||||||
if (!window.chrome.runtime) {
|
|
||||||
window.chrome.runtime = {
|
|
||||||
connect: () => {},
|
|
||||||
sendMessage: () => {},
|
|
||||||
id: undefined,
|
|
||||||
};
|
|
||||||
}
|
|
||||||
})();
|
|
||||||
|
|
||||||
// --- Permissions API ---
|
|
||||||
// Headless reports 'denied' for notification permissions by default,
|
|
||||||
// which is a known fingerprinting signal.
|
|
||||||
(() => {
|
|
||||||
const originalQuery = window.Permissions?.prototype?.query;
|
|
||||||
if (originalQuery) {
|
|
||||||
window.Permissions.prototype.query = function(params) {
|
|
||||||
if (params?.name === 'notifications') {
|
|
||||||
return Promise.resolve({ state: 'prompt', onchange: null });
|
|
||||||
}
|
|
||||||
return originalQuery.call(this, params);
|
|
||||||
};
|
|
||||||
}
|
|
||||||
})();
|
|
||||||
"#
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use crate::tools::builtin::browser::stealth;
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn stealth_js_is_not_empty() {
|
|
||||||
let js = stealth::stealth_js();
|
|
||||||
assert!(js.len() > 100);
|
|
||||||
assert!(js.contains("navigator"));
|
|
||||||
assert!(js.contains("webdriver"));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn stealth_args_are_valid_flags() {
|
|
||||||
for arg in stealth::stealth_args() {
|
|
||||||
assert!(arg.starts_with("--"), "arg should start with --: {}", arg);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -5,13 +5,18 @@ use std::net::{IpAddr, ToSocketAddrs};
|
|||||||
use std::time::Duration;
|
use std::time::Duration;
|
||||||
|
|
||||||
use async_trait::async_trait;
|
use async_trait::async_trait;
|
||||||
|
use futures::StreamExt;
|
||||||
use reqwest::Client;
|
use reqwest::Client;
|
||||||
|
|
||||||
use crate::context::JobContext;
|
use crate::context::JobContext;
|
||||||
use crate::safety::LeakDetector;
|
use crate::safety::LeakDetector;
|
||||||
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
|
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
|
||||||
|
|
||||||
/// Maximum response body size (5 MB). Prevents OOM from unbounded responses.
|
/// Maximum response body size (5 MB).
|
||||||
|
///
|
||||||
|
/// 5 MB is large enough for typical JSON API responses and moderate HTML pages,
|
||||||
|
/// but small enough to prevent OOM from malicious or runaway servers. The WASM
|
||||||
|
/// HTTP wrapper uses the same limit for consistency.
|
||||||
const MAX_RESPONSE_SIZE: usize = 5 * 1024 * 1024;
|
const MAX_RESPONSE_SIZE: usize = 5 * 1024 * 1024;
|
||||||
|
|
||||||
/// Tool for making HTTP requests.
|
/// Tool for making HTTP requests.
|
||||||
@@ -230,18 +235,42 @@ impl Tool for HttpTool {
|
|||||||
.filter_map(|(k, v)| v.to_str().ok().map(|v| (k.to_string(), v.to_string())))
|
.filter_map(|(k, v)| v.to_str().ok().map(|v| (k.to_string(), v.to_string())))
|
||||||
.collect();
|
.collect();
|
||||||
|
|
||||||
// Get response body with size cap to prevent OOM
|
// Pre-check Content-Length header to reject obviously oversized responses
|
||||||
let body_bytes = response.bytes().await.map_err(|e| {
|
// before downloading anything, preventing OOM from malicious servers.
|
||||||
|
if let Some(content_length) = response.headers().get(reqwest::header::CONTENT_LENGTH)
|
||||||
|
&& let Ok(s) = content_length.to_str()
|
||||||
|
&& let Ok(len) = s.parse::<usize>()
|
||||||
|
&& len > MAX_RESPONSE_SIZE
|
||||||
|
{
|
||||||
|
tracing::warn!(
|
||||||
|
url = %parsed_url,
|
||||||
|
content_length = len,
|
||||||
|
max = MAX_RESPONSE_SIZE,
|
||||||
|
"Rejected HTTP response: Content-Length exceeds limit"
|
||||||
|
);
|
||||||
|
return Err(ToolError::ExecutionFailed(format!(
|
||||||
|
"Response Content-Length ({} bytes) exceeds maximum allowed size ({} bytes)",
|
||||||
|
len, MAX_RESPONSE_SIZE
|
||||||
|
)));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Stream the response body with a hard size cap. Even if Content-Length was
|
||||||
|
// absent or lied about the size, we stop reading once we exceed the limit.
|
||||||
|
let mut body = Vec::new();
|
||||||
|
let mut stream = response.bytes_stream();
|
||||||
|
while let Some(chunk) = StreamExt::next(&mut stream).await {
|
||||||
|
let chunk = chunk.map_err(|e| {
|
||||||
ToolError::ExternalService(format!("failed to read response body: {}", e))
|
ToolError::ExternalService(format!("failed to read response body: {}", e))
|
||||||
})?;
|
})?;
|
||||||
|
if body.len() + chunk.len() > MAX_RESPONSE_SIZE {
|
||||||
if body_bytes.len() > MAX_RESPONSE_SIZE {
|
|
||||||
return Err(ToolError::ExecutionFailed(format!(
|
return Err(ToolError::ExecutionFailed(format!(
|
||||||
"Response body too large ({} bytes, max {})",
|
"Response body exceeds maximum allowed size ({} bytes)",
|
||||||
body_bytes.len(),
|
|
||||||
MAX_RESPONSE_SIZE
|
MAX_RESPONSE_SIZE
|
||||||
)));
|
)));
|
||||||
}
|
}
|
||||||
|
body.extend_from_slice(&chunk);
|
||||||
|
}
|
||||||
|
let body_bytes = bytes::Bytes::from(body);
|
||||||
|
|
||||||
let body_text = String::from_utf8_lossy(&body_bytes).into_owned();
|
let body_text = String::from_utf8_lossy(&body_bytes).into_owned();
|
||||||
|
|
||||||
@@ -328,4 +357,10 @@ mod tests {
|
|||||||
// Public
|
// Public
|
||||||
assert!(!is_disallowed_ip(&IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8))));
|
assert!(!is_disallowed_ip(&IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8))));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_max_response_size_is_reasonable() {
|
||||||
|
// MAX_RESPONSE_SIZE should be 5 MB to prevent OOM while allowing typical API responses.
|
||||||
|
assert_eq!(MAX_RESPONSE_SIZE, 5 * 1024 * 1024);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,6 +1,5 @@
|
|||||||
//! Built-in tools that come with the agent.
|
//! Built-in tools that come with the agent.
|
||||||
|
|
||||||
mod browser;
|
|
||||||
mod echo;
|
mod echo;
|
||||||
pub mod extension_tools;
|
pub mod extension_tools;
|
||||||
mod file;
|
mod file;
|
||||||
@@ -12,8 +11,6 @@ pub mod routine;
|
|||||||
pub(crate) mod shell;
|
pub(crate) mod shell;
|
||||||
mod time;
|
mod time;
|
||||||
|
|
||||||
pub use browser::BrowserTool;
|
|
||||||
pub use browser::session::find_chrome;
|
|
||||||
pub use echo::EchoTool;
|
pub use echo::EchoTool;
|
||||||
pub use extension_tools::{
|
pub use extension_tools::{
|
||||||
ToolActivateTool, ToolAuthTool, ToolInstallTool, ToolListTool, ToolRemoveTool, ToolSearchTool,
|
ToolActivateTool, ToolAuthTool, ToolInstallTool, ToolListTool, ToolRemoveTool, ToolSearchTool,
|
||||||
|
|||||||
@@ -426,6 +426,26 @@ impl Tool for ShellTool {
|
|||||||
true // Shell commands should require approval
|
true // Shell commands should require approval
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn requires_approval_for(&self, params: &serde_json::Value) -> bool {
|
||||||
|
let cmd = params
|
||||||
|
.get("command")
|
||||||
|
.and_then(|c| c.as_str().map(String::from))
|
||||||
|
.or_else(|| {
|
||||||
|
params
|
||||||
|
.as_str()
|
||||||
|
.and_then(|s| serde_json::from_str::<serde_json::Value>(s).ok())
|
||||||
|
.and_then(|v| v.get("command").and_then(|c| c.as_str().map(String::from)))
|
||||||
|
});
|
||||||
|
|
||||||
|
if let Some(ref cmd) = cmd
|
||||||
|
&& requires_explicit_approval(cmd)
|
||||||
|
{
|
||||||
|
return true;
|
||||||
|
}
|
||||||
|
|
||||||
|
false
|
||||||
|
}
|
||||||
|
|
||||||
fn requires_sanitization(&self) -> bool {
|
fn requires_sanitization(&self) -> bool {
|
||||||
true // Shell output could contain anything
|
true // Shell output could contain anything
|
||||||
}
|
}
|
||||||
@@ -566,6 +586,34 @@ mod tests {
|
|||||||
assert!(requires_explicit_approval(cmd.as_deref().unwrap()));
|
assert!(requires_explicit_approval(cmd.as_deref().unwrap()));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_requires_approval_for_destructive_command() {
|
||||||
|
let tool = ShellTool::new();
|
||||||
|
// Destructive commands must return true even though shell already
|
||||||
|
// requires base approval -- the distinction matters for auto-approve override.
|
||||||
|
assert!(tool.requires_approval_for(&serde_json::json!({"command": "rm -rf /tmp"})));
|
||||||
|
assert!(tool.requires_approval_for(
|
||||||
|
&serde_json::json!({"command": "git push --force origin main"})
|
||||||
|
));
|
||||||
|
assert!(tool.requires_approval_for(&serde_json::json!({"command": "DROP TABLE users;"})));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_requires_approval_for_safe_command() {
|
||||||
|
let tool = ShellTool::new();
|
||||||
|
// Safe commands should not override auto-approval; only destructive ones do.
|
||||||
|
assert!(!tool.requires_approval_for(&serde_json::json!({"command": "cargo build"})));
|
||||||
|
assert!(!tool.requires_approval_for(&serde_json::json!({"command": "echo hello"})));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_requires_approval_for_string_encoded_args() {
|
||||||
|
let tool = ShellTool::new();
|
||||||
|
// When arguments are string-encoded JSON (rare LLM behavior).
|
||||||
|
let args = serde_json::Value::String(r#"{"command": "rm -rf /tmp/stuff"}"#.to_string());
|
||||||
|
assert!(tool.requires_approval_for(&args));
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_sandbox_policy_builder() {
|
fn test_sandbox_policy_builder() {
|
||||||
let tool = ShellTool::new()
|
let tool = ShellTool::new()
|
||||||
|
|||||||
@@ -14,10 +14,10 @@ use crate::safety::SafetyLayer;
|
|||||||
use crate::secrets::SecretsStore;
|
use crate::secrets::SecretsStore;
|
||||||
use crate::tools::builder::{BuildSoftwareTool, BuilderConfig, LlmSoftwareBuilder};
|
use crate::tools::builder::{BuildSoftwareTool, BuilderConfig, LlmSoftwareBuilder};
|
||||||
use crate::tools::builtin::{
|
use crate::tools::builtin::{
|
||||||
ApplyPatchTool, BrowserTool, CancelJobTool, CreateJobTool, EchoTool, HttpTool, JobStatusTool,
|
ApplyPatchTool, CancelJobTool, CreateJobTool, EchoTool, HttpTool, JobStatusTool, JsonTool,
|
||||||
JsonTool, ListDirTool, ListJobsTool, MemoryReadTool, MemorySearchTool, MemoryTreeTool,
|
ListDirTool, ListJobsTool, MemoryReadTool, MemorySearchTool, MemoryTreeTool, MemoryWriteTool,
|
||||||
MemoryWriteTool, ReadFileTool, ShellTool, TimeTool, ToolActivateTool, ToolAuthTool,
|
ReadFileTool, ShellTool, TimeTool, ToolActivateTool, ToolAuthTool, ToolInstallTool,
|
||||||
ToolInstallTool, ToolListTool, ToolRemoveTool, ToolSearchTool, WriteFileTool,
|
ToolListTool, ToolRemoveTool, ToolSearchTool, WriteFileTool,
|
||||||
};
|
};
|
||||||
use crate::tools::tool::{Tool, ToolDomain};
|
use crate::tools::tool::{Tool, ToolDomain};
|
||||||
use crate::tools::wasm::{
|
use crate::tools::wasm::{
|
||||||
@@ -218,9 +218,8 @@ impl ToolRegistry {
|
|||||||
self.register_sync(Arc::new(WriteFileTool::new()));
|
self.register_sync(Arc::new(WriteFileTool::new()));
|
||||||
self.register_sync(Arc::new(ListDirTool::new()));
|
self.register_sync(Arc::new(ListDirTool::new()));
|
||||||
self.register_sync(Arc::new(ApplyPatchTool::new()));
|
self.register_sync(Arc::new(ApplyPatchTool::new()));
|
||||||
self.register_sync(Arc::new(BrowserTool::new()));
|
|
||||||
|
|
||||||
tracing::info!("Registered 6 development tools (includes browser)");
|
tracing::info!("Registered 5 development tools");
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Register memory tools with a workspace.
|
/// Register memory tools with a workspace.
|
||||||
|
|||||||
@@ -172,6 +172,21 @@ pub trait Tool: Send + Sync {
|
|||||||
false
|
false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Whether this specific invocation should override auto-approval.
|
||||||
|
///
|
||||||
|
/// This method is called after checking `requires_approval()` and finding that
|
||||||
|
/// the tool is auto-approved for this session. Return `true` to force approval
|
||||||
|
/// for this specific invocation despite auto-approval (for example, for
|
||||||
|
/// destructive operations like `rm -rf` or `git push --force`).
|
||||||
|
///
|
||||||
|
/// Return `false` to allow auto-approval to proceed normally.
|
||||||
|
///
|
||||||
|
/// The default returns `false`. Override only if you need parameter-aware
|
||||||
|
/// approval gating.
|
||||||
|
fn requires_approval_for(&self, _params: &serde_json::Value) -> bool {
|
||||||
|
false
|
||||||
|
}
|
||||||
|
|
||||||
/// Maximum time this tool is allowed to run before the caller kills it.
|
/// Maximum time this tool is allowed to run before the caller kills it.
|
||||||
/// Override for long-running tools like sandbox execution.
|
/// Override for long-running tools like sandbox execution.
|
||||||
/// Default: 60 seconds.
|
/// Default: 60 seconds.
|
||||||
@@ -330,4 +345,11 @@ mod tests {
|
|||||||
let err = require_param(¶ms, "data").unwrap_err();
|
let err = require_param(¶ms, "data").unwrap_err();
|
||||||
assert!(err.to_string().contains("missing 'data'"));
|
assert!(err.to_string().contains("missing 'data'"));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_requires_approval_for_default() {
|
||||||
|
let tool = EchoTool;
|
||||||
|
// Default requires_approval_for() returns false, allowing auto-approval.
|
||||||
|
assert!(!tool.requires_approval_for(&serde_json::json!({"message": "hi"})));
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,203 +0,0 @@
|
|||||||
//! Integration test for the browser tool.
|
|
||||||
//!
|
|
||||||
//! Requires Chrome installed. Run with:
|
|
||||||
//! cargo test --test browser_integration -- --nocapture
|
|
||||||
|
|
||||||
use ironclaw::context::JobContext;
|
|
||||||
use ironclaw::tools::Tool;
|
|
||||||
use ironclaw::tools::builtin::{BrowserTool, find_chrome};
|
|
||||||
|
|
||||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
|
||||||
async fn test_browser_navigate_and_screenshot() {
|
|
||||||
// Skip if Chrome/Chromium is not installed (works on macOS, Linux, Windows).
|
|
||||||
if find_chrome().is_none() {
|
|
||||||
eprintln!("Skipping: Chrome not found");
|
|
||||||
return;
|
|
||||||
}
|
|
||||||
|
|
||||||
let tool = BrowserTool::new();
|
|
||||||
let ctx = JobContext::default();
|
|
||||||
|
|
||||||
// 1. Navigate to Wikipedia
|
|
||||||
eprintln!("=== Navigating to Wikipedia...");
|
|
||||||
let nav_result = tool
|
|
||||||
.execute(
|
|
||||||
serde_json::json!({
|
|
||||||
"action": "navigate",
|
|
||||||
"url": "https://en.wikipedia.org/wiki/Mariam_Almheiri"
|
|
||||||
}),
|
|
||||||
&ctx,
|
|
||||||
)
|
|
||||||
.await;
|
|
||||||
|
|
||||||
match &nav_result {
|
|
||||||
Ok(output) => {
|
|
||||||
eprintln!(
|
|
||||||
"Navigation result: {}",
|
|
||||||
serde_json::to_string_pretty(&output.result).unwrap()
|
|
||||||
);
|
|
||||||
let title = output
|
|
||||||
.result
|
|
||||||
.get("title")
|
|
||||||
.and_then(|t| t.as_str())
|
|
||||||
.unwrap_or("");
|
|
||||||
assert!(
|
|
||||||
title.contains("Mariam") || title.contains("Almheiri"),
|
|
||||||
"Page title should mention Mariam Almheiri, got: {}",
|
|
||||||
title
|
|
||||||
);
|
|
||||||
}
|
|
||||||
Err(e) => {
|
|
||||||
eprintln!("Navigation failed: {}", e);
|
|
||||||
panic!("Navigation should succeed");
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 2. Read the accessibility tree
|
|
||||||
eprintln!("\n=== Reading page accessibility tree...");
|
|
||||||
let read_result = tool
|
|
||||||
.execute(serde_json::json!({"action": "read_page"}), &ctx)
|
|
||||||
.await;
|
|
||||||
|
|
||||||
match &read_result {
|
|
||||||
Ok(output) => {
|
|
||||||
let tree = output.result.as_str().unwrap_or("");
|
|
||||||
let line_count = tree.lines().count();
|
|
||||||
eprintln!("Accessibility tree: {} lines", line_count);
|
|
||||||
// Print first 20 lines
|
|
||||||
for line in tree.lines().take(20) {
|
|
||||||
eprintln!(" {}", line);
|
|
||||||
}
|
|
||||||
if line_count > 20 {
|
|
||||||
eprintln!(" ... ({} more lines)", line_count - 20);
|
|
||||||
}
|
|
||||||
assert!(line_count > 3, "Should have some elements on the page");
|
|
||||||
}
|
|
||||||
Err(e) => {
|
|
||||||
eprintln!("Read page failed: {}", e);
|
|
||||||
panic!("Read page should succeed");
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 3. Get page dimensions via eval_js to compute center
|
|
||||||
eprintln!("\n=== Getting page dimensions...");
|
|
||||||
let dims_result = tool
|
|
||||||
.execute(
|
|
||||||
serde_json::json!({
|
|
||||||
"action": "eval_js",
|
|
||||||
"expression": "JSON.stringify({w: window.innerWidth, h: window.innerHeight, scrollH: document.body.scrollHeight})"
|
|
||||||
}),
|
|
||||||
&ctx,
|
|
||||||
)
|
|
||||||
.await;
|
|
||||||
|
|
||||||
let (viewport_w, viewport_h) = match &dims_result {
|
|
||||||
Ok(output) => {
|
|
||||||
let result_str = output
|
|
||||||
.result
|
|
||||||
.get("result")
|
|
||||||
.and_then(|r| r.as_str())
|
|
||||||
.unwrap_or("{}");
|
|
||||||
let dims: serde_json::Value = serde_json::from_str(result_str).unwrap_or_default();
|
|
||||||
let w = dims.get("w").and_then(|v| v.as_f64()).unwrap_or(1920.0);
|
|
||||||
let h = dims.get("h").and_then(|v| v.as_f64()).unwrap_or(1080.0);
|
|
||||||
eprintln!("Viewport: {}x{}", w, h);
|
|
||||||
(w, h)
|
|
||||||
}
|
|
||||||
Err(e) => {
|
|
||||||
eprintln!("eval_js failed: {}", e);
|
|
||||||
(1920.0, 1080.0)
|
|
||||||
}
|
|
||||||
};
|
|
||||||
|
|
||||||
// 4. Scroll to middle of page first
|
|
||||||
eprintln!("\n=== Scrolling to middle of page...");
|
|
||||||
let _ = tool
|
|
||||||
.execute(
|
|
||||||
serde_json::json!({
|
|
||||||
"action": "eval_js",
|
|
||||||
"expression": "window.scrollTo(0, document.body.scrollHeight / 2 - window.innerHeight / 2)"
|
|
||||||
}),
|
|
||||||
&ctx,
|
|
||||||
)
|
|
||||||
.await;
|
|
||||||
|
|
||||||
// Brief wait for scroll to settle
|
|
||||||
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
|
|
||||||
|
|
||||||
// 5. Take full viewport screenshot
|
|
||||||
eprintln!("\n=== Taking viewport screenshot...");
|
|
||||||
let screenshot_result = tool
|
|
||||||
.execute(serde_json::json!({"action": "screenshot"}), &ctx)
|
|
||||||
.await;
|
|
||||||
|
|
||||||
match &screenshot_result {
|
|
||||||
Ok(output) => {
|
|
||||||
let b64 = output
|
|
||||||
.result
|
|
||||||
.get("data")
|
|
||||||
.and_then(|d| d.as_str())
|
|
||||||
.unwrap_or("");
|
|
||||||
eprintln!(
|
|
||||||
"Screenshot: {} base64 chars ({} bytes decoded)",
|
|
||||||
b64.len(),
|
|
||||||
b64.len() * 3 / 4
|
|
||||||
);
|
|
||||||
|
|
||||||
// Save to /tmp for inspection
|
|
||||||
use base64::Engine;
|
|
||||||
if let Ok(bytes) = base64::engine::general_purpose::STANDARD.decode(b64) {
|
|
||||||
let path = "/tmp/ironclaw_browser_test_viewport.png";
|
|
||||||
if std::fs::write(path, &bytes).is_ok() {
|
|
||||||
eprintln!("Saved viewport screenshot to {}", path);
|
|
||||||
}
|
|
||||||
|
|
||||||
// Now crop the center 10x10 using raw PNG manipulation
|
|
||||||
// We'll use eval_js to take a clipped screenshot via CDP directly
|
|
||||||
}
|
|
||||||
}
|
|
||||||
Err(e) => {
|
|
||||||
eprintln!("Screenshot failed: {}", e);
|
|
||||||
panic!("Screenshot should succeed");
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 6. Take a 10x10 screenshot from the center of the viewport using eval_js
|
|
||||||
// We can't directly use the clip param through the current tool API,
|
|
||||||
// so we'll take the viewport screenshot and note the center crop coords.
|
|
||||||
let center_x = (viewport_w / 2.0 - 5.0).max(0.0);
|
|
||||||
let center_y = (viewport_h / 2.0 - 5.0).max(0.0);
|
|
||||||
eprintln!(
|
|
||||||
"\n=== Center 10x10 crop would be at ({}, {}) to ({}, {})",
|
|
||||||
center_x,
|
|
||||||
center_y,
|
|
||||||
center_x + 10.0,
|
|
||||||
center_y + 10.0
|
|
||||||
);
|
|
||||||
|
|
||||||
// 7. Extract some text to verify content loaded
|
|
||||||
eprintln!("\n=== Extracting page text...");
|
|
||||||
let extract_result = tool
|
|
||||||
.execute(
|
|
||||||
serde_json::json!({"action": "extract", "selector": "h1"}),
|
|
||||||
&ctx,
|
|
||||||
)
|
|
||||||
.await;
|
|
||||||
|
|
||||||
match &extract_result {
|
|
||||||
Ok(output) => {
|
|
||||||
let text = output.result.as_str().unwrap_or("");
|
|
||||||
eprintln!("H1 text: {}", text);
|
|
||||||
assert!(
|
|
||||||
text.contains("Mariam") || text.contains("Almheiri"),
|
|
||||||
"H1 should contain the article subject, got: {}",
|
|
||||||
text
|
|
||||||
);
|
|
||||||
}
|
|
||||||
Err(e) => {
|
|
||||||
eprintln!("Extract failed: {}", e);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
eprintln!("\n=== All browser integration tests passed!");
|
|
||||||
}
|
|
||||||
@@ -0,0 +1,22 @@
|
|||||||
|
[package]
|
||||||
|
name = "github-tool"
|
||||||
|
version = "0.1.0"
|
||||||
|
edition = "2021"
|
||||||
|
description = "GitHub integration tool for IronClaw (WASM component)"
|
||||||
|
license = "MIT OR Apache-2.0"
|
||||||
|
publish = false
|
||||||
|
|
||||||
|
[dependencies]
|
||||||
|
serde = { version = "1.0", features = ["derive"] }
|
||||||
|
serde_json = "1.0"
|
||||||
|
wit-bindgen = "0.41.0"
|
||||||
|
|
||||||
|
[lib]
|
||||||
|
crate-type = ["cdylib"]
|
||||||
|
|
||||||
|
[profile.release]
|
||||||
|
opt-level = "s"
|
||||||
|
lto = true
|
||||||
|
strip = true
|
||||||
|
codegen-units = 1
|
||||||
|
|
||||||
@@ -0,0 +1,189 @@
|
|||||||
|
# GitHub Tool for IronClaw
|
||||||
|
|
||||||
|
WASM tool for GitHub integration - manage repos, issues, PRs, and workflows.
|
||||||
|
|
||||||
|
## Features
|
||||||
|
|
||||||
|
- **Repository Info** - Get repo details, list user repos
|
||||||
|
- **Issues** - List, create, and get issue details
|
||||||
|
- **Pull Requests** - List PRs, get PR details, review files, create reviews
|
||||||
|
- **File Content** - Read files from repos
|
||||||
|
- **Workflows** - Trigger GitHub Actions, check run status
|
||||||
|
|
||||||
|
## Setup
|
||||||
|
|
||||||
|
1. Create a GitHub Personal Access Token at <https://github.com/settings/tokens>
|
||||||
|
2. Required scopes: `repo`, `workflow`, `read:org`
|
||||||
|
3. Store the token:
|
||||||
|
|
||||||
|
```
|
||||||
|
ironclaw secret set github_token YOUR_TOKEN
|
||||||
|
```
|
||||||
|
|
||||||
|
## Usage Examples
|
||||||
|
|
||||||
|
### Get Repository Info
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"action": "get_repo",
|
||||||
|
"owner": "nearai",
|
||||||
|
"repo": "ironclaw"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### List Open Issues
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"action": "list_issues",
|
||||||
|
"owner": "nearai",
|
||||||
|
"repo": "ironclaw",
|
||||||
|
"state": "open",
|
||||||
|
"limit": 10
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Create Issue
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"action": "create_issue",
|
||||||
|
"owner": "nearai",
|
||||||
|
"repo": "ironclaw",
|
||||||
|
"title": "Bug: Something is broken",
|
||||||
|
"body": "Detailed description...",
|
||||||
|
"labels": ["bug", "help wanted"]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### List Pull Requests
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"action": "list_pull_requests",
|
||||||
|
"owner": "nearai",
|
||||||
|
"repo": "ironclaw",
|
||||||
|
"state": "open",
|
||||||
|
"limit": 5
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Review PR
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"action": "create_pr_review",
|
||||||
|
"owner": "nearai",
|
||||||
|
"repo": "ironclaw",
|
||||||
|
"pr_number": 42,
|
||||||
|
"body": "LGTM! Great work.",
|
||||||
|
"event": "APPROVE"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Get File Content
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"action": "get_file_content",
|
||||||
|
"owner": "nearai",
|
||||||
|
"repo": "ironclaw",
|
||||||
|
"path": "README.md",
|
||||||
|
"ref": "main"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Trigger Workflow
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"action": "trigger_workflow",
|
||||||
|
"owner": "nearai",
|
||||||
|
"repo": "ironclaw",
|
||||||
|
"workflow_id": "ci.yml",
|
||||||
|
"ref": "main",
|
||||||
|
"inputs": {
|
||||||
|
"environment": "staging"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Check Workflow Runs
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"action": "get_workflow_runs",
|
||||||
|
"owner": "nearai",
|
||||||
|
"repo": "ironclaw",
|
||||||
|
"limit": 5
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### List Workflow Runs (Pagination)
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"action": "get_workflow_runs",
|
||||||
|
"owner": "nearai",
|
||||||
|
"repo": "ironclaw",
|
||||||
|
"limit": 5,
|
||||||
|
"page": 2
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Error Handling
|
||||||
|
|
||||||
|
Errors are returned as strings in the `error` field of the response.
|
||||||
|
|
||||||
|
### Rate Limit Exceeded
|
||||||
|
|
||||||
|
When the GitHub API rate limit is exceeded (and retries fail), you might see:
|
||||||
|
|
||||||
|
```text
|
||||||
|
GitHub API error 429: { "message": "API rate limit exceeded for user ID ...", ... }
|
||||||
|
```
|
||||||
|
|
||||||
|
The tool automatically logs warnings when the rate limit is low (<10 remaining) and retries on 429/5xx errors.
|
||||||
|
|
||||||
|
### Invalid Parameters
|
||||||
|
|
||||||
|
```text
|
||||||
|
Invalid event: 'INVALID'. Must be one of: APPROVE, REQUEST_CHANGES, COMMENT
|
||||||
|
```
|
||||||
|
|
||||||
|
### Missing Token
|
||||||
|
|
||||||
|
```text
|
||||||
|
GitHub token not found in secret store. Set it with: ironclaw secret set github_token <token>...
|
||||||
|
```
|
||||||
|
|
||||||
|
## Troubleshooting
|
||||||
|
|
||||||
|
### "GitHub API error 404: Not Found"
|
||||||
|
|
||||||
|
- Check that the `owner` and `repo` are correct.
|
||||||
|
- Ensure the `github_token` has access to the repository (especially for private repos).
|
||||||
|
- Verify the token scopes include `repo` and `read:org`.
|
||||||
|
|
||||||
|
### "GitHub API error 401: Bad credentials"
|
||||||
|
|
||||||
|
- The token might be invalid or expired.
|
||||||
|
- Update the token: `ironclaw secret set github_token NEW_TOKEN`.
|
||||||
|
|
||||||
|
### Rate Limiting
|
||||||
|
|
||||||
|
- The tool logs a warning when remaining requests drop below 10.
|
||||||
|
- Check logs for "GitHub API rate limit low".
|
||||||
|
- If you hit the limit, wait for the reset time (usually 1 hour).
|
||||||
|
|
||||||
|
## Building
|
||||||
|
|
||||||
|
```bash
|
||||||
|
cd tools-src/github
|
||||||
|
cargo build --target wasm32-wasi --release
|
||||||
|
```
|
||||||
|
|
||||||
|
## License
|
||||||
|
|
||||||
|
MIT/Apache-2.0
|
||||||
@@ -0,0 +1,41 @@
|
|||||||
|
{
|
||||||
|
"capabilities": {
|
||||||
|
"http": {
|
||||||
|
"allowlist": [
|
||||||
|
{
|
||||||
|
"host": "api.github.com",
|
||||||
|
"path_prefix": "/",
|
||||||
|
"methods": [
|
||||||
|
"GET",
|
||||||
|
"POST"
|
||||||
|
]
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"credentials": {
|
||||||
|
"github_token": {
|
||||||
|
"secret_name": "github_token",
|
||||||
|
"location": {
|
||||||
|
"type": "bearer"
|
||||||
|
},
|
||||||
|
"host_patterns": [
|
||||||
|
"api.github.com"
|
||||||
|
]
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"rate_limit": {
|
||||||
|
"requests_per_minute": 60,
|
||||||
|
"requests_per_hour": 3600
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"secrets": {
|
||||||
|
"allowed_names": [
|
||||||
|
"github_token",
|
||||||
|
"github_*"
|
||||||
|
]
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"config": {
|
||||||
|
"default_limit": 30,
|
||||||
|
"max_limit": 100
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,845 @@
|
|||||||
|
//! GitHub WASM Tool for IronClaw.
|
||||||
|
//!
|
||||||
|
//! Provides GitHub integration for reading repos, managing issues,
|
||||||
|
//! reviewing PRs, and triggering workflows.
|
||||||
|
//!
|
||||||
|
//! # Authentication
|
||||||
|
//!
|
||||||
|
//! Store your GitHub Personal Access Token:
|
||||||
|
//! `ironclaw secret set github_token <token>`
|
||||||
|
//!
|
||||||
|
//! Token needs these permissions:
|
||||||
|
//! - repo (for private repos)
|
||||||
|
//! - workflow (for triggering actions)
|
||||||
|
//! - read:org (for org repos)
|
||||||
|
|
||||||
|
wit_bindgen::generate!({
|
||||||
|
world: "sandboxed-tool",
|
||||||
|
path: "../../wit/tool.wit",
|
||||||
|
});
|
||||||
|
|
||||||
|
use serde::Deserialize;
|
||||||
|
|
||||||
|
const MAX_TEXT_LENGTH: usize = 65536;
|
||||||
|
|
||||||
|
/// Validate input length to prevent oversized payloads.
|
||||||
|
fn validate_input_length(s: &str, field_name: &str) -> Result<(), String> {
|
||||||
|
if s.len() > MAX_TEXT_LENGTH {
|
||||||
|
return Err(format!(
|
||||||
|
"Input '{}' exceeds maximum length of {} characters",
|
||||||
|
field_name, MAX_TEXT_LENGTH
|
||||||
|
));
|
||||||
|
}
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Percent-encode a string for safe use in URL path segments.
|
||||||
|
/// Encodes everything except alphanumeric, hyphen, underscore, and dot.
|
||||||
|
fn url_encode_path(s: &str) -> String {
|
||||||
|
let mut out = String::with_capacity(s.len() * 2);
|
||||||
|
for b in s.bytes() {
|
||||||
|
match b {
|
||||||
|
b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'_' | b'.' => {
|
||||||
|
out.push(b as char);
|
||||||
|
}
|
||||||
|
_ => {
|
||||||
|
out.push('%');
|
||||||
|
out.push(char::from(b"0123456789ABCDEF"[(b >> 4) as usize]));
|
||||||
|
out.push(char::from(b"0123456789ABCDEF"[(b & 0xf) as usize]));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
out
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Percent-encode a string for use as a URL query parameter value.
|
||||||
|
/// Currently identical to `url_encode_path`.
|
||||||
|
fn url_encode_query(s: &str) -> String {
|
||||||
|
url_encode_path(s)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Validate that a path segment doesn't contain dangerous characters.
|
||||||
|
/// Returns true if the segment is safe to use.
|
||||||
|
fn validate_path_segment(s: &str) -> bool {
|
||||||
|
!s.is_empty() && !s.contains('/') && !s.contains("..") && !s.contains('?') && !s.contains('#')
|
||||||
|
}
|
||||||
|
|
||||||
|
struct GitHubTool;
|
||||||
|
|
||||||
|
#[derive(Debug, Deserialize)]
|
||||||
|
#[serde(tag = "action")]
|
||||||
|
enum GitHubAction {
|
||||||
|
#[serde(rename = "get_repo")]
|
||||||
|
GetRepo { owner: String, repo: String },
|
||||||
|
#[serde(rename = "list_issues")]
|
||||||
|
ListIssues {
|
||||||
|
owner: String,
|
||||||
|
repo: String,
|
||||||
|
state: Option<String>,
|
||||||
|
page: Option<u32>,
|
||||||
|
limit: Option<u32>,
|
||||||
|
},
|
||||||
|
#[serde(rename = "create_issue")]
|
||||||
|
CreateIssue {
|
||||||
|
owner: String,
|
||||||
|
repo: String,
|
||||||
|
title: String,
|
||||||
|
body: Option<String>,
|
||||||
|
labels: Option<Vec<String>>,
|
||||||
|
},
|
||||||
|
#[serde(rename = "get_issue")]
|
||||||
|
GetIssue {
|
||||||
|
owner: String,
|
||||||
|
repo: String,
|
||||||
|
issue_number: u32,
|
||||||
|
},
|
||||||
|
#[serde(rename = "list_pull_requests")]
|
||||||
|
ListPullRequests {
|
||||||
|
owner: String,
|
||||||
|
repo: String,
|
||||||
|
state: Option<String>,
|
||||||
|
page: Option<u32>,
|
||||||
|
limit: Option<u32>,
|
||||||
|
},
|
||||||
|
#[serde(rename = "get_pull_request")]
|
||||||
|
GetPullRequest {
|
||||||
|
owner: String,
|
||||||
|
repo: String,
|
||||||
|
pr_number: u32,
|
||||||
|
},
|
||||||
|
#[serde(rename = "get_pull_request_files")]
|
||||||
|
GetPullRequestFiles {
|
||||||
|
owner: String,
|
||||||
|
repo: String,
|
||||||
|
pr_number: u32,
|
||||||
|
},
|
||||||
|
#[serde(rename = "create_pr_review")]
|
||||||
|
CreatePrReview {
|
||||||
|
owner: String,
|
||||||
|
repo: String,
|
||||||
|
pr_number: u32,
|
||||||
|
body: String,
|
||||||
|
event: String,
|
||||||
|
},
|
||||||
|
#[serde(rename = "list_repos")]
|
||||||
|
ListRepos {
|
||||||
|
username: String,
|
||||||
|
page: Option<u32>,
|
||||||
|
limit: Option<u32>,
|
||||||
|
},
|
||||||
|
#[serde(rename = "get_file_content")]
|
||||||
|
GetFileContent {
|
||||||
|
owner: String,
|
||||||
|
repo: String,
|
||||||
|
path: String,
|
||||||
|
r#ref: Option<String>,
|
||||||
|
},
|
||||||
|
#[serde(rename = "trigger_workflow")]
|
||||||
|
TriggerWorkflow {
|
||||||
|
owner: String,
|
||||||
|
repo: String,
|
||||||
|
workflow_id: String,
|
||||||
|
r#ref: String,
|
||||||
|
inputs: Option<serde_json::Value>,
|
||||||
|
},
|
||||||
|
#[serde(rename = "get_workflow_runs")]
|
||||||
|
GetWorkflowRuns {
|
||||||
|
owner: String,
|
||||||
|
repo: String,
|
||||||
|
workflow_id: Option<String>,
|
||||||
|
page: Option<u32>,
|
||||||
|
limit: Option<u32>,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
impl exports::near::agent::tool::Guest for GitHubTool {
|
||||||
|
fn execute(req: exports::near::agent::tool::Request) -> exports::near::agent::tool::Response {
|
||||||
|
match execute_inner(&req.params) {
|
||||||
|
Ok(result) => exports::near::agent::tool::Response {
|
||||||
|
output: Some(result),
|
||||||
|
error: None,
|
||||||
|
},
|
||||||
|
Err(e) => exports::near::agent::tool::Response {
|
||||||
|
output: None,
|
||||||
|
error: Some(e),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn schema() -> String {
|
||||||
|
SCHEMA.to_string()
|
||||||
|
}
|
||||||
|
|
||||||
|
fn description() -> String {
|
||||||
|
"GitHub integration for managing repositories, issues, pull requests, \
|
||||||
|
and workflows. Supports reading repo info, listing/creating issues, \
|
||||||
|
reviewing PRs, and triggering GitHub Actions. \
|
||||||
|
Authentication is handled via the 'github_token' secret injected by the host."
|
||||||
|
.to_string()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn execute_inner(params: &str) -> Result<String, String> {
|
||||||
|
let action: GitHubAction =
|
||||||
|
serde_json::from_str(params).map_err(|e| format!("Invalid parameters: {e}"))?;
|
||||||
|
|
||||||
|
// Pre-flight check: ensure token exists in secret store.
|
||||||
|
// We don't use the returned value because the host injects it into the request.
|
||||||
|
let _ = get_github_token()?;
|
||||||
|
|
||||||
|
match action {
|
||||||
|
GitHubAction::GetRepo { owner, repo } => get_repo(&owner, &repo),
|
||||||
|
GitHubAction::ListIssues {
|
||||||
|
owner,
|
||||||
|
repo,
|
||||||
|
state,
|
||||||
|
page,
|
||||||
|
limit,
|
||||||
|
} => list_issues(&owner, &repo, state.as_deref(), page, limit),
|
||||||
|
GitHubAction::CreateIssue {
|
||||||
|
owner,
|
||||||
|
repo,
|
||||||
|
title,
|
||||||
|
body,
|
||||||
|
labels,
|
||||||
|
} => create_issue(&owner, &repo, &title, body.as_deref(), labels),
|
||||||
|
GitHubAction::GetIssue {
|
||||||
|
owner,
|
||||||
|
repo,
|
||||||
|
issue_number,
|
||||||
|
} => get_issue(&owner, &repo, issue_number),
|
||||||
|
GitHubAction::ListPullRequests {
|
||||||
|
owner,
|
||||||
|
repo,
|
||||||
|
state,
|
||||||
|
page,
|
||||||
|
limit,
|
||||||
|
} => list_pull_requests(&owner, &repo, state.as_deref(), page, limit),
|
||||||
|
GitHubAction::GetPullRequest {
|
||||||
|
owner,
|
||||||
|
repo,
|
||||||
|
pr_number,
|
||||||
|
} => get_pull_request(&owner, &repo, pr_number),
|
||||||
|
GitHubAction::GetPullRequestFiles {
|
||||||
|
owner,
|
||||||
|
repo,
|
||||||
|
pr_number,
|
||||||
|
} => get_pull_request_files(&owner, &repo, pr_number),
|
||||||
|
GitHubAction::CreatePrReview {
|
||||||
|
owner,
|
||||||
|
repo,
|
||||||
|
pr_number,
|
||||||
|
body,
|
||||||
|
event,
|
||||||
|
} => create_pr_review(&owner, &repo, pr_number, &body, &event),
|
||||||
|
GitHubAction::ListRepos {
|
||||||
|
username,
|
||||||
|
page,
|
||||||
|
limit,
|
||||||
|
} => list_repos(&username, page, limit),
|
||||||
|
GitHubAction::GetFileContent {
|
||||||
|
owner,
|
||||||
|
repo,
|
||||||
|
path,
|
||||||
|
r#ref,
|
||||||
|
} => get_file_content(&owner, &repo, &path, r#ref.as_deref()),
|
||||||
|
GitHubAction::TriggerWorkflow {
|
||||||
|
owner,
|
||||||
|
repo,
|
||||||
|
workflow_id,
|
||||||
|
r#ref,
|
||||||
|
inputs,
|
||||||
|
} => trigger_workflow(&owner, &repo, &workflow_id, &r#ref, inputs),
|
||||||
|
GitHubAction::GetWorkflowRuns {
|
||||||
|
owner,
|
||||||
|
repo,
|
||||||
|
workflow_id,
|
||||||
|
page,
|
||||||
|
limit,
|
||||||
|
} => get_workflow_runs(&owner, &repo, workflow_id.as_deref(), page, limit),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn get_github_token() -> Result<String, String> {
|
||||||
|
if near::agent::host::secret_exists("github_token") {
|
||||||
|
// Return dummy value since we only need to verify existence.
|
||||||
|
// The actual token is injected by the host.
|
||||||
|
return Ok("present".to_string());
|
||||||
|
}
|
||||||
|
|
||||||
|
Err("GitHub token not found in secret store. Set it with: ironclaw secret set github_token <token>. \
|
||||||
|
Token needs 'repo', 'workflow', and 'read:org' scopes.".into())
|
||||||
|
}
|
||||||
|
|
||||||
|
fn github_request(method: &str, path: &str, body: Option<String>) -> Result<String, String> {
|
||||||
|
let url = format!("https://api.github.com{}", path);
|
||||||
|
|
||||||
|
// Authorization header (Bearer <token>) is injected automatically by the host
|
||||||
|
// via the `http-wrapper` proxy based on the `github_token` secret.
|
||||||
|
let headers = serde_json::json!({
|
||||||
|
"Accept": "application/vnd.github+json",
|
||||||
|
"X-GitHub-Api-Version": "2022-11-28",
|
||||||
|
"User-Agent": "IronClaw-GitHub-Tool"
|
||||||
|
});
|
||||||
|
|
||||||
|
let body_bytes = body.map(|b| b.into_bytes());
|
||||||
|
|
||||||
|
// Simple retry logic for transient errors (max 3 attempts)
|
||||||
|
let max_retries = 3;
|
||||||
|
let mut attempt = 0;
|
||||||
|
|
||||||
|
loop {
|
||||||
|
attempt += 1;
|
||||||
|
|
||||||
|
let response = near::agent::host::http_request(
|
||||||
|
method,
|
||||||
|
&url,
|
||||||
|
&headers.to_string(),
|
||||||
|
body_bytes.as_deref(),
|
||||||
|
None,
|
||||||
|
);
|
||||||
|
|
||||||
|
match response {
|
||||||
|
Ok(resp) => {
|
||||||
|
// Log warning if rate limit is low
|
||||||
|
if let Ok(headers_json) =
|
||||||
|
serde_json::from_str::<serde_json::Value>(&resp.headers_json)
|
||||||
|
{
|
||||||
|
// Header keys are often lowercase in http libs, check case-insensitively if needed,
|
||||||
|
// but usually standard is lowercase/case-insensitive. Let's try lowercase.
|
||||||
|
if let Some(remaining) = headers_json
|
||||||
|
.get("x-ratelimit-remaining")
|
||||||
|
.and_then(|v| v.as_str())
|
||||||
|
{
|
||||||
|
if let Ok(count) = remaining.parse::<u32>() {
|
||||||
|
if count < 10 {
|
||||||
|
near::agent::host::log(
|
||||||
|
near::agent::host::LogLevel::Warn,
|
||||||
|
&format!("GitHub API rate limit low: {} remaining", count),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if resp.status >= 200 && resp.status < 300 {
|
||||||
|
return String::from_utf8(resp.body)
|
||||||
|
.map_err(|e| format!("Invalid UTF-8: {}", e));
|
||||||
|
} else if attempt < max_retries && (resp.status == 429 || resp.status >= 500) {
|
||||||
|
near::agent::host::log(
|
||||||
|
near::agent::host::LogLevel::Warn,
|
||||||
|
&format!(
|
||||||
|
"GitHub API error {} (attempt {}/{}). Retrying...",
|
||||||
|
resp.status, attempt, max_retries
|
||||||
|
),
|
||||||
|
);
|
||||||
|
// Minimal backoff simulation since we can't block easily in WASM without consuming generic budget?
|
||||||
|
// actually std::thread::sleep works in WASMtime if configured, but here we might just spin.
|
||||||
|
// ideally host exposes sleep. For now just retry immediately or rely on host timeout logic?
|
||||||
|
// Let's assume immediate retry for now as simple strategy.
|
||||||
|
continue;
|
||||||
|
} else {
|
||||||
|
let body_str = String::from_utf8_lossy(&resp.body);
|
||||||
|
return Err(format!("GitHub API error {}: {}", resp.status, body_str));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Err(e) => {
|
||||||
|
if attempt < max_retries {
|
||||||
|
near::agent::host::log(
|
||||||
|
near::agent::host::LogLevel::Warn,
|
||||||
|
&format!(
|
||||||
|
"HTTP request failed: {} (attempt {}/{}). Retrying...",
|
||||||
|
e, attempt, max_retries
|
||||||
|
),
|
||||||
|
);
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
return Err(format!(
|
||||||
|
"HTTP request failed after {} attempts: {}",
|
||||||
|
max_retries, e
|
||||||
|
));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// === API Functions ===
|
||||||
|
|
||||||
|
fn get_repo(owner: &str, repo: &str) -> Result<String, String> {
|
||||||
|
if !validate_path_segment(owner) || !validate_path_segment(repo) {
|
||||||
|
return Err("Invalid owner or repo name".into());
|
||||||
|
}
|
||||||
|
let encoded_owner = url_encode_path(owner);
|
||||||
|
let encoded_repo = url_encode_path(repo);
|
||||||
|
github_request(
|
||||||
|
"GET",
|
||||||
|
&format!("/repos/{}/{}", encoded_owner, encoded_repo),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn list_issues(
|
||||||
|
owner: &str,
|
||||||
|
repo: &str,
|
||||||
|
state: Option<&str>,
|
||||||
|
page: Option<u32>,
|
||||||
|
limit: Option<u32>,
|
||||||
|
) -> Result<String, String> {
|
||||||
|
if !validate_path_segment(owner) || !validate_path_segment(repo) {
|
||||||
|
return Err("Invalid owner or repo name".into());
|
||||||
|
}
|
||||||
|
let encoded_owner = url_encode_path(owner);
|
||||||
|
let encoded_repo = url_encode_path(repo);
|
||||||
|
let state = state.unwrap_or("open");
|
||||||
|
let limit = limit.unwrap_or(30).min(100); // Cap at 100
|
||||||
|
let encoded_state = url_encode_query(state);
|
||||||
|
|
||||||
|
let mut path = format!(
|
||||||
|
"/repos/{}/{}/issues?state={}&per_page={}",
|
||||||
|
encoded_owner, encoded_repo, encoded_state, limit
|
||||||
|
);
|
||||||
|
if let Some(p) = page {
|
||||||
|
path.push_str(&format!("&page={}", p));
|
||||||
|
}
|
||||||
|
|
||||||
|
github_request("GET", &path, None)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn create_issue(
|
||||||
|
owner: &str,
|
||||||
|
repo: &str,
|
||||||
|
title: &str,
|
||||||
|
body: Option<&str>,
|
||||||
|
labels: Option<Vec<String>>,
|
||||||
|
) -> Result<String, String> {
|
||||||
|
if !validate_path_segment(owner) || !validate_path_segment(repo) {
|
||||||
|
return Err("Invalid owner or repo name".into());
|
||||||
|
}
|
||||||
|
validate_input_length(title, "title")?;
|
||||||
|
if let Some(b) = body {
|
||||||
|
validate_input_length(b, "body")?;
|
||||||
|
}
|
||||||
|
|
||||||
|
let encoded_owner = url_encode_path(owner);
|
||||||
|
let encoded_repo = url_encode_path(repo);
|
||||||
|
let path = format!("/repos/{}/{}/issues", encoded_owner, encoded_repo);
|
||||||
|
let mut req_body = serde_json::json!({
|
||||||
|
"title": title,
|
||||||
|
});
|
||||||
|
if let Some(body) = body {
|
||||||
|
req_body["body"] = serde_json::json!(body);
|
||||||
|
}
|
||||||
|
if let Some(labels) = labels {
|
||||||
|
req_body["labels"] = serde_json::json!(labels);
|
||||||
|
}
|
||||||
|
github_request("POST", &path, Some(req_body.to_string()))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn get_issue(owner: &str, repo: &str, issue_number: u32) -> Result<String, String> {
|
||||||
|
if !validate_path_segment(owner) || !validate_path_segment(repo) {
|
||||||
|
return Err("Invalid owner or repo name".into());
|
||||||
|
}
|
||||||
|
let encoded_owner = url_encode_path(owner);
|
||||||
|
let encoded_repo = url_encode_path(repo);
|
||||||
|
github_request(
|
||||||
|
"GET",
|
||||||
|
&format!(
|
||||||
|
"/repos/{}/{}/issues/{}",
|
||||||
|
encoded_owner, encoded_repo, issue_number
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn list_pull_requests(
|
||||||
|
owner: &str,
|
||||||
|
repo: &str,
|
||||||
|
state: Option<&str>,
|
||||||
|
page: Option<u32>,
|
||||||
|
limit: Option<u32>,
|
||||||
|
) -> Result<String, String> {
|
||||||
|
if !validate_path_segment(owner) || !validate_path_segment(repo) {
|
||||||
|
return Err("Invalid owner or repo name".into());
|
||||||
|
}
|
||||||
|
let encoded_owner = url_encode_path(owner);
|
||||||
|
let encoded_repo = url_encode_path(repo);
|
||||||
|
let state = state.unwrap_or("open");
|
||||||
|
let limit = limit.unwrap_or(30).min(100); // Cap at 100
|
||||||
|
let encoded_state = url_encode_query(state);
|
||||||
|
|
||||||
|
let mut path = format!(
|
||||||
|
"/repos/{}/{}/pulls?state={}&per_page={}",
|
||||||
|
encoded_owner, encoded_repo, encoded_state, limit
|
||||||
|
);
|
||||||
|
if let Some(p) = page {
|
||||||
|
path.push_str(&format!("&page={}", p));
|
||||||
|
}
|
||||||
|
|
||||||
|
github_request("GET", &path, None)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn get_pull_request(owner: &str, repo: &str, pr_number: u32) -> Result<String, String> {
|
||||||
|
if !validate_path_segment(owner) || !validate_path_segment(repo) {
|
||||||
|
return Err("Invalid owner or repo name".into());
|
||||||
|
}
|
||||||
|
let encoded_owner = url_encode_path(owner);
|
||||||
|
let encoded_repo = url_encode_path(repo);
|
||||||
|
github_request(
|
||||||
|
"GET",
|
||||||
|
&format!(
|
||||||
|
"/repos/{}/{}/pulls/{}",
|
||||||
|
encoded_owner, encoded_repo, pr_number
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn get_pull_request_files(owner: &str, repo: &str, pr_number: u32) -> Result<String, String> {
|
||||||
|
if !validate_path_segment(owner) || !validate_path_segment(repo) {
|
||||||
|
return Err("Invalid owner or repo name".into());
|
||||||
|
}
|
||||||
|
let encoded_owner = url_encode_path(owner);
|
||||||
|
let encoded_repo = url_encode_path(repo);
|
||||||
|
github_request(
|
||||||
|
"GET",
|
||||||
|
&format!(
|
||||||
|
"/repos/{}/{}/pulls/{}/files",
|
||||||
|
encoded_owner, encoded_repo, pr_number
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn create_pr_review(
|
||||||
|
owner: &str,
|
||||||
|
repo: &str,
|
||||||
|
pr_number: u32,
|
||||||
|
body: &str,
|
||||||
|
event: &str,
|
||||||
|
) -> Result<String, String> {
|
||||||
|
if !validate_path_segment(owner) || !validate_path_segment(repo) {
|
||||||
|
return Err("Invalid owner or repo name".into());
|
||||||
|
}
|
||||||
|
validate_input_length(body, "body")?;
|
||||||
|
|
||||||
|
let valid_events = ["APPROVE", "REQUEST_CHANGES", "COMMENT"];
|
||||||
|
if !valid_events.contains(&event) {
|
||||||
|
return Err(format!(
|
||||||
|
"Invalid event: '{}'. Must be one of: {}",
|
||||||
|
event,
|
||||||
|
valid_events.join(", ")
|
||||||
|
));
|
||||||
|
}
|
||||||
|
let encoded_owner = url_encode_path(owner);
|
||||||
|
let encoded_repo = url_encode_path(repo);
|
||||||
|
let path = format!(
|
||||||
|
"/repos/{}/{}/pulls/{}/reviews",
|
||||||
|
encoded_owner, encoded_repo, pr_number
|
||||||
|
);
|
||||||
|
let req_body = serde_json::json!({
|
||||||
|
"body": body,
|
||||||
|
"event": event,
|
||||||
|
});
|
||||||
|
github_request("POST", &path, Some(req_body.to_string()))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn list_repos(username: &str, page: Option<u32>, limit: Option<u32>) -> Result<String, String> {
|
||||||
|
if !validate_path_segment(username) {
|
||||||
|
return Err("Invalid username".into());
|
||||||
|
}
|
||||||
|
let encoded_username = url_encode_path(username);
|
||||||
|
let limit = limit.unwrap_or(30).min(100); // Cap at 100
|
||||||
|
let mut path = format!("/users/{}/repos?per_page={}", encoded_username, limit);
|
||||||
|
if let Some(p) = page {
|
||||||
|
path.push_str(&format!("&page={}", p));
|
||||||
|
}
|
||||||
|
github_request("GET", &path, None)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn get_file_content(
|
||||||
|
owner: &str,
|
||||||
|
repo: &str,
|
||||||
|
path: &str,
|
||||||
|
r#ref: Option<&str>,
|
||||||
|
) -> Result<String, String> {
|
||||||
|
if !validate_path_segment(owner) || !validate_path_segment(repo) {
|
||||||
|
return Err("Invalid owner or repo name".into());
|
||||||
|
}
|
||||||
|
// Validate path segments - reject path traversal attempts and empty segments
|
||||||
|
for segment in path.split('/') {
|
||||||
|
if segment == ".." {
|
||||||
|
return Err("Invalid path: path traversal not allowed".into());
|
||||||
|
}
|
||||||
|
if segment.is_empty() {
|
||||||
|
return Err("Invalid path: empty segment not allowed".into());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Validate ref if provided
|
||||||
|
if let Some(r#ref) = r#ref {
|
||||||
|
if r#ref.contains("..") || r#ref.contains(':') {
|
||||||
|
return Err("Invalid ref: must be a valid branch, tag, or commit SHA".into());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
let encoded_owner = url_encode_path(owner);
|
||||||
|
let encoded_repo = url_encode_path(repo);
|
||||||
|
// Path can contain slashes, so we encode each segment separately
|
||||||
|
let encoded_path = path
|
||||||
|
.split('/')
|
||||||
|
.map(url_encode_path)
|
||||||
|
.collect::<Vec<_>>()
|
||||||
|
.join("/");
|
||||||
|
|
||||||
|
let url_path = if let Some(r#ref) = r#ref {
|
||||||
|
let encoded_ref = url_encode_query(r#ref);
|
||||||
|
format!(
|
||||||
|
"/repos/{}/{}/contents/{}?ref={}",
|
||||||
|
encoded_owner, encoded_repo, encoded_path, encoded_ref
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
format!(
|
||||||
|
"/repos/{}/{}/contents/{}",
|
||||||
|
encoded_owner, encoded_repo, encoded_path
|
||||||
|
)
|
||||||
|
};
|
||||||
|
github_request("GET", &url_path, None)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn trigger_workflow(
|
||||||
|
owner: &str,
|
||||||
|
repo: &str,
|
||||||
|
workflow_id: &str,
|
||||||
|
r#ref: &str,
|
||||||
|
inputs: Option<serde_json::Value>,
|
||||||
|
) -> Result<String, String> {
|
||||||
|
if !validate_path_segment(owner) || !validate_path_segment(repo) {
|
||||||
|
return Err("Invalid owner or repo name".into());
|
||||||
|
}
|
||||||
|
// Validate inputs size if present
|
||||||
|
if let Some(valid_inputs) = &inputs {
|
||||||
|
let inputs_str = valid_inputs.to_string();
|
||||||
|
validate_input_length(&inputs_str, "inputs")?;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Validate workflow_id - must be a safe filename
|
||||||
|
if workflow_id.contains('/') || workflow_id.contains("..") || workflow_id.contains(':') {
|
||||||
|
return Err("Invalid workflow_id: must be a filename or numeric ID".into());
|
||||||
|
}
|
||||||
|
// Validate ref - must be a valid git ref
|
||||||
|
if r#ref.contains("..") || r#ref.contains(':') {
|
||||||
|
return Err("Invalid ref: must be a valid branch, tag, or commit SHA".into());
|
||||||
|
}
|
||||||
|
let encoded_owner = url_encode_path(owner);
|
||||||
|
let encoded_repo = url_encode_path(repo);
|
||||||
|
let encoded_workflow_id = url_encode_path(workflow_id);
|
||||||
|
let path = format!(
|
||||||
|
"/repos/{}/{}/actions/workflows/{}/dispatches",
|
||||||
|
encoded_owner, encoded_repo, encoded_workflow_id
|
||||||
|
);
|
||||||
|
let mut req_body = serde_json::json!({
|
||||||
|
"ref": r#ref,
|
||||||
|
});
|
||||||
|
if let Some(inputs) = inputs {
|
||||||
|
req_body["inputs"] = inputs;
|
||||||
|
}
|
||||||
|
github_request("POST", &path, Some(req_body.to_string()))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn get_workflow_runs(
|
||||||
|
owner: &str,
|
||||||
|
repo: &str,
|
||||||
|
workflow_id: Option<&str>,
|
||||||
|
page: Option<u32>,
|
||||||
|
limit: Option<u32>,
|
||||||
|
) -> Result<String, String> {
|
||||||
|
if !validate_path_segment(owner) || !validate_path_segment(repo) {
|
||||||
|
return Err("Invalid owner or repo name".into());
|
||||||
|
}
|
||||||
|
// Validate workflow_id if provided
|
||||||
|
if let Some(wid) = workflow_id {
|
||||||
|
if wid.contains('/') || wid.contains("..") || wid.contains(':') {
|
||||||
|
return Err("Invalid workflow_id: must be a filename or numeric ID".into());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
let encoded_owner = url_encode_path(owner);
|
||||||
|
let encoded_repo = url_encode_path(repo);
|
||||||
|
let limit = limit.unwrap_or(30).min(100); // Cap at 100
|
||||||
|
let mut path = if let Some(workflow_id) = workflow_id {
|
||||||
|
let encoded_workflow_id = url_encode_path(workflow_id);
|
||||||
|
format!(
|
||||||
|
"/repos/{}/{}/actions/workflows/{}/runs?per_page={}",
|
||||||
|
encoded_owner, encoded_repo, encoded_workflow_id, limit
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
format!(
|
||||||
|
"/repos/{}/{}/actions/runs?per_page={}",
|
||||||
|
encoded_owner, encoded_repo, limit
|
||||||
|
)
|
||||||
|
};
|
||||||
|
if let Some(p) = page {
|
||||||
|
path.push_str(&format!("&page={}", p));
|
||||||
|
}
|
||||||
|
github_request("GET", &path, None)
|
||||||
|
}
|
||||||
|
|
||||||
|
const SCHEMA: &str = r#"{
|
||||||
|
"type": "object",
|
||||||
|
"required": ["action"],
|
||||||
|
"oneOf": [
|
||||||
|
{
|
||||||
|
"properties": {
|
||||||
|
"action": { "const": "get_repo" },
|
||||||
|
"owner": { "type": "string", "description": "Repository owner (user or org)" },
|
||||||
|
"repo": { "type": "string", "description": "Repository name" }
|
||||||
|
},
|
||||||
|
"required": ["action", "owner", "repo"]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"properties": {
|
||||||
|
"action": { "const": "list_issues" },
|
||||||
|
"owner": { "type": "string" },
|
||||||
|
"repo": { "type": "string" },
|
||||||
|
"state": { "type": "string", "enum": ["open", "closed", "all"], "default": "open" },
|
||||||
|
"limit": { "type": "integer", "default": 30 }
|
||||||
|
},
|
||||||
|
"required": ["action", "owner", "repo"]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"properties": {
|
||||||
|
"action": { "const": "create_issue" },
|
||||||
|
"owner": { "type": "string" },
|
||||||
|
"repo": { "type": "string" },
|
||||||
|
"title": { "type": "string" },
|
||||||
|
"body": { "type": "string" },
|
||||||
|
"labels": { "type": "array", "items": { "type": "string" } }
|
||||||
|
},
|
||||||
|
"required": ["action", "owner", "repo", "title"]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"properties": {
|
||||||
|
"action": { "const": "get_issue" },
|
||||||
|
"owner": { "type": "string" },
|
||||||
|
"repo": { "type": "string" },
|
||||||
|
"issue_number": { "type": "integer" }
|
||||||
|
},
|
||||||
|
"required": ["action", "owner", "repo", "issue_number"]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"properties": {
|
||||||
|
"action": { "const": "list_pull_requests" },
|
||||||
|
"owner": { "type": "string" },
|
||||||
|
"repo": { "type": "string" },
|
||||||
|
"state": { "type": "string", "enum": ["open", "closed", "all"], "default": "open" },
|
||||||
|
"limit": { "type": "integer", "default": 30 }
|
||||||
|
},
|
||||||
|
"required": ["action", "owner", "repo"]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"properties": {
|
||||||
|
"action": { "const": "get_pull_request" },
|
||||||
|
"owner": { "type": "string" },
|
||||||
|
"repo": { "type": "string" },
|
||||||
|
"pr_number": { "type": "integer" }
|
||||||
|
},
|
||||||
|
"required": ["action", "owner", "repo", "pr_number"]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"properties": {
|
||||||
|
"action": { "const": "get_pull_request_files" },
|
||||||
|
"owner": { "type": "string" },
|
||||||
|
"repo": { "type": "string" },
|
||||||
|
"pr_number": { "type": "integer" }
|
||||||
|
},
|
||||||
|
"required": ["action", "owner", "repo", "pr_number"]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"properties": {
|
||||||
|
"action": { "const": "create_pr_review" },
|
||||||
|
"owner": { "type": "string" },
|
||||||
|
"repo": { "type": "string" },
|
||||||
|
"pr_number": { "type": "integer" },
|
||||||
|
"body": { "type": "string", "description": "Review comment" },
|
||||||
|
"event": { "type": "string", "enum": ["APPROVE", "REQUEST_CHANGES", "COMMENT"] }
|
||||||
|
},
|
||||||
|
"required": ["action", "owner", "repo", "pr_number", "body", "event"]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"properties": {
|
||||||
|
"action": { "const": "list_repos" },
|
||||||
|
"username": { "type": "string" },
|
||||||
|
"limit": { "type": "integer", "default": 30 }
|
||||||
|
},
|
||||||
|
"required": ["action", "username"]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"properties": {
|
||||||
|
"action": { "const": "get_file_content" },
|
||||||
|
"owner": { "type": "string" },
|
||||||
|
"repo": { "type": "string" },
|
||||||
|
"path": { "type": "string", "description": "File path in repo" },
|
||||||
|
"ref": { "type": "string", "description": "Branch/commit (default: default branch)" }
|
||||||
|
},
|
||||||
|
"required": ["action", "owner", "repo", "path"]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"properties": {
|
||||||
|
"action": { "const": "trigger_workflow" },
|
||||||
|
"owner": { "type": "string" },
|
||||||
|
"repo": { "type": "string" },
|
||||||
|
"workflow_id": { "type": "string", "description": "Workflow filename or ID" },
|
||||||
|
"ref": { "type": "string", "description": "Branch to run on" },
|
||||||
|
"inputs": { "type": "object" }
|
||||||
|
},
|
||||||
|
"required": ["action", "owner", "repo", "workflow_id", "ref"]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"properties": {
|
||||||
|
"action": { "const": "get_workflow_runs" },
|
||||||
|
"owner": { "type": "string" },
|
||||||
|
"repo": { "type": "string" },
|
||||||
|
"workflow_id": { "type": "string" },
|
||||||
|
"limit": { "type": "integer", "default": 30 }
|
||||||
|
},
|
||||||
|
"required": ["action", "owner", "repo"]
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}"#;
|
||||||
|
|
||||||
|
export!(GitHubTool);
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_url_encode_path() {
|
||||||
|
assert_eq!(url_encode_path("foo-bar_123.baz"), "foo-bar_123.baz");
|
||||||
|
assert_eq!(url_encode_path("foo bar"), "foo%20bar");
|
||||||
|
assert_eq!(url_encode_path("foo/bar"), "foo%2Fbar");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_validate_path_segment() {
|
||||||
|
assert!(validate_path_segment("foo"));
|
||||||
|
assert!(!validate_path_segment(""));
|
||||||
|
assert!(!validate_path_segment("foo/bar"));
|
||||||
|
assert!(!validate_path_segment(".."));
|
||||||
|
// Empty segments are handled in get_file_content logic, not here
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_validate_event_in_create_pr_review() {
|
||||||
|
let valid = ["APPROVE", "REQUEST_CHANGES", "COMMENT"];
|
||||||
|
// Ensure valid inputs are accepted
|
||||||
|
for event in valid {
|
||||||
|
assert!(valid.contains(&event));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_input_length_validation() {
|
||||||
|
assert!(validate_input_length("short", "test").is_ok());
|
||||||
|
|
||||||
|
let long = "a".repeat(MAX_TEXT_LENGTH + 1);
|
||||||
|
assert!(validate_input_length(&long, "test").is_err());
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -136,7 +136,7 @@ fn parse_message(v: &serde_json::Value) -> Message {
|
|||||||
date: get_header(payload, "Date"),
|
date: get_header(payload, "Date"),
|
||||||
body: extract_body(payload),
|
body: extract_body(payload),
|
||||||
snippet: v["snippet"].as_str().unwrap_or("").to_string(),
|
snippet: v["snippet"].as_str().unwrap_or("").to_string(),
|
||||||
is_unread: label_ids.contains(&"UNREAD".to_string()),
|
is_unread: label_ids.iter().any(|l| l == "UNREAD"),
|
||||||
label_ids,
|
label_ids,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -198,7 +198,7 @@ pub fn list_messages(
|
|||||||
to: get_header(payload, "To"),
|
to: get_header(payload, "To"),
|
||||||
date: get_header(payload, "Date"),
|
date: get_header(payload, "Date"),
|
||||||
snippet: msg["snippet"].as_str().unwrap_or("").to_string(),
|
snippet: msg["snippet"].as_str().unwrap_or("").to_string(),
|
||||||
is_unread: label_ids.contains(&"UNREAD".to_string()),
|
is_unread: label_ids.iter().any(|l| l == "UNREAD"),
|
||||||
label_ids,
|
label_ids,
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,7 +6,7 @@
|
|||||||
//! # Capabilities Required
|
//! # Capabilities Required
|
||||||
//!
|
//!
|
||||||
//! - HTTP: `www.googleapis.com/calendar/v3/*` (GET, POST, PUT, PATCH, DELETE)
|
//! - HTTP: `www.googleapis.com/calendar/v3/*` (GET, POST, PUT, PATCH, DELETE)
|
||||||
//! - Secrets: `google_calendar_token` (OAuth 2.0 token, injected automatically)
|
//! - Secrets: `google_oauth_token` (OAuth 2.0 token, injected automatically)
|
||||||
//!
|
//!
|
||||||
//! # Supported Actions
|
//! # Supported Actions
|
||||||
//!
|
//!
|
||||||
|
|||||||
@@ -269,8 +269,13 @@ pub fn replace_text(
|
|||||||
|
|
||||||
let parsed = batch_update_raw(document_id, vec![request])?;
|
let parsed = batch_update_raw(document_id, vec![request])?;
|
||||||
|
|
||||||
let occurrences = parsed["replies"][0]["replaceAllText"]["occurrencesChanged"]
|
let first_reply = parsed["replies"].as_array().and_then(|arr| arr.first());
|
||||||
|
let occurrences = first_reply
|
||||||
|
.map(|r| {
|
||||||
|
r["replaceAllText"]["occurrencesChanged"]
|
||||||
.as_i64()
|
.as_i64()
|
||||||
|
.unwrap_or(0)
|
||||||
|
})
|
||||||
.unwrap_or(0);
|
.unwrap_or(0);
|
||||||
|
|
||||||
Ok(ReplaceResult {
|
Ok(ReplaceResult {
|
||||||
|
|||||||
@@ -330,7 +330,13 @@ pub fn add_sheet(spreadsheet_id: &str, title: &str) -> Result<AddSheetResult, St
|
|||||||
|
|
||||||
let parsed = batch_update(spreadsheet_id, requests)?;
|
let parsed = batch_update(spreadsheet_id, requests)?;
|
||||||
|
|
||||||
let reply = &parsed["replies"][0]["addSheet"]["properties"];
|
let reply = parsed["replies"]
|
||||||
|
.as_array()
|
||||||
|
.and_then(|arr| arr.first())
|
||||||
|
.map(|r| &r["addSheet"]["properties"]);
|
||||||
|
|
||||||
|
let reply = reply.ok_or_else(|| "No reply from batch update".to_string())?;
|
||||||
|
|
||||||
Ok(AddSheetResult {
|
Ok(AddSheetResult {
|
||||||
sheet: SheetInfo {
|
sheet: SheetInfo {
|
||||||
sheet_id: reply["sheetId"].as_i64().unwrap_or(0),
|
sheet_id: reply["sheetId"].as_i64().unwrap_or(0),
|
||||||
|
|||||||
Reference in New Issue
Block a user