Compare commits

..
Author SHA1 Message Date
Illia PolosukhinandClaude Opus 4.6 09e6a7e6d8 fix: review fixes for libSQL backend (shared connections, panics, indexes)
- Replace .expect() with proper error propagation in 3 call sites
- Share Arc<Database> between backend and stores instead of single Connection
- Add connect-per-operation pattern to LibSqlSecretsStore and LibSqlWasmToolStore
- Wrap store() INSERT + SELECT-back in a transaction
- Add ~22 missing indexes for parity with PostgreSQL schema
- Add 18 leak_detection_patterns seed rows matching PostgreSQL V2 migration
- Fix super:: import to use crate:: style
- Gate mask_password_in_url behind #[cfg(feature = "postgres")]
- Rewrite secrets store init with or_else chain for runtime backend selection

Co-Authored-By: Claude Opus 4.6 <[email protected]>
2026-02-13 17:03:58 -08:00
ZakiandClaude Opus 4.6 5814d77b16 fix: add missing JobContext fields and resolve fmt/clippy warnings
Add total_tokens_used and max_tokens fields to JobContext in
libsql_backend.rs, apply cargo fmt, and fix clippy warnings.

Co-Authored-By: Claude Opus 4.6 <[email protected]>
2026-02-13 08:24:03 -08:00
ZakiandClaude Opus 4.6 83773af997 merge: Resolve conflicts with main, add user-scoped DB methods
Merge main's security hardening (user-scoped job/conversation access,
cargo-dist config, CI improvements) into turso branch.

Add Database trait methods for user-scoped operations:
- list_sandbox_jobs_for_user
- sandbox_job_summary_for_user
- sandbox_job_belongs_to_user
- conversation_belongs_to_user

Implemented in both postgres and libsql backends.

Co-Authored-By: Claude Opus 4.6 <[email protected]>
2026-02-13 07:30:27 -08:00
ZakiandClaude Opus 4.6 7474fd4c52 fix: address PR review feedback for libSQL backend
- P0: Switch libsql_backend to connection-per-operation pattern to fix
  shared Connection concurrency issue across tokio tasks
- P0: Wrap secrets store INSERT+SELECT in transaction to fix TOCTOU race
- P0: Document encryption-at-rest limitations and json_patch divergence
- P1: Fix get_opt_text removing .filter(|s| !s.is_empty()) that conflated
  empty strings with NULL
- P1: Replace datetime('now') with fmt_ts(&Utc::now()) for consistent
  RFC 3339 timestamps across all queries
- P2: Use explicit _rowid column in FTS5 triggers and joins for stability
  across VACUUM operations
- P2: Add tracing::warn when embedding provided but vector search disabled
  in hybrid_search
- Extract shared connect_from_config() helper to deduplicate DB connection
  logic across main.rs, cli/config.rs, and cli/mcp.rs

Co-Authored-By: Claude Opus 4.6 <[email protected]>
2026-02-12 07:02:21 -08:00
ZakiandClaude Opus 4.6 64b6f559fd feat: enable onboarding wizard for libSQL builds
Refactor the setup wizard to work with both postgres and libsql feature
flags. Previously the wizard was gated behind #[cfg(feature = "postgres")]
only, so libsql-only builds would print an error on `ironclaw onboard`.

- Add libsql fields to Settings (database_backend, libsql_path, libsql_url)
- Split wizard database/migration/secrets methods into feature-gated variants
- Add step_database_libsql() with local path and Turso remote replica prompts
- Update setup/mod.rs and main.rs feature gates to any(postgres, libsql)
- Extend check_onboard_needed() to detect libsql database presence

Co-Authored-By: Claude Opus 4.6 <[email protected]>
2026-02-12 06:45:42 -08:00
ZakiandClaude Opus 4.6 0de6f6aabb feat: add libSQL/Turso database backend with full feature parity
Introduce a Database trait abstraction (~60 async methods) enabling
compile-time backend selection between PostgreSQL and libSQL/Turso.
Convert all modules from concrete Store to Arc<dyn Database>, add
LibSqlSecretsStore and LibSqlWasmToolStore implementations, wire
libsql stores throughout CLI and main entry points, and make the
setup wizard backend-agnostic.

Key changes:
- src/db/: Database trait, PostgresDatabase adapter, LibSqlBackend
  with native SQLite-dialect SQL, and idempotent migration system
- src/secrets/store.rs: LibSqlSecretsStore (all 8 trait methods)
- src/tools/wasm/storage.rs: LibSqlWasmToolStore (all 7 trait methods)
- src/main.rs, cli/tool.rs, cli/mcp.rs: backend-conditional wiring
- src/setup/channels.rs: SecretsContext uses Arc<dyn SecretsStore>
- Feature-gate postgres-only tests and examples

Co-Authored-By: Claude Opus 4.6 <[email protected]>
2026-02-12 04:35:05 -08:00
134 changed files with 4686 additions and 12079 deletions
-97
View File
@@ -1,97 +0,0 @@
---
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
-81
View File
@@ -1,81 +0,0 @@
---
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.
-245
View File
@@ -1,245 +0,0 @@
---
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.
-170
View File
@@ -1,170 +0,0 @@
---
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.
+1
View File
@@ -39,6 +39,7 @@ permissions:
# If there's a prerelease-style suffix to the version, then the release(s)
# will be marked as a prerelease.
on:
pull_request:
push:
tags:
- '**[0-9]+.[0-9]+.[0-9]+*'
-9
View File
@@ -1,15 +1,6 @@
.env
.env.local
.env.*
!.env.example
# Claude Code worktrees
.claude/worktrees/
# Sidecar tool data
.sidecar/
.todos/
target/
-54
View File
@@ -7,60 +7,6 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
## [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
### Other
+1 -49
View File
@@ -198,9 +198,8 @@ When designing new features or systems, always prefer generic/extensible archite
### Error Handling
- Use `thiserror` for error types in `error.rs`
- Never use `.unwrap()` or `.expect()` in production code (tests are fine)
- Never use `.unwrap()` in production code (tests are fine)
- Map errors with context: `.map_err(|e| SomeError::Variant { reason: e.to_string() })?`
- Before committing, grep for `.unwrap()` and `.expect(` in changed files to catch violations mechanically
### Async
- All I/O is async with tokio
@@ -630,22 +629,6 @@ RUST_LOG=ironclaw::agent=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
- Use `crate::` imports, not `super::`
@@ -654,37 +637,6 @@ for that module's behavior. When modifying code in a module that has a spec:
- Keep functions focused, extract helpers when logic is reused
- Comments for non-obvious logic only
## Review & Fix Discipline
Hard-won lessons from code review -- follow these when fixing bugs or addressing review feedback.
### Fix the pattern, not just the instance
When a reviewer flags a bug (e.g., TOCTOU race in INSERT + SELECT-back), search the entire codebase for all instances of that same pattern. A fix in `SecretsStore::create()` that doesn't also fix `WasmToolStore::store()` is half a fix.
### Propagate architectural fixes to satellite types
If a core type changes its concurrency model (e.g., `LibSqlBackend` switches to connection-per-operation), every type that was handed a resource from the old model (e.g., `LibSqlSecretsStore`, `LibSqlWasmToolStore` holding a single `Connection`) must also be updated. Grep for the old type across the codebase.
### Schema translation is more than DDL
When translating a database schema between backends (PostgreSQL to libSQL, etc.), check for:
- **Indexes** -- diff `CREATE INDEX` statements between the two schemas
- **Seed data** -- check for `INSERT INTO` in migrations (e.g., `leak_detection_patterns`)
- **Semantic differences** -- document where SQL functions behave differently (e.g., `json_patch` vs `jsonb_set`)
### Feature flag testing
When adding feature-gated code, test compilation with each feature in isolation:
```bash
cargo check # default features
cargo check --no-default-features --features libsql # libsql only
cargo check --all-features # all features
```
Dead code behind the wrong `#[cfg]` gate will only show up when building with a single feature.
### Mechanical verification before committing
Run these checks on changed files before committing:
- `grep -rnE '\.unwrap\(|\.expect\(' <files>` -- no panics in production
- `grep -rn 'super::' <files>` -- use `crate::` imports
- If you fixed a pattern bug, `grep` for other instances of that pattern across `src/`
## Workspace & Memory System
Inspired by [OpenClaw](https://github.com/openclaw/openclaw), the workspace provides persistent memory for agents with a flexible filesystem-like structure.
Generated
+12 -4
View File
@@ -2210,11 +2210,11 @@ dependencies = [
"hyper 1.8.1",
"hyper-util",
"rustls",
"rustls-native-certs",
"rustls-pki-types",
"tokio",
"tokio-rustls",
"tower-service",
"webpki-roots",
]
[[package]]
@@ -2490,7 +2490,7 @@ dependencies = [
[[package]]
name = "ironclaw"
version = "0.5.0"
version = "0.1.3"
dependencies = [
"aes-gcm",
"aho-corasick",
@@ -2548,7 +2548,6 @@ dependencies = [
"tower-http 0.6.8",
"tracing",
"tracing-subscriber",
"url",
"urlencoding",
"uuid",
"wasmparser 0.220.1",
@@ -4034,7 +4033,6 @@ dependencies = [
"pin-project-lite",
"quinn",
"rustls",
"rustls-native-certs",
"rustls-pki-types",
"serde",
"serde_json",
@@ -4052,6 +4050,7 @@ dependencies = [
"wasm-bindgen-futures",
"wasm-streams",
"web-sys",
"webpki-roots",
]
[[package]]
@@ -6238,6 +6237,15 @@ dependencies = [
"wasm-bindgen",
]
[[package]]
name = "webpki-roots"
version = "1.0.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "12bed680863276c63889429bfd6cab3b99943659923822de1c8a39c49e4d722c"
dependencies = [
"rustls-pki-types",
]
[[package]]
name = "which"
version = "4.4.2"
+6 -17
View File
@@ -1,16 +1,8 @@
[workspace]
exclude = [
"channels-src/telegram",
"channels-src/slack",
"channels-src/whatsapp",
"tools-src/gmail",
]
[package]
name = "ironclaw"
version = "0.5.0"
version = "0.1.3"
edition = "2024"
rust-version = "1.92"
rust-version = "1.85"
description = "Secure personal AI assistant that protects your data and expands its capabilities on the fly"
authors = ["NEAR AI <[email protected]>"]
license = "MIT OR Apache-2.0"
@@ -30,7 +22,7 @@ tokio-stream = { version = "0.1", features = ["sync"] }
futures = "0.3"
# HTTP client
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls-native-roots", "stream"] }
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls", "stream"] }
# Serialization
serde = { version = "1", features = ["derive"] }
@@ -92,8 +84,7 @@ fs4 = "0.6"
# Secrecy for sensitive values
secrecy = { version = "0.10", features = ["serde"] }
# URL parsing and encoding
url = "2"
# URL encoding for OAuth flow
urlencoding = "2"
# Open URLs in browser
@@ -147,7 +138,7 @@ pretty_assertions = "1"
tempfile = "3"
[features]
default = ["postgres", "libsql"]
default = ["postgres"]
postgres = [
"dep:deadpool-postgres",
"dep:tokio-postgres",
@@ -191,13 +182,11 @@ windows-archive = ".tar.gz"
# The archive format to use for non-windows builds (defaults .tar.xz)
unix-archive = ".tar.gz"
# Which actions to run on pull requests
pr-run-mode = "skip"
pr-run-mode = "upload"
# Path that installers should place binaries in
install-path = "CARGO_HOME"
# Whether to install an updater program
install-updater = true
# Cache intermediate build artifacts to speed up the release pipelines
cache-builds = true
[workspace.metadata.dist.github-custom-runners]
aarch64-unknown-linux-gnu = "ubuntu-24.04-arm"
-46
View File
@@ -1,46 +0,0 @@
# Multi-stage Dockerfile for the IronClaw agent (cloud deployment).
#
# Build:
# docker build --platform linux/amd64 -t ironclaw:latest .
#
# Run:
# docker run --env-file .env -p 3000:3000 ironclaw:latest
# Stage 1: Build
FROM rust:1.92-slim-bookworm AS builder
RUN apt-get update && apt-get install -y --no-install-recommends \
pkg-config libssl-dev cmake gcc g++ \
&& rm -rf /var/lib/apt/lists/*
WORKDIR /app
# Copy manifests first for layer caching
COPY Cargo.toml Cargo.lock ./
# Copy source and build artifacts
COPY src/ src/
COPY migrations/ migrations/
COPY wit/ wit/
RUN cargo build --release --bin ironclaw
# Stage 2: Runtime
FROM debian:bookworm-slim
RUN apt-get update && apt-get install -y --no-install-recommends \
ca-certificates libssl3 \
&& rm -rf /var/lib/apt/lists/*
COPY --from=builder /app/target/release/ironclaw /usr/local/bin/ironclaw
COPY --from=builder /app/migrations /app/migrations
# Non-root user
RUN useradd -m -u 1000 -s /bin/bash ironclaw
USER ironclaw
EXPOSE 3000
ENV RUST_LOG=ironclaw=info
ENTRYPOINT ["ironclaw"]
+2 -2
View File
@@ -9,7 +9,7 @@
# The image includes common development tools so workers can build software,
# run tests, and execute shell commands.
FROM rust:1.92-bookworm AS builder
FROM rust:1.85-bookworm AS builder
WORKDIR /build
COPY . .
@@ -40,7 +40,7 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
ENV RUSTUP_HOME=/usr/local/rustup \
CARGO_HOME=/usr/local/cargo \
PATH=/usr/local/cargo/bin:$PATH
RUN curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | sh -s -- -y --default-toolchain 1.92.0 \
RUN curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | sh -s -- -y --default-toolchain 1.85.0 \
&& chmod -R a+r /usr/local/rustup /usr/local/cargo
# Install Claude Code CLI (for claude-bridge mode)
+18 -14
View File
@@ -112,7 +112,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
| `pairing` | ✅ | ✅ | - | list/approve for channel DM pairing |
| `nodes` | ✅ | ❌ | P3 | Device management |
| `plugins` | ✅ | ❌ | P3 | Plugin management |
| `hooks` | ✅ | | P2 | Lifecycle hooks |
| `hooks` | ✅ | | P2 | Lifecycle hooks |
| `cron` | ✅ | ❌ | P2 | Scheduled jobs |
| `webhooks` | ✅ | ❌ | P3 | Webhook config |
| `message send` | ✅ | ❌ | P2 | Send to channels |
@@ -133,7 +133,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|---------|----------|----------|-------|
| Pi agent runtime | ✅ | | IronClaw uses custom runtime |
| RPC-based execution | ✅ | ✅ | Orchestrator/worker pattern |
| Multi-provider failover | ✅ | | `FailoverProvider` tries providers sequentially on retryable errors |
| Multi-provider failover | ✅ | | Provider fallback chains |
| Per-sender sessions | ✅ | ✅ | |
| Global sessions | ✅ | ❌ | Optional shared context |
| Session pruning | ✅ | ❌ | Auto cleanup old sessions |
@@ -164,7 +164,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
| AWS Bedrock | ✅ | ❌ | P3 | |
| Google Gemini | ✅ | ❌ | P3 | |
| OpenRouter | ✅ | ❌ | P3 | |
| Ollama (local) | ✅ | | - | via `rig::providers::ollama` (full support) |
| Ollama (local) | ✅ | | P2 | Local models |
| node-llama-cpp | ✅ | | - | N/A for Rust |
| llama.cpp (native) | ❌ | 🔮 | P3 | Rust bindings |
@@ -173,8 +173,8 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
| Feature | OpenClaw | IronClaw | Notes |
|---------|----------|----------|-------|
| Auto-discovery | ✅ | ❌ | |
| Failover chains | ✅ | | `FailoverProvider` with configurable `fallback_model` |
| Cooldown management | ✅ | | Lock-free per-provider cooldown in `FailoverProvider` |
| Failover chains | ✅ | | Provider fallback |
| Cooldown management | ✅ | | Skip failed providers |
| Per-session model override | ✅ | ✅ | Model selector in TUI |
| 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 |
| Timezone support | ✅ | ✅ | - | Via cron expressions |
| One-shot/recurring jobs | ✅ | ✅ | - | Manual + cron triggers |
| `beforeInbound` hook | ✅ | | P2 | |
| `beforeOutbound` hook | ✅ | | P2 | |
| `beforeToolCall` hook | ✅ | | P2 | |
| `beforeInbound` hook | ✅ | | P2 | |
| `beforeOutbound` hook | ✅ | | P2 | |
| `beforeToolCall` hook | ✅ | | P2 | |
| `onMessage` hook | ✅ | ✅ | - | Routines with event trigger |
| `onSessionStart` hook | ✅ | | P2 | |
| `onSessionEnd` hook | ✅ | | P2 | |
| `onSessionStart` hook | ✅ | | P2 | |
| `onSessionEnd` hook | ✅ | | P2 | |
| `transcribeAudio` hook | ✅ | ❌ | P3 | |
| `transformResponse` hook | ✅ | | P2 | |
| `transformResponse` hook | ✅ | | P2 | |
| Bundled hooks | ✅ | ❌ | P2 | |
| Plugin hooks | ✅ | ❌ | P3 | |
| Workspace hooks | ✅ | ❌ | P2 | Inline code |
@@ -419,11 +419,15 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
- ❌ Slack channel (real implementation)
- ✅ Telegram channel (WASM, DM pairing, caption, /start)
- ❌ WhatsApp channel
- Multi-provider failover (`FailoverProvider` with retryable error classification)
- Hooks system (beforeInbound, beforeToolCall, beforeOutbound, onSessionStart, onSessionEnd, transformResponse)
- Multi-provider failover
- Hooks system (beforeInbound, beforeToolCall, etc.)
### P2 - Medium Priority
-Media handling (images, PDFs)
-Cron job scheduling
- ❌ Web Control UI
- ❌ WebChat channel
- 🚧 Media handling (caption support; no image/PDF processing)
- ❌ CLI subcommands (config, status, memory, doctor)
- ❌ Ollama/local model support
- ❌ Configuration hot-reload
- ❌ Webhook trigger endpoint in web gateway
+36 -36
View File
@@ -181,42 +181,42 @@ External content passes through multiple security layers:
## Architecture
```
┌────────────────────────────────────────────────────────────────┐
│ Channels │
│ ┌──────┐ ┌──────┐ ┌─────────────┐ ┌─────────────┐ │
│ │ REPL │ │ HTTP │ │WASM Channels│ │ Web Gateway │ │
│ └──┬───┘ └──┬───┘ └──────┬──────┘ │ (SSE + WS) │ │
│ │ │ │ └──────┬──────┘ │
│ └─────────┴──────────────┴────────────────┘ │
│ │ │
│ ┌─────────▼─────────┐ │
│ │ Agent Loop │ Intent routing │
│ └────┬─────────┬───┘ │
│ │ │ │
│ ┌──────────▼───┐ ┌──▼──────────────┐ │
│ │ Scheduler │ │ Routines Engine │ │
│ │(parallel jobs)│ │(cron, event, wh) │ │
│ └──────┬───────┘ └────────┬─────────┘ │
│ │ │ │
│ ┌─────────────┼───────────────────┘ │
│ │ │ │
│ ┌───▼────┐ ┌────▼────────────────┐ │
│ │ Local │ │ Orchestrator │ │
│ │Workers │ │ ┌───────────────┐ │ │
│ │(in-proc)│ │ │ Docker Sandbox│ │ │
│ └───┬────┘ │ │ Containers │ │ │
│ │ │ │ ┌───────────┐ │ │ │
│ │ │ │ │Worker / CC│ │ │ │
│ │ │ │ └───────────┘ │ │ │
│ │ │ └───────────────┘ │ │
│ │ └─────────┬───────────┘ │
│ └──────────────────┤ │
│ │ │
│ ┌───────────▼──────────┐ │
│ │ Tool Registry │ │
│ │ Built-in, MCP, WASM │ │
│ └──────────────────────┘ │
└────────────────────────────────────────────────────────────────┘
┌────────────────────────────────────────────────────────────────────
│ Channels
│ ┌──────┐ ┌──────┐ ┌─────────────┐ ┌─────────────┐
│ │ REPL │ │ HTTP │ │WASM Channels│ │ Web Gateway │
│ └──┬───┘ └──┬───┘ └──────┬──────┘ │ (SSE + WS) │
│ │ │ │ └──────┬──────┘
│ └─────────┴──────────────┴────────────────┘
│ │
│ ┌─────────▼─────────┐
│ │ Agent Loop │ Intent routing
│ └────┬─────────┬───┘
│ │ │
│ ┌──────────▼───┐ ┌──▼──────────────┐
│ │ Scheduler │ │ Routines Engine │
│ │(parallel jobs)│ │(cron, event, wh) │
│ └──────┬───────┘ └────────┬─────────┘
│ │ │
│ ┌─────────────┼───────────────────┘
│ │ │
│ ┌───▼────┐ ┌────▼────────────────┐
│ │ Local │ │ Orchestrator │
│ │Workers │ │ ┌───────────────┐ │
│ │(in-proc)│ │ │ Docker Sandbox│ │
│ └───┬────┘ │ │ Containers │ │
│ │ │ │ ┌───────────┐ │ │
│ │ │ │ │Worker / CC│ │ │
│ │ │ │ └───────────┘ │ │
│ │ │ └───────────────┘ │
│ │ └─────────┬───────────┘
│ └──────────────────┤
│ │
│ ┌───────────▼──────────┐
│ │ Tool Registry │
│ │ Built-in, MCP, WASM │
│ └──────────────────────┘
└────────────────────────────────────────────────────────────────────
```
### Core Components
-23
View File
@@ -1,23 +0,0 @@
[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
-121
View File
@@ -1,121 +0,0 @@
# 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
@@ -1,39 +0,0 @@
{
"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
}
}
-476
View File
@@ -1,476 +0,0 @@
//! 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");
}
}
+2 -14
View File
@@ -338,13 +338,7 @@ fn emit_message(
team_id,
};
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()
});
let metadata_json = serde_json::to_string(&metadata).unwrap_or_else(|_| "{}".to_string());
// Strip @ mentions of the bot from the text for cleaner messages
let cleaned_text = strip_bot_mention(&text);
@@ -372,13 +366,7 @@ fn strip_bot_mention(text: &str) -> String {
/// Create a JSON HTTP response.
fn json_response(status: u16, value: serde_json::Value) -> OutgoingHttpResponse {
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 body = serde_json::to_vec(&value).unwrap_or_default();
let headers = serde_json::json!({"Content-Type": "application/json"});
OutgoingHttpResponse {
+19 -9
View File
@@ -285,7 +285,11 @@ impl Guest for TelegramChannel {
}
// Persist dm_policy and allow_from for DM pairing in handle_message
let dm_policy = config.dm_policy.as_deref().unwrap_or("pairing").to_string();
let dm_policy = config
.dm_policy
.as_deref()
.unwrap_or("pairing")
.to_string();
let _ = channel_host::workspace_write(DM_POLICY_PATH, &dm_policy);
let allow_from_json = serde_json::to_string(&config.allow_from.unwrap_or_default())
@@ -840,8 +844,8 @@ fn send_pairing_reply(chat_id: i64, code: &str) -> Result<(), String> {
"parse_mode": "Markdown",
});
let payload_bytes =
serde_json::to_vec(&payload).map_err(|e| format!("Failed to serialize payload: {}", e))?;
let payload_bytes = serde_json::to_vec(&payload)
.map_err(|e| format!("Failed to serialize payload: {}", e))?;
let headers = serde_json::json!({
"Content-Type": "application/json"
@@ -911,10 +915,15 @@ fn handle_message(message: TelegramMessage) {
let is_private = message.chat.chat_type == "private";
// Owner validation: when owner_id is set, only that user can message
let owner_id_str = channel_host::workspace_read(OWNER_ID_PATH).filter(|s| !s.is_empty());
let owner_configured = channel_host::workspace_read(OWNER_ID_PATH)
.map(|s| !s.is_empty())
.unwrap_or(false);
if let Some(ref id_str) = owner_id_str {
if let Ok(owner_id) = id_str.parse::<i64>() {
if owner_configured {
if let Ok(owner_id) = channel_host::workspace_read(OWNER_ID_PATH)
.unwrap()
.parse::<i64>()
{
if from.id != owner_id {
channel_host::log(
channel_host::LogLevel::Debug,
@@ -928,8 +937,8 @@ fn handle_message(message: TelegramMessage) {
}
} else if is_private {
// No owner_id: apply dm_policy for private chats
let dm_policy =
channel_host::workspace_read(DM_POLICY_PATH).unwrap_or_else(|| "pairing".to_string());
let dm_policy = channel_host::workspace_read(DM_POLICY_PATH)
.unwrap_or_else(|| "pairing".to_string());
if dm_policy != "open" {
// Build effective allow list: config allow_from + pairing store
@@ -992,7 +1001,8 @@ fn handle_message(message: TelegramMessage) {
if !respond_to_all {
let has_command = content.starts_with('/');
let bot_username = channel_host::workspace_read(BOT_USERNAME_PATH).unwrap_or_default();
let bot_username = channel_host::workspace_read(BOT_USERNAME_PATH)
.unwrap_or_default();
let has_bot_mention = if bot_username.is_empty() {
content.contains('@')
} else {
+6 -23
View File
@@ -254,19 +254,10 @@ struct WhatsAppChannel;
impl Guest for WhatsAppChannel {
fn on_start(config_json: String) -> Result<ChannelConfig, String> {
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(),
reply_to_message: default_reply_to_message(),
}
}
};
let config: WhatsAppConfig = serde_json::from_str(&config_json).unwrap_or(WhatsAppConfig {
api_version: default_api_version(),
reply_to_message: default_reply_to_message(),
});
channel_host::log(
channel_host::LogLevel::Info,
@@ -276,9 +267,6 @@ 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
Ok(ChannelConfig {
display_name: "WhatsApp".to_string(),
@@ -339,16 +327,11 @@ impl Guest for WhatsAppChannel {
let metadata: WhatsAppMessageMetadata = serde_json::from_str(&response.metadata_json)
.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
// Host will replace {WHATSAPP_ACCESS_TOKEN} with actual token in Authorization header
let api_url = format!(
"https://graph.facebook.com/{}/{}/messages",
api_version, metadata.phone_number_id
"https://graph.facebook.com/v18.0/{}/messages",
metadata.phone_number_id
);
// Build sendMessage payload
-13
View File
@@ -1,13 +0,0 @@
[Unit]
Description=Cloud SQL Auth Proxy
After=network.target
[Service]
Type=simple
DynamicUser=yes
ExecStart=/usr/local/bin/cloud-sql-proxy ironclaw-prod:us-central1:ironclaw-db --port=5432
Restart=always
RestartSec=5
[Install]
WantedBy=multi-user.target
-27
View File
@@ -1,27 +0,0 @@
# WARNING: Replace all CHANGE_ME values before deploying.
# Do not use placeholder passwords in production.
DATABASE_URL=postgres://ironclaw:CHANGE_ME@localhost:5432/ironclaw
# NEAR AI
NEARAI_SESSION_TOKEN=CHANGE_ME
NEARAI_MODEL=claude-3-5-sonnet-20241022
NEARAI_BASE_URL=https://cloud-api.near.ai
NEARAI_AUTH_URL=https://private.near.ai
NEARAI_API_MODE=chat_completions
# Agent
AGENT_NAME=ironclaw
CLI_ENABLED=false
# Web Gateway
GATEWAY_ENABLED=true
# 0.0.0.0 binds to all interfaces (required for Docker --network=host).
# Use 127.0.0.1 if running outside Docker or for local-only access.
GATEWAY_HOST=0.0.0.0
GATEWAY_PORT=3000
GATEWAY_AUTH_TOKEN=CHANGE_ME
# Disabled for initial deploy
SANDBOX_ENABLED=false
HEARTBEAT_ENABLED=false
EMBEDDING_ENABLED=false
-20
View File
@@ -1,20 +0,0 @@
[Unit]
Description=IronClaw AI Assistant
After=cloud-sql-proxy.service docker.service
Requires=cloud-sql-proxy.service
[Service]
Type=simple
ExecStartPre=/usr/bin/docker pull us-central1-docker.pkg.dev/ironclaw-prod/ironclaw/agent:latest
ExecStart=/usr/bin/docker run --rm \
--name ironclaw \
--env-file /opt/ironclaw/.env \
--network=host \
us-central1-docker.pkg.dev/ironclaw-prod/ironclaw/agent:latest \
--no-onboard
ExecStop=/usr/bin/docker stop ironclaw
Restart=always
RestartSec=10
[Install]
WantedBy=multi-user.target
-68
View File
@@ -1,68 +0,0 @@
#!/usr/bin/env bash
# VM bootstrap script for IronClaw on GCP Compute Engine.
#
# Run on a fresh Debian 12 VM after SSH:
# sudo bash setup.sh
#
# Prerequisites:
# - VM has the ironclaw-vm service account attached
# - Cloud SQL Auth Proxy accessible via IAM
# - Artifact Registry image pushed
set -euo pipefail
# Must run as root
if [ "$(id -u)" -ne 0 ]; then
echo "ERROR: This script must be run as root (sudo bash setup.sh)"
exit 1
fi
echo "==> Installing Docker"
apt-get update
apt-get install -y docker.io
systemctl enable docker
systemctl start docker
echo "==> Installing Cloud SQL Auth Proxy"
curl -fsSL -o /usr/local/bin/cloud-sql-proxy \
https://storage.googleapis.com/cloud-sql-connectors/cloud-sql-proxy/v2.14.3/cloud-sql-proxy.linux.amd64
chmod +x /usr/local/bin/cloud-sql-proxy
echo "==> Installing systemd services"
cp /tmp/deploy/cloud-sql-proxy.service /etc/systemd/system/
cp /tmp/deploy/ironclaw.service /etc/systemd/system/
systemctl daemon-reload
echo "==> Starting Cloud SQL Auth Proxy"
systemctl enable cloud-sql-proxy
systemctl start cloud-sql-proxy
echo "==> Configuring Docker registry auth"
# The VM service account provides Artifact Registry access
gcloud auth configure-docker us-central1-docker.pkg.dev --quiet
echo "==> Creating config directory"
# Owned by root, readable only by root. Docker reads --env-file as root
# before dropping to uid 1000 (ironclaw) inside the container.
mkdir -p /opt/ironclaw
chmod 700 /opt/ironclaw
if [ ! -f /opt/ironclaw/.env ]; then
echo "WARNING: /opt/ironclaw/.env does not exist."
echo "Create it with your configuration before starting IronClaw."
echo "See deploy/env.example for the required variables."
echo ""
echo "Then run: systemctl enable ironclaw && systemctl start ironclaw"
else
chmod 600 /opt/ironclaw/.env
echo "==> Starting IronClaw"
systemctl enable ironclaw
systemctl start ironclaw
fi
echo "==> Setup complete"
echo ""
echo "Verify with:"
echo " systemctl status cloud-sql-proxy"
echo " systemctl status ironclaw"
echo " docker logs ironclaw"
+1
View File
@@ -81,6 +81,7 @@ async fn main() -> anyhow::Result<()> {
let session = create_session_manager(SessionConfig {
auth_base_url: config.llm.nearai.auth_base_url.clone(),
session_path: config.llm.nearai.session_path.clone(),
..Default::default()
})
.await;
let llm = create_llm_provider(&config.llm, session)?;
+165 -348
View File
@@ -22,14 +22,13 @@ use crate::context::JobContext;
use crate::db::Database;
use crate::error::Error;
use crate::extensions::ExtensionManager;
use crate::hooks::HookRegistry;
use crate::llm::{ChatMessage, LlmProvider, Reasoning, ReasoningContext, RespondResult};
use crate::safety::SafetyLayer;
use crate::tools::ToolRegistry;
use crate::workspace::Workspace;
/// Collapse a tool output string into a single-line preview for display.
pub(crate) fn truncate_for_preview(output: &str, max_chars: usize) -> String {
fn truncate_for_preview(output: &str, max_chars: usize) -> String {
let collapsed: String = output
.chars()
.take(max_chars + 50)
@@ -38,14 +37,8 @@ pub(crate) fn truncate_for_preview(output: &str, max_chars: usize) -> String {
.split_whitespace()
.collect::<Vec<_>>()
.join(" ");
// char_indices gives us byte offsets at char boundaries, so the slice is always valid UTF-8.
if collapsed.chars().count() > max_chars {
let byte_offset = collapsed
.char_indices()
.nth(max_chars)
.map(|(i, _)| i)
.unwrap_or(collapsed.len());
format!("{}...", &collapsed[..byte_offset])
if collapsed.len() > max_chars {
format!("{}...", &collapsed[..max_chars])
} else {
collapsed
}
@@ -68,14 +61,10 @@ enum AgenticLoopResult {
pub struct AgentDeps {
pub store: Option<Arc<dyn Database>>,
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 tools: Arc<ToolRegistry>,
pub workspace: Option<Arc<Workspace>>,
pub extension_manager: Option<Arc<ExtensionManager>>,
pub hooks: Arc<HookRegistry>,
}
/// The main agent that coordinates all components.
@@ -118,7 +107,6 @@ impl Agent {
deps.safety.clone(),
deps.tools.clone(),
deps.store.clone(),
deps.hooks.clone(),
));
Self {
@@ -144,11 +132,6 @@ impl Agent {
&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> {
&self.deps.safety
}
@@ -161,10 +144,6 @@ impl Agent {
self.deps.workspace.as_ref()
}
fn hooks(&self) -> &Arc<HookRegistry> {
&self.deps.hooks
}
/// Run the agent main loop.
pub async fn run(self) -> Result<(), Error> {
// Start channels
@@ -316,7 +295,7 @@ impl Agent {
Some(spawn_heartbeat(
config,
workspace.clone(),
self.cheap_llm().clone(),
self.llm().clone(),
Some(notify_tx),
))
} else {
@@ -432,32 +411,10 @@ impl Agent {
match self.handle_message(&message).await {
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
.channels
.respond(&message, OutgoingResponse::text(response))
.await;
}
}
let _ = self
.channels
.respond(&message, OutgoingResponse::text(response))
.await;
}
Ok(Some(_)) => {
// Empty response, nothing to send (e.g. approval handled via send_status)
@@ -503,33 +460,7 @@ impl Agent {
async fn handle_message(&self, message: &IncomingMessage) -> Result<Option<String>, Error> {
// Parse submission type first
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
}
}
let submission = SubmissionParser::parse(&message.content);
// Hydrate thread from DB if it's a historical thread not in memory
if let Some(ref external_thread_id) = message.thread_id {
@@ -723,17 +654,19 @@ impl Agent {
}
// Restore response chain from conversation metadata
if let Some(store) = self.store()
&& let Ok(Some(metadata)) = store.get_conversation_metadata(thread_uuid).await
&& let Some(rid) = metadata
.get("last_response_id")
.and_then(|v| v.as_str())
.map(String::from)
{
thread.last_response_id = Some(rid.clone());
self.llm()
.seed_response_chain(&thread_uuid.to_string(), rid);
tracing::debug!("Restored response chain for thread {}", thread_uuid);
if let Some(store) = self.store() {
if let Ok(Some(metadata)) = store.get_conversation_metadata(thread_uuid).await {
if let Some(rid) = metadata
.get("last_response_id")
.and_then(|v| v.as_str())
.map(String::from)
{
thread.last_response_id = Some(rid.clone());
self.llm()
.seed_response_chain(&thread_uuid.to_string(), rid);
tracing::debug!("Restored response chain for thread {}", thread_uuid);
}
}
}
// Insert into session and register with session manager
@@ -938,27 +871,6 @@ impl Agent {
// Complete, fail, or request approval
match result {
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);
self.persist_response_chain(thread);
let _ = self
@@ -1042,12 +954,13 @@ impl Agent {
return;
}
if let Some(ref resp) = response
&& let Err(e) = store
if let Some(ref resp) = response {
if let Err(e) = store
.add_conversation_message(thread_id, "assistant", resp)
.await
{
tracing::warn!("Failed to persist assistant message: {}", e);
{
tracing::warn!("Failed to persist assistant message: {}", e);
}
}
});
}
@@ -1145,14 +1058,14 @@ impl Agent {
// Check if interrupted
{
let sess = session.lock().await;
if let Some(thread) = sess.threads.get(&thread_id)
&& thread.state == ThreadState::Interrupted
{
return Err(crate::error::JobError::ContextError {
id: thread_id,
reason: "Interrupted".to_string(),
if let Some(thread) = sess.threads.get(&thread_id) {
if thread.state == ThreadState::Interrupted {
return Err(crate::error::JobError::ContextError {
id: thread_id,
reason: "Interrupted".to_string(),
}
.into());
}
.into());
}
}
@@ -1227,90 +1140,66 @@ impl Agent {
// Record tool calls in the thread
{
let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&thread_id)
&& let Some(turn) = thread.last_turn_mut()
{
for tc in &tool_calls {
turn.record_tool_call(&tc.name, tc.arguments.clone());
if let Some(thread) = sess.threads.get_mut(&thread_id) {
if let Some(turn) = thread.last_turn_mut() {
for tc in &tool_calls {
turn.record_tool_call(&tc.name, tc.arguments.clone());
}
}
}
}
// Execute each tool (with approval checking and hook interception)
for mut tc in tool_calls {
// Execute each tool (with approval checking)
for tc in tool_calls {
// Check if tool requires approval
if let Some(tool) = self.tools().get(&tc.name).await
&& tool.requires_approval()
{
// Check if auto-approved for this session
let mut is_auto_approved = {
let sess = session.lock().await;
sess.is_tool_auto_approved(&tc.name)
};
// Let the tool inspect the specific parameters and
// override auto-approval (e.g. destructive shell commands).
if is_auto_approved && tool.requires_approval_for(&tc.arguments) {
tracing::info!(
tool = %tc.name,
"Tool requires explicit approval for these parameters despite auto-approve"
);
is_auto_approved = false;
}
if !is_auto_approved {
// Need approval - store pending request and return
let pending = PendingApproval {
request_id: Uuid::new_v4(),
tool_name: tc.name.clone(),
parameters: tc.arguments.clone(),
description: tool.description().to_string(),
tool_call_id: tc.id.clone(),
context_messages: context_messages.clone(),
if let Some(tool) = self.tools().get(&tc.name).await {
if tool.requires_approval() {
// Check if auto-approved for this session
let mut is_auto_approved = {
let sess = session.lock().await;
sess.is_tool_auto_approved(&tc.name)
};
return Ok(AgenticLoopResult::NeedApproval { pending });
}
}
// 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
);
// For shell commands, override auto-approval for
// destructive patterns that should always require
// explicit per-invocation approval.
if is_auto_approved && tc.name == "shell" {
if let Some(cmd) = 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))
})
{
if crate::tools::builtin::shell::requires_explicit_approval(
&cmd,
) {
tracing::info!(
"Shell command '{}' requires explicit approval despite auto-approve",
cmd.chars().take(80).collect::<String>()
);
is_auto_approved = false;
}
}
},
_ => {} // Continue, fail-open errors already logged
}
if !is_auto_approved {
// Need approval - store pending request and return
let pending = PendingApproval {
request_id: Uuid::new_v4(),
tool_name: tc.name.clone(),
parameters: tc.arguments.clone(),
description: tool.description().to_string(),
tool_call_id: tc.id.clone(),
context_messages: context_messages.clone(),
};
return Ok(AgenticLoopResult::NeedApproval { pending });
}
}
}
@@ -1341,34 +1230,34 @@ impl Agent {
)
.await;
if let Ok(ref output) = tool_result
&& !output.is_empty()
{
let _ = self
.channels
.send_status(
&message.channel,
StatusUpdate::ToolResult {
name: tc.name.clone(),
preview: output.clone(),
},
&message.metadata,
)
.await;
if let Ok(ref output) = tool_result {
if !output.is_empty() {
let _ = self
.channels
.send_status(
&message.channel,
StatusUpdate::ToolResult {
name: tc.name.clone(),
preview: truncate_for_preview(output, 200),
},
&message.metadata,
)
.await;
}
}
// Record result in thread
{
let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&thread_id)
&& let Some(turn) = thread.last_turn_mut()
{
match &tool_result {
Ok(output) => {
turn.record_tool_result(serde_json::json!(output));
}
Err(e) => {
turn.record_tool_error(e.to_string());
if let Some(thread) = sess.threads.get_mut(&thread_id) {
if let Some(turn) = thread.last_turn_mut() {
match &tool_result {
Ok(output) => {
turn.record_tool_result(serde_json::json!(output));
}
Err(e) => {
turn.record_tool_error(e.to_string());
}
}
}
}
@@ -1579,9 +1468,6 @@ impl Agent {
session: Arc<Mutex<Session>>,
thread_id: Uuid,
) -> 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 mut mgr = undo_mgr.lock().await;
@@ -1589,6 +1475,7 @@ impl Agent {
return Ok(SubmissionResult::ok_with_message("Nothing to undo."));
}
let mut sess = session.lock().await;
let thread = sess
.threads
.get_mut(&thread_id)
@@ -1599,10 +1486,12 @@ impl Agent {
let current_turn = thread.turn_number();
if let Some(checkpoint) = mgr.undo(current_turn, current_messages) {
// Extract values before consuming the reference
let turn_number = checkpoint.turn_number;
let messages = checkpoint.messages.clone();
let undo_count = mgr.undo_count();
// Restore thread from checkpoint
thread.restore_from_messages(checkpoint.messages);
thread.restore_from_messages(messages);
Ok(SubmissionResult::ok_with_message(format!(
"Undone to turn {}. {} undo(s) remaining.",
turn_number, undo_count
@@ -1617,9 +1506,6 @@ impl Agent {
session: Arc<Mutex<Session>>,
thread_id: Uuid,
) -> 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 mut mgr = undo_mgr.lock().await;
@@ -1627,15 +1513,12 @@ impl Agent {
return Ok(SubmissionResult::ok_with_message("Nothing to redo."));
}
let thread = sess
.threads
.get_mut(&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) {
if let Some(checkpoint) = mgr.redo() {
let mut sess = session.lock().await;
let thread = sess
.threads
.get_mut(&thread_id)
.ok_or_else(|| Error::from(crate::error::JobError::NotFound { id: thread_id }))?;
thread.restore_from_messages(checkpoint.messages);
Ok(SubmissionResult::ok_with_message(format!(
"Redone to turn {}.",
@@ -1757,17 +1640,17 @@ impl Agent {
};
// Verify request ID if provided
if let Some(req_id) = request_id
&& req_id != pending.request_id
{
// Put it back and return error
let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&thread_id) {
thread.await_approval(pending);
if let Some(req_id) = request_id {
if req_id != pending.request_id {
// Put it back and return error
let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&thread_id) {
thread.await_approval(pending);
}
return Ok(SubmissionResult::error(
"Request ID mismatch. Use the correct request ID.",
));
}
return Ok(SubmissionResult::error(
"Request ID mismatch. Use the correct request ID.",
));
}
if approved {
@@ -1821,20 +1704,20 @@ impl Agent {
)
.await;
if let Ok(ref output) = tool_result
&& !output.is_empty()
{
let _ = self
.channels
.send_status(
&message.channel,
StatusUpdate::ToolResult {
name: pending.tool_name.clone(),
preview: output.clone(),
},
&message.metadata,
)
.await;
if let Ok(ref output) = tool_result {
if !output.is_empty() {
let _ = self
.channels
.send_status(
&message.channel,
StatusUpdate::ToolResult {
name: pending.tool_name.clone(),
preview: truncate_for_preview(output, 200),
},
&message.metadata,
)
.await;
}
}
// Build context including the tool result
@@ -1843,15 +1726,15 @@ impl Agent {
// Record result in thread
{
let mut sess = session.lock().await;
if let Some(thread) = sess.threads.get_mut(&thread_id)
&& let Some(turn) = thread.last_turn_mut()
{
match &tool_result {
Ok(output) => {
turn.record_tool_result(serde_json::json!(output));
}
Err(e) => {
turn.record_tool_error(e.to_string());
if let Some(thread) = sess.threads.get_mut(&thread_id) {
if let Some(turn) = thread.last_turn_mut() {
match &tool_result {
Ok(output) => {
turn.record_tool_result(serde_json::json!(output));
}
Err(e) => {
turn.record_tool_error(e.to_string());
}
}
}
}
@@ -2211,15 +2094,15 @@ impl Agent {
}
// Persist new job to database (fire-and-forget)
if let Some(store) = self.store()
&& let Ok(ctx) = self.context_manager.get_context(job_id).await
{
let store = store.clone();
tokio::spawn(async move {
if let Err(e) = store.save_job(&ctx).await {
tracing::warn!("Failed to persist new job {}: {}", job_id, e);
}
});
if let Some(store) = self.store() {
if let Ok(ctx) = self.context_manager.get_context(job_id).await {
let store = store.clone();
tokio::spawn(async move {
if let Err(e) = store.save_job(&ctx).await {
tracing::warn!("Failed to persist new job {}: {}", job_id, e);
}
});
}
}
// Schedule for execution
@@ -2299,10 +2182,10 @@ impl Agent {
let mut output = String::from("Jobs:\n");
for job_id in jobs {
if let Ok(ctx) = self.context_manager.get_context(job_id).await
&& ctx.user_id == user_id
{
output.push_str(&format!(" {} - {} ({:?})\n", job_id, ctx.title, ctx.state));
if let Ok(ctx) = self.context_manager.get_context(job_id).await {
if ctx.user_id == user_id {
output.push_str(&format!(" {} - {} ({:?})\n", job_id, ctx.title, ctx.state));
}
}
}
@@ -2753,70 +2636,4 @@ mod tests {
assert!(detect_auth_awaiting("tool_activate", &result).is_none());
}
// --- truncate_for_preview tests ---
use super::truncate_for_preview;
#[test]
fn test_truncate_short_input() {
assert_eq!(truncate_for_preview("hello", 10), "hello");
}
#[test]
fn test_truncate_empty_input() {
assert_eq!(truncate_for_preview("", 10), "");
}
#[test]
fn test_truncate_exact_length() {
assert_eq!(truncate_for_preview("hello", 5), "hello");
}
#[test]
fn test_truncate_over_limit() {
let result = truncate_for_preview("hello world, this is long", 10);
assert!(result.ends_with("..."));
// "hello worl" = 10 chars + "..."
assert_eq!(result, "hello worl...");
}
#[test]
fn test_truncate_collapses_newlines() {
let result = truncate_for_preview("line1\nline2\nline3", 100);
assert!(!result.contains('\n'));
assert_eq!(result, "line1 line2 line3");
}
#[test]
fn test_truncate_collapses_whitespace() {
let result = truncate_for_preview("hello world", 100);
assert_eq!(result, "hello world");
}
#[test]
fn test_truncate_multibyte_utf8() {
// Each emoji is 4 bytes. Truncating at char boundary must not panic.
let input = "😀😁😂🤣😃😄😅😆😉😊";
let result = truncate_for_preview(input, 5);
assert!(result.ends_with("..."));
// First 5 chars = 5 emoji
assert_eq!(result, "😀😁😂🤣😃...");
}
#[test]
fn test_truncate_cjk_characters() {
// CJK chars are 3 bytes each in UTF-8.
let input = "你好世界测试数据很长的字符串";
let result = truncate_for_preview(input, 4);
assert_eq!(result, "你好世界...");
}
#[test]
fn test_truncate_mixed_multibyte_and_ascii() {
let input = "hello 世界 foo";
let result = truncate_for_preview(input, 8);
// 'h','e','l','l','o',' ','世','界' = 8 chars
assert_eq!(result, "hello 世界...");
}
}
-1
View File
@@ -26,7 +26,6 @@ pub mod task;
pub mod undo;
pub mod worker;
pub(crate) use agent_loop::truncate_for_preview;
pub use agent_loop::{Agent, AgentDeps};
pub use compaction::{CompactionResult, ContextCompactor};
pub use context_monitor::{CompactionStrategy, ContextBreakdown, ContextMonitor};
+3 -2
View File
@@ -103,9 +103,10 @@ impl RoutineEngine {
if let Trigger::Event {
channel: Some(ch), ..
} = &routine.trigger
&& ch != &message.channel
{
continue;
if ch != &message.channel {
continue;
}
}
// Regex match
-5
View File
@@ -14,7 +14,6 @@ use crate::config::AgentConfig;
use crate::context::{ContextManager, JobContext, JobState};
use crate::db::Database;
use crate::error::{Error, JobError};
use crate::hooks::HookRegistry;
use crate::llm::LlmProvider;
use crate::safety::SafetyLayer;
use crate::tools::ToolRegistry;
@@ -50,7 +49,6 @@ pub struct Scheduler {
safety: Arc<SafetyLayer>,
tools: Arc<ToolRegistry>,
store: Option<Arc<dyn Database>>,
hooks: Arc<HookRegistry>,
/// Running jobs (main LLM-driven jobs).
jobs: Arc<RwLock<HashMap<Uuid, ScheduledJob>>>,
/// Running sub-tasks (tool executions, background tasks).
@@ -66,7 +64,6 @@ impl Scheduler {
safety: Arc<SafetyLayer>,
tools: Arc<ToolRegistry>,
store: Option<Arc<dyn Database>>,
hooks: Arc<HookRegistry>,
) -> Self {
Self {
config,
@@ -75,7 +72,6 @@ impl Scheduler {
safety,
tools,
store,
hooks,
jobs: Arc::new(RwLock::new(HashMap::new())),
subtasks: Arc::new(RwLock::new(HashMap::new())),
}
@@ -122,7 +118,6 @@ impl Scheduler {
safety: self.safety.clone(),
tools: self.tools.clone(),
store: self.store.clone(),
hooks: self.hooks.clone(),
timeout: self.config.job_timeout,
use_planning: self.config.use_planning,
};
+18 -18
View File
@@ -119,25 +119,25 @@ impl SelfRepair for DefaultSelfRepair {
let mut stuck_jobs = Vec::new();
for job_id in stuck_ids {
if let Ok(ctx) = self.context_manager.get_context(job_id).await
&& ctx.state == JobState::Stuck
{
let stuck_duration = ctx
.started_at
.map(|start| {
let now = Utc::now();
let duration = now.signed_duration_since(start);
Duration::from_secs(duration.num_seconds().max(0) as u64)
})
.unwrap_or_default();
if let Ok(ctx) = self.context_manager.get_context(job_id).await {
if ctx.state == JobState::Stuck {
let stuck_duration = ctx
.started_at
.map(|start| {
let now = Utc::now();
let duration = now.signed_duration_since(start);
Duration::from_secs(duration.num_seconds().max(0) as u64)
})
.unwrap_or_default();
stuck_jobs.push(StuckJob {
job_id,
last_activity: ctx.started_at.unwrap_or(ctx.created_at),
stuck_duration,
last_error: None,
repair_attempts: ctx.repair_attempts,
});
stuck_jobs.push(StuckJob {
job_id,
last_activity: ctx.started_at.unwrap_or(ctx.created_at),
stuck_duration,
last_error: None,
repair_attempts: ctx.repair_attempts,
});
}
}
}
+5 -5
View File
@@ -346,11 +346,11 @@ impl Thread {
let mut turn = Turn::new(turn_number, &msg.content);
// Check if next is assistant response
if let Some(next) = iter.peek()
&& next.role == crate::llm::Role::Assistant
{
let response = iter.next().expect("peeked");
turn.complete(&response.content);
if let Some(next) = iter.peek() {
if next.role == crate::llm::Role::Assistant {
let response = iter.next().expect("peeked");
turn.complete(&response.content);
}
}
self.turns.push(turn);
+8 -61
View File
@@ -11,7 +11,6 @@ use uuid::Uuid;
use crate::agent::session::Session;
use crate::agent::undo::UndoManager;
use crate::hooks::HookRegistry;
/// Key for mapping external thread IDs to internal ones.
#[derive(Clone, Hash, Eq, PartialEq)]
@@ -26,7 +25,6 @@ pub struct SessionManager {
sessions: RwLock<HashMap<String, Arc<Mutex<Session>>>>,
thread_map: RwLock<HashMap<ThreadKey, Uuid>>,
undo_managers: RwLock<HashMap<Uuid, Arc<Mutex<UndoManager>>>>,
hooks: Option<Arc<HookRegistry>>,
}
impl SessionManager {
@@ -36,16 +34,9 @@ impl SessionManager {
sessions: RwLock::new(HashMap::new()),
thread_map: 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.
pub async fn get_or_create_session(&self, user_id: &str) -> Arc<Mutex<Session>> {
// Fast path: check if session exists
@@ -63,28 +54,8 @@ impl SessionManager {
return Arc::clone(session);
}
let new_session = Session::new(user_id);
let session_id = new_session.id.to_string();
let session = Arc::new(Mutex::new(new_session));
let session = Arc::new(Mutex::new(Session::new(user_id)));
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
}
@@ -202,8 +173,8 @@ impl SessionManager {
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);
// Find stale sessions (user_id + session_id)
let stale_sessions: Vec<(String, String)> = {
// Find stale session user_ids
let stale_users: Vec<String> = {
let sessions = self.sessions.read().await;
sessions
.iter()
@@ -211,7 +182,7 @@ impl SessionManager {
// Try to lock; skip if contended (someone is actively using it)
let sess = session.try_lock().ok()?;
if sess.last_active_at < cutoff {
Some((user_id.clone(), sess.id.to_string()))
Some(user_id.clone())
} else {
None
}
@@ -219,11 +190,6 @@ impl SessionManager {
.collect()
};
let stale_users: Vec<String> = stale_sessions
.iter()
.map(|(user_id, _)| user_id.clone())
.collect();
if stale_users.is_empty() {
return 0;
}
@@ -233,30 +199,11 @@ impl SessionManager {
{
let sessions = self.sessions.read().await;
for user_id in &stale_users {
if let Some(session) = sessions.get(user_id)
&& let Ok(sess) = session.try_lock()
{
stale_thread_ids.extend(sess.threads.keys());
}
}
}
// 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);
if let Some(session) = sessions.get(user_id) {
if let Ok(sess) = session.try_lock() {
stale_thread_ids.extend(sess.threads.keys());
}
});
}
}
}
+14 -13
View File
@@ -93,26 +93,27 @@ impl SubmissionParser {
// /thread <uuid> - switch thread
if let Some(rest) = lower.strip_prefix("/thread ") {
let rest = rest.trim();
if rest != "new"
&& let Ok(id) = Uuid::parse_str(rest)
{
return Submission::SwitchThread { thread_id: id };
if rest != "new" {
if let Ok(id) = Uuid::parse_str(rest) {
return Submission::SwitchThread { thread_id: id };
}
}
}
// /resume <uuid> - resume from checkpoint
if let Some(rest) = lower.strip_prefix("/resume ")
&& let Ok(id) = Uuid::parse_str(rest.trim())
{
return Submission::Resume { checkpoint_id: id };
if let Some(rest) = lower.strip_prefix("/resume ") {
if let Ok(id) = Uuid::parse_str(rest.trim()) {
return Submission::Resume { checkpoint_id: id };
}
}
// Try structured JSON approval (from web gateway's /api/chat/approval endpoint)
if trimmed.starts_with('{')
&& let Ok(submission) = serde_json::from_str::<Submission>(trimmed)
&& matches!(submission, Submission::ExecApproval { .. })
{
return submission;
if trimmed.starts_with('{') {
if let Ok(submission) = serde_json::from_str::<Submission>(trimmed) {
if matches!(submission, Submission::ExecApproval { .. }) {
return submission;
}
}
}
// Approval responses (simple yes/no/always for pending approvals)
+16 -136
View File
@@ -43,10 +43,6 @@ impl Checkpoint {
}
/// 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 {
/// Stack of past checkpoints (for undo).
undo_stack: VecDeque<Checkpoint>,
@@ -72,14 +68,6 @@ impl UndoManager {
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.
///
/// This clears the redo stack since we're creating a new history branch.
@@ -92,23 +80,24 @@ impl UndoManager {
// Clear redo stack (new branch of history)
self.redo_stack.clear();
// Create and push checkpoint
let checkpoint = Checkpoint::new(turn_number, messages, description);
self.push_undo(checkpoint);
self.undo_stack.push_back(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.
///
/// 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.
/// The current state should be saved to redo stack before calling this.
pub fn undo(
&mut self,
current_turn: usize,
current_messages: Vec<ChatMessage>,
) -> Option<Checkpoint> {
) -> Option<&Checkpoint> {
if self.undo_stack.is_empty() {
return None;
}
@@ -121,8 +110,9 @@ impl UndoManager {
);
self.redo_stack.push(current);
// Pop and return the most recent checkpoint
self.undo_stack.pop_back()
// Return the most recent checkpoint without removing it
// (we keep it so multiple undos can work)
self.undo_stack.back()
}
/// Pop the last checkpoint from the undo stack.
@@ -131,29 +121,7 @@ impl UndoManager {
}
/// Redo: restore a previously undone state.
///
/// 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);
pub fn redo(&mut self) -> Option<Checkpoint> {
self.redo_stack.pop()
}
@@ -246,16 +214,14 @@ mod tests {
assert!(manager.can_undo());
assert!(!manager.can_redo());
// Undo - returns owned Checkpoint now
// Undo
let current = vec![ChatMessage::user("Hello"), ChatMessage::assistant("Hi")];
let checkpoint = manager.undo(2, current);
assert!(checkpoint.is_some());
let checkpoint = checkpoint.unwrap();
assert_eq!(checkpoint.turn_number, 1);
assert!(manager.can_redo());
// Redo - now requires current state parameters
let restored = manager.redo(checkpoint.turn_number, checkpoint.messages);
// Redo
let restored = manager.redo();
assert!(restored.is_some());
}
@@ -283,90 +249,4 @@ mod tests {
assert!(restored.is_some());
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);
}
}
+50 -91
View File
@@ -12,7 +12,6 @@ use crate::agent::task::TaskOutput;
use crate::context::{ContextManager, JobState};
use crate::db::Database;
use crate::error::Error;
use crate::hooks::HookRegistry;
use crate::llm::{
ActionPlan, ChatMessage, LlmProvider, Reasoning, ReasoningContext, RespondResult, ToolSelection,
};
@@ -30,7 +29,6 @@ pub struct WorkerDeps {
pub safety: Arc<SafetyLayer>,
pub tools: Arc<ToolRegistry>,
pub store: Option<Arc<dyn Database>>,
pub hooks: Arc<HookRegistry>,
pub timeout: Duration,
pub use_planning: bool,
}
@@ -229,11 +227,11 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
}
// Check for cancellation
if let Ok(ctx) = self.context_manager().get_context(self.job_id).await
&& ctx.state == JobState::Cancelled
{
tracing::info!("Worker for job {} detected cancellation", self.job_id);
return Ok(());
if let Ok(ctx) = self.context_manager().get_context(self.job_id).await {
if ctx.state == JobState::Cancelled {
tracing::info!("Worker for job {} detected cancellation", self.job_id);
return Ok(());
}
}
iteration += 1;
@@ -301,7 +299,6 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
parameters: tc.arguments.clone(),
reasoning: String::new(),
alternatives: vec![],
tool_call_id: tc.id.clone(),
};
self.process_tool_result(reason_ctx, &selection, result)
@@ -354,11 +351,23 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
.map(|selection| {
let tool_name = selection.tool_name.clone();
let params = selection.parameters.clone();
let deps = self.deps.clone();
let tools = self.tools().clone();
let context_manager = self.context_manager().clone();
let safety = self.safety().clone();
let job_id = self.job_id;
let store = self.deps.store.clone();
async move {
let result = Self::execute_tool_inner(&deps, job_id, &tool_name, &params).await;
let result = Self::execute_tool_inner(
tools,
context_manager,
safety,
store,
job_id,
&tool_name,
&params,
)
.await;
ToolExecResult { result }
}
})
@@ -369,18 +378,20 @@ 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.
async fn execute_tool_inner(
deps: &WorkerDeps,
tools: Arc<ToolRegistry>,
context_manager: Arc<ContextManager>,
safety: Arc<SafetyLayer>,
store: Option<Arc<dyn Database>>,
job_id: Uuid,
tool_name: &str,
params: &serde_json::Value,
) -> Result<String, Error> {
let tool =
deps.tools
.get(tool_name)
.await
.ok_or_else(|| crate::error::ToolError::NotFound {
name: tool_name.to_string(),
})?;
let tool = tools
.get(tool_name)
.await
.ok_or_else(|| crate::error::ToolError::NotFound {
name: tool_name.to_string(),
})?;
// Tools requiring approval are blocked in autonomous jobs
if tool.requires_approval() {
@@ -390,46 +401,8 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
.into());
}
// Fetch job context early so we have the real user_id for hooks
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(),
}
};
// Get job context for the tool
let job_ctx = context_manager.get_context(job_id).await?;
if job_ctx.state == JobState::Cancelled {
return Err(crate::error::ToolError::ExecutionFailed {
name: tool_name.to_string(),
@@ -439,7 +412,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
}
// Validate tool parameters
let validation = deps.safety.validator().validate_tool_params(&params);
let validation = safety.validator().validate_tool_params(params);
if !validation.is_valid {
let details = validation
.errors
@@ -504,8 +477,8 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
Ok(Ok(output)) => {
let output_str = serde_json::to_string_pretty(&output.result)
.ok()
.map(|s| deps.safety.sanitize_tool_output(tool_name, &s).content);
deps.context_manager
.map(|s| safety.sanitize_tool_output(tool_name, &s).content);
context_manager
.update_memory(job_id, |mem| {
let rec = mem.create_action(tool_name, params.clone()).succeed(
output_str.clone(),
@@ -518,8 +491,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
.await
.ok()
}
Ok(Err(e)) => deps
.context_manager
Ok(Err(e)) => context_manager
.update_memory(job_id, |mem| {
let rec = mem
.create_action(tool_name, params.clone())
@@ -529,8 +501,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
})
.await
.ok(),
Err(_) => deps
.context_manager
Err(_) => context_manager
.update_memory(job_id, |mem| {
let rec = mem
.create_action(tool_name, params.clone())
@@ -543,7 +514,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
};
// Persist action to database (fire-and-forget)
if let (Some(action), Some(store)) = (action, deps.store.clone()) {
if let (Some(action), Some(store)) = (action, store) {
tokio::spawn(async move {
if let Err(e) = store.save_action(job_id, &action).await {
tracing::warn!("Failed to persist action for job {}: {}", job_id, e);
@@ -594,7 +565,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
);
reason_ctx.messages.push(ChatMessage::tool_result(
&selection.tool_call_id,
"tool_call_id",
&selection.tool_name,
wrapped,
));
@@ -626,7 +597,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
}
reason_ctx.messages.push(ChatMessage::tool_result(
&selection.tool_call_id,
"tool_call_id",
&selection.tool_name,
format!("Error: {}", e),
));
@@ -676,15 +647,12 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
.execute_tool(&action.tool_name, &action.parameters)
.await;
// Create a synthetic ToolSelection for process_tool_result.
// Plan actions don't originate from an LLM tool_call response so
// there is no real tool_call_id; generate a unique one.
// Create a synthetic ToolSelection for process_tool_result
let selection = ToolSelection {
tool_name: action.tool_name.clone(),
parameters: action.parameters.clone(),
reasoning: action.reasoning.clone(),
alternatives: vec![],
tool_call_id: format!("plan_{}_{}", self.job_id, i),
};
// Process the result
@@ -729,7 +697,16 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
tool_name: &str,
params: &serde_json::Value,
) -> Result<String, Error> {
Self::execute_tool_inner(&self.deps, self.job_id, tool_name, params).await
Self::execute_tool_inner(
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> {
@@ -797,26 +774,8 @@ impl From<TaskOutput> for Result<String, Error> {
#[cfg(test)]
mod tests {
use crate::llm::ToolSelection;
use crate::util::llm_signals_completion;
#[test]
fn test_tool_selection_preserves_call_id() {
let selection = ToolSelection {
tool_name: "memory_search".to_string(),
parameters: serde_json::json!({"query": "test"}),
reasoning: "Need to search memory".to_string(),
alternatives: vec![],
tool_call_id: "call_abc123".to_string(),
};
assert_eq!(selection.tool_call_id, "call_abc123");
assert_ne!(
selection.tool_call_id, "tool_call_id",
"tool_call_id must not be the hardcoded placeholder string"
);
}
#[test]
fn test_completion_positive_signals() {
assert!(llm_signals_completion("The job is complete."));
-208
View File
@@ -1,208 +0,0 @@
//! 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);
}
}
+172 -311
View File
@@ -1,145 +1,147 @@
//! Bootstrap helpers for IronClaw.
//! Bootstrap configuration for IronClaw.
//!
//! The only setting that truly needs disk persistence before the database is
//! available is `DATABASE_URL` (chicken-and-egg: can't connect to DB without
//! it). Everything else is auto-detected or read from env vars.
//! These are the only settings that MUST live on disk because they're needed
//! before the database connection is established. Everything else lives in the
//! `settings` table in PostgreSQL.
//!
//! File: `~/.ironclaw/.env` (standard dotenvy format)
//! File: `~/.ironclaw/bootstrap.json`
use std::path::PathBuf;
/// Path to the IronClaw-specific `.env` file: `~/.ironclaw/.env`.
pub fn ironclaw_env_path() -> PathBuf {
dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".ironclaw")
.join(".env")
use serde::{Deserialize, Serialize};
use crate::settings::KeySource;
/// Minimal config needed to connect to the database and decrypt secrets.
///
/// This is the only JSON file IronClaw reads from disk at startup.
/// All other configuration lives in the `settings` table in PostgreSQL.
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BootstrapConfig {
/// Database connection URL (postgres://...).
#[serde(default)]
pub database_url: Option<String>,
/// Database connection pool size.
#[serde(default)]
pub database_pool_size: Option<usize>,
/// Source for the secrets master key.
#[serde(default)]
pub secrets_master_key_source: KeySource,
/// Whether onboarding wizard has been completed.
#[serde(default)]
pub onboard_completed: bool,
}
/// Load env vars from `~/.ironclaw/.env` (in addition to the standard `.env`).
///
/// Call this **after** `dotenvy::dotenv()` so that the standard `./.env`
/// takes priority over `~/.ironclaw/.env`. dotenvy never overwrites
/// existing env vars, so the effective priority is:
///
/// explicit env vars > `./.env` > `~/.ironclaw/.env`
///
/// If `~/.ironclaw/.env` doesn't exist but the legacy `bootstrap.json` does,
/// extracts `DATABASE_URL` from it and writes the `.env` file (one-time
/// upgrade from the old config format).
pub fn load_ironclaw_env() {
let path = ironclaw_env_path();
if !path.exists() {
// One-time upgrade: extract DATABASE_URL from legacy bootstrap.json
migrate_bootstrap_json_to_env(&path);
}
if path.exists() {
let _ = dotenvy::from_path(&path);
}
}
/// If `bootstrap.json` exists, pull `database_url` out of it and write `.env`.
fn migrate_bootstrap_json_to_env(env_path: &std::path::Path) {
let ironclaw_dir = env_path
.parent()
.unwrap_or_else(|| std::path::Path::new("."));
let bootstrap_path = ironclaw_dir.join("bootstrap.json");
if !bootstrap_path.exists() {
return;
}
let content = match std::fs::read_to_string(&bootstrap_path) {
Ok(c) => c,
Err(_) => return,
};
// Minimal parse: just grab database_url from the JSON
let parsed: serde_json::Value = match serde_json::from_str(&content) {
Ok(v) => v,
Err(_) => return,
};
if let Some(url) = parsed.get("database_url").and_then(|v| v.as_str()) {
if let Some(parent) = env_path.parent()
&& let Err(e) = std::fs::create_dir_all(parent)
{
eprintln!("Warning: failed to create {}: {}", parent.display(), e);
return;
impl Default for BootstrapConfig {
fn default() -> Self {
Self {
database_url: None,
database_pool_size: None,
secrets_master_key_source: KeySource::None,
onboard_completed: false,
}
if let Err(e) = std::fs::write(env_path, format!("DATABASE_URL=\"{}\"\n", url)) {
eprintln!("Warning: failed to migrate bootstrap.json to .env: {}", e);
return;
}
}
impl BootstrapConfig {
/// Default bootstrap file path: `~/.ironclaw/bootstrap.json`.
pub fn default_path() -> PathBuf {
dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".ironclaw")
.join("bootstrap.json")
}
/// Legacy settings.json path (for migration detection).
pub fn legacy_settings_path() -> PathBuf {
dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".ironclaw")
.join("settings.json")
}
/// Load from the default path, falling back to legacy settings.json,
/// then to defaults if neither exists.
pub fn load() -> Self {
let bootstrap_path = Self::default_path();
if bootstrap_path.exists() {
return Self::load_from(&bootstrap_path);
}
rename_to_migrated(&bootstrap_path);
eprintln!(
"Migrated DATABASE_URL from bootstrap.json to {}",
env_path.display()
);
// Fall back to legacy settings.json (extract just the 4 bootstrap fields)
let legacy_path = Self::legacy_settings_path();
if legacy_path.exists() {
return Self::load_from_legacy(&legacy_path);
}
Self::default()
}
/// Load from a specific path.
pub fn load_from(path: &PathBuf) -> Self {
match std::fs::read_to_string(path) {
Ok(data) => serde_json::from_str(&data).unwrap_or_default(),
Err(_) => Self::default(),
}
}
/// Extract bootstrap fields from a legacy settings.json.
fn load_from_legacy(path: &PathBuf) -> Self {
match std::fs::read_to_string(path) {
Ok(data) => {
// The legacy Settings struct is a superset; serde will ignore extra fields.
serde_json::from_str(&data).unwrap_or_default()
}
Err(_) => Self::default(),
}
}
/// Save to the default path.
pub fn save(&self) -> std::io::Result<()> {
self.save_to(&Self::default_path())
}
/// Save to a specific path.
pub fn save_to(&self, path: &PathBuf) -> std::io::Result<()> {
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
let json = serde_json::to_string_pretty(self)
.map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e.to_string()))?;
std::fs::write(path, json)
}
}
/// Write database bootstrap vars to `~/.ironclaw/.env`.
/// One-time migration from disk config files to the database settings table.
///
/// 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.
/// On first boot after upgrade, checks if:
/// 1. `~/.ironclaw/settings.json` exists
/// 2. The DB settings table is empty for this user
///
/// Creates the parent directory if it doesn't exist.
/// Values are double-quoted so that `#` (common in URL-encoded passwords)
/// and other shell-special characters are preserved by dotenvy.
pub fn save_bootstrap_env(vars: &[(&str, &str)]) -> std::io::Result<()> {
let path = ironclaw_env_path();
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
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.
///
/// Only runs when a `settings.json` exists on disk AND the DB has no settings
/// yet. After the wizard writes directly to the DB, this path is only hit by
/// users upgrading from the old disk-only configuration.
///
/// After syncing, renames `settings.json` to `.migrated` so it won't trigger again.
/// If both conditions hold, migrates settings, MCP servers, and session data
/// to the database, writes `bootstrap.json`, and renames old files to `.migrated`.
pub async fn migrate_disk_to_db(
store: &dyn crate::db::Database,
user_id: &str,
) -> Result<(), MigrationError> {
let ironclaw_dir = dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".ironclaw");
let legacy_settings_path = ironclaw_dir.join("settings.json");
let legacy_settings_path = BootstrapConfig::legacy_settings_path();
if !legacy_settings_path.exists() {
tracing::debug!("No legacy settings.json found, skipping disk-to-DB migration");
return Ok(());
}
// If DB already has settings, this is not a first boot, the wizard already
// wrote directly to the DB. Just clean up the stale file.
// Only migrate if DB is empty for this user
let has_settings = store.has_settings(user_id).await.map_err(|e| {
MigrationError::Database(format!("Failed to check existing settings: {}", e))
})?;
if has_settings {
tracing::info!("DB already has settings, renaming stale settings.json");
rename_to_migrated(&legacy_settings_path);
tracing::debug!(
"DB already has settings for user '{}', skipping migration",
user_id
);
return Ok(());
}
@@ -158,14 +160,22 @@ pub async fn migrate_disk_to_db(
tracing::info!("Migrated {} settings to database", db_map.len());
}
// 2. Write DATABASE_URL to ~/.ironclaw/.env
if let Some(ref url) = settings.database_url {
save_database_url(url)
.map_err(|e| MigrationError::Io(format!("Failed to write .env: {}", e)))?;
tracing::info!("Wrote DATABASE_URL to {}", ironclaw_env_path().display());
}
// 2. Write bootstrap.json with the 4 essential fields
let bootstrap = BootstrapConfig {
database_url: settings.database_url.clone(),
database_pool_size: settings.database_pool_size,
secrets_master_key_source: settings.secrets_master_key_source,
onboard_completed: settings.onboard_completed,
};
bootstrap
.save()
.map_err(|e| MigrationError::Io(format!("Failed to write bootstrap.json: {}", e)))?;
tracing::info!("Wrote bootstrap.json");
// 3. Migrate mcp-servers.json if it exists
let ironclaw_dir = dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".ironclaw");
let mcp_path = ironclaw_dir.join("mcp-servers.json");
if mcp_path.exists() {
match std::fs::read_to_string(&mcp_path) {
@@ -201,7 +211,7 @@ pub async fn migrate_disk_to_db(
Ok(content) => match serde_json::from_str::<serde_json::Value>(&content) {
Ok(value) => {
store
.set_setting(user_id, "nearai.session_token", &value)
.set_setting(user_id, "nearai.session", &value)
.await
.map_err(|e| {
MigrationError::Database(format!(
@@ -226,19 +236,12 @@ pub async fn migrate_disk_to_db(
// 5. Rename settings.json to .migrated (don't delete, safety net)
rename_to_migrated(&legacy_settings_path);
// 6. Clean up old bootstrap.json if it exists (superseded by .env)
let old_bootstrap = ironclaw_dir.join("bootstrap.json");
if old_bootstrap.exists() {
rename_to_migrated(&old_bootstrap);
tracing::info!("Renamed old bootstrap.json to .migrated");
}
tracing::info!("Disk-to-DB migration complete");
Ok(())
}
/// Rename a file to `<name>.migrated` as a safety net.
fn rename_to_migrated(path: &std::path::Path) {
fn rename_to_migrated(path: &PathBuf) {
let mut migrated = path.as_os_str().to_owned();
migrated.push(".migrated");
if let Err(e) = std::fs::rename(path, &migrated) {
@@ -261,204 +264,62 @@ mod tests {
use tempfile::tempdir;
#[test]
fn test_save_and_load_database_url() {
fn test_bootstrap_save_load() {
let dir = tempdir().unwrap();
let env_path = dir.path().join(".env");
let path = dir.path().join("bootstrap.json");
// Write in the quoted format that save_database_url uses
let url = "postgres://localhost:5432/ironclaw_test";
std::fs::write(&env_path, format!("DATABASE_URL=\"{}\"\n", url)).unwrap();
let config = BootstrapConfig {
database_url: Some("postgres://localhost/test".to_string()),
database_pool_size: Some(5),
secrets_master_key_source: KeySource::Keychain,
onboard_completed: true,
};
// Verify the content is a valid dotenv line (quoted)
let content = std::fs::read_to_string(&env_path).unwrap();
config.save_to(&path).unwrap();
let loaded = BootstrapConfig::load_from(&path);
assert_eq!(
content,
"DATABASE_URL=\"postgres://localhost:5432/ironclaw_test\"\n"
loaded.database_url,
Some("postgres://localhost/test".to_string())
);
// Verify dotenvy can parse it (strips quotes automatically)
let parsed: Vec<(String, String)> = dotenvy::from_path_iter(&env_path)
.unwrap()
.filter_map(|r| r.ok())
.collect();
assert_eq!(parsed.len(), 1);
assert_eq!(parsed[0].0, "DATABASE_URL");
assert_eq!(parsed[0].1, url);
assert_eq!(loaded.database_pool_size, Some(5));
assert_eq!(loaded.secrets_master_key_source, KeySource::Keychain);
assert!(loaded.onboard_completed);
}
#[test]
fn test_save_database_url_with_hash_in_password() {
fn test_bootstrap_from_legacy_settings() {
let dir = tempdir().unwrap();
let env_path = dir.path().join(".env");
let path = dir.path().join("settings.json");
// URLs with # in the password are common (URL-encoded special chars).
// Without quoting, dotenvy treats # as a comment delimiter.
let url = "postgres://user:p%23ss@localhost:5432/ironclaw";
std::fs::write(&env_path, format!("DATABASE_URL=\"{}\"\n", url)).unwrap();
let parsed: Vec<(String, String)> = dotenvy::from_path_iter(&env_path)
.unwrap()
.filter_map(|r| r.ok())
.collect();
assert_eq!(parsed.len(), 1);
assert_eq!(parsed[0].0, "DATABASE_URL");
assert_eq!(parsed[0].1, url);
}
#[test]
fn test_save_database_url_creates_parent_dirs() {
let dir = tempdir().unwrap();
let nested = dir.path().join("deep").join("nested");
let env_path = nested.join(".env");
// Parent doesn't exist yet
assert!(!nested.exists());
// The global function uses a fixed path, so we test the logic directly
std::fs::create_dir_all(&nested).unwrap();
std::fs::write(&env_path, "DATABASE_URL=postgres://test\n").unwrap();
assert!(env_path.exists());
let content = std::fs::read_to_string(&env_path).unwrap();
assert!(content.contains("DATABASE_URL=postgres://test"));
}
#[test]
fn test_ironclaw_env_path() {
let path = ironclaw_env_path();
assert!(path.ends_with(".ironclaw/.env"));
}
#[test]
fn test_migrate_bootstrap_json_to_env() {
let dir = tempdir().unwrap();
let env_path = dir.path().join(".env");
let bootstrap_path = dir.path().join("bootstrap.json");
// Write a legacy bootstrap.json
let bootstrap_json = serde_json::json!({
"database_url": "postgres://localhost/ironclaw_upgrade",
"database_pool_size": 5,
// Write a legacy settings.json with many extra fields
let legacy = serde_json::json!({
"database_url": "postgres://localhost/ironclaw",
"database_pool_size": 10,
"secrets_master_key_source": "keychain",
"onboard_completed": true
"onboard_completed": true,
"selected_model": "claude-3-5-sonnet",
"agent": { "name": "testbot", "max_parallel_jobs": 3 },
"heartbeat": { "enabled": true }
});
std::fs::write(
&bootstrap_path,
serde_json::to_string_pretty(&bootstrap_json).unwrap(),
)
.unwrap();
std::fs::write(&path, serde_json::to_string_pretty(&legacy).unwrap()).unwrap();
assert!(!env_path.exists());
assert!(bootstrap_path.exists());
// Run the migration
migrate_bootstrap_json_to_env(&env_path);
// .env should now exist with DATABASE_URL
assert!(env_path.exists());
let content = std::fs::read_to_string(&env_path).unwrap();
let config = BootstrapConfig::load_from_legacy(&path);
assert_eq!(
content,
"DATABASE_URL=\"postgres://localhost/ironclaw_upgrade\"\n"
config.database_url,
Some("postgres://localhost/ironclaw".to_string())
);
// bootstrap.json should be renamed to .migrated
assert!(!bootstrap_path.exists());
assert!(dir.path().join("bootstrap.json.migrated").exists());
assert_eq!(config.database_pool_size, Some(10));
assert_eq!(config.secrets_master_key_source, KeySource::Keychain);
assert!(config.onboard_completed);
}
#[test]
fn test_migrate_bootstrap_json_no_database_url() {
let dir = tempdir().unwrap();
let env_path = dir.path().join(".env");
let bootstrap_path = dir.path().join("bootstrap.json");
// bootstrap.json with no database_url
let bootstrap_json = serde_json::json!({
"onboard_completed": false
});
std::fs::write(
&bootstrap_path,
serde_json::to_string_pretty(&bootstrap_json).unwrap(),
)
.unwrap();
migrate_bootstrap_json_to_env(&env_path);
// .env should NOT be created
assert!(!env_path.exists());
// bootstrap.json should remain (no migration happened)
assert!(bootstrap_path.exists());
}
#[test]
fn test_migrate_bootstrap_json_missing() {
let dir = tempdir().unwrap();
let env_path = dir.path().join(".env");
// No bootstrap.json at all
migrate_bootstrap_json_to_env(&env_path);
// Nothing should happen
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"));
fn test_bootstrap_defaults() {
let config = BootstrapConfig::default();
assert!(config.database_url.is_none());
assert!(config.database_pool_size.is_none());
assert_eq!(config.secrets_master_key_source, KeySource::None);
assert!(!config.onboard_completed);
}
}
+9 -31
View File
@@ -33,16 +33,9 @@ use termimad::MadSkin;
use tokio::sync::mpsc;
use tokio_stream::wrappers::ReceiverStream;
use crate::agent::truncate_for_preview;
use crate::channels::{Channel, IncomingMessage, MessageStream, OutgoingResponse, StatusUpdate};
use crate::error::ChannelError;
/// Max characters for tool result previews in the terminal.
const CLI_TOOL_RESULT_MAX: usize = 200;
/// Max characters for thinking/status messages in the terminal.
const CLI_STATUS_MAX: usize = 200;
/// Slash commands available in the REPL.
const SLASH_COMMANDS: &[&str] = &[
"/help",
@@ -184,8 +177,6 @@ pub struct ReplChannel {
debug_mode: Arc<AtomicBool>,
/// Whether we're currently streaming (chunks have been printed without a trailing newline).
is_streaming: Arc<AtomicBool>,
/// When true, the one-liner startup banner is suppressed (boot screen shown instead).
suppress_banner: Arc<AtomicBool>,
}
impl ReplChannel {
@@ -195,7 +186,6 @@ impl ReplChannel {
single_message: None,
debug_mode: Arc::new(AtomicBool::new(false)),
is_streaming: Arc::new(AtomicBool::new(false)),
suppress_banner: Arc::new(AtomicBool::new(false)),
}
}
@@ -205,15 +195,9 @@ impl ReplChannel {
single_message: Some(message),
debug_mode: 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 {
self.debug_mode.load(Ordering::Relaxed)
}
@@ -273,12 +257,11 @@ impl Channel for ReplChannel {
let (tx, rx) = mpsc::channel(32);
let single_message = self.single_message.clone();
let debug_mode = Arc::clone(&self.debug_mode);
let suppress_banner = Arc::clone(&self.suppress_banner);
std::thread::spawn(move || {
// Single message mode: send it and return
if let Some(msg) = single_message {
let incoming = IncomingMessage::new("repl", "default", &msg);
let incoming = IncomingMessage::new("repl", "user", &msg);
let _ = tx.blocking_send(incoming);
return;
}
@@ -308,10 +291,8 @@ impl Channel for ReplChannel {
}
let _ = rl.load_history(&hist_path);
if !suppress_banner.load(Ordering::Relaxed) {
println!("\x1b[1mIronClaw\x1b[0m /help for commands, /quit to exit");
println!();
}
println!("\x1b[1mIronClaw\x1b[0m /help for commands, /quit to exit");
println!();
loop {
let prompt = if debug_mode.load(Ordering::Relaxed) {
@@ -348,21 +329,21 @@ impl Channel for ReplChannel {
_ => {}
}
let msg = IncomingMessage::new("repl", "default", line);
let msg = IncomingMessage::new("repl", "user", line);
if tx.blocking_send(msg).is_err() {
break;
}
}
Err(ReadlineError::Interrupted) => {
// Ctrl+C: send /interrupt
let msg = IncomingMessage::new("repl", "default", "/interrupt");
let msg = IncomingMessage::new("repl", "user", "/interrupt");
if tx.blocking_send(msg).is_err() {
break;
}
}
Err(ReadlineError::Eof) => {
// Ctrl+D: send /quit so the agent loop runs graceful shutdown
let msg = IncomingMessage::new("repl", "default", "/quit");
let msg = IncomingMessage::new("repl", "user", "/quit");
let _ = tx.blocking_send(msg);
break;
}
@@ -419,8 +400,7 @@ impl Channel for ReplChannel {
match status {
StatusUpdate::Thinking(msg) => {
let display = truncate_for_preview(&msg, CLI_STATUS_MAX);
eprintln!(" \x1b[90m\u{25CB} {display}\x1b[0m");
eprintln!(" \x1b[90m\u{25CB} {msg}\x1b[0m");
}
StatusUpdate::ToolStarted { name } => {
eprintln!(" \x1b[33m\u{25CB} {name}\x1b[0m");
@@ -433,8 +413,7 @@ impl Channel for ReplChannel {
}
}
StatusUpdate::ToolResult { name: _, preview } => {
let display = truncate_for_preview(&preview, CLI_TOOL_RESULT_MAX);
eprintln!(" \x1b[90m{display}\x1b[0m");
eprintln!(" \x1b[90m{preview}\x1b[0m");
}
StatusUpdate::StreamChunk(chunk) => {
// Print separator on the false-to-true transition
@@ -459,8 +438,7 @@ impl Channel for ReplChannel {
}
StatusUpdate::Status(msg) => {
if debug || msg.contains("approval") || msg.contains("Approval") {
let display = truncate_for_preview(&msg, CLI_STATUS_MAX);
eprintln!(" \x1b[90m{display}\x1b[0m");
eprintln!(" \x1b[90m{msg}\x1b[0m");
}
}
StatusUpdate::ApprovalNeeded {
+43 -112
View File
@@ -76,9 +76,6 @@ struct ChannelStoreData {
credentials: HashMap<String, String>,
/// Pairing store for DM pairing (guest access control).
pairing_store: Arc<PairingStore>,
/// Dedicated tokio runtime for HTTP requests, lazily initialized.
/// Reused across multiple `http_request` calls within one execution.
http_runtime: Option<tokio::runtime::Runtime>,
}
impl ChannelStoreData {
@@ -99,7 +96,6 @@ impl ChannelStoreData {
table: ResourceTable::new(),
credentials,
pairing_store,
http_runtime: None,
}
}
@@ -138,13 +134,13 @@ impl ChannelStoreData {
if result.contains('{') && result.contains('}') {
// Only warn if it looks like an unresolved placeholder (not JSON braces)
let brace_pattern = regex::Regex::new(r"\{[A-Z_]+\}").ok();
if let Some(re) = brace_pattern
&& re.is_match(&result)
{
tracing::warn!(
context = %context,
"String may contain unresolved credential placeholders"
);
if let Some(re) = brace_pattern {
if re.is_match(&result) {
tracing::warn!(
context = %context,
"String may contain unresolved credential placeholders"
);
}
}
}
@@ -287,25 +283,10 @@ impl near::agent::channel_host::Host for ChannelStoreData {
.map(|h| h.max_response_bytes)
.unwrap_or(10 * 1024 * 1024);
// Make the HTTP request using a dedicated single-threaded runtime.
// We're inside spawn_blocking, so we can't rely on the main runtime's
// I/O driver (it may be busy with WASM compilation or other startup work).
// A dedicated runtime gives us our own I/O driver and avoids contention.
// The runtime is lazily created and reused across calls within one execution.
if self.http_runtime.is_none() {
self.http_runtime = Some(
tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.map_err(|e| format!("Failed to create HTTP runtime: {e}"))?,
);
}
let rt = self.http_runtime.as_ref().expect("just initialized");
let result = rt.block_on(async {
let client = reqwest::Client::builder()
.connect_timeout(std::time::Duration::from_secs(10))
.build()
.map_err(|e| format!("Failed to build HTTP client: {e}"))?;
// Make the HTTP request using blocking I/O
// We're already in a spawn_blocking context, so we can use block_on
let result = tokio::runtime::Handle::current().block_on(async {
let client = reqwest::Client::new();
let mut request = match method.to_uppercase().as_str() {
"GET" => client.get(&url),
@@ -327,9 +308,9 @@ impl near::agent::channel_host::Host for ChannelStoreData {
request = request.body(body_bytes);
}
// Send request with caller-specified timeout (default 30s, max 5min).
let timeout_ms = timeout_ms.unwrap_or(30_000).min(300_000) as u64;
let timeout = std::time::Duration::from_millis(timeout_ms);
// Send request with caller-specified timeout (default 30s).
// Cap at callback_timeout to prevent outliving the host wrapper.
let timeout = std::time::Duration::from_millis(timeout_ms.unwrap_or(30_000) as u64);
let response = request.timeout(timeout).send().await.map_err(|e| {
// Walk the full error chain so we get the actual root cause
// (DNS, TLS, connection refused, etc.) instead of just
@@ -357,13 +338,13 @@ impl near::agent::channel_host::Host for ChannelStoreData {
// Enforce max response body size to prevent memory exhaustion.
let max_response = max_response_bytes;
if let Some(cl) = response.content_length()
&& cl as usize > max_response
{
return Err(format!(
"Response body too large: {} bytes exceeds limit of {} bytes",
cl, max_response
));
if let Some(cl) = response.content_length() {
if cl as usize > max_response {
return Err(format!(
"Response body too large: {} bytes exceeds limit of {} bytes",
cl, max_response
));
}
}
let body = response
.bytes()
@@ -814,21 +795,7 @@ impl WasmChannel {
.await;
match result {
Ok(Ok((config, mut host_state))) => {
// Surface WASM guest logs (errors/warnings from webhook setup, etc.)
for entry in host_state.take_logs() {
match entry.level {
crate::tools::wasm::LogLevel::Error => {
tracing::error!(channel = %self.name, "{}", entry.message);
}
crate::tools::wasm::LogLevel::Warn => {
tracing::warn!(channel = %self.name, "{}", entry.message);
}
_ => {
tracing::debug!(channel = %self.name, "{}", entry.message);
}
}
}
Ok(Ok((config, _host_state))) => {
tracing::info!(
channel = %self.name,
display_name = %config.display_name,
@@ -1528,8 +1495,8 @@ impl WasmChannel {
match result {
Ok(emitted_messages) => {
// Process any emitted messages
if !emitted_messages.is_empty()
&& let Err(e) = Self::dispatch_emitted_messages(
if !emitted_messages.is_empty() {
if let Err(e) = Self::dispatch_emitted_messages(
&channel_name,
emitted_messages,
&message_tx,
@@ -1541,6 +1508,7 @@ impl WasmChannel {
"Failed to dispatch emitted messages from poll"
);
}
}
}
Err(e) => {
tracing::warn!(
@@ -1770,22 +1738,22 @@ impl Channel for WasmChannel {
*self.endpoints.write().await = endpoints;
// Start polling if configured
if let Some(poll_config) = &config.poll
&& poll_config.enabled
{
let interval = self
.capabilities
.validate_poll_interval(poll_config.interval_ms)
.map_err(|e| ChannelError::StartupFailed {
name: self.name.clone(),
reason: e,
})?;
if let Some(poll_config) = &config.poll {
if poll_config.enabled {
let interval = self
.capabilities
.validate_poll_interval(poll_config.interval_ms)
.map_err(|e| ChannelError::StartupFailed {
name: self.name.clone(),
reason: e,
})?;
// Create shutdown channel for polling and store the sender to keep it alive
let (poll_shutdown_tx, poll_shutdown_rx) = oneshot::channel();
*self.poll_shutdown_tx.write().await = Some(poll_shutdown_tx);
// Create shutdown channel for polling and store the sender to keep it alive
let (poll_shutdown_tx, poll_shutdown_rx) = oneshot::channel();
*self.poll_shutdown_tx.write().await = Some(poll_shutdown_tx);
self.start_polling(Duration::from_millis(interval as u64), poll_shutdown_rx);
self.start_polling(Duration::from_millis(interval as u64), poll_shutdown_rx);
}
}
tracing::info!(
@@ -2648,52 +2616,15 @@ mod tests {
assert_eq!(store.redact_credentials(input), input);
}
/// Verify that WASM HTTP host functions work using a dedicated
/// current-thread runtime inside spawn_blocking.
/// Verify that the block_on-inside-spawn_blocking pattern used by the WASM
/// channel HTTP host function doesn't deadlock or panic.
#[tokio::test]
async fn test_dedicated_runtime_inside_spawn_blocking() {
async fn test_block_on_inside_spawn_blocking_does_not_deadlock() {
let result = tokio::task::spawn_blocking(|| {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("failed to build runtime");
rt.block_on(async { 42 })
tokio::runtime::Handle::current().block_on(async { 42 })
})
.await
.expect("spawn_blocking panicked");
assert_eq!(result, 42);
}
/// Verify a real HTTP request works using the dedicated-runtime pattern.
/// This catches DNS, TLS, and I/O driver issues that trivial tests miss.
#[tokio::test]
#[ignore] // requires network
async fn test_dedicated_runtime_real_http() {
let result = tokio::task::spawn_blocking(|| {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("failed to build runtime");
rt.block_on(async {
let client = reqwest::Client::builder()
.connect_timeout(std::time::Duration::from_secs(10))
.build()
.expect("failed to build client");
let resp = client
.get("https://api.telegram.org/bot000/getMe")
.timeout(std::time::Duration::from_secs(10))
.send()
.await;
match resp {
Ok(r) => r.status().as_u16(),
Err(e) if e.is_timeout() => panic!("request timed out: {e}"),
Err(e) => panic!("unexpected error: {e}"),
}
})
})
.await
.expect("spawn_blocking panicked");
// 404 because "000" is not a valid bot token
assert_eq!(result, 404);
}
}
+12 -10
View File
@@ -25,21 +25,23 @@ pub async fn auth_middleware(
next: Next,
) -> Response {
// Try Authorization header first (constant-time comparison)
if let Some(auth_header) = headers.get("authorization")
&& let Ok(value) = auth_header.to_str()
&& let Some(token) = value.strip_prefix("Bearer ")
&& bool::from(token.as_bytes().ct_eq(auth.token.as_bytes()))
{
return next.run(request).await;
if let Some(auth_header) = headers.get("authorization") {
if let Ok(value) = auth_header.to_str() {
if let Some(token) = value.strip_prefix("Bearer ") {
if bool::from(token.as_bytes().ct_eq(auth.token.as_bytes())) {
return next.run(request).await;
}
}
}
}
// Fall back to query parameter for SSE EventSource (constant-time comparison)
if let Some(query) = request.uri().query() {
for pair in query.split('&') {
if let Some(token) = pair.strip_prefix("token=")
&& bool::from(token.as_bytes().ct_eq(auth.token.as_bytes()))
{
return next.run(request).await;
if let Some(token) = pair.strip_prefix("token=") {
if bool::from(token.as_bytes().ct_eq(auth.token.as_bytes())) {
return next.run(request).await;
}
}
}
}
+8 -8
View File
@@ -473,10 +473,10 @@ pub async fn chat_completions_handler(
if let Some(mt) = req.max_tokens {
tool_req = tool_req.with_max_tokens(mt);
}
if let Some(ref tc) = req.tool_choice
&& let Some(choice) = normalize_tool_choice(tc)
{
tool_req = tool_req.with_tool_choice(choice);
if let Some(ref tc) = req.tool_choice {
if let Some(choice) = normalize_tool_choice(tc) {
tool_req = tool_req.with_tool_choice(choice);
}
}
let resp = llm
@@ -591,10 +591,10 @@ async fn handle_streaming(
if let Some(mt) = req.max_tokens {
tool_req = tool_req.with_max_tokens(mt);
}
if let Some(ref tc) = req.tool_choice
&& let Some(choice) = normalize_tool_choice(tc)
{
tool_req = tool_req.with_tool_choice(choice);
if let Some(ref tc) = req.tool_choice {
if let Some(choice) = normalize_tool_choice(tc) {
tool_req = tool_req.with_tool_choice(choice);
}
}
LlmResult::WithTools(
llm.complete_with_tools(tool_req)
+151 -150
View File
@@ -525,10 +525,10 @@ pub async fn clear_auth_mode(state: &GatewayState) {
if let Some(ref sm) = state.session_manager {
let session = sm.get_or_create_session(&state.user_id).await;
let mut sess = session.lock().await;
if let Some(thread_id) = sess.active_thread
&& let Some(thread) = sess.threads.get_mut(&thread_id)
{
thread.pending_auth = None;
if let Some(thread_id) = sess.active_thread {
if let Some(thread) = sess.threads.get_mut(&thread_id) {
thread.pending_auth = None;
}
}
}
}
@@ -626,69 +626,69 @@ async fn chat_history_handler(
// Verify the thread belongs to the authenticated user before returning any data.
// In-memory threads are already scoped by user via session_manager, but DB
// lookups could expose another user's conversation if the UUID is guessed.
if query.thread_id.is_some()
&& let Some(ref store) = state.store
{
let owned = store
.conversation_belongs_to_user(thread_id, &state.user_id)
.await
.unwrap_or(false);
if !owned && !sess.threads.contains_key(&thread_id) {
return Err((StatusCode::NOT_FOUND, "Thread not found".to_string()));
if query.thread_id.is_some() {
if let Some(ref store) = state.store {
let owned = store
.conversation_belongs_to_user(thread_id, &state.user_id)
.await
.unwrap_or(false);
if !owned && !sess.threads.contains_key(&thread_id) {
return Err((StatusCode::NOT_FOUND, "Thread not found".to_string()));
}
}
}
// For paginated requests (before cursor set), always go to DB
if before_cursor.is_some()
&& let Some(ref store) = state.store
{
let (messages, has_more) = store
.list_conversation_messages_paginated(thread_id, before_cursor, limit as i64)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
if before_cursor.is_some() {
if let Some(ref store) = state.store {
let (messages, has_more) = store
.list_conversation_messages_paginated(thread_id, before_cursor, limit as i64)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
let oldest_timestamp = messages.first().map(|m| m.created_at.to_rfc3339());
let turns = build_turns_from_db_messages(&messages);
return Ok(Json(HistoryResponse {
thread_id,
turns,
has_more,
oldest_timestamp,
}));
let oldest_timestamp = messages.first().map(|m| m.created_at.to_rfc3339());
let turns = build_turns_from_db_messages(&messages);
return Ok(Json(HistoryResponse {
thread_id,
turns,
has_more,
oldest_timestamp,
}));
}
}
// Try in-memory first (freshest data for active threads)
if let Some(thread) = sess.threads.get(&thread_id)
&& !thread.turns.is_empty()
{
let turns: Vec<TurnInfo> = thread
.turns
.iter()
.map(|t| TurnInfo {
turn_number: t.turn_number,
user_input: t.user_input.clone(),
response: t.response.clone(),
state: format!("{:?}", t.state),
started_at: t.started_at.to_rfc3339(),
completed_at: t.completed_at.map(|dt| dt.to_rfc3339()),
tool_calls: t
.tool_calls
.iter()
.map(|tc| ToolCallInfo {
name: tc.name.clone(),
has_result: tc.result.is_some(),
has_error: tc.error.is_some(),
})
.collect(),
})
.collect();
if let Some(thread) = sess.threads.get(&thread_id) {
if !thread.turns.is_empty() {
let turns: Vec<TurnInfo> = thread
.turns
.iter()
.map(|t| TurnInfo {
turn_number: t.turn_number,
user_input: t.user_input.clone(),
response: t.response.clone(),
state: format!("{:?}", t.state),
started_at: t.started_at.to_rfc3339(),
completed_at: t.completed_at.map(|dt| dt.to_rfc3339()),
tool_calls: t
.tool_calls
.iter()
.map(|tc| ToolCallInfo {
name: tc.name.clone(),
has_result: tc.result.is_some(),
has_error: tc.error.is_some(),
})
.collect(),
})
.collect();
return Ok(Json(HistoryResponse {
thread_id,
turns,
has_more: false,
oldest_timestamp: None,
}));
return Ok(Json(HistoryResponse {
thread_id,
turns,
has_more: false,
oldest_timestamp: None,
}));
}
}
// Fall back to DB for historical threads not in memory (paginated)
@@ -738,12 +738,12 @@ fn build_turns_from_db_messages(messages: &[crate::history::ConversationMessage]
};
// Check if next message is an assistant response
if let Some(next) = iter.peek()
&& next.role == "assistant"
{
let assistant_msg = iter.next().expect("peeked");
turn.response = Some(assistant_msg.content.clone());
turn.completed_at = Some(assistant_msg.created_at.to_rfc3339());
if let Some(next) = iter.peek() {
if next.role == "assistant" {
let assistant_msg = iter.next().expect("peeked");
turn.response = Some(assistant_msg.content.clone());
turn.completed_at = Some(assistant_msg.created_at.to_rfc3339());
}
}
// Incomplete turn (user message without response)
@@ -1126,65 +1126,65 @@ async fn jobs_detail_handler(
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
// Try sandbox job from DB first, scoped to the authenticated user.
if let Some(ref store) = state.store
&& let Ok(Some(job)) = store.get_sandbox_job(job_id).await
{
if job.user_id != state.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
let browse_id = std::path::Path::new(&job.project_dir)
.file_name()
.map(|n| n.to_string_lossy().to_string())
.unwrap_or_else(|| job.id.to_string());
if let Some(ref store) = state.store {
if let Ok(Some(job)) = store.get_sandbox_job(job_id).await {
if job.user_id != state.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
let browse_id = std::path::Path::new(&job.project_dir)
.file_name()
.map(|n| n.to_string_lossy().to_string())
.unwrap_or_else(|| job.id.to_string());
let ui_state = match job.status.as_str() {
"creating" => "pending",
"running" => "in_progress",
s => s,
};
let ui_state = match job.status.as_str() {
"creating" => "pending",
"running" => "in_progress",
s => s,
};
let elapsed_secs = job.started_at.map(|start| {
let end = job.completed_at.unwrap_or_else(chrono::Utc::now);
(end - start).num_seconds().max(0) as u64
});
// Synthesize transitions from timestamps.
let mut transitions = Vec::new();
if let Some(started) = job.started_at {
transitions.push(TransitionInfo {
from: "creating".to_string(),
to: "running".to_string(),
timestamp: started.to_rfc3339(),
reason: None,
let elapsed_secs = job.started_at.map(|start| {
let end = job.completed_at.unwrap_or_else(chrono::Utc::now);
(end - start).num_seconds().max(0) as u64
});
}
if let Some(completed) = job.completed_at {
transitions.push(TransitionInfo {
from: "running".to_string(),
to: job.status.clone(),
timestamp: completed.to_rfc3339(),
reason: job.failure_reason.clone(),
});
}
return Ok(Json(JobDetailResponse {
id: job.id,
title: job.task.clone(),
description: String::new(),
state: ui_state.to_string(),
user_id: job.user_id.clone(),
created_at: job.created_at.to_rfc3339(),
started_at: job.started_at.map(|dt| dt.to_rfc3339()),
completed_at: job.completed_at.map(|dt| dt.to_rfc3339()),
elapsed_secs,
project_dir: Some(job.project_dir.clone()),
browse_url: Some(format!("/projects/{}/", browse_id)),
job_mode: {
let mode = store.get_sandbox_job_mode(job.id).await.ok().flatten();
mode.filter(|m| m != "worker")
},
transitions,
}));
// Synthesize transitions from timestamps.
let mut transitions = Vec::new();
if let Some(started) = job.started_at {
transitions.push(TransitionInfo {
from: "creating".to_string(),
to: "running".to_string(),
timestamp: started.to_rfc3339(),
reason: None,
});
}
if let Some(completed) = job.completed_at {
transitions.push(TransitionInfo {
from: "running".to_string(),
to: job.status.clone(),
timestamp: completed.to_rfc3339(),
reason: job.failure_reason.clone(),
});
}
return Ok(Json(JobDetailResponse {
id: job.id,
title: job.task.clone(),
description: String::new(),
state: ui_state.to_string(),
user_id: job.user_id.clone(),
created_at: job.created_at.to_rfc3339(),
started_at: job.started_at.map(|dt| dt.to_rfc3339()),
completed_at: job.completed_at.map(|dt| dt.to_rfc3339()),
elapsed_secs,
project_dir: Some(job.project_dir.clone()),
browse_url: Some(format!("/projects/{}/", browse_id)),
job_mode: {
let mode = store.get_sandbox_job_mode(job.id).await.ok().flatten();
mode.filter(|m| m != "worker")
},
transitions,
}));
}
}
Err((StatusCode::NOT_FOUND, "Job not found".to_string()))
@@ -1198,35 +1198,35 @@ async fn jobs_cancel_handler(
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
// Try sandbox job cancellation, scoped to the authenticated user.
if let Some(ref store) = state.store
&& let Ok(Some(job)) = store.get_sandbox_job(job_id).await
{
if job.user_id != state.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
if job.status == "running" || job.status == "creating" {
// Stop the container if we have a job manager.
if let Some(ref jm) = state.job_manager
&& let Err(e) = jm.stop_job(job_id).await
{
tracing::warn!(job_id = %job_id, error = %e, "Failed to stop container during cancellation");
if let Some(ref store) = state.store {
if let Ok(Some(job)) = store.get_sandbox_job(job_id).await {
if job.user_id != state.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
store
.update_sandbox_job_status(
job_id,
"failed",
Some(false),
Some("Cancelled by user"),
None,
Some(chrono::Utc::now()),
)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
if job.status == "running" || job.status == "creating" {
// Stop the container if we have a job manager.
if let Some(ref jm) = state.job_manager {
if let Err(e) = jm.stop_job(job_id).await {
tracing::warn!(job_id = %job_id, error = %e, "Failed to stop container during cancellation");
}
}
store
.update_sandbox_job_status(
job_id,
"failed",
Some(false),
Some("Cancelled by user"),
None,
Some(chrono::Utc::now()),
)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
}
return Ok(Json(serde_json::json!({
"status": "cancelled",
"job_id": job_id,
})));
}
return Ok(Json(serde_json::json!({
"status": "cancelled",
"job_id": job_id,
})));
}
Err((StatusCode::NOT_FOUND, "Job not found".to_string()))
@@ -1334,13 +1334,14 @@ async fn jobs_prompt_handler(
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
// Verify user owns this job.
if let Some(ref store) = state.store
&& !store
if let Some(ref store) = state.store {
if !store
.sandbox_job_belongs_to_user(job_id, &state.user_id)
.await
.unwrap_or(false)
{
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
{
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
}
let content = body
+6 -15
View File
@@ -12,7 +12,6 @@ let loadingOlder = false;
let jobEvents = new Map(); // job_id -> Array of events
let jobListRefreshTimer = null;
const JOB_EVENTS_CAP = 500;
const MEMORY_SEARCH_QUERY_MAX_LENGTH = 100;
// --- Auth ---
@@ -1002,12 +1001,9 @@ function buildBreadcrumb(path) {
}
function searchMemory(query) {
const normalizedQuery = normalizeSearchQuery(query);
if (!normalizedQuery) return;
apiFetch('/api/memory/search', {
method: 'POST',
body: { query: normalizedQuery, limit: 20 },
body: { query, limit: 20 },
}).then((data) => {
const tree = document.getElementById('memory-tree');
tree.innerHTML = '';
@@ -1018,23 +1014,18 @@ function searchMemory(query) {
for (const result of data.results) {
const item = document.createElement('div');
item.className = 'search-result';
const snippet = snippetAround(result.content, normalizedQuery, 120);
const snippet = snippetAround(result.content, query, 120);
item.innerHTML = '<div class="path">' + escapeHtml(result.path) + '</div>'
+ '<div class="snippet">' + highlightQuery(snippet, normalizedQuery) + '</div>';
+ '<div class="snippet">' + highlightQuery(snippet, query) + '</div>';
item.addEventListener('click', () => readMemoryFile(result.path));
tree.appendChild(item);
}
}).catch(() => {});
}
function normalizeSearchQuery(query) {
return (typeof query === 'string' ? query : '').slice(0, MEMORY_SEARCH_QUERY_MAX_LENGTH);
}
function snippetAround(text, query, len) {
const normalizedQuery = normalizeSearchQuery(query);
const lower = text.toLowerCase();
const idx = lower.indexOf(normalizedQuery.toLowerCase());
const idx = lower.indexOf(query.toLowerCase());
if (idx < 0) return text.substring(0, len);
const start = Math.max(0, idx - Math.floor(len / 2));
const end = Math.min(text.length, start + len);
@@ -1047,11 +1038,11 @@ function snippetAround(text, query, len) {
function highlightQuery(text, query) {
if (!query) return escapeHtml(text);
const escaped = escapeHtml(text);
const normalizedQuery = normalizeSearchQuery(query);
const queryEscaped = normalizedQuery.replace(/[.*+?^${}()|[\]\\]/g, '\\$&');
const queryEscaped = query.replace(/[.*+?^${}()|[\]\\]/g, '\\$&');
const re = new RegExp('(' + queryEscaped + ')', 'gi');
return escaped.replace(re, '<mark>$1</mark>');
}
// --- Logs ---
const LOG_MAX_ENTRIES = 2000;
+1 -5
View File
@@ -5,11 +5,7 @@
<meta name="viewport" content="width=device-width, initial-scale=1.0">
<title>IronClaw</title>
<link rel="stylesheet" href="/style.css">
<script
src="https://cdn.jsdelivr.net/npm/[email protected]/lib/marked.umd.min.js"
integrity="sha384-pN9zSKOnTZwXRtYZAu0PBPEgR2B7DOC1aeLxQ33oJ0oy5iN1we6gm57xldM2irDG"
crossorigin="anonymous"
></script>
<script src="https://cdn.jsdelivr.net/npm/marked/marked.min.js"></script>
</head>
<body>
<!-- Auth Screen -->
+60 -28
View File
@@ -48,6 +48,8 @@ pub enum ConfigCommand {
/// Connects to the database to read/write settings. Falls back to disk
/// if the database is not available.
pub async fn run_config_command(cmd: ConfigCommand) -> anyhow::Result<()> {
let _ = dotenvy::dotenv();
// Try to connect to the DB for settings access
let db: Option<Arc<dyn crate::db::Database>> = match connect_db().await {
Ok(d) => Some(d),
@@ -90,7 +92,7 @@ async fn load_settings(store: Option<&dyn crate::db::Database>) -> Settings {
_ => {}
}
}
Settings::default()
Settings::load()
}
/// List all settings.
@@ -108,10 +110,10 @@ async fn list_settings(
println!();
for (key, value) in all {
if let Some(ref f) = filter
&& !key.starts_with(f)
{
continue;
if let Some(ref f) = filter {
if !key.starts_with(f) {
continue;
}
}
let display_value = if value.len() > 60 {
@@ -153,17 +155,19 @@ async fn set_setting(
.set(path, value)
.map_err(|e| anyhow::anyhow!("{}", e))?;
let store = store.ok_or_else(|| {
anyhow::anyhow!("Database connection required to save settings. Check DATABASE_URL.")
})?;
let json_value = match serde_json::from_str::<serde_json::Value>(value) {
Ok(v) => v,
Err(_) => serde_json::Value::String(value.to_string()),
};
store
.set_setting(DEFAULT_USER_ID, path, &json_value)
.await
.map_err(|e| anyhow::anyhow!("Failed to save to database: {}", e))?;
// Save to DB if available, otherwise disk
if let Some(store) = store {
let json_value = match serde_json::from_str::<serde_json::Value>(value) {
Ok(v) => v,
Err(_) => serde_json::Value::String(value.to_string()),
};
store
.set_setting(DEFAULT_USER_ID, path, &json_value)
.await
.map_err(|e| anyhow::anyhow!("Failed to save to database: {}", e))?;
} else {
settings.save()?;
}
println!("Set {} = {}", path, value);
Ok(())
@@ -176,13 +180,17 @@ async fn reset_setting(store: Option<&dyn crate::db::Database>, path: &str) -> a
.get(path)
.ok_or_else(|| anyhow::anyhow!("Unknown setting: {}", path))?;
let store = store.ok_or_else(|| {
anyhow::anyhow!("Database connection required to reset settings. Check DATABASE_URL.")
})?;
store
.delete_setting(DEFAULT_USER_ID, path)
.await
.map_err(|e| anyhow::anyhow!("Failed to delete setting from database: {}", e))?;
// Delete from DB (falling back to default) or reset on disk
if let Some(store) = store {
store
.delete_setting(DEFAULT_USER_ID, path)
.await
.map_err(|e| anyhow::anyhow!("Failed to delete setting from database: {}", e))?;
} else {
let mut settings = Settings::load();
settings.reset(path).map_err(|e| anyhow::anyhow!("{}", e))?;
settings.save()?;
}
println!("Reset {} to default: {}", path, default_value);
Ok(())
@@ -192,13 +200,37 @@ async fn reset_setting(store: Option<&dyn crate::db::Database>, path: &str) -> a
fn show_path(has_db: bool) -> anyhow::Result<()> {
if has_db {
println!("Settings stored in: database (settings table)");
println!(
"Bootstrap config: {}",
crate::bootstrap::BootstrapConfig::default_path().display()
);
} else {
println!("Settings stored in: PostgreSQL (not connected, using defaults)");
let path = Settings::default_path();
println!("Settings stored in: {} (disk fallback)", path.display());
if path.exists() {
let metadata = std::fs::metadata(&path)?;
println!(" Size: {} bytes", metadata.len());
if let Ok(modified) = metadata.modified() {
use std::time::SystemTime;
let duration = SystemTime::now()
.duration_since(modified)
.unwrap_or_default();
let secs = duration.as_secs();
if secs < 60 {
println!(" Modified: {} seconds ago", secs);
} else if secs < 3600 {
println!(" Modified: {} minutes ago", secs / 60);
} else if secs < 86400 {
println!(" Modified: {} hours ago", secs / 3600);
} else {
println!(" Modified: {} days ago", secs / 86400);
}
}
} else {
println!(" (does not exist, using defaults)");
}
}
println!(
"Env config: {}",
crate::bootstrap::ironclaw_env_path().display()
);
Ok(())
}
-1
View File
@@ -12,7 +12,6 @@
mod config;
mod mcp;
pub mod memory;
pub mod oauth_defaults;
mod pairing;
pub mod status;
mod tool;
-342
View File
@@ -1,342 +0,0 @@
//! Shared OAuth infrastructure: built-in credentials, callback server, landing pages.
//!
//! Every OAuth flow in the codebase (WASM tool auth, MCP server auth, NEAR AI login)
//! uses the same callback port, landing page, and listener logic from this module.
//!
//! # Built-in Credentials
//!
//! Many CLI tools (gcloud, rclone, gdrive) ship with default OAuth credentials
//! so users don't need to register their own OAuth app. Google explicitly
//! documents that client_secret for "Desktop App" / "Installed App" types
//! is NOT actually secret.
//!
//! Default credentials are hardcoded below. They can be overridden at:
//!
//! - **Compile time**: Set IRONCLAW_GOOGLE_CLIENT_ID / IRONCLAW_GOOGLE_CLIENT_SECRET
//! env vars before building to replace the hardcoded defaults.
//! - **Runtime**: Users can set GOOGLE_OAUTH_CLIENT_ID / GOOGLE_OAUTH_CLIENT_SECRET
//! env vars, which take priority over built-in defaults.
use std::time::Duration;
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::net::TcpListener;
// ── Built-in credentials ────────────────────────────────────────────────
pub struct OAuthCredentials {
pub client_id: &'static str,
pub client_secret: &'static str,
}
/// Google OAuth "Desktop App" credentials, shared across all Google tools.
/// Compile-time env vars override the hardcoded defaults below.
const GOOGLE_CLIENT_ID: &str = match option_env!("IRONCLAW_GOOGLE_CLIENT_ID") {
Some(v) => v,
None => "564604149681-efo25d43rs85v0tibdepsmdv5dsrhhr0.apps.googleusercontent.com",
};
const GOOGLE_CLIENT_SECRET: &str = match option_env!("IRONCLAW_GOOGLE_CLIENT_SECRET") {
Some(v) => v,
None => "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",
};
/// Returns built-in OAuth credentials for a provider, keyed by secret_name.
///
/// The secret_name comes from the tool's capabilities.json `auth.secret_name` field.
/// Returns `None` if no built-in credentials are configured for that provider.
pub fn builtin_credentials(secret_name: &str) -> Option<OAuthCredentials> {
match secret_name {
"google_oauth_token" => Some(OAuthCredentials {
client_id: GOOGLE_CLIENT_ID,
client_secret: GOOGLE_CLIENT_SECRET,
}),
_ => None,
}
}
// ── Shared callback server ──────────────────────────────────────────────
/// Fixed port for all OAuth callbacks.
///
/// Every redirect URI registered with providers must use this port:
/// `http://localhost:9876/callback` (or `/auth/callback` for NEAR AI).
pub const OAUTH_CALLBACK_PORT: u16 = 9876;
/// Error from the OAuth callback listener.
#[derive(Debug, thiserror::Error)]
pub enum OAuthCallbackError {
#[error("Port {0} is in use (another auth flow running?): {1}")]
PortInUse(u16, String),
#[error("Authorization denied by user")]
Denied,
#[error("Timed out waiting for authorization")]
Timeout,
#[error("IO error: {0}")]
Io(String),
}
/// Bind the OAuth callback listener on the fixed port.
///
/// Binds to IPv4 `127.0.0.1` first because callback URLs use `127.0.0.1`
/// explicitly (e.g., NEAR AI redirects to `http://127.0.0.1:9876/auth/callback`).
/// Falls back to IPv6 `[::1]` only if IPv4 binding fails for a reason other
/// than `AddrInUse`. If the port is already occupied, fails immediately.
pub async fn bind_callback_listener() -> Result<TcpListener, OAuthCallbackError> {
let ipv4_addr = format!("127.0.0.1:{}", OAUTH_CALLBACK_PORT);
match TcpListener::bind(&ipv4_addr).await {
Ok(listener) => return Ok(listener),
Err(e) if e.kind() == std::io::ErrorKind::AddrInUse => {
return Err(OAuthCallbackError::PortInUse(
OAUTH_CALLBACK_PORT,
e.to_string(),
));
}
Err(_) => {
// IPv4 not available, fall back to IPv6
}
}
TcpListener::bind(format!("[::1]:{}", OAUTH_CALLBACK_PORT))
.await
.map_err(|e| {
if e.kind() == std::io::ErrorKind::AddrInUse {
OAuthCallbackError::PortInUse(OAUTH_CALLBACK_PORT, e.to_string())
} else {
OAuthCallbackError::Io(e.to_string())
}
})
}
/// Wait for an OAuth callback and extract a query parameter value.
///
/// Listens for a GET request matching `path_prefix` (e.g., "/callback" or "/auth/callback"),
/// extracts the value of `param_name` (e.g., "code" or "token"), and shows a branded
/// landing page using `display_name` (e.g., "Google", "Notion", "NEAR AI").
///
/// Times out after 5 minutes.
pub async fn wait_for_callback(
listener: TcpListener,
path_prefix: &str,
param_name: &str,
display_name: &str,
) -> Result<String, OAuthCallbackError> {
let path_prefix = path_prefix.to_string();
let param_name = param_name.to_string();
let display_name = display_name.to_string();
tokio::time::timeout(Duration::from_secs(300), async move {
loop {
let (mut socket, _) = listener
.accept()
.await
.map_err(|e| OAuthCallbackError::Io(e.to_string()))?;
let mut reader = BufReader::new(&mut socket);
let mut request_line = String::new();
reader
.read_line(&mut request_line)
.await
.map_err(|e| OAuthCallbackError::Io(e.to_string()))?;
if let Some(path) = request_line.split_whitespace().nth(1)
&& path.starts_with(&path_prefix)
&& let Some(query) = path.split('?').nth(1)
{
// Check for error first
if query.contains("error=") {
let html = landing_html(&display_name, false);
let response = format!(
"HTTP/1.1 400 Bad Request\r\n\
Content-Type: text/html; charset=utf-8\r\n\
Connection: close\r\n\
\r\n\
{}",
html
);
let _ = socket.write_all(response.as_bytes()).await;
return Err(OAuthCallbackError::Denied);
}
// Look for the target parameter
for param in query.split('&') {
let parts: Vec<&str> = param.splitn(2, '=').collect();
if parts.len() == 2 && parts[0] == param_name {
let value = urlencoding::decode(parts[1])
.unwrap_or_else(|_| parts[1].into())
.into_owned();
let html = landing_html(&display_name, true);
let response = format!(
"HTTP/1.1 200 OK\r\n\
Content-Type: text/html; charset=utf-8\r\n\
Connection: close\r\n\
\r\n\
{}",
html
);
let _ = socket.write_all(response.as_bytes()).await;
let _ = socket.shutdown().await;
return Ok(value);
}
}
}
// Not the callback we're looking for
let response = "HTTP/1.1 404 Not Found\r\nConnection: close\r\n\r\n";
let _ = socket.write_all(response.as_bytes()).await;
}
})
.await
.map_err(|_| OAuthCallbackError::Timeout)?
}
/// Escape a string for safe interpolation into HTML content.
fn html_escape(s: &str) -> String {
let mut out = String::with_capacity(s.len());
for c in s.chars() {
match c {
'&' => out.push_str("&amp;"),
'<' => out.push_str("&lt;"),
'>' => out.push_str("&gt;"),
'"' => out.push_str("&quot;"),
'\'' => out.push_str("&#x27;"),
_ => out.push(c),
}
}
out
}
/// HTML landing page shown in the browser after an OAuth redirect.
pub fn landing_html(provider_name: &str, success: bool) -> String {
let safe_name = html_escape(provider_name);
let (icon, heading, subtitle, accent) = if success {
(
r##"<div style="width:64px;height:64px;border-radius:50%;background:#22c55e;display:flex;align-items:center;justify-content:center;margin:0 auto 24px">
<svg width="32" height="32" viewBox="0 0 24 24" fill="none" stroke="#fff" stroke-width="3" stroke-linecap="round" stroke-linejoin="round"><polyline points="20 6 9 17 4 12"/></svg>
</div>"##,
format!("{} Connected", safe_name),
"You can close this window and return to your terminal.",
"#22c55e",
)
} else {
(
r##"<div style="width:64px;height:64px;border-radius:50%;background:#ef4444;display:flex;align-items:center;justify-content:center;margin:0 auto 24px">
<svg width="32" height="32" viewBox="0 0 24 24" fill="none" stroke="#fff" stroke-width="3" stroke-linecap="round" stroke-linejoin="round"><line x1="18" y1="6" x2="6" y2="18"/><line x1="6" y1="6" x2="18" y2="18"/></svg>
</div>"##,
"Authorization Failed".to_string(),
"The request was denied. You can close this window and try again.",
"#ef4444",
)
};
format!(
r#"<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="utf-8">
<meta name="viewport" content="width=device-width,initial-scale=1">
<title>IronClaw - {heading}</title>
<style>
* {{ margin:0; padding:0; box-sizing:border-box }}
body {{
font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, Helvetica, Arial, sans-serif;
background: #0a0a0a;
color: #e5e5e5;
display: flex;
justify-content: center;
align-items: center;
min-height: 100vh;
}}
.card {{
text-align: center;
padding: 48px 40px;
max-width: 420px;
border: 1px solid #262626;
border-radius: 16px;
background: #141414;
}}
h1 {{
font-size: 22px;
font-weight: 600;
margin-bottom: 8px;
color: #fafafa;
}}
p {{
font-size: 14px;
color: #a3a3a3;
line-height: 1.5;
}}
.accent {{ color: {accent}; }}
.brand {{
margin-top: 32px;
font-size: 12px;
color: #525252;
letter-spacing: 0.5px;
text-transform: uppercase;
}}
</style>
</head>
<body>
<div class="card">
{icon}
<h1>{heading}</h1>
<p>{subtitle}</p>
<div class="brand">IronClaw</div>
</div>
</body>
</html>"#,
heading = heading,
icon = icon,
subtitle = subtitle,
accent = accent,
)
}
#[cfg(test)]
mod tests {
use crate::cli::oauth_defaults::{builtin_credentials, landing_html};
#[test]
fn test_unknown_provider_returns_none() {
assert!(builtin_credentials("unknown_token").is_none());
}
#[test]
fn test_google_returns_based_on_compile_env() {
let creds = builtin_credentials("google_oauth_token");
assert!(creds.is_some());
let creds = creds.unwrap();
assert!(!creds.client_id.is_empty());
assert!(!creds.client_secret.is_empty());
}
#[test]
fn test_landing_html_success_contains_key_elements() {
let html = landing_html("Google", true);
assert!(html.contains("Google Connected"));
assert!(html.contains("charset"));
assert!(html.contains("IronClaw"));
assert!(html.contains("#22c55e")); // green accent
assert!(!html.contains("Failed"));
}
#[test]
fn test_landing_html_escapes_provider_name() {
let html = landing_html("<script>alert(1)</script>", true);
assert!(!html.contains("<script>"));
assert!(html.contains("&lt;script&gt;"));
}
#[test]
fn test_landing_html_error_contains_key_elements() {
let html = landing_html("Notion", false);
assert!(html.contains("Authorization Failed"));
assert!(html.contains("charset"));
assert!(html.contains("IronClaw"));
assert!(html.contains("#ef4444")); // red accent
assert!(!html.contains("Connected"));
}
}
+24 -44
View File
@@ -9,7 +9,7 @@ use crate::settings::Settings;
/// Run the status command, printing system health info.
pub async fn run_status_command() -> anyhow::Result<()> {
let settings = Settings::default();
let settings = Settings::load();
println!("IronClaw Status");
println!("===============\n");
@@ -22,36 +22,16 @@ pub async fn run_status_command() -> anyhow::Result<()> {
);
// Database
let db_url_set = settings.database_url.is_some() || std::env::var("DATABASE_URL").is_ok();
print!(" Database: ");
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 {
Ok(()) => println!("connected (PostgreSQL)"),
Err(e) => println!("error ({})", e),
}
} else {
println!("not configured");
}
if db_url_set {
// Try to connect
match check_database().await {
Ok(()) => println!("connected"),
Err(e) => println!("error ({})", e),
}
} else {
println!("not configured");
}
// Session / Auth
@@ -63,17 +43,15 @@ pub async fn run_status_command() -> anyhow::Result<()> {
println!("not found (run `ironclaw onboard`)");
}
// Secrets (auto-detect from env only; skip keychain probe to avoid
// triggering macOS system password dialogs on a simple status check)
// Secrets
print!(" Secrets: ");
if std::env::var("SECRETS_MASTER_KEY").is_ok() {
println!("configured (env)");
let secrets_configured = settings.secrets_master_key_source != crate::settings::KeySource::None
|| std::env::var("SECRETS_MASTER_KEY").is_ok()
|| crate::secrets::keychain::has_master_key().await;
if secrets_configured {
println!("configured ({:?})", settings.secrets_master_key_source);
} else {
// 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)");
println!("not configured");
}
// Embeddings
@@ -151,18 +129,20 @@ pub async fn run_status_command() -> anyhow::Result<()> {
Err(_) => println!("none configured"),
}
// Config path
println!(
"\n Config: {}",
crate::bootstrap::ironclaw_env_path().display()
);
// Settings path
println!("\n Settings: {}", Settings::default_path().display());
Ok(())
}
#[cfg(feature = "postgres")]
async fn check_database() -> anyhow::Result<()> {
let url = std::env::var("DATABASE_URL").map_err(|_| anyhow::anyhow!("DATABASE_URL not set"))?;
let _ = dotenvy::dotenv();
let settings = Settings::load();
let url = std::env::var("DATABASE_URL")
.ok()
.or(settings.database_url)
.ok_or_else(|| anyhow::anyhow!("no URL"))?;
let config: deadpool_postgres::Config = deadpool_postgres::Config {
url: Some(url),
+153 -175
View File
@@ -423,11 +423,11 @@ async fn extract_crate_name(cargo_toml: &Path) -> anyhow::Result<String> {
// Simple TOML parsing for [package] name
for line in content.lines() {
let line = line.trim();
if line.starts_with("name")
&& let Some((_, value)) = line.split_once('=')
{
let name = value.trim().trim_matches('"').trim_matches('\'');
return Ok(name.to_string());
if line.starts_with("name") {
if let Some((_, value)) = line.split_once('=') {
let name = value.trim().trim_matches('"').trim_matches('\'');
return Ok(name.to_string());
}
}
}
@@ -491,10 +491,10 @@ async fn list_tools(dir: Option<PathBuf>, verbose: bool) -> anyhow::Result<()> {
if has_caps {
let caps_path = path.with_extension("capabilities.json");
if let Ok(content) = fs::read_to_string(&caps_path).await
&& let Ok(caps) = CapabilitiesFile::from_json(&content)
{
print_capabilities_summary(&caps);
if let Ok(content) = fs::read_to_string(&caps_path).await {
if let Ok(caps) = CapabilitiesFile::from_json(&content) {
print_capabilities_summary(&caps);
}
}
}
println!();
@@ -607,16 +607,16 @@ fn print_capabilities_summary(caps: &CapabilitiesFile) {
}
}
if let Some(ref secrets) = caps.secrets
&& !secrets.allowed_names.is_empty()
{
parts.push(format!("secrets: {}", secrets.allowed_names.len()));
if let Some(ref secrets) = caps.secrets {
if !secrets.allowed_names.is_empty() {
parts.push(format!("secrets: {}", secrets.allowed_names.len()));
}
}
if let Some(ref ws) = caps.workspace
&& !ws.allowed_prefixes.is_empty()
{
parts.push("workspace: read".to_string());
if let Some(ref ws) = caps.workspace {
if !ws.allowed_prefixes.is_empty() {
parts.push("workspace: read".to_string());
}
}
if !parts.is_empty() {
@@ -653,30 +653,30 @@ fn print_capabilities_detail(caps: &CapabilitiesFile) {
}
}
if let Some(ref secrets) = caps.secrets
&& !secrets.allowed_names.is_empty()
{
println!(" Secrets (existence check only):");
for name in &secrets.allowed_names {
println!(" {}", name);
if let Some(ref secrets) = caps.secrets {
if !secrets.allowed_names.is_empty() {
println!(" Secrets (existence check only):");
for name in &secrets.allowed_names {
println!(" {}", name);
}
}
}
if let Some(ref tool_invoke) = caps.tool_invoke
&& !tool_invoke.aliases.is_empty()
{
println!(" Tool aliases:");
for (alias, real_name) in &tool_invoke.aliases {
println!(" {} -> {}", alias, real_name);
if let Some(ref tool_invoke) = caps.tool_invoke {
if !tool_invoke.aliases.is_empty() {
println!(" Tool aliases:");
for (alias, real_name) in &tool_invoke.aliases {
println!(" {} -> {}", alias, real_name);
}
}
}
if let Some(ref ws) = caps.workspace
&& !ws.allowed_prefixes.is_empty()
{
println!(" Workspace read prefixes:");
for prefix in &ws.allowed_prefixes {
println!(" {}", prefix);
if let Some(ref ws) = caps.workspace {
if !ws.allowed_prefixes.is_empty() {
println!(" Workspace read prefixes:");
for prefix in &ws.allowed_prefixes {
println!(" {}", prefix);
}
}
}
}
@@ -802,100 +802,48 @@ async fn auth_tool(name: String, dir: Option<PathBuf>, user_id: String) -> anyho
}
// Check for environment variable
if let Some(ref env_var) = auth.env_var
&& let Ok(token) = std::env::var(env_var)
&& !token.is_empty()
{
println!(" Found {} in environment.", env_var);
println!();
if let Some(ref env_var) = auth.env_var {
if let Ok(token) = std::env::var(env_var) {
if !token.is_empty() {
println!(" Found {} in environment.", env_var);
println!();
// Validate if endpoint is provided
if let Some(ref validation) = auth.validation_endpoint {
print!(" Validating token...");
std::io::stdout().flush()?;
// Validate if endpoint is provided
if let Some(ref validation) = auth.validation_endpoint {
print!(" Validating token...");
std::io::stdout().flush()?;
match validate_token(&token, validation, &auth.secret_name).await {
Ok(()) => {
println!("");
}
Err(e) => {
println!("");
println!(" Validation failed: {}", e);
println!();
println!(" Falling back to manual entry...");
return auth_tool_manual(secrets_store.as_ref(), &user_id, &auth).await;
match validate_token(&token, validation, &auth.secret_name).await {
Ok(()) => {
println!("");
}
Err(e) => {
println!("");
println!(" Validation failed: {}", e);
println!();
println!(" Falling back to manual entry...");
return auth_tool_manual(secrets_store.as_ref(), &user_id, &auth).await;
}
}
}
// Save the token
save_token(secrets_store.as_ref(), &user_id, &auth, &token).await?;
print_success(display_name);
return Ok(());
}
}
// Save the token
save_token(secrets_store.as_ref(), &user_id, &auth, &token, None, None).await?;
print_success(display_name);
return Ok(());
}
// Check for OAuth configuration
if let Some(ref oauth) = auth.oauth {
// For providers with shared tokens (e.g., all Google tools share google_oauth_token),
// combine scopes from all installed tools so one auth covers everything.
let combined = combine_provider_scopes(&tools_dir, &auth.secret_name, oauth).await;
if combined.scopes.len() > oauth.scopes.len() {
let extra = combined.scopes.len() - oauth.scopes.len();
println!(
" Including scopes from {} other installed tool(s) sharing this credential.",
extra
);
println!();
}
return auth_tool_oauth(secrets_store.as_ref(), &user_id, &auth, &combined).await;
return auth_tool_oauth(secrets_store.as_ref(), &user_id, &auth, oauth).await;
}
// Fall back to manual entry
auth_tool_manual(secrets_store.as_ref(), &user_id, &auth).await
}
/// Scan the tools directory for all capabilities files sharing the same secret_name
/// and combine their OAuth scopes. This way, authing any Google tool requests scopes
/// for ALL installed Google tools, so one login covers everything.
async fn combine_provider_scopes(
tools_dir: &Path,
secret_name: &str,
base_oauth: &crate::tools::wasm::OAuthConfigSchema,
) -> crate::tools::wasm::OAuthConfigSchema {
let mut all_scopes: std::collections::HashSet<String> =
base_oauth.scopes.iter().cloned().collect();
if let Ok(mut entries) = tokio::fs::read_dir(tools_dir).await {
while let Ok(Some(entry)) = entries.next_entry().await {
let path = entry.path();
if path.extension().and_then(|e| e.to_str()) != Some("json") {
continue;
}
let name = path
.file_name()
.and_then(|n| n.to_str())
.unwrap_or_default();
if !name.ends_with(".capabilities.json") {
continue;
}
if let Ok(content) = tokio::fs::read_to_string(&path).await
&& let Ok(caps) = CapabilitiesFile::from_json(&content)
&& let Some(auth) = &caps.auth
&& auth.secret_name == secret_name
&& let Some(oauth) = &auth.oauth
{
all_scopes.extend(oauth.scopes.iter().cloned());
}
}
}
let mut combined = base_oauth.clone();
combined.scopes = all_scopes.into_iter().collect();
combined.scopes.sort(); // deterministic ordering
combined
}
/// OAuth browser-based login flow.
async fn auth_tool_oauth(
store: &(dyn SecretsStore + Send + Sync),
@@ -906,14 +854,12 @@ async fn auth_tool_oauth(
use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD};
use rand::RngCore;
use sha2::{Digest, Sha256};
use crate::cli::oauth_defaults::{self, OAUTH_CALLBACK_PORT};
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::net::TcpListener;
let display_name = auth.display_name.as_deref().unwrap_or(&auth.secret_name);
// Get client_id: capabilities file > runtime env var > built-in defaults
let builtin = oauth_defaults::builtin_credentials(&auth.secret_name);
// Get client_id from config or env
let client_id = oauth
.client_id
.clone()
@@ -923,32 +869,41 @@ async fn auth_tool_oauth(
.as_ref()
.and_then(|env| std::env::var(env).ok())
})
.or_else(|| builtin.as_ref().map(|c| c.client_id.to_string()))
.ok_or_else(|| {
anyhow::anyhow!(
"OAuth client_id not configured.\n\
Set {} env var, or build with IRONCLAW_GOOGLE_CLIENT_ID.",
oauth.client_id_env.as_deref().unwrap_or("the client_id")
Set it in the capabilities file or via environment variable."
)
})?;
// Get client_secret: capabilities file > runtime env var > built-in defaults
let client_secret = oauth
.client_secret
.clone()
.or_else(|| {
oauth
.client_secret_env
.as_ref()
.and_then(|env| std::env::var(env).ok())
})
.or_else(|| builtin.as_ref().map(|c| c.client_secret.to_string()));
// Get client_secret if provided
let client_secret = oauth.client_secret.clone().or_else(|| {
oauth
.client_secret_env
.as_ref()
.and_then(|env| std::env::var(env).ok())
});
println!(" Starting OAuth authentication...");
println!();
let listener = oauth_defaults::bind_callback_listener().await?;
let redirect_uri = format!("http://localhost:{}/callback", OAUTH_CALLBACK_PORT);
// Find an available port for the callback
let mut listener = None;
let mut port = 0;
for p in 9876..=9886 {
match TcpListener::bind(format!("127.0.0.1:{}", p)).await {
Ok(l) => {
listener = Some(l);
port = p;
break;
}
Err(_) => continue,
}
}
let listener = listener.ok_or_else(|| anyhow::anyhow!("Could not find available port"))?;
let redirect_uri = format!("http://localhost:{}/callback", port);
// Generate PKCE verifier and challenge
let (code_verifier, code_challenge) = if oauth.use_pkce {
@@ -1007,8 +962,65 @@ async fn auth_tool_oauth(
println!(" Waiting for authorization...");
let code =
oauth_defaults::wait_for_callback(listener, "/callback", "code", display_name).await?;
// Wait for callback with timeout
let timeout = std::time::Duration::from_secs(300);
let code = tokio::time::timeout(timeout, async {
loop {
let (mut socket, _) = listener.accept().await?;
let mut reader = BufReader::new(&mut socket);
let mut request_line = String::new();
reader.read_line(&mut request_line).await?;
// Parse GET /callback?code=xxx HTTP/1.1
if let Some(path) = request_line.split_whitespace().nth(1) {
if path.starts_with("/callback") {
if let Some(query) = path.split('?').nth(1) {
for param in query.split('&') {
let parts: Vec<&str> = param.splitn(2, '=').collect();
if parts.len() == 2 && parts[0] == "code" {
let code = urlencoding::decode(parts[1])
.unwrap_or_else(|_| parts[1].into())
.into_owned();
// Send success response
let response = format!(
"HTTP/1.1 200 OK\r\n\
Content-Type: text/html\r\n\
\r\n\
<!DOCTYPE html><html><body style=\"font-family: sans-serif; \
display: flex; justify-content: center; align-items: center; \
height: 100vh; margin: 0; background: #191919; color: white;\">\
<div style=\"text-align: center;\">\
<h1> {} Connected!</h1>\
<p>You can close this window.</p>\
</div></body></html>",
display_name
);
let _ = socket.write_all(response.as_bytes()).await;
let _ = socket.shutdown().await;
return Ok::<_, anyhow::Error>(code);
}
}
// Check for error
if query.contains("error=") {
let response =
"HTTP/1.1 400 Bad Request\r\n\r\nAuthorization denied";
let _ = socket.write_all(response.as_bytes()).await;
return Err(anyhow::anyhow!("Authorization denied by user"));
}
}
}
}
let response = "HTTP/1.1 404 Not Found\r\n\r\n";
let _ = socket.write_all(response.as_bytes()).await;
}
})
.await
.map_err(|_| anyhow::anyhow!("Timed out waiting for authorization"))??;
println!();
println!(" Exchanging code for token...");
@@ -1059,19 +1071,8 @@ async fn auth_tool_oauth(
)
})?;
let refresh_token = token_data.get("refresh_token").and_then(|v| v.as_str());
let expires_in = token_data.get("expires_in").and_then(|v| v.as_u64());
// Save the token (with refresh token and expiry if provided)
save_token(
store,
user_id,
auth,
access_token,
refresh_token,
expires_in,
)
.await?;
// Save the token
save_token(store, user_id, auth, access_token).await?;
// Extract any additional info for display
let workspace_name = token_data
@@ -1173,8 +1174,8 @@ async fn auth_tool_manual(
}
}
// Save the token (manual path: no refresh token or expiry)
save_token(store, user_id, auth, &token, None, None).await?;
// Save the token
save_token(store, user_id, auth, &token).await?;
print_success(display_name);
Ok(())
}
@@ -1265,16 +1266,11 @@ async fn validate_token(
}
/// Save token to secrets store.
///
/// Optionally stores a refresh token (as `{secret_name}_refresh_token`) and
/// sets `expires_at` on the access token so the runtime can auto-refresh.
async fn save_token(
store: &(dyn SecretsStore + Send + Sync),
user_id: &str,
auth: &crate::tools::wasm::AuthCapabilitySchema,
token: &str,
refresh_token: Option<&str>,
expires_in: Option<u64>,
) -> anyhow::Result<()> {
let mut params = CreateSecretParams::new(&auth.secret_name, token);
@@ -1282,29 +1278,11 @@ async fn save_token(
params = params.with_provider(provider);
}
if let Some(secs) = expires_in {
let expires_at = chrono::Utc::now() + chrono::Duration::seconds(secs as i64);
params = params.with_expiry(expires_at);
}
store
.create(user_id, params)
.await
.map_err(|e| anyhow::anyhow!("Failed to save token: {}", e))?;
// Store refresh token separately (no expiry, it's long-lived)
if let Some(rt) = refresh_token {
let refresh_name = format!("{}_refresh_token", auth.secret_name);
let mut refresh_params = CreateSecretParams::new(&refresh_name, rt);
if let Some(ref provider) = auth.provider {
refresh_params = refresh_params.with_provider(provider);
}
store
.create(user_id, refresh_params)
.await
.map_err(|e| anyhow::anyhow!("Failed to save refresh token: {}", e))?;
}
Ok(())
}
+75 -159
View File
@@ -1,13 +1,11 @@
//! Configuration for IronClaw.
//!
//! Settings are loaded with priority: env var > database > default.
//! `DATABASE_URL` lives in `~/.ironclaw/.env` (loaded via dotenvy early
//! in startup). Everything else comes from env vars, the DB settings
//! table, or auto-detection.
//! The database replaces the old `settings.json` file for all settings
//! except the 4 bootstrap fields (database_url, pool_size, secrets key
//! source, onboard_completed) which live in `~/.ironclaw/bootstrap.json`.
use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::OnceLock;
use std::time::Duration;
use secrecy::{ExposeSecret, SecretString};
@@ -15,13 +13,6 @@ use secrecy::{ExposeSecret, SecretString};
use crate::error::ConfigError;
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.
#[derive(Debug, Clone)]
pub struct Config {
@@ -49,9 +40,9 @@ impl Config {
pub async fn from_db(
store: &dyn crate::db::Database,
user_id: &str,
bootstrap: &crate::bootstrap::BootstrapConfig,
) -> Result<Self, ConfigError> {
let _ = dotenvy::dotenv();
crate::bootstrap::load_ironclaw_env();
// Load all settings from DB into a Settings struct
let db_settings = match store.get_all_settings(user_id).await {
@@ -62,7 +53,7 @@ impl Config {
}
};
Self::build(&db_settings).await
Self::build(bootstrap, &db_settings).await
}
/// Load configuration from environment variables only (no database).
@@ -70,20 +61,20 @@ impl Config {
/// Used during early startup before the database is connected,
/// and by CLI commands that don't have DB access.
/// Falls back to legacy `settings.json` on disk if present.
///
/// Loads both `./.env` (standard, higher priority) and `~/.ironclaw/.env`
/// (lower priority) via dotenvy, which never overwrites existing vars.
pub async fn from_env() -> Result<Self, ConfigError> {
let _ = dotenvy::dotenv();
crate::bootstrap::load_ironclaw_env();
let bootstrap = crate::bootstrap::BootstrapConfig::load();
let settings = Settings::load();
Self::build(&settings).await
Self::build(&bootstrap, &settings).await
}
/// Build config from settings (shared by from_env and from_db).
async fn build(settings: &Settings) -> Result<Self, ConfigError> {
/// Build config from bootstrap + settings (shared by from_env and from_db).
async fn build(
bootstrap: &crate::bootstrap::BootstrapConfig,
settings: &Settings,
) -> Result<Self, ConfigError> {
Ok(Self {
database: DatabaseConfig::resolve()?,
database: DatabaseConfig::resolve(bootstrap)?,
llm: LlmConfig::resolve(settings)?,
embeddings: EmbeddingsConfig::resolve(settings)?,
tunnel: TunnelConfig::resolve(settings)?,
@@ -91,7 +82,7 @@ impl Config {
agent: AgentConfig::resolve(settings)?,
safety: SafetyConfig::resolve()?,
wasm: WasmConfig::resolve()?,
secrets: SecretsConfig::resolve().await?,
secrets: SecretsConfig::resolve(bootstrap).await?,
builder: BuilderModeConfig::resolve()?,
heartbeat: HeartbeatConfig::resolve(settings)?,
routines: RoutineConfig::resolve()?,
@@ -116,13 +107,13 @@ impl TunnelConfig {
let public_url = optional_env("TUNNEL_URL")?
.or_else(|| settings.tunnel.public_url.clone().filter(|s| !s.is_empty()));
if let Some(ref url) = public_url
&& !url.starts_with("https://")
{
return Err(ConfigError::InvalidValue {
key: "TUNNEL_URL".to_string(),
message: "must start with https:// (webhooks require HTTPS)".to_string(),
});
if let Some(ref url) = public_url {
if !url.starts_with("https://") {
return Err(ConfigError::InvalidValue {
key: "TUNNEL_URL".to_string(),
message: "must start with https:// (webhooks require HTTPS)".to_string(),
});
}
}
Ok(Self { public_url })
@@ -153,15 +144,6 @@ pub enum DatabaseBackend {
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 {
type Err = String;
@@ -197,7 +179,7 @@ pub struct DatabaseConfig {
}
impl DatabaseConfig {
fn resolve() -> Result<Self, ConfigError> {
fn resolve(bootstrap: &crate::bootstrap::BootstrapConfig) -> Result<Self, ConfigError> {
let backend: DatabaseBackend = if let Some(b) = optional_env("DATABASE_BACKEND")? {
b.parse().map_err(|e| ConfigError::InvalidValue {
key: "DATABASE_BACKEND".to_string(),
@@ -209,8 +191,8 @@ impl DatabaseConfig {
// PostgreSQL URL is required only when using the postgres backend.
// For libsql backend, default to an empty placeholder.
// DATABASE_URL is loaded from ~/.ironclaw/.env via dotenvy early in startup.
let url = optional_env("DATABASE_URL")?
.or_else(|| bootstrap.database_url.clone())
.or_else(|| {
if backend == DatabaseBackend::LibSql {
Some("unused://libsql".to_string())
@@ -223,7 +205,15 @@ impl DatabaseConfig {
hint: "Run 'ironclaw onboard' or set DATABASE_URL environment variable".to_string(),
})?;
let pool_size = parse_optional_env("DATABASE_POOL_SIZE", 10)?;
let pool_size = optional_env("DATABASE_POOL_SIZE")?
.map(|s| s.parse())
.transpose()
.map_err(|e| ConfigError::InvalidValue {
key: "DATABASE_POOL_SIZE".to_string(),
message: format!("must be a positive integer: {e}"),
})?
.or(bootstrap.database_pool_size)
.unwrap_or(10);
let libsql_path = optional_env("LIBSQL_PATH")?.map(PathBuf::from).or_else(|| {
if backend == DatabaseBackend::LibSql {
@@ -397,9 +387,6 @@ impl std::str::FromStr for NearAiApiMode {
pub struct NearAiConfig {
/// Model to use (e.g., "claude-3-5-sonnet-20241022", "gpt-4o")
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)
pub base_url: String,
/// Base URL for auth/refresh endpoints (default: https://private.near.ai)
@@ -410,44 +397,16 @@ pub struct NearAiConfig {
pub api_mode: NearAiApiMode,
/// API key for cloud-api (required for chat_completions mode)
pub api_key: Option<SecretString>,
/// Optional fallback model for failover (default: None).
/// When set, a secondary provider is created with this model and wrapped
/// in a `FailoverProvider` so transient errors on the primary model
/// automatically fall through to the fallback.
pub fallback_model: Option<String>,
/// Maximum number of retries for transient errors (default: 3).
/// With the default of 3, the provider makes up to 4 total attempts
/// (1 initial + 3 retries) before giving up.
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 {
fn resolve(settings: &Settings) -> Result<Self, ConfigError> {
// Determine backend: env var > settings > default (NearAi)
// Determine backend (default: NearAi)
let backend: LlmBackend = if let Some(b) = optional_env("LLM_BACKEND")? {
b.parse().map_err(|e| ConfigError::InvalidValue {
key: "LLM_BACKEND".to_string(),
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 {
LlmBackend::NearAi
};
@@ -473,7 +432,6 @@ impl LlmConfig {
"fireworks::accounts/fireworks/models/llama4-maverick-instruct-basic"
.to_string()
}),
cheap_model: optional_env("NEARAI_CHEAP_MODEL")?,
base_url: optional_env("NEARAI_BASE_URL")?
.unwrap_or_else(|| "https://cloud-api.near.ai".to_string()),
auth_base_url: optional_env("NEARAI_AUTH_URL")?
@@ -483,10 +441,6 @@ impl LlmConfig {
.unwrap_or_else(default_session_path),
api_mode,
api_key: nearai_api_key,
fallback_model: optional_env("NEARAI_FALLBACK_MODEL")?,
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
@@ -519,7 +473,6 @@ impl LlmConfig {
let ollama = if backend == LlmBackend::Ollama {
let base_url = optional_env("OLLAMA_BASE_URL")?
.or_else(|| settings.ollama_base_url.clone())
.unwrap_or_else(|| "http://localhost:11434".to_string());
let model = optional_env("OLLAMA_MODEL")?.unwrap_or_else(|| "llama3".to_string());
Some(OllamaConfig { base_url, model })
@@ -528,9 +481,8 @@ impl LlmConfig {
};
let openai_compatible = if backend == LlmBackend::OpenAiCompatible {
let base_url = optional_env("LLM_BASE_URL")?
.or_else(|| settings.openai_compatible_base_url.clone())
.ok_or_else(|| ConfigError::MissingRequired {
let base_url =
optional_env("LLM_BASE_URL")?.ok_or_else(|| ConfigError::MissingRequired {
key: "LLM_BASE_URL".to_string(),
hint: "Set LLM_BASE_URL when LLM_BACKEND=openai_compatible".to_string(),
})?;
@@ -900,41 +852,52 @@ 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 {
/// Auto-detect secrets master key from env var, then OS keychain.
///
/// Sequential probe: SECRETS_MASTER_KEY env var first, then OS keychain.
/// No saved "source" needed; just try each source in order.
async fn resolve() -> Result<Self, ConfigError> {
async fn resolve(bootstrap: &crate::bootstrap::BootstrapConfig) -> Result<Self, ConfigError> {
use crate::settings::KeySource;
let (master_key, source) = if let Some(env_key) = optional_env("SECRETS_MASTER_KEY")? {
(Some(SecretString::from(env_key)), KeySource::Env)
} else {
// Probe the OS keychain; if a key is stored, use it
match crate::secrets::keychain::get_master_key().await {
Ok(key_bytes) => {
let key_hex: String = key_bytes.iter().map(|b| format!("{:02x}", b)).collect();
(Some(SecretString::from(key_hex)), KeySource::Keychain)
match bootstrap.secrets_master_key_source {
KeySource::Keychain => {
// Try to load from OS keychain (async on Linux)
match crate::secrets::keychain::get_master_key().await {
Ok(key_bytes) => {
let key_hex: String =
key_bytes.iter().map(|b| format!("{:02x}", b)).collect();
(Some(SecretString::from(key_hex)), KeySource::Keychain)
}
Err(_) => {
// Keychain configured but key not found
// This might happen if keychain was cleared
tracing::warn!(
"Secrets configured for keychain but key not found. \
Run 'ironclaw onboard' to reconfigure."
);
(None, KeySource::None)
}
}
}
Err(_) => (None, KeySource::None),
KeySource::Env => {
tracing::warn!(
"Secrets configured for env var but SECRETS_MASTER_KEY not set."
);
(None, KeySource::None)
}
KeySource::None => (None, KeySource::None),
}
};
let enabled = master_key.is_some();
if let Some(ref key) = master_key
&& key.expose_secret().len() < 32
{
return Err(ConfigError::InvalidValue {
key: "SECRETS_MASTER_KEY".to_string(),
message: "must be at least 32 bytes for AES-256-GCM".to_string(),
});
if let Some(ref key) = master_key {
if key.expose_secret().len() < 32 {
return Err(ConfigError::InvalidValue {
key: "SECRETS_MASTER_KEY".to_string(),
message: "must be at least 32 bytes for AES-256-GCM".to_string(),
});
}
}
Ok(Self {
@@ -1388,64 +1351,17 @@ 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
fn optional_env(key: &str) -> Result<Option<String>, ConfigError> {
// Check real env vars first (always win over injected secrets)
match std::env::var(key) {
Ok(val) if val.is_empty() => {}
Ok(val) => return Ok(Some(val)),
Err(std::env::VarError::NotPresent) => {}
Err(e) => {
return Err(ConfigError::ParseError(format!(
"failed to read {key}: {e}"
)));
}
Ok(val) if val.is_empty() => Ok(None),
Ok(val) => Ok(Some(val)),
Err(std::env::VarError::NotPresent) => Ok(None),
Err(e) => Err(ConfigError::ParseError(format!(
"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>
+8 -8
View File
@@ -772,10 +772,10 @@ impl Database for LibSqlBackend {
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
{
if let Ok(id_str) = row.get::<String>(0)
&& let Ok(id) = id_str.parse()
{
ids.push(id);
if let Ok(id_str) = row.get::<String>(0) {
if let Ok(id) = id_str.parse() {
ids.push(id);
}
}
}
Ok(ids)
@@ -2199,10 +2199,10 @@ impl Database for LibSqlBackend {
e.content_preview = None;
}
// Update to latest timestamp
if let (Some(existing), Some(new)) = (&e.updated_at, &updated_at)
&& new > existing
{
e.updated_at = Some(*new);
if let (Some(existing), Some(new)) = (&e.updated_at, &updated_at) {
if new > existing {
e.updated_at = Some(*new);
}
}
})
.or_insert(WorkspaceEntry {
-3
View File
@@ -40,9 +40,6 @@ pub enum Error {
#[error("Workspace error: {0}")]
Workspace(#[from] WorkspaceError),
#[error("Hook error: {0}")]
Hook(#[from] crate::hooks::HookError),
#[error("Orchestrator error: {0}")]
Orchestrator(#[from] OrchestratorError),
+5 -1
View File
@@ -20,6 +20,10 @@ impl CostEstimator {
// Default tool costs (in USD or equivalent)
tool_costs.insert("http".to_string(), dec!(0.0001)); // API call
tool_costs.insert("marketplace".to_string(), dec!(0.01)); // Gas costs
tool_costs.insert("ecommerce".to_string(), dec!(0.001)); // API call
tool_costs.insert("taskrabbit".to_string(), dec!(0.0)); // Cost comes from task itself
tool_costs.insert("restaurant".to_string(), dec!(0.001)); // API call
tool_costs.insert("echo".to_string(), dec!(0.0)); // Free
tool_costs.insert("time".to_string(), dec!(0.0)); // Free
tool_costs.insert("json".to_string(), dec!(0.0)); // Free
@@ -70,7 +74,7 @@ mod tests {
let estimator = CostEstimator::new();
assert_eq!(estimator.estimate_tool("echo"), dec!(0.0));
assert_eq!(estimator.estimate_tool("http"), dec!(0.0001));
assert_eq!(estimator.estimate_tool("marketplace"), dec!(0.01));
assert!(estimator.estimate_tool("unknown") > dec!(0.0));
}
+4
View File
@@ -16,6 +16,10 @@ impl TimeEstimator {
// Default tool durations
tool_durations.insert("http".to_string(), Duration::from_secs(5));
tool_durations.insert("marketplace".to_string(), Duration::from_secs(10));
tool_durations.insert("ecommerce".to_string(), Duration::from_secs(8));
tool_durations.insert("taskrabbit".to_string(), Duration::from_secs(30)); // Just API, not task itself
tool_durations.insert("restaurant".to_string(), Duration::from_secs(5));
tool_durations.insert("echo".to_string(), Duration::from_millis(10));
tool_durations.insert("time".to_string(), Duration::from_millis(1));
tool_durations.insert("json".to_string(), Duration::from_millis(5));
+6 -5
View File
@@ -144,11 +144,12 @@ impl SuccessEvaluator for RuleBasedEvaluator {
// Check for critical errors
for action in actions.iter().filter(|a| !a.success) {
if let Some(ref error) = action.error
&& (error.to_lowercase().contains("critical")
|| error.to_lowercase().contains("fatal"))
{
issues.push(format!("Critical error in {}: {}", action.tool_name, error));
if let Some(ref error) = action.error {
if error.to_lowercase().contains("critical")
|| error.to_lowercase().contains("fatal")
{
issues.push(format!("Critical error in {}: {}", action.tool_name, error));
}
}
}
+27 -27
View File
@@ -492,13 +492,13 @@ impl ExtensionManager {
}
// Check Content-Length header before downloading the full body
if let Some(len) = response.content_length()
&& len as usize > MAX_WASM_SIZE
{
return Err(ExtensionError::InstallFailed(format!(
"WASM binary too large ({} bytes, max {} bytes)",
len, MAX_WASM_SIZE
)));
if let Some(len) = response.content_length() {
if len as usize > MAX_WASM_SIZE {
return Err(ExtensionError::InstallFailed(format!(
"WASM binary too large ({} bytes, max {} bytes)",
len, MAX_WASM_SIZE
)));
}
}
let bytes = response
@@ -768,27 +768,27 @@ impl ExtensionManager {
};
// Check env var first
if let Some(ref env_var) = auth.env_var
&& let Ok(value) = std::env::var(env_var)
{
// Store the env var value as a secret
let params =
CreateSecretParams::new(&auth.secret_name, &value).with_provider(name.to_string());
self.secrets
.create(&self.user_id, params)
.await
.map_err(|e| ExtensionError::AuthFailed(e.to_string()))?;
if let Some(ref env_var) = auth.env_var {
if let Ok(value) = std::env::var(env_var) {
// Store the env var value as a secret
let params = CreateSecretParams::new(&auth.secret_name, &value)
.with_provider(name.to_string());
self.secrets
.create(&self.user_id, params)
.await
.map_err(|e| ExtensionError::AuthFailed(e.to_string()))?;
return Ok(AuthResult {
name: name.to_string(),
kind: ExtensionKind::WasmTool,
auth_url: None,
callback_type: None,
instructions: None,
setup_url: None,
awaiting_token: false,
status: "authenticated".to_string(),
});
return Ok(AuthResult {
name: name.to_string(),
kind: ExtensionKind::WasmTool,
auth_url: None,
callback_type: None,
instructions: None,
setup_url: None,
awaiting_token: false,
status: "authenticated".to_string(),
});
}
}
// Check if already authenticated
-5
View File
@@ -36,11 +36,6 @@ pub struct Store {
#[cfg(feature = "postgres")]
impl Store {
/// Wrap an existing pool (useful when the caller already has a connection).
pub fn from_pool(pool: Pool) -> Self {
Self { pool }
}
/// Create a new store and connect to the database.
pub async fn new(config: &DatabaseConfig) -> Result<Self, DatabaseError> {
let mut cfg = Config::new();
-199
View File
@@ -1,199 +0,0 @@
//! 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>;
}
-19
View File
@@ -1,19 +0,0 @@
//! 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;
-555
View File
@@ -1,555 +0,0 @@
//! 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(&current_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(&current_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());
}
}
-3
View File
@@ -39,7 +39,6 @@
//! - **Continuous learning** - Improve estimates from historical data
pub mod agent;
pub mod boot_screen;
pub mod bootstrap;
pub mod channels;
pub mod cli;
@@ -51,7 +50,6 @@ pub mod estimation;
pub mod evaluation;
pub mod extensions;
pub mod history;
pub mod hooks;
pub mod llm;
pub mod orchestrator;
pub mod pairing;
@@ -61,7 +59,6 @@ pub mod secrets;
pub mod settings;
pub mod setup;
pub mod tools;
pub mod tracing_fmt;
pub mod util;
pub mod worker;
pub mod workspace;
-1017
View File
File diff suppressed because it is too large Load Diff
+12 -127
View File
@@ -8,16 +8,13 @@
//! - **OpenAI-compatible**: Any endpoint that speaks the OpenAI API
mod costs;
pub mod failover;
mod nearai;
mod nearai_chat;
mod provider;
mod reasoning;
mod retry;
mod rig_adapter;
pub mod session;
pub use failover::{CooldownConfig, FailoverProvider};
pub use nearai::{ModelInfo, NearAiProvider};
pub use nearai_chat::NearAiChatProvider;
pub use provider::{
@@ -36,7 +33,7 @@ use std::sync::Arc;
use rig::client::CompletionClient;
use secrecy::ExposeSecret;
use crate::config::{LlmBackend, LlmConfig, NearAiApiMode, NearAiConfig};
use crate::config::{LlmBackend, LlmConfig, NearAiApiMode};
use crate::error::LlmError;
/// Create an LLM provider based on configuration.
@@ -49,7 +46,7 @@ pub fn create_llm_provider(
session: Arc<SessionManager>,
) -> Result<Arc<dyn LlmProvider>, LlmError> {
match config.backend {
LlmBackend::NearAi => create_llm_provider_with_config(&config.nearai, session),
LlmBackend::NearAi => create_nearai_provider(config, session),
LlmBackend::OpenAi => create_openai_provider(config),
LlmBackend::Anthropic => create_anthropic_provider(config),
LlmBackend::Ollama => create_ollama_provider(config),
@@ -57,28 +54,21 @@ pub fn create_llm_provider(
}
}
/// Create an LLM provider from a `NearAiConfig` directly.
///
/// This is useful when constructing additional providers for failover,
/// where only the model name differs from the primary config.
pub fn create_llm_provider_with_config(
config: &NearAiConfig,
fn create_nearai_provider(
config: &LlmConfig,
session: Arc<SessionManager>,
) -> Result<Arc<dyn LlmProvider>, LlmError> {
match config.api_mode {
match config.nearai.api_mode {
NearAiApiMode::Responses => {
tracing::info!(
model = %config.model,
"Using Responses API (chat-api) with session auth"
);
Ok(Arc::new(NearAiProvider::new(config.clone(), session)))
tracing::info!("Using NEAR AI Responses API (chat-api) with session auth");
Ok(Arc::new(NearAiProvider::new(
config.nearai.clone(),
session,
)))
}
NearAiApiMode::ChatCompletions => {
tracing::info!(
model = %config.model,
"Using Chat Completions API (cloud-api) with API key auth"
);
Ok(Arc::new(NearAiChatProvider::new(config.clone())?))
tracing::info!("Using NEAR AI Chat Completions API (cloud-api) with API key auth");
Ok(Arc::new(NearAiChatProvider::new(config.nearai.clone())?))
}
}
}
@@ -183,108 +173,3 @@ fn create_openai_compatible_provider(config: &LlmConfig) -> Result<Arc<dyn LlmPr
);
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());
}
}
+96 -138
View File
@@ -19,7 +19,6 @@ use crate::llm::provider::{
ChatMessage, CompletionRequest, CompletionResponse, FinishReason, LlmProvider, Role, ToolCall,
ToolCompletionRequest, ToolCompletionResponse,
};
use crate::llm::retry::{is_retryable_status, retry_backoff_delay};
use crate::llm::session::SessionManager;
/// Information about an available model from NEAR AI API.
@@ -210,20 +209,20 @@ impl NearAiProvider {
data: Option<Vec<ModelEntry>>,
}
if let Ok(resp) = serde_json::from_str::<ModelsResponse>(&response_text)
&& let Some(entries) = resp.models.or(resp.data)
{
let models: Vec<ModelInfo> = entries
.into_iter()
.filter_map(|e| {
e.get_name().map(|name| ModelInfo {
name,
provider: None,
if let Ok(resp) = serde_json::from_str::<ModelsResponse>(&response_text) {
if let Some(entries) = resp.models.or(resp.data) {
let models: Vec<ModelInfo> = entries
.into_iter()
.filter_map(|e| {
e.get_name().map(|name| ModelInfo {
name,
provider: None,
})
})
})
.collect();
if !models.is_empty() {
return Ok(models);
.collect();
if !models.is_empty() {
return Ok(models);
}
}
}
@@ -271,139 +270,88 @@ impl NearAiProvider {
}
}
/// Inner request implementation with retry logic for transient errors.
///
/// Retries on HTTP 429, 500, 502, 503, 504 with exponential backoff.
/// Does not retry on client errors (400, 401, 403, 404) or parse errors.
/// Inner request implementation without retry logic.
async fn send_request_inner<T: Serialize + std::fmt::Debug, R: for<'de> Deserialize<'de>>(
&self,
path: &str,
body: &T,
) -> Result<R, LlmError> {
let url = self.api_url(path);
let max_retries = self.config.max_retries;
let token = self.session.get_token().await?;
for attempt in 0..=max_retries {
let token = self.session.get_token().await?;
tracing::debug!("Sending request to NEAR AI: {}", url);
tracing::debug!("Request body: {:?}", body);
tracing::debug!(
"Sending request to NEAR AI: {} (attempt {})",
url,
attempt + 1
);
tracing::debug!("Request body: {:?}", body);
let response = self
.client
.post(&url)
.header("Authorization", format!("Bearer {}", token.expose_secret()))
.header("Content-Type", "application/json")
.json(body)
.send()
.await
.map_err(|e| {
tracing::error!("NEAR AI request failed: {}", e);
e
})?;
let response = self
.client
.post(&url)
.header("Authorization", format!("Bearer {}", token.expose_secret()))
.header("Content-Type", "application/json")
.json(body)
.send()
.await;
let status = response.status();
let response_text = response.text().await.unwrap_or_default();
let response = match response {
Ok(r) => r,
Err(e) => {
tracing::error!("NEAR AI request failed: {}", e);
// Network errors (timeout, connection refused) are transient
if attempt < max_retries {
let delay = retry_backoff_delay(attempt);
tracing::warn!(
"NEAR AI request error (attempt {}/{}), retrying in {:?}: {}",
attempt + 1,
max_retries + 1,
delay,
e,
);
tokio::time::sleep(delay).await;
continue;
}
return Err(e.into());
}
};
tracing::debug!("NEAR AI response status: {}", status);
tracing::debug!("NEAR AI response body: {}", response_text);
let status = response.status();
let response_text = response.text().await.unwrap_or_default();
if !status.is_success() {
// Check for session expiration (401 with specific message patterns)
if status.as_u16() == 401 {
let is_session_expired = response_text.to_lowercase().contains("session")
&& (response_text.to_lowercase().contains("expired")
|| response_text.to_lowercase().contains("invalid"));
tracing::debug!("NEAR AI response status: {}", status);
tracing::debug!("NEAR AI response body: {}", response_text);
if !status.is_success() {
let status_code = status.as_u16();
// Check for session expiration (401 with specific message patterns)
if status_code == 401 {
let lower = response_text.to_lowercase();
let is_session_expired = lower.contains("session")
&& (lower.contains("expired") || lower.contains("invalid"));
if is_session_expired {
return Err(LlmError::SessionExpired {
provider: "nearai".to_string(),
});
}
// Generic 401 -- not retryable
return Err(LlmError::AuthFailed {
if is_session_expired {
return Err(LlmError::SessionExpired {
provider: "nearai".to_string(),
});
}
// Check if this is a transient error worth retrying
if is_retryable_status(status_code) && attempt < max_retries {
let delay = retry_backoff_delay(attempt);
tracing::warn!(
"NEAR AI returned HTTP {} (attempt {}/{}), retrying in {:?}",
status_code,
attempt + 1,
max_retries + 1,
delay,
);
tokio::time::sleep(delay).await;
continue;
}
// Non-retryable error or exhausted retries
if let Ok(error) = serde_json::from_str::<NearAiErrorResponse>(&response_text) {
if status_code == 429 {
return Err(LlmError::RateLimited {
provider: "nearai".to_string(),
retry_after: None,
});
}
return Err(LlmError::RequestFailed {
provider: "nearai".to_string(),
reason: error.error,
});
}
return Err(LlmError::RequestFailed {
// Generic 401 without session expiration indication
return Err(LlmError::AuthFailed {
provider: "nearai".to_string(),
reason: format!("HTTP {}: {}", status, response_text),
});
}
// Success -- parse the response
return match serde_json::from_str::<R>(&response_text) {
Ok(parsed) => Ok(parsed),
Err(e) => {
tracing::debug!("Response is not expected JSON format: {}", e);
tracing::debug!("Will try alternative parsing in caller");
Err(LlmError::InvalidResponse {
// Try to parse as JSON error
if let Ok(error) = serde_json::from_str::<NearAiErrorResponse>(&response_text) {
if status.as_u16() == 429 {
return Err(LlmError::RateLimited {
provider: "nearai".to_string(),
reason: format!("Parse error: {}. Raw: {}", e, response_text),
})
retry_after: None,
});
}
};
return Err(LlmError::RequestFailed {
provider: "nearai".to_string(),
reason: error.error,
});
}
return Err(LlmError::RequestFailed {
provider: "nearai".to_string(),
reason: format!("HTTP {}: {}", status, response_text),
});
}
// This is unreachable because the loop always returns, but the compiler
// cannot prove that. Return a generic error as a safety net.
Err(LlmError::RequestFailed {
provider: "nearai".to_string(),
reason: "retry loop exited unexpectedly".to_string(),
})
// Try to parse as our expected type
match serde_json::from_str::<R>(&response_text) {
Ok(parsed) => Ok(parsed),
Err(e) => {
tracing::debug!("Response is not expected JSON format: {}", e);
tracing::debug!("Will try alternative parsing in caller");
Err(LlmError::InvalidResponse {
provider: "nearai".to_string(),
reason: format!("Parse error: {}. Raw: {}", e, response_text),
})
}
}
}
}
@@ -508,7 +456,7 @@ impl LlmProvider for NearAiProvider {
Err(e) => return Err(e),
};
tracing::debug!("NEAR AI response: output_items={}", response.output.len());
tracing::debug!("NEAR AI response: {:?}", response);
// Extract text from response output
// Try multiple formats since API response shape may vary
@@ -516,6 +464,11 @@ impl LlmProvider for NearAiProvider {
.output
.iter()
.filter_map(|item| {
tracing::debug!(
"Processing output item: type={}, text={:?}",
item.item_type,
item.text
);
if item.item_type == "message" {
// First check for direct text field on item
if let Some(ref text) = item.text {
@@ -526,6 +479,11 @@ impl LlmProvider for NearAiProvider {
contents
.iter()
.filter_map(|c| {
tracing::debug!(
"Content item: type={}, text={:?}",
c.content_type,
c.text
);
// Accept various content types that might contain text
match c.content_type.as_str() {
"output_text" | "text" => c.text.clone(),
@@ -736,21 +694,21 @@ impl LlmProvider for NearAiProvider {
}
}
}
} else if item.item_type == "function_call"
&& let (Some(name), Some(call_id)) = (&item.name, &item.call_id)
{
// Parse arguments JSON string into Value
let arguments = item
.arguments
.as_ref()
.and_then(|s| serde_json::from_str(s).ok())
.unwrap_or(serde_json::Value::Object(Default::default()));
} else if item.item_type == "function_call" {
if let (Some(name), Some(call_id)) = (&item.name, &item.call_id) {
// Parse arguments JSON string into Value
let arguments = item
.arguments
.as_ref()
.and_then(|s| serde_json::from_str(s).ok())
.unwrap_or(serde_json::Value::Object(Default::default()));
tool_calls.push(ToolCall {
id: call_id.clone(),
name: name.clone(),
arguments,
});
tool_calls.push(ToolCall {
id: call_id.clone(),
name: name.clone(),
arguments,
});
}
}
}
+41 -95
View File
@@ -16,7 +16,6 @@ use crate::llm::provider::{
ChatMessage, CompletionRequest, CompletionResponse, FinishReason, LlmProvider, ModelMetadata,
Role, ToolCall, ToolCompletionRequest, ToolCompletionResponse,
};
use crate::llm::retry::{is_retryable_status, retry_backoff_delay};
/// NEAR AI Chat Completions API provider.
pub struct NearAiChatProvider {
@@ -63,116 +62,63 @@ impl NearAiChatProvider {
.unwrap_or_default()
}
/// Send a request to the chat completions API with retry on transient errors.
///
/// Retries on HTTP 429, 500, 502, 503, 504 with exponential backoff.
/// Does not retry on client errors (400, 401, 403, 404) or parse errors.
/// Send a request to the chat completions API.
async fn send_request<T: Serialize, R: for<'de> Deserialize<'de>>(
&self,
body: &T,
) -> Result<R, LlmError> {
let url = self.api_url("chat/completions");
let max_retries = self.config.max_retries;
for attempt in 0..=max_retries {
tracing::debug!(
"Sending request to NEAR AI Chat: {} (attempt {})",
url,
attempt + 1,
);
tracing::debug!("Sending request to NEAR AI Chat: {}", url);
if tracing::enabled!(tracing::Level::DEBUG)
&& let Ok(json) = serde_json::to_string(body)
{
tracing::debug!("NEAR AI Chat request body: {}", json);
}
// Log the request body for debugging tool call issues
if let Ok(json) = serde_json::to_string(body) {
tracing::debug!("NEAR AI Chat request body: {}", json);
}
let response = self
.client
.post(&url)
.header("Authorization", format!("Bearer {}", self.api_key()))
.header("Content-Type", "application/json")
.json(body)
.send()
.await;
let response = match response {
Ok(r) => r,
Err(e) => {
tracing::error!("NEAR AI Chat request failed: {}", e);
if attempt < max_retries {
let delay = retry_backoff_delay(attempt);
tracing::warn!(
"NEAR AI Chat request error (attempt {}/{}), retrying in {:?}: {}",
attempt + 1,
max_retries + 1,
delay,
e,
);
tokio::time::sleep(delay).await;
continue;
}
return Err(LlmError::RequestFailed {
provider: "nearai_chat".to_string(),
reason: e.to_string(),
});
let response = self
.client
.post(&url)
.header("Authorization", format!("Bearer {}", self.api_key()))
.header("Content-Type", "application/json")
.json(body)
.send()
.await
.map_err(|e| {
tracing::error!("NEAR AI Chat request failed: {}", e);
LlmError::RequestFailed {
provider: "nearai_chat".to_string(),
reason: e.to_string(),
}
};
let status = response.status();
let response_text = response.text().await.unwrap_or_default();
})?;
tracing::debug!("NEAR AI Chat response status: {}", status);
tracing::debug!("NEAR AI Chat response body: {}", response_text);
let status = response.status();
let response_text = response.text().await.unwrap_or_default();
if !status.is_success() {
let status_code = status.as_u16();
// Auth errors are not retryable
if status_code == 401 {
return Err(LlmError::AuthFailed {
provider: "nearai_chat".to_string(),
});
}
tracing::debug!("NEAR AI Chat response status: {}", status);
tracing::debug!("NEAR AI Chat response body: {}", response_text);
// Transient errors: retry with backoff
if is_retryable_status(status_code) && attempt < max_retries {
let delay = retry_backoff_delay(attempt);
tracing::warn!(
"NEAR AI Chat returned HTTP {} (attempt {}/{}), retrying in {:?}",
status_code,
attempt + 1,
max_retries + 1,
delay,
);
tokio::time::sleep(delay).await;
continue;
}
// Non-retryable or exhausted retries
if status_code == 429 {
return Err(LlmError::RateLimited {
provider: "nearai_chat".to_string(),
retry_after: None,
});
}
return Err(LlmError::RequestFailed {
if !status.is_success() {
if status.as_u16() == 401 {
return Err(LlmError::AuthFailed {
provider: "nearai_chat".to_string(),
reason: format!("HTTP {}: {}", status, response_text),
});
}
// Success — parse the response
return serde_json::from_str(&response_text).map_err(|e| LlmError::InvalidResponse {
if status.as_u16() == 429 {
return Err(LlmError::RateLimited {
provider: "nearai_chat".to_string(),
retry_after: None,
});
}
return Err(LlmError::RequestFailed {
provider: "nearai_chat".to_string(),
reason: format!("JSON parse error: {}. Raw: {}", e, response_text),
reason: format!("HTTP {}: {}", status, response_text),
});
}
// Safety net: unreachable because the loop always returns
Err(LlmError::RequestFailed {
serde_json::from_str(&response_text).map_err(|e| LlmError::InvalidResponse {
provider: "nearai_chat".to_string(),
reason: "retry loop exited unexpectedly".to_string(),
reason: format!("JSON parse error: {}. Raw: {}", e, response_text),
})
}
@@ -449,10 +395,10 @@ fn flatten_tool_messages(messages: Vec<ChatCompletionMessage>) -> Vec<ChatComple
if let (true, Some(calls)) = (msg.role == "assistant", &msg.tool_calls) {
// Convert assistant tool_calls into descriptive text
let mut parts: Vec<String> = Vec::new();
if let Some(ref text) = msg.content
&& !text.is_empty()
{
parts.push(text.clone());
if let Some(ref text) = msg.content {
if !text.is_empty() {
parts.push(text.clone());
}
}
for tc in calls {
parts.push(format!(
+15 -21
View File
@@ -113,12 +113,6 @@ pub struct ToolSelection {
pub reasoning: String,
/// Alternative tools considered.
pub alternatives: Vec<String>,
/// The tool call ID from the LLM response.
///
/// OpenAI-compatible providers assign each tool call a unique ID that must
/// be echoed back in the corresponding tool result message. Without this,
/// the provider cannot match results to their originating calls.
pub tool_call_id: String,
}
/// Token usage from a single LLM call.
@@ -250,7 +244,6 @@ impl Reasoning {
parameters: tool_call.arguments,
reasoning: reasoning.clone(),
alternatives: vec![],
tool_call_id: tool_call.id,
})
.collect();
@@ -588,20 +581,21 @@ fn recover_tool_calls_from_content(
}
// Try JSON first: {"name":"x","arguments":{}}
if let Ok(parsed) = serde_json::from_str::<serde_json::Value>(inner)
&& let Some(name) = parsed.get("name").and_then(|v| v.as_str())
&& tool_names.contains(name)
{
let arguments = parsed
.get("arguments")
.cloned()
.unwrap_or(serde_json::Value::Object(Default::default()));
calls.push(ToolCall {
id: format!("recovered_{}", calls.len()),
name: name.to_string(),
arguments,
});
continue;
if let Ok(parsed) = serde_json::from_str::<serde_json::Value>(inner) {
if let Some(name) = parsed.get("name").and_then(|v| v.as_str()) {
if tool_names.contains(name) {
let arguments = parsed
.get("arguments")
.cloned()
.unwrap_or(serde_json::Value::Object(Default::default()));
calls.push(ToolCall {
id: format!("recovered_{}", calls.len()),
name: name.to_string(),
arguments,
});
continue;
}
}
}
// Bare tool name (e.g. "<tool_call>tool_list</tool_call>")
-96
View File
@@ -1,96 +0,0 @@
//! Shared retry helpers for LLM providers.
//!
//! Provides exponential backoff with jitter and retryable status classification
//! used by both `NearAiProvider` and `NearAiChatProvider`.
use std::time::Duration;
use rand::Rng;
/// Returns `true` if the HTTP status code is transient and worth retrying.
pub(crate) fn is_retryable_status(status: u16) -> bool {
matches!(status, 429 | 500 | 502 | 503 | 504)
}
/// Calculate exponential backoff delay with random jitter.
///
/// Base delay is 1 second, doubled each attempt, with +/-25% jitter.
/// - attempt 0: ~1s (0.75s - 1.25s)
/// - attempt 1: ~2s (1.5s - 2.5s)
/// - attempt 2: ~4s (3.0s - 5.0s)
pub(crate) fn retry_backoff_delay(attempt: u32) -> Duration {
let base_ms: u64 = 1000u64.saturating_mul(2u64.saturating_pow(attempt));
let jitter_range = base_ms / 4; // 25%
let jitter = if jitter_range > 0 {
let offset = rand::thread_rng().gen_range(0..=jitter_range * 2);
offset as i64 - jitter_range as i64
} else {
0
};
let delay_ms = (base_ms as i64 + jitter).max(100) as u64;
Duration::from_millis(delay_ms)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_is_retryable_status() {
// Transient errors should be retryable
assert!(is_retryable_status(429));
assert!(is_retryable_status(500));
assert!(is_retryable_status(502));
assert!(is_retryable_status(503));
assert!(is_retryable_status(504));
// Client errors should not be retryable
assert!(!is_retryable_status(400));
assert!(!is_retryable_status(401));
assert!(!is_retryable_status(403));
assert!(!is_retryable_status(404));
assert!(!is_retryable_status(422));
// Success codes should not be retryable
assert!(!is_retryable_status(200));
assert!(!is_retryable_status(201));
}
#[test]
fn test_retry_backoff_delay_exponential_growth() {
// Run multiple samples to verify the range, accounting for jitter
for _ in 0..20 {
let d0 = retry_backoff_delay(0);
let d1 = retry_backoff_delay(1);
let d2 = retry_backoff_delay(2);
// Attempt 0: base 1000ms, jitter +/-250ms -> [750, 1250]
assert!(d0.as_millis() >= 750, "attempt 0 too low: {:?}", d0);
assert!(d0.as_millis() <= 1250, "attempt 0 too high: {:?}", d0);
// Attempt 1: base 2000ms, jitter +/-500ms -> [1500, 2500]
assert!(d1.as_millis() >= 1500, "attempt 1 too low: {:?}", d1);
assert!(d1.as_millis() <= 2500, "attempt 1 too high: {:?}", d1);
// Attempt 2: base 4000ms, jitter +/-1000ms -> [3000, 5000]
assert!(d2.as_millis() >= 3000, "attempt 2 too low: {:?}", d2);
assert!(d2.as_millis() <= 5000, "attempt 2 too high: {:?}", d2);
}
}
#[test]
fn test_retry_backoff_delay_minimum() {
// Even at attempt 0, delay should be at least 100ms (the minimum floor)
for _ in 0..20 {
let delay = retry_backoff_delay(0);
assert!(delay.as_millis() >= 100);
}
}
#[test]
fn test_retry_backoff_delay_no_overflow() {
// Very high attempt numbers should not panic from overflow
let delay = retry_backoff_delay(30);
assert!(delay.as_millis() >= 100);
}
}
+186 -54
View File
@@ -31,6 +31,8 @@ pub struct SessionConfig {
pub auth_base_url: String,
/// Path to session file (e.g., ~/.ironclaw/session.json).
pub session_path: PathBuf,
/// Port range for OAuth callback server.
pub callback_port_range: (u16, u16),
}
impl Default for SessionConfig {
@@ -38,6 +40,7 @@ impl Default for SessionConfig {
Self {
auth_base_url: "https://private.near.ai".to_string(),
session_path: default_session_path(),
callback_port_range: (9876, 9886),
}
}
}
@@ -80,16 +83,16 @@ impl SessionManager {
};
// Try to load existing session synchronously during construction
if let Ok(data) = std::fs::read_to_string(&manager.config.session_path)
&& let Ok(session) = serde_json::from_str::<SessionData>(&data)
{
// We can't await here, so we use try_write
if let Ok(mut guard) = manager.token.try_write() {
*guard = Some(SecretString::from(session.session_token));
tracing::info!(
"Loaded session token from {}",
manager.config.session_path.display()
);
if let Ok(data) = std::fs::read_to_string(&manager.config.session_path) {
if let Ok(session) = serde_json::from_str::<SessionData>(&data) {
// We can't await here, so we use try_write
if let Ok(mut guard) = manager.token.try_write() {
*guard = Some(SecretString::from(session.session_token));
tracing::info!(
"Loaded session token from {}",
manager.config.session_path.display()
);
}
}
}
@@ -219,21 +222,38 @@ impl SessionManager {
/// Start the OAuth login flow.
///
/// 1. Bind the fixed callback port
/// 1. Find an available port for the callback server
/// 2. Print the auth URL and attempt to open browser
/// 3. Wait for OAuth callback with session token
/// 4. Save and return the token
async fn initiate_login(&self) -> Result<(), LlmError> {
use crate::cli::oauth_defaults::{self, OAUTH_CALLBACK_PORT};
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::net::TcpListener;
let listener = oauth_defaults::bind_callback_listener()
.await
.map_err(|e| LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: e.to_string(),
})?;
// Find an available port
let mut listener = None;
let mut port = 0;
let callback_url = format!("http://127.0.0.1:{}", OAUTH_CALLBACK_PORT);
for p in self.config.callback_port_range.0..=self.config.callback_port_range.1 {
match TcpListener::bind(format!("127.0.0.1:{}", p)).await {
Ok(l) => {
listener = Some(l);
port = p;
break;
}
Err(_) => continue,
}
}
let listener = listener.ok_or_else(|| LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: format!(
"Could not find available port in range {}-{}",
self.config.callback_port_range.0, self.config.callback_port_range.1
),
})?;
let callback_url = format!("http://127.0.0.1:{}", port);
// Show auth provider menu
println!();
@@ -313,16 +333,138 @@ impl SessionManager {
println!();
println!("Waiting for authentication...");
// The NEAR AI API redirects to: {frontend_callback}/auth/callback?token=X&...
let session_token =
oauth_defaults::wait_for_callback(listener, "/auth/callback", "token", "NEAR AI")
.await
.map_err(|e| LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: e.to_string(),
// Wait for callback with timeout
// The API redirects to: {frontend_callback}/auth/callback?token=X&session_id=X&expires_at=X&is_new_user=X
let timeout = std::time::Duration::from_secs(300); // 5 minutes
let selected_provider = auth_provider.to_string();
let (session_token, auth_provider) = tokio::time::timeout(timeout, async move {
loop {
let (mut socket, _) = listener.accept().await.map_err(|e| {
LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: format!("Failed to accept connection: {}", e),
}
})?;
let auth_provider = Some(auth_provider.to_string());
let mut reader = BufReader::new(&mut socket);
let mut request_line = String::new();
reader.read_line(&mut request_line).await.map_err(|e| {
LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: format!("Failed to read request: {}", e),
}
})?;
// Parse GET /auth/callback?token=xxx&session_id=xxx&expires_at=xxx&is_new_user=xxx HTTP/1.1
if let Some(path) = request_line.split_whitespace().nth(1) {
if path.starts_with("/auth/callback") {
// Parse query parameters
if let Some(query) = path.split('?').nth(1) {
let mut token = None;
for param in query.split('&') {
let parts: Vec<&str> = param.splitn(2, '=').collect();
if parts.len() == 2 && parts[0] == "token" {
token = Some(
urlencoding::decode(parts[1])
.unwrap_or_else(|_| parts[1].into())
.into_owned(),
);
}
}
if let Some(token) = token {
// Send success response with nice styling
let response = concat!(
"HTTP/1.1 200 OK\r\n",
"Content-Type: text/html; charset=utf-8\r\n",
"Connection: close\r\n",
"\r\n",
"<!DOCTYPE html>\n",
"<html>\n",
"<head>\n",
" <meta charset=\"utf-8\">\n",
" <title>NEAR AI - Authentication Successful</title>\n",
" <style>\n",
" * { margin: 0; padding: 0; box-sizing: border-box; }\n",
" body {\n",
" font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', Roboto, sans-serif;\n",
" background: linear-gradient(135deg, #1a1a2e 0%, #16213e 100%);\n",
" min-height: 100vh;\n",
" display: flex;\n",
" align-items: center;\n",
" justify-content: center;\n",
" color: #fff;\n",
" }\n",
" .container {\n",
" text-align: center;\n",
" padding: 3rem;\n",
" background: rgba(255,255,255,0.05);\n",
" border-radius: 16px;\n",
" backdrop-filter: blur(10px);\n",
" border: 1px solid rgba(255,255,255,0.1);\n",
" max-width: 400px;\n",
" }\n",
" .checkmark {\n",
" width: 80px;\n",
" height: 80px;\n",
" background: linear-gradient(135deg, #00d9a5 0%, #00b386 100%);\n",
" border-radius: 50%;\n",
" display: flex;\n",
" align-items: center;\n",
" justify-content: center;\n",
" margin: 0 auto 1.5rem;\n",
" font-size: 40px;\n",
" }\n",
" h1 {\n",
" font-size: 1.5rem;\n",
" font-weight: 600;\n",
" margin-bottom: 0.75rem;\n",
" }\n",
" p {\n",
" color: rgba(255,255,255,0.7);\n",
" font-size: 0.95rem;\n",
" line-height: 1.5;\n",
" }\n",
" .brand {\n",
" margin-top: 2rem;\n",
" padding-top: 1.5rem;\n",
" border-top: 1px solid rgba(255,255,255,0.1);\n",
" font-size: 0.8rem;\n",
" color: rgba(255,255,255,0.4);\n",
" }\n",
" </style>\n",
"</head>\n",
"<body>\n",
" <div class=\"container\">\n",
" <div class=\"checkmark\">&#10003;</div>\n",
" <h1>Authentication Successful</h1>\n",
" <p>You can close this window and return to the terminal.</p>\n",
" <div class=\"brand\">NEAR AI Agent</div>\n",
" </div>\n",
"</body>\n",
"</html>"
);
let _ = socket.write_all(response.as_bytes()).await;
let _ = socket.shutdown().await;
return Ok::<_, LlmError>((token, Some(selected_provider.clone())));
}
}
}
}
// Not the callback we're looking for, send 404
let response = "HTTP/1.1 404 Not Found\r\nConnection: close\r\n\r\n";
let _ = socket.write_all(response.as_bytes()).await;
}
})
.await
.map_err(|_| LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: "Authentication timed out after 5 minutes".to_string(),
})??;
// Save the token
self.save_session(&session_token, auth_provider.as_deref())
@@ -428,30 +570,17 @@ impl SessionManager {
})?;
let user_id = self.user_id.read().await.clone();
let value = if let Some(value) = store
let value = store
.get_setting(&user_id, "nearai.session_token")
.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 {
provider: "nearai".to_string(),
reason: format!("DB query failed: {}", e),
})?
.ok_or(LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: "No session in DB".to_string(),
})?
};
})?
.ok_or_else(|| LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: "No session in DB".to_string(),
})?;
let session: SessionData =
serde_json::from_value(value).map_err(|e| LlmError::SessionRenewalFailed {
@@ -513,14 +642,15 @@ pub async fn create_session_manager(config: SessionConfig) -> Arc<SessionManager
let manager = SessionManager::new_async(config).await;
// Check for legacy env var and migrate if present and no file token
if !manager.has_token().await
&& let Ok(token) = std::env::var("NEARAI_SESSION_TOKEN")
&& !token.is_empty()
{
tracing::info!("Migrating session token from NEARAI_SESSION_TOKEN env var to file");
manager.set_token(SecretString::from(token.clone())).await;
if let Err(e) = manager.save_session(&token, None).await {
tracing::warn!("Failed to save migrated session: {}", e);
if !manager.has_token().await {
if let Ok(token) = std::env::var("NEARAI_SESSION_TOKEN") {
if !token.is_empty() {
tracing::info!("Migrating session token from NEARAI_SESSION_TOKEN env var to file");
manager.set_token(SecretString::from(token.clone())).await;
if let Err(e) = manager.save_session(&token, None).await {
tracing::warn!("Failed to save migrated session: {}", e);
}
}
}
}
@@ -541,6 +671,7 @@ mod tests {
let config = SessionConfig {
auth_base_url: "https://example.com".to_string(),
session_path: session_path.clone(),
callback_port_range: (9900, 9910),
};
let manager = SessionManager::new_async(config.clone()).await;
@@ -581,6 +712,7 @@ mod tests {
let config = SessionConfig {
auth_base_url: "https://example.com".to_string(),
session_path: dir.path().join("nonexistent.json"),
callback_port_range: (9900, 9910),
};
let manager = SessionManager::new_async(config).await;
+118 -236
View File
@@ -22,11 +22,7 @@ use ironclaw::{
config::Config,
context::ContextManager,
extensions::ExtensionManager,
hooks::HookRegistry,
llm::{
CooldownConfig, FailoverProvider, LlmProvider, SessionConfig, create_cheap_llm_provider,
create_llm_provider, create_llm_provider_with_config, create_session_manager,
},
llm::{SessionConfig, create_llm_provider, create_session_manager},
orchestrator::{
ContainerJobConfig, ContainerJobManager, OrchestratorApi, TokenStore,
api::OrchestratorState,
@@ -49,6 +45,7 @@ use ironclaw::secrets::PostgresSecretsStore;
use ironclaw::secrets::SecretsCrypto;
#[cfg(any(feature = "postgres", feature = "libsql"))]
use ironclaw::setup::{SetupConfig, SetupWizard};
#[tokio::main]
async fn main() -> anyhow::Result<()> {
let cli = Cli::parse();
@@ -93,6 +90,7 @@ async fn main() -> anyhow::Result<()> {
.init();
// Memory commands need database (and optionally embeddings)
let _ = dotenvy::dotenv();
let config = Config::from_env()
.await
.map_err(|e| anyhow::anyhow!("{}", e))?;
@@ -101,6 +99,7 @@ async fn main() -> anyhow::Result<()> {
let session = ironclaw::llm::create_session_manager(ironclaw::llm::SessionConfig {
auth_base_url: config.llm.nearai.auth_base_url.clone(),
session_path: config.llm.nearai.session_path.clone(),
..Default::default()
})
.await;
@@ -153,6 +152,7 @@ async fn main() -> anyhow::Result<()> {
return run_pairing_command(pairing_cmd.clone()).map_err(|e| anyhow::anyhow!("{}", e));
}
Some(Command::Status) => {
let _ = dotenvy::dotenv();
tracing_subscriber::fmt()
.with_env_filter(
EnvFilter::try_from_default_env().unwrap_or_else(|_| EnvFilter::new("warn")),
@@ -243,10 +243,8 @@ async fn main() -> anyhow::Result<()> {
skip_auth,
channels_only,
}) => {
// Load .env files before running onboarding wizard.
// Standard ./.env first (higher priority), then ~/.ironclaw/.env.
// Load .env before running onboarding wizard
let _ = dotenvy::dotenv();
ironclaw::bootstrap::load_ironclaw_env();
#[cfg(any(feature = "postgres", feature = "libsql"))]
{
@@ -269,23 +267,23 @@ async fn main() -> anyhow::Result<()> {
}
}
// Load .env files early so DATABASE_URL (and any other vars) are
// available to all subsequent env-based config resolution.
// Standard ./.env first (higher priority), then ~/.ironclaw/.env.
// Load .env if present
let _ = dotenvy::dotenv();
ironclaw::bootstrap::load_ironclaw_env();
// Enhanced first-run detection
#[cfg(any(feature = "postgres", feature = "libsql"))]
if !cli.no_onboard
&& let Some(reason) = check_onboard_needed()
{
println!("Onboarding needed: {}", reason);
println!();
let mut wizard = SetupWizard::new();
wizard.run().await?;
if !cli.no_onboard {
if let Some(reason) = check_onboard_needed().await {
println!("Onboarding needed: {}", reason);
println!();
let mut wizard = SetupWizard::new();
wizard.run().await?;
}
}
// Load bootstrap config (4 fields that must live on disk)
let bootstrap = ironclaw::bootstrap::BootstrapConfig::load();
// Load initial config from env + disk (before DB is available)
let mut config = match Config::from_env().await {
Ok(c) => c,
@@ -305,20 +303,18 @@ async fn main() -> anyhow::Result<()> {
let session_config = SessionConfig {
auth_base_url: config.llm.nearai.auth_base_url.clone(),
session_path: config.llm.nearai.session_path.clone(),
..Default::default()
};
let session = create_session_manager(session_config).await;
// Session-based auth is only needed for NEAR AI backend without an API key.
// 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()
{
// Ensure we're authenticated before proceeding (only needed for NEAR AI backend)
if config.llm.backend == ironclaw::config::LlmBackend::NearAi {
session.ensure_authenticated().await?;
}
// Initialize tracing
let env_filter = EnvFilter::try_from_default_env()
.unwrap_or_else(|_| EnvFilter::new("ironclaw=info,tower_http=warn"));
.unwrap_or_else(|_| EnvFilter::new("ironclaw=info,tower_http=debug"));
// Create log broadcaster before tracing init so the WebLogLayer can capture all events.
// This gets wired to the gateway's /api/logs/events SSE endpoint later.
@@ -326,11 +322,7 @@ async fn main() -> anyhow::Result<()> {
tracing_subscriber::registry()
.with(env_filter)
.with(
tracing_subscriber::fmt::layer()
.with_target(false)
.with_writer(ironclaw::tracing_fmt::TruncatingStderr::default()),
)
.with(tracing_subscriber::fmt::layer().with_target(false))
.with(WebLogLayer::new(Arc::clone(&log_broadcaster)))
.init();
@@ -338,10 +330,7 @@ async fn main() -> anyhow::Result<()> {
let repl_channel = if let Some(ref msg) = cli.message {
Some(ReplChannel::with_message(msg.clone()))
} else if config.channels.cli.enabled {
let repl = ReplChannel::new();
// Suppress the one-liner banner; boot screen will be shown instead.
repl.suppress_banner();
Some(repl)
Some(ReplChannel::new())
} else {
None
};
@@ -429,7 +418,7 @@ async fn main() -> anyhow::Result<()> {
}
// Reload config from DB now that we have a connection.
match Config::from_db(db.as_ref(), "default").await {
match Config::from_db(db.as_ref(), "default", &bootstrap).await {
Ok(db_config) => {
config = db_config;
tracing::info!("Configuration reloaded from database");
@@ -450,112 +439,10 @@ async fn main() -> anyhow::Result<()> {
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)
let llm = create_llm_provider(&config.llm, session.clone())?;
tracing::info!("LLM provider initialized: {}", llm.model_name());
// Wrap in failover if a fallback model is configured
let llm: Arc<dyn LlmProvider> =
if let Some(fallback_model) = config.llm.nearai.fallback_model.as_ref() {
if fallback_model == &config.llm.nearai.model {
tracing::warn!(
"fallback_model is the same as primary model, failover may not be effective"
);
}
let mut fallback_config = config.llm.nearai.clone();
fallback_config.model = fallback_model.clone();
let fallback = create_llm_provider_with_config(&fallback_config, session.clone())?;
tracing::info!(
primary = %llm.model_name(),
fallback = %fallback.model_name(),
"LLM failover enabled"
);
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 {
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
let safety = Arc::new(SafetyLayer::new(&config.safety));
tracing::info!("Safety layer initialized");
@@ -629,6 +516,49 @@ async fn main() -> anyhow::Result<()> {
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());
// Create WASM tool runtime (sync, just builds the wasmtime engine)
@@ -649,10 +579,7 @@ async fn main() -> anyhow::Result<()> {
// Both register into the shared ToolRegistry (RwLock-based) so concurrent writes are safe.
let wasm_tools_future = async {
if let Some(ref runtime) = wasm_tool_runtime {
let mut loader = WasmToolLoader::new(Arc::clone(runtime), Arc::clone(&tools));
if let Some(ref secrets) = secrets_store {
loader = loader.with_secrets_store(Arc::clone(secrets));
}
let loader = WasmToolLoader::new(Arc::clone(runtime), Arc::clone(&tools));
// Load installed tools from ~/.ironclaw/tools/
match loader.load_from_dir(&config.wasm.tools_dir).await {
@@ -898,14 +825,12 @@ async fn main() -> anyhow::Result<()> {
// Initialize channel manager
let mut channels = ChannelManager::new();
let mut channel_names: Vec<String> = Vec::new();
if let Some(repl) = repl_channel {
channels.add(Box::new(repl));
if cli.message.is_some() {
tracing::info!("Single message mode");
} else {
channel_names.push("repl".to_string());
tracing::info!("REPL mode enabled");
}
}
@@ -977,13 +902,13 @@ async fn main() -> anyhow::Result<()> {
// Inject owner_id for Telegram so the bot only responds
// to the bound user account.
if channel_name == "telegram"
&& let Some(owner_id) = config.channels.telegram_owner_id
{
config_updates.insert(
"owner_id".to_string(),
serde_json::json!(owner_id),
);
if channel_name == "telegram" {
if let Some(owner_id) = config.channels.telegram_owner_id {
config_updates.insert(
"owner_id".to_string(),
serde_json::json!(owner_id),
);
}
}
if !config_updates.is_empty() {
@@ -1041,7 +966,6 @@ async fn main() -> anyhow::Result<()> {
}
}
channel_names.push(channel_name.clone());
channels.add(Box::new(SharedWasmChannel::new(channel_arc)));
}
@@ -1075,24 +999,23 @@ async fn main() -> anyhow::Result<()> {
// Extract its routes for the unified server; the channel itself just
// provides the mpsc stream.
let mut webhook_server_addr: Option<std::net::SocketAddr> = None;
if !cli.cli_only
&& let Some(ref http_config) = config.channels.http
{
let http_channel = HttpChannel::new(http_config.clone());
webhook_routes.push(http_channel.routes());
let (host, port) = http_channel.addr();
webhook_server_addr = Some(
format!("{}:{}", host, port)
.parse()
.expect("HttpConfig host:port must be a valid SocketAddr"),
);
channel_names.push("http".to_string());
channels.add(Box::new(http_channel));
tracing::info!(
"HTTP channel enabled on {}:{}",
http_config.host,
http_config.port
);
if !cli.cli_only {
if let Some(ref http_config) = config.channels.http {
let http_channel = HttpChannel::new(http_config.clone());
webhook_routes.push(http_channel.routes());
let (host, port) = http_channel.addr();
webhook_server_addr = Some(
format!("{}:{}", host, port)
.parse()
.expect("HttpConfig host:port must be a valid SocketAddr"),
);
channels.add(Box::new(http_channel));
tracing::info!(
"HTTP channel enabled on {}:{}",
http_config.host,
http_config.port
);
}
}
// Start the unified webhook server if any routes were registered.
@@ -1149,11 +1072,8 @@ async fn main() -> anyhow::Result<()> {
// Create context manager (shared between job tools and agent)
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)
let session_manager = Arc::new(SessionManager::new().with_hooks(hooks.clone()));
let session_manager = Arc::new(SessionManager::new());
// Register job tools (sandbox deps auto-injected when container_job_manager is available)
tools.register_job_tools(
@@ -1163,7 +1083,6 @@ async fn main() -> anyhow::Result<()> {
);
// Add web gateway channel if configured
let mut gateway_url: Option<String> = None;
if let Some(ref gw_config) = config.channels.gateway {
let mut gw = GatewayChannel::new(gw_config.clone());
if let Some(ref ws) = workspace {
@@ -1196,39 +1115,29 @@ async fn main() -> anyhow::Result<()> {
}
}
gateway_url = Some(format!(
"http://{}:{}/?token={}",
gw_config.host,
gw_config.port,
gw.auth_token()
));
tracing::info!(
"Web gateway enabled on {}:{}",
gw_config.host,
gw_config.port
);
tracing::info!("Web UI: http://{}:{}/", gw_config.host, gw_config.port);
tracing::info!(
"Web UI: http://{}:{}/?token={}",
gw_config.host,
gw_config.port,
gw.auth_token()
);
channel_names.push("gateway".to_string());
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
let deps = AgentDeps {
store: db,
llm,
cheap_llm,
safety,
tools,
workspace,
extension_manager,
hooks,
};
let agent = Agent::new(
config.agent.clone(),
@@ -1242,38 +1151,6 @@ async fn main() -> anyhow::Result<()> {
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)
agent.run().await?;
@@ -1289,11 +1166,13 @@ async fn main() -> anyhow::Result<()> {
/// Check if onboarding is needed and return the reason.
///
/// Returns `Some(reason)` if onboarding should be triggered, `None` otherwise.
/// Called after `load_ironclaw_env()`, so DATABASE_URL from `~/.ironclaw/.env`
/// is already in the environment.
#[cfg(any(feature = "postgres", feature = "libsql"))]
fn check_onboard_needed() -> Option<&'static str> {
let has_db = std::env::var("DATABASE_URL").is_ok()
async fn check_onboard_needed() -> Option<&'static str> {
let bootstrap = ironclaw::bootstrap::BootstrapConfig::load();
// Database not configured (and not in env)
let has_db = bootstrap.database_url.is_some()
|| std::env::var("DATABASE_URL").is_ok()
|| std::env::var("LIBSQL_PATH").is_ok()
|| ironclaw::config::default_libsql_path().exists();
@@ -1301,16 +1180,19 @@ fn check_onboard_needed() -> Option<&'static str> {
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");
}
// Secrets not configured (and not in env)
if bootstrap.secrets_master_key_source == ironclaw::settings::KeySource::None
&& std::env::var("SECRETS_MASTER_KEY").is_err()
&& !ironclaw::secrets::keychain::has_master_key().await
{
// Only require secrets setup if user hasn't explicitly disabled it
// For now, we don't require it for first run
}
// First run (onboarding never completed and no session)
let session_path = ironclaw::llm::session::default_session_path();
if !bootstrap.onboard_completed && !session_path.exists() {
return Some("First run");
}
None
+12 -12
View File
@@ -202,7 +202,7 @@ async fn report_complete(
State(state): State<OrchestratorState>,
Path(job_id): Path<Uuid>,
Json(report): Json<CompletionReport>,
) -> Result<Json<serde_json::Value>, StatusCode> {
) -> Result<StatusCode, StatusCode> {
if report.success {
tracing::info!(
job_id = %job_id,
@@ -223,7 +223,7 @@ async fn report_complete(
};
let _ = state.job_manager.complete_job(job_id, result).await;
Ok(Json(serde_json::json!({"status": "ok"})))
Ok(StatusCode::OK)
}
// -- Sandbox job event handlers --
@@ -339,16 +339,16 @@ async fn get_prompt_handler(
Path(job_id): Path<Uuid>,
) -> Result<(StatusCode, Json<serde_json::Value>), StatusCode> {
let mut queue = state.prompt_queue.lock().await;
if let Some(prompts) = queue.get_mut(&job_id)
&& let Some(prompt) = prompts.pop_front()
{
return Ok((
StatusCode::OK,
Json(serde_json::json!({
"content": prompt.content,
"done": prompt.done,
})),
));
if let Some(prompts) = queue.get_mut(&job_id) {
if let Some(prompt) = prompts.pop_front() {
return Ok((
StatusCode::OK,
Json(serde_json::json!({
"content": prompt.content,
"done": prompt.done,
})),
));
}
}
// Return 204 with an empty body. The Json wrapper requires some value
+38 -38
View File
@@ -229,17 +229,17 @@ impl ContainerJobManager {
.unwrap_or_else(|| PathBuf::from("."))
.join(".ironclaw")
.join("projects");
if let Ok(canonical_base) = projects_base.canonicalize()
&& !canonical.starts_with(&canonical_base)
{
return Err(OrchestratorError::ContainerCreationFailed {
job_id,
reason: format!(
"project directory {} is outside allowed base {}",
canonical.display(),
canonical_base.display()
),
});
if let Ok(canonical_base) = projects_base.canonicalize() {
if !canonical.starts_with(&canonical_base) {
return Err(OrchestratorError::ContainerCreationFailed {
job_id,
reason: format!(
"project directory {} is outside allowed base {}",
canonical.display(),
canonical_base.display()
),
});
}
}
binds.push(format!("{}:/workspace:rw", canonical.display()));
env_vec.push("IRONCLAW_WORKSPACE=/workspace".to_string());
@@ -442,36 +442,36 @@ impl ContainerJobManager {
let containers = self.containers.read().await;
containers.get(&job_id).map(|h| h.container_id.clone())
};
if let Some(cid) = container_id
&& !cid.is_empty()
{
match connect_docker().await {
Ok(docker) => {
if let Err(e) = docker
.stop_container(
&cid,
Some(bollard::container::StopContainerOptions { t: 5 }),
)
.await
{
tracing::warn!(job_id = %job_id, error = %e, "Failed to stop completed container");
if let Some(cid) = container_id {
if !cid.is_empty() {
match connect_docker().await {
Ok(docker) => {
if let Err(e) = docker
.stop_container(
&cid,
Some(bollard::container::StopContainerOptions { t: 5 }),
)
.await
{
tracing::warn!(job_id = %job_id, error = %e, "Failed to stop completed container");
}
if let Err(e) = docker
.remove_container(
&cid,
Some(bollard::container::RemoveContainerOptions {
force: true,
..Default::default()
}),
)
.await
{
tracing::warn!(job_id = %job_id, error = %e, "Failed to remove completed container");
}
}
if let Err(e) = docker
.remove_container(
&cid,
Some(bollard::container::RemoveContainerOptions {
force: true,
..Default::default()
}),
)
.await
{
tracing::warn!(job_id = %job_id, error = %e, "Failed to remove completed container");
Err(e) => {
tracing::warn!(job_id = %job_id, error = %e, "Failed to connect to Docker for container cleanup");
}
}
Err(e) => {
tracing::warn!(job_id = %job_id, error = %e, "Failed to connect to Docker for container cleanup");
}
}
}
self.token_store.revoke(job_id).await;
+4 -4
View File
@@ -147,10 +147,10 @@ impl LeakDetector {
// Build prefix matcher for patterns that start with a known prefix
let mut prefixes = Vec::new();
for (idx, pattern) in patterns.iter().enumerate() {
if let Some(prefix) = extract_literal_prefix(pattern.regex.as_str())
&& prefix.len() >= 3
{
prefixes.push((prefix, idx));
if let Some(prefix) = extract_literal_prefix(pattern.regex.as_str()) {
if prefix.len() >= 3 {
prefixes.push((prefix, idx));
}
}
}
+7 -6
View File
@@ -494,10 +494,10 @@ impl ContainerRunner {
/// 3. `~/.docker/run/docker.sock` (Docker Desktop on macOS)
pub async fn connect_docker() -> Result<Docker> {
// First try bollard defaults (checks DOCKER_HOST, then /var/run/docker.sock)
if let Ok(docker) = Docker::connect_with_local_defaults()
&& docker.ping().await.is_ok()
{
return Ok(docker);
if let Ok(docker) = Docker::connect_with_local_defaults() {
if docker.ping().await.is_ok() {
return Ok(docker);
}
}
// Try Docker Desktop socket (macOS)
@@ -507,9 +507,10 @@ pub async fn connect_docker() -> Result<Docker> {
let sock_str = desktop_sock.to_string_lossy();
if let Ok(docker) =
Docker::connect_with_socket(&sock_str, 120, bollard::API_DEFAULT_VERSION)
&& docker.ping().await.is_ok()
{
return Ok(docker);
if docker.ping().await.is_ok() {
return Ok(docker);
}
}
}
}
+9 -9
View File
@@ -259,11 +259,11 @@ async fn handle_connect(
let decision = state.decider.decide(&network_req).await;
if !decision.is_allowed()
&& let NetworkDecision::Deny { reason } = decision
{
tracing::info!("Proxy: blocked CONNECT {} - {}", host, reason);
return error_response(StatusCode::FORBIDDEN, reason);
if !decision.is_allowed() {
if let NetworkDecision::Deny { reason } = decision {
tracing::info!("Proxy: blocked CONNECT {} - {}", host, reason);
return error_response(StatusCode::FORBIDDEN, reason);
}
}
tracing::debug!("Proxy: allowing CONNECT to {}", host);
@@ -294,10 +294,10 @@ async fn forward_request(
// Copy headers (except hop-by-hop headers)
for (name, value) in req.headers() {
if !is_hop_by_hop_header(name.as_str())
&& let Ok(v) = value.to_str()
{
builder = builder.header(name.as_str(), v);
if !is_hop_by_hop_header(name.as_str()) {
if let Ok(v) = value.to_str() {
builder = builder.header(name.as_str(), v);
}
}
}
+5 -4
View File
@@ -109,11 +109,12 @@ impl NetworkPolicyDecider for DefaultPolicyDecider {
async fn decide(&self, request: &NetworkRequest) -> NetworkDecision {
// First check if the domain is allowed
let validation = self.allowlist.is_allowed(&request.host);
if !validation.is_allowed()
&& let crate::sandbox::proxy::allowlist::DomainValidationResult::Denied(reason) =
if !validation.is_allowed() {
if let crate::sandbox::proxy::allowlist::DomainValidationResult::Denied(reason) =
validation
{
return NetworkDecision::Deny { reason };
{
return NetworkDecision::Deny { reason };
}
}
// Check if we need to inject credentials
+1 -1
View File
@@ -261,7 +261,7 @@ pub use platform::{delete_master_key, get_master_key, has_master_key, store_mast
/// Parse a hex string to bytes.
fn hex_to_bytes(hex: &str) -> Result<Vec<u8>, SecretError> {
if !hex.len().is_multiple_of(2) {
if hex.len() % 2 != 0 {
return Err(SecretError::KeychainError(
"Invalid hex string length".to_string(),
));
+22 -59
View File
@@ -153,10 +153,10 @@ impl SecretsStore for PostgresSecretsStore {
let secret = row_to_secret(&r);
// Check expiration
if let Some(expires_at) = secret.expires_at
&& expires_at < Utc::now()
{
return Err(SecretError::Expired);
if let Some(expires_at) = secret.expires_at {
if expires_at < Utc::now() {
return Err(SecretError::Expired);
}
}
Ok(secret)
@@ -276,10 +276,10 @@ impl SecretsStore for PostgresSecretsStore {
}
// Simple glob: * matches any suffix
if let Some(prefix) = pattern.strip_suffix('*')
&& secret_name.starts_with(prefix)
{
return Ok(true);
if let Some(prefix) = pattern.strip_suffix('*') {
if secret_name.starts_with(prefix) {
return Ok(true);
}
}
}
@@ -432,10 +432,10 @@ impl SecretsStore for LibSqlSecretsStore {
Some(row) => {
let secret = libsql_row_to_secret(&row)?;
if let Some(expires_at) = secret.expires_at
&& expires_at < Utc::now()
{
return Err(SecretError::Expired);
if let Some(expires_at) = secret.expires_at {
if expires_at < Utc::now() {
return Err(SecretError::Expired);
}
}
Ok(secret)
@@ -541,10 +541,10 @@ impl SecretsStore for LibSqlSecretsStore {
return Ok(true);
}
if let Some(prefix) = pattern.strip_suffix('*')
&& secret_name.starts_with(prefix)
{
return Ok(true);
if let Some(prefix) = pattern.strip_suffix('*') {
if secret_name.starts_with(prefix) {
return Ok(true);
}
}
}
@@ -695,21 +695,12 @@ pub mod testing {
}
async fn get(&self, user_id: &str, name: &str) -> Result<Secret, SecretError> {
let secret = self
.secrets
self.secrets
.read()
.await
.get(&(user_id.to_string(), name.to_string()))
.cloned()
.ok_or_else(|| SecretError::NotFound(name.to_string()))?;
if let Some(expires_at) = secret.expires_at
&& expires_at < Utc::now()
{
return Err(SecretError::Expired);
}
Ok(secret)
.ok_or_else(|| SecretError::NotFound(name.to_string()))
}
async fn get_decrypted(
@@ -770,10 +761,10 @@ pub mod testing {
if pattern == secret_name {
return Ok(true);
}
if let Some(prefix) = pattern.strip_suffix('*')
&& secret_name.starts_with(prefix)
{
return Ok(true);
if let Some(prefix) = pattern.strip_suffix('*') {
if secret_name.starts_with(prefix) {
return Ok(true);
}
}
}
Ok(false)
@@ -898,34 +889,6 @@ mod tests {
);
}
#[tokio::test]
async fn test_expired_secret_returns_error() {
let store = test_store();
let expires_at = chrono::Utc::now() - chrono::Duration::hours(1);
let params = CreateSecretParams::new("expired_key", "value").with_expiry(expires_at);
store.create("user1", params).await.unwrap();
let result = store.get("user1", "expired_key").await;
assert!(result.is_err());
assert!(matches!(
result.unwrap_err(),
crate::secrets::SecretError::Expired
));
}
#[tokio::test]
async fn test_non_expired_secret_succeeds() {
let store = test_store();
let expires_at = chrono::Utc::now() + chrono::Duration::hours(1);
let params = CreateSecretParams::new("fresh_key", "value").with_expiry(expires_at);
store.create("user1", params).await.unwrap();
let result = store.get("user1", "fresh_key").await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_user_isolation() {
let store = test_store();
+83 -75
View File
@@ -40,18 +40,8 @@ pub struct Settings {
#[serde(default)]
pub secrets_master_key_source: KeySource,
// === Step 3: Inference Provider ===
/// 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 3: NEAR AI Auth ===
// Session stored separately in session.json
// === Step 4: Model Selection ===
/// Currently selected model.
@@ -509,16 +499,20 @@ impl Default for BuilderSettings {
}
impl Settings {
/// Get the default settings file path (~/.ironclaw/settings.json).
pub fn default_path() -> PathBuf {
dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".ironclaw")
.join("settings.json")
}
/// Reconstruct Settings from a flat key-value map (as stored in the DB).
///
/// Each key is a dotted path (e.g., "agent.name"), value is a JSONB value.
/// Missing keys get their default value.
pub fn from_db_map(map: &std::collections::HashMap<String, serde_json::Value>) -> Self {
// 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.
// Start with defaults, then overlay each DB setting
let mut settings = Self::default();
for (key, value) in map {
@@ -527,23 +521,17 @@ impl Settings {
serde_json::Value::String(s) => s.clone(),
serde_json::Value::Bool(b) => b.to_string(),
serde_json::Value::Number(n) => n.to_string(),
serde_json::Value::Null => continue, // null means default, skip
serde_json::Value::Null => "null".to_string(),
other => other.to_string(),
};
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!(
"Failed to apply DB setting '{}' = '{}': {}",
key,
value_str,
e
);
}
if let Err(e) = settings.set(key, &value_str) {
tracing::warn!(
"Failed to apply DB setting '{}' = '{}': {}",
key,
value_str,
e
);
}
}
@@ -564,27 +552,50 @@ impl Settings {
map
}
/// Get the default settings file path (~/.ironclaw/settings.json).
pub fn default_path() -> std::path::PathBuf {
dirs::home_dir()
.unwrap_or_else(|| std::path::PathBuf::from("."))
.join(".ironclaw")
.join("settings.json")
}
/// Load settings from disk, returning default if not found.
pub fn load() -> Self {
Self::load_from(&Self::default_path())
}
/// Load settings from a specific path (used by bootstrap legacy migration).
pub fn load_from(path: &std::path::Path) -> Self {
/// Load settings from a specific path.
pub fn load_from(path: &PathBuf) -> Self {
match std::fs::read_to_string(path) {
Ok(data) => serde_json::from_str(&data).unwrap_or_default(),
Err(_) => Self::default(),
}
}
/// Save settings to disk.
pub fn save(&self) -> std::io::Result<()> {
self.save_to(&Self::default_path())
}
/// Save settings to a specific path.
pub fn save_to(&self, path: &PathBuf) -> std::io::Result<()> {
// Ensure parent directory exists
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
let json = serde_json::to_string_pretty(self)
.map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e.to_string()))?;
std::fs::write(path, json)
}
/// Get the selected model, falling back to the provided default.
pub fn model_or(&self, default: &str) -> String {
self.selected_model
.clone()
.unwrap_or_else(|| default.to_string())
}
/// Set the selected model and save.
pub fn set_model(&mut self, model: &str) -> std::io::Result<()> {
self.selected_model = Some(model.to_string());
self.save()
}
/// Get a setting value by dotted path (e.g., "agent.max_parallel_jobs").
pub fn get(&self, path: &str) -> Option<String> {
let json = serde_json::to_value(self).ok()?;
@@ -769,22 +780,42 @@ fn collect_settings(
#[cfg(test)]
mod tests {
use super::*;
use tempfile::tempdir;
#[test]
fn test_db_map_round_trip() {
fn test_settings_save_load() {
let dir = tempdir().unwrap();
let path = dir.path().join("settings.json");
let settings = Settings {
selected_model: Some("claude-3-5-sonnet-20241022".to_string()),
..Default::default()
};
let map = settings.to_db_map();
let restored = Settings::from_db_map(&map);
settings.save_to(&path).unwrap();
let loaded = Settings::load_from(&path);
assert_eq!(
restored.selected_model,
loaded.selected_model,
Some("claude-3-5-sonnet-20241022".to_string())
);
}
#[test]
fn test_model_or_default() {
let settings = Settings::default();
assert_eq!(
settings.model_or("default-model"),
"default-model".to_string()
);
let settings = Settings {
selected_model: Some("my-model".to_string()),
..Default::default()
};
assert_eq!(settings.model_or("default-model"), "my-model".to_string());
}
#[test]
fn test_get_setting() {
let settings = Settings::default();
@@ -855,13 +886,16 @@ mod tests {
}
#[test]
fn test_telegram_owner_id_db_round_trip() {
fn test_telegram_owner_id_round_trip() {
let dir = tempdir().unwrap();
let path = dir.path().join("settings.json");
let mut settings = Settings::default();
settings.channels.telegram_owner_id = Some(123456789);
settings.save_to(&path).unwrap();
let map = settings.to_db_map();
let restored = Settings::from_db_map(&map);
assert_eq!(restored.channels.telegram_owner_id, Some(123456789));
let loaded = Settings::load_from(&path);
assert_eq!(loaded.channels.telegram_owner_id, Some(123456789));
}
#[test]
@@ -878,30 +912,4 @@ mod tests {
.unwrap();
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())
);
}
}
-539
View File
@@ -1,539 +0,0 @@
# 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`
+145 -184
View File
@@ -15,27 +15,11 @@ use serde::Deserialize;
#[cfg(feature = "postgres")]
use crate::secrets::SecretsCrypto;
use crate::secrets::{CreateSecretParams, SecretsStore};
use crate::settings::{Settings, TunnelSettings};
use crate::settings::Settings;
use crate::setup::prompts::{
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.
pub struct SecretsContext {
store: Arc<dyn SecretsStore>,
@@ -61,39 +45,32 @@ impl SecretsContext {
}
/// Save a secret to the database.
pub async fn save_secret(
&self,
name: &str,
value: &SecretString,
) -> Result<(), ChannelSetupError> {
pub async fn save_secret(&self, name: &str, value: &SecretString) -> Result<(), String> {
let params = CreateSecretParams::new(name, value.expose_secret());
self.store
.create(&self.user_id, params)
.await
.map_err(|e| ChannelSetupError::Secrets(format!("Failed to save secret: {}", e)))?;
.map_err(|e| format!("Failed to save secret: {}", e))?;
Ok(())
}
/// Check if a secret exists.
pub async fn secret_exists(&self, name: &str) -> bool {
match self.store.exists(&self.user_id, name).await {
Ok(exists) => exists,
Err(e) => {
tracing::warn!(secret = name, error = %e, "Failed to check if secret exists, assuming absent");
false
}
}
self.store
.exists(&self.user_id, name)
.await
.unwrap_or(false)
}
/// Read a secret from the database (decrypted).
pub async fn get_secret(&self, name: &str) -> Result<SecretString, ChannelSetupError> {
pub async fn get_secret(&self, name: &str) -> Result<SecretString, String> {
let decrypted = self
.store
.get_decrypted(&self.user_id, name)
.await
.map_err(|e| ChannelSetupError::Secrets(format!("Failed to read secret: {}", e)))?;
.map_err(|e| format!("Failed to read secret: {}", e))?;
Ok(SecretString::from(decrypted.expose().to_string()))
}
}
@@ -130,6 +107,7 @@ struct TelegramGetUpdatesResponse {
#[derive(Debug, Deserialize)]
struct TelegramUpdate {
#[allow(dead_code)]
update_id: i64,
message: Option<TelegramUpdateMessage>,
}
@@ -153,10 +131,7 @@ struct TelegramUpdateUser {
/// 2. Entering the bot token
/// 3. Validating the token
/// 4. Saving the token to the database
pub async fn setup_telegram(
secrets: &SecretsContext,
settings: &Settings,
) -> Result<TelegramSetupResult, ChannelSetupError> {
pub async fn setup_telegram(secrets: &SecretsContext) -> Result<TelegramSetupResult, String> {
println!("Telegram Setup:");
println!();
print_info("To create a Telegram bot:");
@@ -168,10 +143,10 @@ pub async fn setup_telegram(
// Check if token already exists
if secrets.secret_exists("telegram_bot_token").await {
print_info("Existing Telegram token found in database.");
if !confirm("Replace existing token?", false)? {
if !confirm("Replace existing token?", false).map_err(|e| e.to_string())? {
// Still offer to configure webhook secret and owner binding
let webhook_secret = setup_telegram_webhook_secret(secrets, &settings.tunnel).await?;
let owner_id = bind_telegram_owner_flow(secrets, settings).await?;
let webhook_secret = setup_telegram_webhook_secret(secrets).await?;
let owner_id = bind_telegram_owner_flow(secrets).await?;
return Ok(TelegramSetupResult {
enabled: true,
bot_username: None,
@@ -181,48 +156,47 @@ pub async fn setup_telegram(
}
}
loop {
let token = secret_input("Bot token (from @BotFather)")?;
let token = secret_input("Bot token (from @BotFather)").map_err(|e| e.to_string())?;
// Validate the token
print_info("Validating bot token...");
// Validate the token
print_info("Validating bot token...");
match validate_telegram_token(&token).await {
Ok(username) => {
print_success(&format!(
"Bot validated: @{}",
username.as_deref().unwrap_or("unknown")
));
match validate_telegram_token(&token).await {
Ok(username) => {
print_success(&format!(
"Bot validated: @{}",
username.as_deref().unwrap_or("unknown")
));
// Save to database
secrets.save_secret("telegram_bot_token", &token).await?;
print_success("Token saved to database");
// Save to database
secrets.save_secret("telegram_bot_token", &token).await?;
print_success("Token saved to database");
// Bind bot to owner's Telegram account
let owner_id = bind_telegram_owner(&token).await?;
// Bind bot to owner's Telegram account
let owner_id = bind_telegram_owner(&token).await?;
// Offer webhook secret configuration
let webhook_secret =
setup_telegram_webhook_secret(secrets, &settings.tunnel).await?;
// Offer webhook secret configuration
let webhook_secret = setup_telegram_webhook_secret(secrets).await?;
return Ok(TelegramSetupResult {
enabled: true,
bot_username: username,
webhook_secret,
owner_id,
});
}
Err(e) => {
print_error(&format!("Token validation failed: {}", e));
Ok(TelegramSetupResult {
enabled: true,
bot_username: username,
webhook_secret,
owner_id,
})
}
Err(e) => {
print_error(&format!("Token validation failed: {}", e));
if !confirm("Try again?", true)? {
return Ok(TelegramSetupResult {
enabled: false,
bot_username: None,
webhook_secret: None,
owner_id: None,
});
}
if confirm("Try again?", true).map_err(|e| e.to_string())? {
Box::pin(setup_telegram(secrets)).await
} else {
Ok(TelegramSetupResult {
enabled: false,
bot_username: None,
webhook_secret: None,
owner_id: None,
})
}
}
}
@@ -232,14 +206,14 @@ pub async fn setup_telegram(
///
/// Polls `getUpdates` until a message arrives, then captures the sender's user ID.
/// Returns `None` if the user declines or the flow times out.
async fn bind_telegram_owner(token: &SecretString) -> Result<Option<i64>, ChannelSetupError> {
async fn bind_telegram_owner(token: &SecretString) -> Result<Option<i64>, String> {
println!();
print_info("Account Binding (recommended):");
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.");
println!();
if !confirm("Bind bot to your Telegram account?", true)? {
if !confirm("Bind bot to your Telegram account?", true).map_err(|e| e.to_string())? {
print_info("Skipping account binding. Bot will accept messages from all users.");
return Ok(None);
}
@@ -250,16 +224,14 @@ async fn bind_telegram_owner(token: &SecretString) -> Result<Option<i64>, Channe
let client = Client::builder()
.timeout(std::time::Duration::from_secs(35))
.build()
.map_err(|e| ChannelSetupError::Network(format!("Failed to create HTTP client: {}", e)))?;
.map_err(|e| format!("Failed to create HTTP client: {}", e))?;
// Clear any existing webhook so getUpdates works
let delete_url = format!(
"https://api.telegram.org/bot{}/deleteWebhook",
token.expose_secret()
);
if let Err(e) = client.post(&delete_url).send().await {
tracing::warn!("Failed to delete webhook (getUpdates may not work): {e}");
}
let _ = client.post(&delete_url).send().await;
let updates_url = format!(
"https://api.telegram.org/bot{}/getUpdates",
@@ -274,56 +246,49 @@ async fn bind_telegram_owner(token: &SecretString) -> Result<Option<i64>, Channe
.query(&[("timeout", "30"), ("allowed_updates", "[\"message\"]")])
.send()
.await
.map_err(|e| ChannelSetupError::Network(format!("getUpdates request failed: {}", e)))?;
.map_err(|e| format!("getUpdates request failed: {}", e))?;
if !response.status().is_success() {
return Err(ChannelSetupError::Network(format!(
"getUpdates returned status {}",
response.status()
)));
return Err(format!("getUpdates returned status {}", response.status()));
}
let body: TelegramGetUpdatesResponse = response.json().await.map_err(|e| {
ChannelSetupError::Network(format!("Failed to parse getUpdates response: {}", e))
})?;
let body: TelegramGetUpdatesResponse = response
.json()
.await
.map_err(|e| format!("Failed to parse getUpdates response: {}", e))?;
if !body.ok {
return Err(ChannelSetupError::Network(
"Telegram API returned error for getUpdates".to_string(),
));
return Err("Telegram API returned error for getUpdates".to_string());
}
// Find the first message with a sender
for update in &body.result {
if let Some(ref msg) = update.message
&& let Some(ref from) = msg.from
{
let display_name = from
.username
.as_ref()
.map(|u| format!("@{}", u))
.unwrap_or_else(|| from.first_name.clone());
if let Some(ref msg) = update.message {
if let Some(ref from) = msg.from {
let display_name = from
.username
.as_ref()
.map(|u| format!("@{}", u))
.unwrap_or_else(|| from.first_name.clone());
print_success(&format!(
"Received message from {} (ID: {})",
display_name, from.id
));
print_success(&format!(
"Received message from {} (ID: {})",
display_name, from.id
));
// Acknowledge the update so it doesn't pile up
let ack_url = format!(
"https://api.telegram.org/bot{}/getUpdates",
token.expose_secret()
);
if let Err(e) = client
.get(&ack_url)
.query(&[("offset", &(update.update_id + 1).to_string())])
.send()
.await
{
tracing::warn!("Failed to acknowledge Telegram update: {e}");
// Acknowledge the update so it doesn't pile up
let ack_url = format!(
"https://api.telegram.org/bot{}/getUpdates",
token.expose_secret()
);
let _ = client
.get(&ack_url)
.query(&[("offset", &(update.update_id + 1).to_string())])
.send()
.await;
return Ok(Some(from.id));
}
return Ok(Some(from.id));
}
}
}
@@ -336,13 +301,12 @@ async fn bind_telegram_owner(token: &SecretString) -> Result<Option<i64>, Channe
/// Bind flow when the token already exists (reads from secrets store).
///
/// Retrieves the saved bot token and delegates to `bind_telegram_owner`.
async fn bind_telegram_owner_flow(
secrets: &SecretsContext,
settings: &Settings,
) -> Result<Option<i64>, ChannelSetupError> {
async fn bind_telegram_owner_flow(secrets: &SecretsContext) -> Result<Option<i64>, String> {
// Check current settings first
let settings = Settings::load();
if settings.channels.telegram_owner_id.is_some() {
print_info("Bot is already bound to a Telegram account.");
if !confirm("Re-bind to a different account?", false)? {
if !confirm("Re-bind to a different account?", false).map_err(|e| e.to_string())? {
return Ok(settings.channels.telegram_owner_id);
}
}
@@ -357,10 +321,12 @@ async fn bind_telegram_owner_flow(
///
/// This is shared across all channels that need webhook endpoints.
/// Returns the tunnel URL if configured.
pub fn setup_tunnel(settings: &Settings) -> Result<Option<String>, ChannelSetupError> {
pub fn setup_tunnel() -> Result<Option<String>, String> {
// Check if already configured
let settings = Settings::load();
if let Some(ref url) = settings.tunnel.public_url {
print_info(&format!("Existing tunnel configured: {}", url));
if !confirm("Change tunnel configuration?", false)? {
if !confirm("Change tunnel configuration?", false).map_err(|e| e.to_string())? {
return Ok(Some(url.clone()));
}
}
@@ -380,24 +346,30 @@ pub fn setup_tunnel(settings: &Settings) -> Result<Option<String>, ChannelSetupE
print_info("Security comes from provider-specific secrets (e.g., Telegram webhook secret).");
println!();
if !confirm("Configure a tunnel?", false)? {
if !confirm("Configure a tunnel?", false).map_err(|e| e.to_string())? {
return Ok(None);
}
let tunnel_url = input("Tunnel URL (e.g., https://abc123.ngrok.io)")?;
let tunnel_url =
input("Tunnel URL (e.g., https://abc123.ngrok.io)").map_err(|e| e.to_string())?;
// Validate URL format
if !tunnel_url.starts_with("https://") {
print_error("URL must start with https:// (webhooks require HTTPS)");
return Err(ChannelSetupError::Validation(
"Invalid tunnel URL: must use HTTPS".to_string(),
));
return Err("Invalid tunnel URL: must use HTTPS".to_string());
}
// Remove trailing slash if present
let tunnel_url = tunnel_url.trim_end_matches('/').to_string();
print_success(&format!("Tunnel URL configured: {}", tunnel_url));
// Save to settings
let mut settings = Settings::load();
settings.tunnel.public_url = Some(tunnel_url.clone());
settings
.save()
.map_err(|e| format!("Failed to save settings: {}", e))?;
print_success(&format!("Tunnel URL saved: {}", tunnel_url));
print_info("");
print_info("Make sure your tunnel is running before starting the agent.");
print_info("You can also set TUNNEL_URL environment variable to override.");
@@ -408,11 +380,10 @@ pub fn setup_tunnel(settings: &Settings) -> Result<Option<String>, ChannelSetupE
/// Set up Telegram webhook secret for signature validation.
///
/// Returns the webhook secret if configured.
async fn setup_telegram_webhook_secret(
secrets: &SecretsContext,
tunnel: &TunnelSettings,
) -> Result<Option<String>, ChannelSetupError> {
if tunnel.public_url.is_none() {
async fn setup_telegram_webhook_secret(secrets: &SecretsContext) -> Result<Option<String>, String> {
// Check if tunnel is configured
let settings = Settings::load();
if settings.tunnel.public_url.is_none() {
print_info("");
print_info("No tunnel configured. Telegram will use polling mode (30s+ delay).");
print_info("Run setup again to configure a tunnel for instant delivery.");
@@ -424,7 +395,7 @@ async fn setup_telegram_webhook_secret(
print_info("A webhook secret adds an extra layer of security by validating");
print_info("that requests actually come from Telegram's servers.");
if !confirm("Generate a webhook secret?", true)? {
if !confirm("Generate a webhook secret?", true).map_err(|e| e.to_string())? {
return Ok(None);
}
@@ -443,13 +414,11 @@ async fn setup_telegram_webhook_secret(
/// Validate a Telegram bot token by calling the getMe API.
///
/// Returns the bot's username if valid.
pub async fn validate_telegram_token(
token: &SecretString,
) -> Result<Option<String>, ChannelSetupError> {
pub async fn validate_telegram_token(token: &SecretString) -> Result<Option<String>, String> {
let client = Client::builder()
.timeout(std::time::Duration::from_secs(10))
.build()
.map_err(|e| ChannelSetupError::Network(format!("Failed to create HTTP client: {}", e)))?;
.map_err(|e| format!("Failed to create HTTP client: {}", e))?;
let url = format!(
"https://api.telegram.org/bot{}/getMe",
@@ -460,26 +429,21 @@ pub async fn validate_telegram_token(
.get(&url)
.send()
.await
.map_err(|e| ChannelSetupError::Network(format!("Request failed: {}", e)))?;
.map_err(|e| format!("Request failed: {}", e))?;
if !response.status().is_success() {
return Err(ChannelSetupError::Network(format!(
"API returned status {}",
response.status()
)));
return Err(format!("API returned status {}", response.status()));
}
let body: TelegramGetMeResponse = response
.json()
.await
.map_err(|e| ChannelSetupError::Network(format!("Failed to parse response: {}", e)))?;
.map_err(|e| format!("Failed to parse response: {}", e))?;
if body.ok {
Ok(body.result.and_then(|u| u.username))
} else {
Err(ChannelSetupError::Network(
"Telegram API returned error".to_string(),
))
Err("Telegram API returned error".to_string())
}
}
@@ -492,34 +456,38 @@ pub struct HttpSetupResult {
}
/// Set up HTTP webhook channel.
pub async fn setup_http(secrets: &SecretsContext) -> Result<HttpSetupResult, ChannelSetupError> {
pub async fn setup_http(secrets: &SecretsContext) -> Result<HttpSetupResult, String> {
println!("HTTP Webhook Setup:");
println!();
print_info("The HTTP webhook allows external services to send messages to the agent.");
println!();
let port_str = optional_input("Port", Some("default: 8080"))?;
let port_str = optional_input("Port", Some("default: 8080")).map_err(|e| e.to_string())?;
let port: u16 = port_str
.as_deref()
.unwrap_or("8080")
.parse()
.map_err(|e| ChannelSetupError::Validation(format!("Invalid port: {}", e)))?;
.map_err(|e| format!("Invalid port: {}", e))?;
if port < 1024 {
print_info("Note: Ports below 1024 may require root privileges");
}
let host =
optional_input("Host", Some("default: 0.0.0.0"))?.unwrap_or_else(|| "0.0.0.0".to_string());
let host = optional_input("Host", Some("default: 0.0.0.0"))
.map_err(|e| e.to_string())?
.unwrap_or_else(|| "0.0.0.0".to_string());
// Generate a webhook secret
if confirm("Generate a webhook secret for authentication?", true)? {
if confirm("Generate a webhook secret for authentication?", true).map_err(|e| e.to_string())? {
let secret = generate_webhook_secret();
secrets
.save_secret("http_webhook_secret", &SecretString::from(secret))
.save_secret("http_webhook_secret", &SecretString::from(secret.clone()))
.await?;
print_success("Webhook secret generated and saved to database");
print_info("Retrieve it later with: ironclaw secret get http_webhook_secret");
print_info(&format!(
"Secret: {} (store this for your webhook clients)",
secret
));
}
print_success(&format!("HTTP webhook will listen on {}:{}", host, port));
@@ -533,7 +501,11 @@ pub async fn setup_http(secrets: &SecretsContext) -> Result<HttpSetupResult, Cha
/// Generate a random webhook secret.
pub fn generate_webhook_secret() -> String {
generate_secret_with_length(32)
use rand::RngCore;
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.
@@ -551,7 +523,7 @@ pub async fn setup_wasm_channel(
secrets: &SecretsContext,
channel_name: &str,
setup: &crate::channels::wasm::SetupSchema,
) -> Result<WasmChannelSetupResult, ChannelSetupError> {
) -> Result<WasmChannelSetupResult, String> {
println!("{} Setup:", channel_name);
println!();
@@ -562,7 +534,7 @@ pub async fn setup_wasm_channel(
"Existing {} found in database.",
secret_config.name
));
if !confirm("Replace existing value?", false)? {
if !confirm("Replace existing value?", false).map_err(|e| e.to_string())? {
continue;
}
}
@@ -570,7 +542,8 @@ pub async fn setup_wasm_channel(
// Get the value from user or auto-generate
let value = if secret_config.optional {
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 !v.is_empty() {
@@ -597,21 +570,18 @@ pub async fn setup_wasm_channel(
}
} else {
// Required secret
let input_value = secret_input(&secret_config.prompt)?;
let input_value = secret_input(&secret_config.prompt).map_err(|e| e.to_string())?;
// Validate if pattern is provided
if let Some(ref pattern) = secret_config.validation {
let re = regex::Regex::new(pattern).map_err(|e| {
ChannelSetupError::Validation(format!("Invalid validation pattern: {}", e))
})?;
let re = regex::Regex::new(pattern)
.map_err(|e| format!("Invalid validation pattern: {}", e))?;
if !re.is_match(input_value.expose_secret()) {
print_error(&format!(
"Value does not match expected format: {}",
pattern
));
return Err(ChannelSetupError::Validation(
"Validation failed".to_string(),
));
return Err("Validation failed".to_string());
}
}
@@ -623,11 +593,14 @@ pub async fn setup_wasm_channel(
print_success(&format!("{} saved to database", secret_config.name));
}
// TODO: Substitute secrets into the validation URL and make a
// GET request to verify the configured credentials actually work.
// Optionally validate the configuration
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!(
"Validation endpoint configured: {} (validation not yet implemented)",
"Validation endpoint configured: {} (validation skipped)",
validation_endpoint
));
}
@@ -651,23 +624,11 @@ fn generate_secret_with_length(length: usize) -> String {
#[cfg(test)]
mod tests {
use crate::setup::channels::generate_webhook_secret;
use super::*;
#[test]
fn test_generate_webhook_secret() {
let secret = generate_webhook_secret();
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);
}
}
+2 -3
View File
@@ -3,7 +3,7 @@
//! Provides a guided setup experience for:
//! 1. Database connection
//! 2. Security (secrets master key)
//! 3. Inference provider selection
//! 3. NEAR AI authentication
//! 4. Model selection
//! 5. Embeddings
//! 6. Channel configuration (HTTP, Telegram, etc.)
@@ -24,8 +24,7 @@ mod prompts;
mod wizard;
pub use channels::{
ChannelSetupError, SecretsContext, setup_http, setup_telegram, setup_tunnel,
validate_telegram_token,
SecretsContext, setup_http, setup_telegram, setup_tunnel, validate_telegram_token,
};
pub use prompts::{
confirm, input, optional_input, print_error, print_header, print_info, print_step,
+4 -10
View File
@@ -21,7 +21,6 @@ use secrecy::SecretString;
/// Display a numbered menu and get user selection.
///
/// Returns the index (0-based) of the selected option.
/// Pressing Enter without input selects the first option (index 0).
///
/// # Example
///
@@ -55,11 +54,10 @@ pub fn select_one(prompt: &str, options: &[&str]) -> io::Result<usize> {
}
// Parse number
if let Ok(num) = input.parse::<usize>()
&& num >= 1
&& num <= options.len()
{
return Ok(num - 1);
if let Ok(num) = input.parse::<usize>() {
if num >= 1 && num <= options.len() {
return Ok(num - 1);
}
}
writeln!(
@@ -85,10 +83,6 @@ pub fn select_one(prompt: &str, options: &[&str]) -> io::Result<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 selected: Vec<bool> = options.iter().map(|(_, s)| *s).collect();
let mut cursor_pos = 0;
+123 -938
View File
File diff suppressed because it is too large Load Diff
+27 -27
View File
@@ -326,20 +326,20 @@ impl TestHarness {
}
// Verify expected output
if let Some(ref expected) = test.expected_output
&& &actual != expected
{
return TestResult {
name: test.name.clone(),
passed: false,
duration,
error: Some(format!(
"Output mismatch:\nExpected: {}\nActual: {}",
serde_json::to_string_pretty(expected).unwrap_or_default(),
serde_json::to_string_pretty(&actual).unwrap_or_default()
)),
actual_output: Some(actual),
};
if let Some(ref expected) = test.expected_output {
if &actual != expected {
return TestResult {
name: test.name.clone(),
passed: false,
duration,
error: Some(format!(
"Output mismatch:\nExpected: {}\nActual: {}",
serde_json::to_string_pretty(expected).unwrap_or_default(),
serde_json::to_string_pretty(&actual).unwrap_or_default()
)),
actual_output: Some(actual),
};
}
}
// Verify expected fields
@@ -357,19 +357,19 @@ impl TestHarness {
};
}
if let Some(ref expected_value) = field.value
&& field_value != Some(expected_value)
{
return TestResult {
name: test.name.clone(),
passed: false,
duration,
error: Some(format!(
"Field '{}' mismatch: expected {:?}, got {:?}",
field.path, expected_value, field_value
)),
actual_output: Some(actual),
};
if let Some(ref expected_value) = field.value {
if field_value != Some(expected_value) {
return TestResult {
name: test.name.clone(),
passed: false,
duration,
error: Some(format!(
"Field '{}' mismatch: expected {:?}, got {:?}",
field.path, expected_value, field_value
)),
actual_output: Some(actual),
};
}
}
}
}
+7 -2
View File
@@ -3,7 +3,7 @@
use async_trait::async_trait;
use crate::context::JobContext;
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
use crate::tools::tool::{Tool, ToolError, ToolOutput};
/// Simple echo tool for testing.
pub struct EchoTool;
@@ -38,7 +38,12 @@ impl Tool for EchoTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let message = require_str(&params, "message")?;
let message = params
.get("message")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'message' parameter".to_string())
})?;
Ok(ToolOutput::text(message, start.elapsed()))
}
+136
View File
@@ -0,0 +1,136 @@
//! E-commerce tool for shopping and price comparison.
use async_trait::async_trait;
use crate::context::JobContext;
use crate::tools::tool::{Tool, ToolError, ToolOutput};
/// Tool for e-commerce operations (Amazon, price comparison, etc.).
pub struct EcommerceTool {
// TODO: Add API clients
}
impl EcommerceTool {
/// Create a new e-commerce tool.
pub fn new() -> Self {
Self {}
}
}
impl Default for EcommerceTool {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl Tool for EcommerceTool {
fn name(&self) -> &str {
"ecommerce"
}
fn description(&self) -> &str {
"Search products, compare prices, and find deals across e-commerce platforms."
}
fn parameters_schema(&self) -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": {
"action": {
"type": "string",
"enum": ["search", "get_product", "compare_prices", "track_price"],
"description": "The e-commerce action to perform"
},
"query": {
"type": "string",
"description": "Search query (for search action)"
},
"product_id": {
"type": "string",
"description": "Product ID or ASIN (for get_product, compare_prices)"
},
"platform": {
"type": "string",
"enum": ["amazon", "ebay", "walmart", "all"],
"description": "E-commerce platform to search"
},
"max_price": {
"type": "number",
"description": "Maximum price filter"
},
"category": {
"type": "string",
"description": "Product category filter"
}
},
"required": ["action"]
})
}
async fn execute(
&self,
params: serde_json::Value,
_ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let action = params
.get("action")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'action' parameter".to_string())
})?;
// TODO: Implement actual e-commerce API integrations
let result = match action {
"search" => {
let query = params.get("query").and_then(|v| v.as_str()).unwrap_or("");
serde_json::json!({
"query": query,
"results": [],
"message": "E-commerce integration not yet implemented"
})
}
"get_product" => {
let product_id = params
.get("product_id")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'product_id' parameter".to_string())
})?;
serde_json::json!({
"product_id": product_id,
"found": false,
"message": "E-commerce integration not yet implemented"
})
}
"compare_prices" => {
serde_json::json!({
"prices": [],
"message": "E-commerce integration not yet implemented"
})
}
"track_price" => {
serde_json::json!({
"tracking": false,
"message": "E-commerce integration not yet implemented"
})
}
_ => {
return Err(ToolError::InvalidParameters(format!(
"unknown action: {}",
action
)));
}
};
Ok(ToolOutput::success(result, start.elapsed()))
}
fn requires_sanitization(&self) -> bool {
true // External e-commerce data
}
}
+17 -5
View File
@@ -9,7 +9,7 @@ use async_trait::async_trait;
use crate::context::JobContext;
use crate::extensions::{ExtensionKind, ExtensionManager};
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
use crate::tools::tool::{Tool, ToolError, ToolOutput};
// ── tool_search ──────────────────────────────────────────────────────────
@@ -133,7 +133,10 @@ impl Tool for ToolInstallTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let name = require_str(&params, "name")?;
let name = params
.get("name")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("name is required".to_string()))?;
let url = params.get("url").and_then(|v| v.as_str());
@@ -207,7 +210,10 @@ impl Tool for ToolAuthTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let name = require_str(&params, "name")?;
let name = params
.get("name")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("name is required".to_string()))?;
let result = self
.manager
@@ -300,7 +306,10 @@ impl Tool for ToolActivateTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let name = require_str(&params, "name")?;
let name = params
.get("name")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("name is required".to_string()))?;
match self.manager.activate(name).await {
Ok(result) => {
@@ -462,7 +471,10 @@ impl Tool for ToolRemoveTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let name = require_str(&params, "name")?;
let name = params
.get("name")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("name is required".to_string()))?;
let message = self
.manager
+25 -7
View File
@@ -11,7 +11,7 @@ use async_trait::async_trait;
use tokio::fs;
use crate::context::JobContext;
use crate::tools::tool::{Tool, ToolDomain, ToolError, ToolOutput, require_str};
use crate::tools::tool::{Tool, ToolDomain, ToolError, ToolOutput};
use crate::workspace::paths as ws_paths;
/// Well-known workspace filenames that must go through memory_write, not write_file.
@@ -203,7 +203,10 @@ impl Tool for ReadFileTool {
params: serde_json::Value,
_ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let path_str = require_str(&params, "path")?;
let path_str = params
.get("path")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'path' parameter".into()))?;
let offset = params.get("offset").and_then(|v| v.as_u64()).unwrap_or(0) as usize;
let limit = params.get("limit").and_then(|v| v.as_u64());
@@ -325,7 +328,10 @@ impl Tool for WriteFileTool {
params: serde_json::Value,
_ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let path_str = require_str(&params, "path")?;
let path_str = params
.get("path")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'path' parameter".into()))?;
// Reject workspace paths: these live in the database, not on disk.
if is_workspace_path(path_str) {
@@ -336,7 +342,10 @@ impl Tool for WriteFileTool {
)));
}
let content = require_str(&params, "content")?;
let content = params
.get("content")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'content' parameter".into()))?;
let start = std::time::Instant::now();
@@ -641,11 +650,20 @@ impl Tool for ApplyPatchTool {
params: serde_json::Value,
_ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let path_str = require_str(&params, "path")?;
let path_str = params
.get("path")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'path' parameter".into()))?;
let old_string = require_str(&params, "old_string")?;
let old_string = params
.get("old_string")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'old_string' parameter".into()))?;
let new_string = require_str(&params, "new_string")?;
let new_string = params
.get("new_string")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'new_string' parameter".into()))?;
let replace_all = params
.get("replace_all")
+27 -54
View File
@@ -5,18 +5,13 @@ use std::net::{IpAddr, ToSocketAddrs};
use std::time::Duration;
use async_trait::async_trait;
use futures::StreamExt;
use reqwest::Client;
use crate::context::JobContext;
use crate::safety::LeakDetector;
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
use crate::tools::tool::{Tool, ToolError, ToolOutput};
/// 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.
/// Maximum response body size (5 MB). Prevents OOM from unbounded responses.
const MAX_RESPONSE_SIZE: usize = 5 * 1024 * 1024;
/// Tool for making HTTP requests.
@@ -59,12 +54,12 @@ fn validate_url(url: &str) -> Result<reqwest::Url, ToolError> {
}
// Check literal IP addresses
if let Ok(ip) = host.parse::<IpAddr>()
&& is_disallowed_ip(&ip)
{
return Err(ToolError::NotAuthorized(
"private or local IPs are not allowed".to_string(),
));
if let Ok(ip) = host.parse::<IpAddr>() {
if is_disallowed_ip(&ip) {
return Err(ToolError::NotAuthorized(
"private or local IPs are not allowed".to_string(),
));
}
}
// Resolve hostname and check all resolved IPs against the blocklist.
@@ -159,9 +154,17 @@ impl Tool for HttpTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let method = require_str(&params, "method")?;
let method = params
.get("method")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'method' parameter".to_string())
})?;
let url = require_str(&params, "url")?;
let url = params
.get("url")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'url' parameter".to_string()))?;
let parsed_url = validate_url(url)?;
// Parse headers
@@ -235,43 +238,19 @@ impl Tool for HttpTool {
.filter_map(|(k, v)| v.to_str().ok().map(|v| (k.to_string(), v.to_string())))
.collect();
// Pre-check Content-Length header to reject obviously oversized responses
// 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"
);
// Get response body with size cap to prevent OOM
let body_bytes = response.bytes().await.map_err(|e| {
ToolError::ExternalService(format!("failed to read response body: {}", e))
})?;
if body_bytes.len() > MAX_RESPONSE_SIZE {
return Err(ToolError::ExecutionFailed(format!(
"Response Content-Length ({} bytes) exceeds maximum allowed size ({} bytes)",
len, MAX_RESPONSE_SIZE
"Response body too large ({} bytes, max {})",
body_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))
})?;
if body.len() + chunk.len() > MAX_RESPONSE_SIZE {
return Err(ToolError::ExecutionFailed(format!(
"Response body exceeds maximum allowed size ({} bytes)",
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();
// Try to parse as JSON, fall back to string
@@ -357,10 +336,4 @@ mod tests {
// Public
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);
}
}
+31 -17
View File
@@ -18,7 +18,7 @@ use crate::context::{ContextManager, JobContext, JobState};
use crate::db::Database;
use crate::history::SandboxJobRecord;
use crate::orchestrator::job_manager::{ContainerJobManager, JobMode};
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
use crate::tools::tool::{Tool, ToolError, ToolOutput};
/// Tool for creating a new job.
///
@@ -158,18 +158,18 @@ impl CreateJobTool {
});
// Persist the job mode to DB
if mode == JobMode::ClaudeCode
&& let Some(store) = self.store.clone()
{
let job_id_copy = job_id;
tokio::spawn(async move {
if let Err(e) = store
.update_sandbox_job_mode(job_id_copy, "claude_code")
.await
{
tracing::warn!(job_id = %job_id_copy, "Failed to set job mode: {}", e);
}
});
if mode == JobMode::ClaudeCode {
if let Some(store) = self.store.clone() {
let job_id_copy = job_id;
tokio::spawn(async move {
if let Err(e) = store
.update_sandbox_job_mode(job_id_copy, "claude_code")
.await
{
tracing::warn!(job_id = %job_id_copy, "Failed to set job mode: {}", e);
}
});
}
}
// Create the container job with the pre-determined job_id.
@@ -467,9 +467,17 @@ impl Tool for CreateJobTool {
params: serde_json::Value,
ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let title = require_str(&params, "title")?;
let title = params
.get("title")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'title' parameter".into()))?;
let description = require_str(&params, "description")?;
let description = params
.get("description")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'description' parameter".into())
})?;
if self.sandbox_enabled() {
let wait = params.get("wait").and_then(|v| v.as_bool()).unwrap_or(true);
@@ -627,7 +635,10 @@ impl Tool for JobStatusTool {
let start = std::time::Instant::now();
let requester_id = ctx.user_id.clone();
let job_id_str = require_str(&params, "job_id")?;
let job_id_str = params
.get("job_id")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'job_id' parameter".into()))?;
let job_id = Uuid::parse_str(job_id_str).map_err(|_| {
ToolError::InvalidParameters(format!("invalid job ID format: {}", job_id_str))
@@ -709,7 +720,10 @@ impl Tool for CancelJobTool {
let start = std::time::Instant::now();
let requester_id = ctx.user_id.clone();
let job_id_str = require_str(&params, "job_id")?;
let job_id_str = params
.get("job_id")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'job_id' parameter".into()))?;
let job_id = Uuid::parse_str(job_id_str).map_err(|_| {
ToolError::InvalidParameters(format!("invalid job ID format: {}", job_id_str))
+10 -3
View File
@@ -3,7 +3,7 @@
use async_trait::async_trait;
use crate::context::JobContext;
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_param, require_str};
use crate::tools::tool::{Tool, ToolError, ToolOutput};
/// Tool for JSON manipulation (parse, query, transform).
pub struct JsonTool;
@@ -46,9 +46,16 @@ impl Tool for JsonTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let operation = require_str(&params, "operation")?;
let operation = params
.get("operation")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'operation' parameter".to_string())
})?;
let data = require_param(&params, "data")?;
let data = params
.get("data")
.ok_or_else(|| ToolError::InvalidParameters("missing 'data' parameter".to_string()))?;
let result = match operation {
"parse" => {
+160
View File
@@ -0,0 +1,160 @@
//! NEAR AI Marketplace tool.
use async_trait::async_trait;
use rust_decimal::Decimal;
use crate::context::JobContext;
use crate::tools::tool::{Tool, ToolError, ToolOutput};
/// Tool for interacting with the NEAR AI marketplace.
pub struct MarketplaceTool {
// TODO: Add marketplace client
}
impl MarketplaceTool {
/// Create a new marketplace tool.
pub fn new() -> Self {
Self {}
}
}
impl Default for MarketplaceTool {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl Tool for MarketplaceTool {
fn name(&self) -> &str {
"marketplace"
}
fn description(&self) -> &str {
"Interact with the NEAR AI marketplace: search jobs, submit bids, deliver work."
}
fn parameters_schema(&self) -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": {
"action": {
"type": "string",
"enum": ["search_jobs", "get_job", "submit_bid", "accept_job", "submit_work", "get_status"],
"description": "The marketplace action to perform"
},
"job_id": {
"type": "string",
"description": "Job ID (for get_job, submit_bid, accept_job, submit_work)"
},
"query": {
"type": "string",
"description": "Search query (for search_jobs)"
},
"category": {
"type": "string",
"description": "Job category filter (for search_jobs)"
},
"bid_amount": {
"type": "number",
"description": "Bid amount in NEAR (for submit_bid)"
},
"work_url": {
"type": "string",
"description": "URL to submitted work (for submit_work)"
},
"work_description": {
"type": "string",
"description": "Description of completed work (for submit_work)"
}
},
"required": ["action"]
})
}
async fn execute(
&self,
params: serde_json::Value,
_ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let action = params
.get("action")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'action' parameter".to_string())
})?;
// TODO: Implement actual marketplace integration
let result = match action {
"search_jobs" => {
// Placeholder response
serde_json::json!({
"jobs": [],
"total": 0,
"message": "Marketplace integration not yet implemented"
})
}
"get_job" => {
let job_id = params
.get("job_id")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'job_id' parameter".to_string())
})?;
serde_json::json!({
"job_id": job_id,
"status": "not_found",
"message": "Marketplace integration not yet implemented"
})
}
"submit_bid" => {
serde_json::json!({
"success": false,
"message": "Marketplace integration not yet implemented"
})
}
"accept_job" => {
serde_json::json!({
"success": false,
"message": "Marketplace integration not yet implemented"
})
}
"submit_work" => {
serde_json::json!({
"success": false,
"message": "Marketplace integration not yet implemented"
})
}
"get_status" => {
serde_json::json!({
"connected": false,
"message": "Marketplace integration not yet implemented"
})
}
_ => {
return Err(ToolError::InvalidParameters(format!(
"unknown action: {}",
action
)));
}
};
Ok(ToolOutput::success(result, start.elapsed()))
}
fn estimated_cost(&self, params: &serde_json::Value) -> Option<Decimal> {
// Bidding has a cost
if params.get("action").and_then(|v| v.as_str()) == Some("submit_bid") {
Some(Decimal::new(1, 2)) // 0.01 NEAR gas cost
} else {
None
}
}
fn requires_sanitization(&self) -> bool {
true // External marketplace data
}
}
+15 -4
View File
@@ -17,7 +17,7 @@ use std::sync::Arc;
use async_trait::async_trait;
use crate::context::JobContext;
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
use crate::tools::tool::{Tool, ToolError, ToolOutput};
use crate::workspace::{Workspace, paths};
/// Identity files that the LLM must not overwrite via tool calls.
@@ -81,7 +81,10 @@ impl Tool for MemorySearchTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let query = require_str(&params, "query")?;
let query = params
.get("query")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'query' parameter".to_string()))?;
let limit = params
.get("limit")
@@ -173,7 +176,12 @@ impl Tool for MemoryWriteTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let content = require_str(&params, "content")?;
let content = params
.get("content")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'content' parameter".to_string())
})?;
if content.trim().is_empty() {
return Err(ToolError::InvalidParameters(
@@ -329,7 +337,10 @@ impl Tool for MemoryReadTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let path = require_str(&params, "path")?;
let path = params
.get("path")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'path' parameter".to_string()))?;
let doc = self
.workspace
+8
View File
@@ -1,17 +1,22 @@
//! Built-in tools that come with the agent.
mod echo;
mod ecommerce;
pub mod extension_tools;
mod file;
mod http;
mod job;
mod json;
mod marketplace;
mod memory;
mod restaurant;
pub mod routine;
pub(crate) mod shell;
mod taskrabbit;
mod time;
pub use echo::EchoTool;
pub use ecommerce::EcommerceTool;
pub use extension_tools::{
ToolActivateTool, ToolAuthTool, ToolInstallTool, ToolListTool, ToolRemoveTool, ToolSearchTool,
};
@@ -19,9 +24,12 @@ pub use file::{ApplyPatchTool, ListDirTool, ReadFileTool, WriteFileTool};
pub use http::HttpTool;
pub use job::{CancelJobTool, CreateJobTool, JobStatusTool, ListJobsTool};
pub use json::JsonTool;
pub use marketplace::MarketplaceTool;
pub use memory::{MemoryReadTool, MemorySearchTool, MemoryTreeTool, MemoryWriteTool};
pub use restaurant::RestaurantTool;
pub use routine::{
RoutineCreateTool, RoutineDeleteTool, RoutineHistoryTool, RoutineListTool, RoutineUpdateTool,
};
pub use shell::ShellTool;
pub use taskrabbit::TaskRabbitTool;
pub use time::TimeTool;
+172
View File
@@ -0,0 +1,172 @@
//! Restaurant reservation tool.
use async_trait::async_trait;
use crate::context::JobContext;
use crate::tools::tool::{Tool, ToolError, ToolOutput};
/// Tool for restaurant reservations (OpenTable, Resy, etc.).
pub struct RestaurantTool {
// TODO: Add reservation API clients
}
impl RestaurantTool {
/// Create a new restaurant tool.
pub fn new() -> Self {
Self {}
}
}
impl Default for RestaurantTool {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl Tool for RestaurantTool {
fn name(&self) -> &str {
"restaurant"
}
fn description(&self) -> &str {
"Search restaurants, check availability, and make reservations via OpenTable, Resy, etc."
}
fn parameters_schema(&self) -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": {
"action": {
"type": "string",
"enum": ["search", "check_availability", "make_reservation", "cancel_reservation", "get_reservation"],
"description": "The restaurant action to perform"
},
"query": {
"type": "string",
"description": "Search query (cuisine type, restaurant name, etc.)"
},
"location": {
"type": "object",
"properties": {
"city": { "type": "string" },
"neighborhood": { "type": "string" },
"latitude": { "type": "number" },
"longitude": { "type": "number" }
},
"description": "Location to search near"
},
"date": {
"type": "string",
"description": "Reservation date (YYYY-MM-DD)"
},
"time": {
"type": "string",
"description": "Preferred time (HH:MM)"
},
"party_size": {
"type": "integer",
"description": "Number of guests"
},
"restaurant_id": {
"type": "string",
"description": "Restaurant ID (for check_availability, make_reservation)"
},
"reservation_id": {
"type": "string",
"description": "Reservation ID (for cancel_reservation, get_reservation)"
},
"guest_name": {
"type": "string",
"description": "Name for the reservation"
},
"guest_phone": {
"type": "string",
"description": "Phone number for the reservation"
},
"guest_email": {
"type": "string",
"description": "Email for the reservation"
}
},
"required": ["action"]
})
}
async fn execute(
&self,
params: serde_json::Value,
_ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let action = params
.get("action")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'action' parameter".to_string())
})?;
// TODO: Implement actual restaurant reservation API integrations
let result = match action {
"search" => {
let query = params.get("query").and_then(|v| v.as_str()).unwrap_or("");
serde_json::json!({
"query": query,
"restaurants": [],
"message": "Restaurant integration not yet implemented"
})
}
"check_availability" => {
let restaurant_id = params
.get("restaurant_id")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters(
"missing 'restaurant_id' parameter".to_string(),
)
})?;
serde_json::json!({
"restaurant_id": restaurant_id,
"available_times": [],
"message": "Restaurant integration not yet implemented"
})
}
"make_reservation" => {
serde_json::json!({
"success": false,
"message": "Restaurant integration not yet implemented"
})
}
"cancel_reservation" => {
serde_json::json!({
"cancelled": false,
"message": "Restaurant integration not yet implemented"
})
}
"get_reservation" => {
let reservation_id = params.get("reservation_id").and_then(|v| v.as_str());
serde_json::json!({
"reservation_id": reservation_id,
"found": false,
"message": "Restaurant integration not yet implemented"
})
}
_ => {
return Err(ToolError::InvalidParameters(format!(
"unknown action: {}",
action
)));
}
};
Ok(ToolOutput::success(result, start.elapsed()))
}
fn requires_sanitization(&self) -> bool {
true // External restaurant data
}
}
+25 -7
View File
@@ -20,7 +20,7 @@ use crate::agent::routine::{
use crate::agent::routine_engine::RoutineEngine;
use crate::context::JobContext;
use crate::db::Database;
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
use crate::tools::tool::{Tool, ToolError, ToolOutput};
// ==================== routine_create ====================
@@ -106,16 +106,25 @@ impl Tool for RoutineCreateTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let name = require_str(&params, "name")?;
let name = params
.get("name")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'name'".to_string()))?;
let description = params
.get("description")
.and_then(|v| v.as_str())
.unwrap_or("");
let trigger_type = require_str(&params, "trigger_type")?;
let trigger_type = params
.get("trigger_type")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'trigger_type'".to_string()))?;
let prompt = require_str(&params, "prompt")?;
let prompt = params
.get("prompt")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'prompt'".to_string()))?;
// Build trigger
let trigger = match trigger_type {
@@ -399,7 +408,10 @@ impl Tool for RoutineUpdateTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let name = require_str(&params, "name")?;
let name = params
.get("name")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'name'".to_string()))?;
let mut routine = self
.store
@@ -502,7 +514,10 @@ impl Tool for RoutineDeleteTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let name = require_str(&params, "name")?;
let name = params
.get("name")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'name'".to_string()))?;
let routine = self
.store
@@ -580,7 +595,10 @@ impl Tool for RoutineHistoryTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let name = require_str(&params, "name")?;
let name = params
.get("name")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'name'".to_string()))?;
let limit = params
.get("limit")
+11 -98
View File
@@ -30,7 +30,7 @@ use tokio::process::Command;
use crate::context::JobContext;
use crate::sandbox::{SandboxManager, SandboxPolicy};
use crate::tools::tool::{Tool, ToolDomain, ToolError, ToolOutput, require_str};
use crate::tools::tool::{Tool, ToolDomain, ToolError, ToolOutput};
/// Maximum output size before truncation (64KB).
const MAX_OUTPUT_SIZE: usize = 64 * 1024;
@@ -343,12 +343,12 @@ impl ShellTool {
// Use sandbox if configured; fail-closed (never silently fall through
// to unsandboxed execution when sandbox was intended).
if let Some(ref sandbox) = self.sandbox
&& (sandbox.is_initialized() || sandbox.config().enabled)
{
return self
.execute_sandboxed(sandbox, cmd, &cwd, timeout_duration)
.await;
if let Some(ref sandbox) = self.sandbox {
if sandbox.is_initialized() || sandbox.config().enabled {
return self
.execute_sandboxed(sandbox, cmd, &cwd, timeout_duration)
.await;
}
}
// Only execute directly when no sandbox was configured at all.
@@ -401,7 +401,10 @@ impl Tool for ShellTool {
params: serde_json::Value,
_ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let command = require_str(&params, "command")?;
let command = params
.get("command")
.and_then(|v| v.as_str())
.ok_or_else(|| ToolError::InvalidParameters("missing 'command' parameter".into()))?;
let workdir = params.get("workdir").and_then(|v| v.as_str());
let timeout = params.get("timeout").and_then(|v| v.as_u64());
@@ -426,26 +429,6 @@ impl Tool for ShellTool {
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 {
true // Shell output could contain anything
}
@@ -544,76 +527,6 @@ mod tests {
));
}
/// Replicate the extraction logic from agent_loop.rs to prove it works
/// when `arguments` is a `serde_json::Value::Object` (the common case
/// that was previously broken because `Value::Object.as_str()` returns None).
#[test]
fn test_destructive_command_extraction_from_object_args() {
let arguments = serde_json::json!({"command": "rm -rf /tmp/stuff"});
let cmd = arguments
.get("command")
.and_then(|c| c.as_str().map(String::from))
.or_else(|| {
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)))
});
assert_eq!(cmd.as_deref(), Some("rm -rf /tmp/stuff"));
assert!(requires_explicit_approval(cmd.as_deref().unwrap()));
}
/// Verify extraction still works when `arguments` is a JSON string
/// (rare, but possible if the LLM provider returns string-encoded JSON).
#[test]
fn test_destructive_command_extraction_from_string_args() {
let arguments =
serde_json::Value::String(r#"{"command": "git push --force origin main"}"#.to_string());
let cmd = arguments
.get("command")
.and_then(|c| c.as_str().map(String::from))
.or_else(|| {
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)))
});
assert_eq!(cmd.as_deref(), Some("git push --force origin main"));
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]
fn test_sandbox_policy_builder() {
let tool = ShellTool::new()
+157
View File
@@ -0,0 +1,157 @@
//! TaskRabbit tool for real-world task delegation.
use async_trait::async_trait;
use rust_decimal::Decimal;
use crate::context::JobContext;
use crate::tools::tool::{Tool, ToolError, ToolOutput};
/// Tool for delegating real-world tasks via TaskRabbit.
pub struct TaskRabbitTool {
// TODO: Add TaskRabbit API client
}
impl TaskRabbitTool {
/// Create a new TaskRabbit tool.
pub fn new() -> Self {
Self {}
}
}
impl Default for TaskRabbitTool {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl Tool for TaskRabbitTool {
fn name(&self) -> &str {
"taskrabbit"
}
fn description(&self) -> &str {
"Delegate real-world tasks to TaskRabbit taskers (delivery, assembly, cleaning, etc.)."
}
fn parameters_schema(&self) -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": {
"action": {
"type": "string",
"enum": ["search_taskers", "get_quote", "book_task", "get_status", "cancel_task"],
"description": "The TaskRabbit action to perform"
},
"task_type": {
"type": "string",
"enum": ["delivery", "assembly", "moving", "cleaning", "handyman", "other"],
"description": "Type of task"
},
"description": {
"type": "string",
"description": "Detailed description of the task"
},
"location": {
"type": "object",
"properties": {
"address": { "type": "string" },
"city": { "type": "string" },
"state": { "type": "string" },
"zip": { "type": "string" }
},
"description": "Location for the task"
},
"scheduled_time": {
"type": "string",
"description": "ISO 8601 datetime for when the task should be performed"
},
"budget": {
"type": "number",
"description": "Maximum budget for the task in USD"
},
"task_id": {
"type": "string",
"description": "Task ID (for get_status, cancel_task)"
}
},
"required": ["action"]
})
}
async fn execute(
&self,
params: serde_json::Value,
_ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let action = params
.get("action")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'action' parameter".to_string())
})?;
// TODO: Implement actual TaskRabbit API integration
let result = match action {
"search_taskers" => {
serde_json::json!({
"taskers": [],
"message": "TaskRabbit integration not yet implemented"
})
}
"get_quote" => {
serde_json::json!({
"quotes": [],
"message": "TaskRabbit integration not yet implemented"
})
}
"book_task" => {
serde_json::json!({
"booked": false,
"message": "TaskRabbit integration not yet implemented"
})
}
"get_status" => {
let task_id = params.get("task_id").and_then(|v| v.as_str());
serde_json::json!({
"task_id": task_id,
"status": "unknown",
"message": "TaskRabbit integration not yet implemented"
})
}
"cancel_task" => {
serde_json::json!({
"cancelled": false,
"message": "TaskRabbit integration not yet implemented"
})
}
_ => {
return Err(ToolError::InvalidParameters(format!(
"unknown action: {}",
action
)));
}
};
Ok(ToolOutput::success(result, start.elapsed()))
}
fn estimated_cost(&self, params: &serde_json::Value) -> Option<Decimal> {
// Booking a task has associated costs
if params.get("action").and_then(|v| v.as_str()) == Some("book_task") {
params
.get("budget")
.and_then(|v| v.as_f64())
.map(|b| Decimal::try_from(b).unwrap_or_default())
} else {
None
}
}
fn requires_sanitization(&self) -> bool {
true // External TaskRabbit data
}
}
+25 -5
View File
@@ -4,7 +4,7 @@ use async_trait::async_trait;
use chrono::{DateTime, Utc};
use crate::context::JobContext;
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
use crate::tools::tool::{Tool, ToolError, ToolOutput};
/// Tool for getting current time and date operations.
pub struct TimeTool;
@@ -52,7 +52,12 @@ impl Tool for TimeTool {
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let operation = require_str(&params, "operation")?;
let operation = params
.get("operation")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'operation' parameter".to_string())
})?;
let result = match operation {
"now" => {
@@ -64,7 +69,12 @@ impl Tool for TimeTool {
})
}
"parse" => {
let timestamp = require_str(&params, "timestamp")?;
let timestamp = params
.get("timestamp")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'timestamp' parameter".to_string())
})?;
let dt: DateTime<Utc> = timestamp.parse().map_err(|e| {
ToolError::InvalidParameters(format!("invalid timestamp: {}", e))
@@ -77,9 +87,19 @@ impl Tool for TimeTool {
})
}
"diff" => {
let ts1 = require_str(&params, "timestamp")?;
let ts1 = params
.get("timestamp")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'timestamp' parameter".to_string())
})?;
let ts2 = require_str(&params, "timestamp2")?;
let ts2 = params
.get("timestamp2")
.and_then(|v| v.as_str())
.ok_or_else(|| {
ToolError::InvalidParameters("missing 'timestamp2' parameter".to_string())
})?;
let dt1: DateTime<Utc> = ts1.parse().map_err(|e| {
ToolError::InvalidParameters(format!("invalid timestamp: {}", e))

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