Compare commits

..
22 Commits
Author SHA1 Message Date
github-actions[bot]GitHubgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
c1ca3bb91c chore: release v0.4.0 (#124)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-02-17 16:35:49 +00:00
e499795b8c fix: undo() peeks without popping, breaking repeated undo and leaking redo stack (#71)
* fix: undo() peeks without popping, breaking repeated undo and leaking redo stack

undo() used self.undo_stack.back() (peek) instead of pop_back(), so
repeated undo always returned the same checkpoint while pushing to
the redo stack unboundedly.

Additionally, redo() did not save the current state to the undo stack,
breaking the undo/redo cycle.

Changes:
- undo(): change back() to pop_back(), return owned Checkpoint
- redo(): accept current_turn/current_messages params, save current
  state to undo stack before popping from redo stack
- Update process_undo/process_redo callers in agent_loop.rs
- Add tests for repeated undo, undo/redo cycling, stack size invariant

* fix: standardize lock ordering and extract push_undo helper

Address review feedback:
- Standardize lock order (Session before UndoManager) in process_undo
  and process_redo to match process_user_input and prevent deadlocks
- Extract push_undo() helper to deduplicate push-and-trim logic shared
  by checkpoint() and redo()

* docs: add move-semantics notes and stack invariant to UndoManager

Address review feedback requesting documentation about the ownership
semantics of undo/redo parameters and the stack size invariant.

---------

Co-authored-by: Yi LIU <[email protected]>
Co-authored-by: firat.sertgoz <[email protected]>
2026-02-17 16:35:19 +00:00
5e1da4827a fix: check Content-Length before downloading HTTP response body (#74)
* fix: check Content-Length before downloading HTTP response body

The HTTP tool previously downloaded the entire response body into memory
before checking the size limit, allowing a malicious server to cause OOM.
Now the Content-Length header is checked first to reject obviously
oversized responses, and the body is streamed with a hard size cap so
reading stops as soon as the limit is exceeded.

* fix: check chunk size before allocation and fix Content-Length parsing

Address review feedback:
- Check body.len() + chunk.len() before extend_from_slice to prevent
  OOM from a single oversized chunk
- Use let-chain for Content-Length parsing instead of unwrap_or to
  gracefully handle invalid headers

* docs: document MAX_RESPONSE_SIZE rationale and add tracing on rejection

Address review feedback: explain why 5 MB was chosen for the response
size limit and log a warning when Content-Length causes early rejection.

---------

Co-authored-by: Yi LIU <[email protected]>
2026-02-17 16:33:42 +00:00
d04af5cd75 web: add integrity check for marked CDN and cap highlight regex input (#109)
* web: add integrity check for marked CDN and cap highlight regex input

* web: normalize memory search query before snippet+highlight matching

* web: place memory query length constant with top-level config

---------

Co-authored-by: Clawyered <[email protected]>
Co-authored-by: Illia Polosukhin <[email protected]>
2026-02-17 16:33:08 +00:00
956037c4d3 llm: fallback to legacy nearai.session key when loading DB session (#111)
* llm: fallback to legacy nearai.session when loading DB session

* llm: simplify session fallback load with if-let form

---------

Co-authored-by: Clawyered <[email protected]>
Co-authored-by: Illia Polosukhin <[email protected]>
2026-02-17 16:32:52 +00:00
68a1851c19 feat: add cooldown management to FailoverProvider (#114)
Track per-provider failure state with lock-free atomics and temporarily
skip providers that have repeatedly failed with retryable errors. This
reduces latency when a provider is known to be down, instead of
wasting time on every request trying all providers sequentially.

- Add CooldownConfig (duration + threshold) and ProviderCooldown (atomics)
- Rewrite try_providers() to skip cooled-down providers, with a safety
  net that always tries the oldest-cooled provider if all are down
- Add 2 env vars: LLM_FAILOVER_COOLDOWN_SECS, LLM_FAILOVER_THRESHOLD
- Add MultiCallMockProvider and 7 new test cases
- Mark "Cooldown management" as complete in FEATURE_PARITY.md

Co-authored-by: Claude Opus 4.6 <[email protected]>
2026-02-17 08:00:32 +00:00
github-actions[bot]GitHubgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
dfa105539b chore: release v0.4.0 (#122)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-02-17 07:53:03 +00:00
8929baf76a feat: add review and fix-issue project commands (#104)
* feat: add review and fix-issue project commands

Add 4 Claude Code project commands adapted from global skills,
tailored to IronClaw's build/test/lint workflow and conventions:

- review-pr: Paranoid architect PR review across 6 lenses
- review-crate: Deep Rust crate audit (vulnerabilities, bugs, unfinished work)
- respond-pr: Triage and address PR review comments
- fix-issue: End-to-end GitHub issue resolution with branch/plan/implement flow

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* fix: address PR review feedback on project commands

- Add headRefOid to gh pr view and resolve {owner}/{repo} in review-pr.md
  so Step 6 line comments actually work (Gemini + Copilot)
- Add --paginate to gh api calls in respond-pr.md for large PRs (Gemini + Copilot)
- Use gh repo view --json defaultBranchRef instead of hardcoded main/master
  fallback in fix-issue.md (Gemini)
- Narrow allowed-tools in all four commands to match repo convention of
  specific subcommands (Bash(cargo fmt:*) style) instead of broad wildcards (Copilot)
- Clarify >20 files guidance in review-pr.md: read all, process in priority order (Copilot)
- Make cargo audit mandatory with install hint in review-crate.md (Gemini)

Co-Authored-By: Claude Opus 4.6 <[email protected]>

---------

Co-authored-by: Claude Opus 4.6 <[email protected]>
2026-02-17 07:52:42 +00:00
e07dfab449 chore: remove accidentally committed .sidecar and .todos directories (#123)
These are local tool data directories (Sidecar) that should not be
tracked. Added both to .gitignore to prevent future accidents.

Co-authored-by: Claude Opus 4.6 <[email protected]>
2026-02-17 07:29:18 +00:00
6783cba4e4 feat: move per-invocation approval check into Tool trait (#119)
* feat: move per-invocation approval check into Tool trait (#94)

Move shell-specific destructive command detection out of agent_loop.rs
into a new `requires_approval_for(params)` method on the Tool trait.
ShellTool overrides it to check for destructive patterns (rm -rf, git
push --force, etc.) while the default delegates to `requires_approval()`.

This follows the project's tool architecture principle of keeping
tool-specific logic out of the main agent codebase, and enables other
tools to implement per-invocation gating without modifying the agent loop.

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* fix: requires_approval_for default should return false, not self.requires_approval()

The previous default broke auto-approval for all tools: since
requires_approval_for() delegated to requires_approval(), any
auto-approved tool would have its auto-approval immediately overridden
on every invocation. The correct semantic is:

- requires_approval(): "Does this tool use the approval system?"
- requires_approval_for(params): "Should this invocation override auto-approval?"

The default for the latter must be false (allow auto-approval).
ShellTool's fallback for safe commands is also changed to false.

Co-Authored-By: Claude Opus 4.6 <[email protected]>

---------

Co-authored-by: Claude Opus 4.6 <[email protected]>
2026-02-17 06:42:13 +00:00
63302ab406 feat: add polished boot screen on CLI startup (#118)
* feat: add polished boot screen on CLI startup

Replace the minimal one-liner REPL banner with an ANSI-styled status
panel that summarizes the agent's runtime state after initialization:
model, database, tool count, enabled features, active channels, and
the gateway URL. The boot screen is shown only in interactive CLI mode
(skipped for single-message -m mode).

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* fix: address PR review feedback on boot screen

- Stop logging gateway auth token in tracing::info! (security)
- Use info.agent_name instead of hardcoded "IronClaw" in header
- Display embeddings provider in features line: "embeddings (openai)"
- Add Display impl for DatabaseBackend, simplify main.rs match

Co-Authored-By: Claude Opus 4.6 <[email protected]>

---------

Co-authored-by: Claude Opus 4.6 <[email protected]>
2026-02-17 05:50:57 +00:00
7c553b0973 feat: Add lifecycle hooks system with 6 interception points (#18)
* feat: Add lifecycle hooks system with 6 interception points

Implement extensible hook infrastructure for intercepting and transforming
agent operations at well-defined points in the lifecycle:

- BeforeInbound: intercept/modify/reject incoming user messages
- BeforeToolCall: intercept/modify/reject tool executions (chat + job)
- BeforeOutbound: intercept/modify/suppress outgoing responses
- TransformResponse: transform final response before completing a turn
- OnSessionStart: fire-and-forget notification on new session creation
- OnSessionEnd: fire-and-forget notification on session pruning

Hooks execute in priority order with modification chaining, reject
short-circuits, configurable failure modes (FailOpen/FailClosed),
and per-hook timeouts. Empty registry is zero-cost (all hooks pass
through immediately).

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* fix: enforce hook fail-closed semantics

* Merge upstream/main into feat/hooks-system-clean

Resolve merge conflicts:
- FEATURE_PARITY.md: Keep both upstream cron/routines status and hooks status
- src/error.rs: Keep both Hook and Orchestrator/Worker error variants

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* fix: resolve CI test failures in pairing store and wizard

- Fix pairing store truncate bug: record_failed_approve used
  .truncate(true) which wiped the file before reading, causing rate
  limiting to never accumulate past 1 attempt. Changed to
  .truncate(false) to preserve existing data.

- Fix wizard test: skip test_install_missing_bundled_channels when
  telegram WASM artifact specifically isn't available, not just when
  all channels are empty (whatsapp may exist without telegram).

- Add workspace exclude for subcrate directories to prevent cargo
  from discovering them as workspace members during builds.

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* fix: address PR #18 review comments

- Remove duplicate maybe_hydrate_thread call (rebase artifact)
- Fix RwLock held across async hook execution in HookRegistry::run()
- Add tracing::warn for silent JSON parse failures in hook modifications
- Refactor execute_tool_inner to accept &WorkerDeps instead of 8 Arc params
- Use real user_id from JobContext instead of job_id UUID in BeforeToolCall hook

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* fix: cargo fmt + remove tracked worktree breaking CI

- Apply rustfmt formatting (method chain line breaks, match arm style)
- Remove .claude/worktrees/ from git tracking (caused submodule error in CI)
- Add .claude/worktrees/ to .gitignore

Co-Authored-By: Claude Opus 4.6 <[email protected]>

---------

Co-authored-by: Firat Sertgoz <[email protected]>
Co-authored-by: Claude Opus 4.6 <[email protected]>
2026-02-17 05:40:05 +00:00
github-actions[bot]GitHubgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
5e44185e48 chore: release v0.3.0 (#117)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-02-17 05:39:12 +00:00
72623c9e5b feat: direct api key and cheap model (#116)
* feat: Support direct API key auth and cheap model routing

Allow using IronClaw with any OpenAI-compatible API provider (e.g.
Anthropic Claude) via API key, without requiring NEAR AI session auth.

Changes:
- Skip session authentication in chat_completions mode (API key auth)
- Skip first-run onboard check when NEARAI_API_KEY is configured
- Add `cheap_model` config field (NEARAI_CHEAP_MODEL env var) for a
  secondary lightweight model used for heartbeat, routing, evaluation
- Add `create_cheap_llm_provider()` factory in llm module
- Add `cheap_llm` to AgentDeps with fallback to main model
- Route heartbeat through cheap model to reduce costs
- Fix wizard compilation for new config field

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* fix: address PR #20 review feedback

- Check API key presence (not api_mode) for auth skip (ilblackdragon)
- Add Settings::load() call in check_onboard_needed (ilblackdragon)
- Warn and ignore cheap_model for non-NearAi backends (ilblackdragon)
- Add unit tests for create_cheap_llm_provider (ilblackdragon)
- Minor formatting cleanup in cheap provider match arm

Co-Authored-By: Claude Opus 4.6 <[email protected]>

---------

Co-authored-by: Samuel Barbosa <[email protected]>
Co-authored-by: Claude Opus 4.6 <[email protected]>
2026-02-17 01:24:27 +00:00
github-actions[bot]GitHubgithub-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
6895adbcc9 chore: release v0.2.0 (#60)
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
2026-02-16 22:28:14 +00:00
Vlad Frolov f1480f471b ci: Explicitly enable cargo-dist caching for binary artifacts building 2026-02-16 21:34:49 +01:00
Vlad Frolov 9db949746f ci: Skip building binary artifacts on every PR 2026-02-16 21:29:03 +01:00
61a123a746 Add GitHub tool and Discord channel (#34)
* Add GitHub tool for IronClaw - manage repos, issues, PRs, and workflows

* Add Discord channel for IronClaw - slash commands and button interactions

* Security fixes: URL encoding, secret validation, Discord button handler

- Add URL encoding for all path segments and query parameters (P1)
- Add path segment validation to prevent path traversal
- Add secret_exists check for better error messages (P2)
- Fix http_request signature to use 5 args (P2)
- Fix Discord button handler to check member field (P2)
- Fix typo in Discord slash command format (P2)
- Add github.capabilities.json and discord.capabilities.json (Blocker)
- Add Cargo.toml for Discord channel (Blocker)
- Add limit caps (max 100) for all list operations (P3)
- Remove debug logging

* Apply Copilot review fixes

Security & Code Quality:
- Use secret_get instead of workspace_read for GitHub token
- Remove manual Authorization header (host injects via capabilities)
- Add validation for file paths (reject path traversal)
- Add validation for workflow_id and git refs
- Fix url_encode_query comment
- Add release profile optimizations to Cargo.toml files
- Fix package names to match conventions (github-tool, discord-channel)
- Add metadata fields to Cargo.toml
- Fix rate limits to be consistent (60/min, 3600/hr)
- Fix Discord user_name to filter empty global_name
- Fix Discord metadata serialization error handling
- Update Discord README to clarify which secrets are used by host vs WASM
- Better formatting for Discord command option values

* applied all PR change requests and comments

* cleaned up workspace

* Adding validation for empty path segments and event enum in GitHub tool

* addedvalidation for events and vaidation to reject empty file path in github tools and implemented safe UTF-8 trunacating

* added codegen units and updated truncating logic also update capabilities.json as requested by copilot review

* added codegen units and updated truncating logic also update capabilities.json as requested by copilot review

* fixed message trucating and remove url_encode alias, also appled all requested changes from last PR comment

---------

Co-authored-by: root <root@cafx>
Co-authored-by: Peni <[email protected]>
Co-authored-by: Illia Polosukhin <[email protected]>
Co-authored-by: firat.sertgoz <[email protected]>
2026-02-16 15:58:13 +04:00
0e981429ee feat: mark Ollama + OpenAI-compatible as implemented (#102)
Co-authored-by: BroccoliFin <[email protected]>
2026-02-16 03:38:06 +00:00
Illia PolosukhinandClaude Opus 4.6 1b38a64e15 docs: add module specification rules to CLAUDE.md
Any agent working on a module with a README.md spec must read it first,
keep code and spec in sync, and treat the spec as the tiebreaker when
they disagree.

Co-Authored-By: Claude Opus 4.6 <[email protected]>
2026-02-15 00:33:41 -08:00
Illia PolosukhinandClaude Opus 4.6 2e5f8b60d5 docs: add setup/onboarding specification (src/setup/README.md)
Authoritative specification for the 7-step onboarding wizard. Documents
the full flow, settings persistence (two-layer architecture), platform
caveats (macOS keychain dialogs, URL passwords), secrets context, and
a modification checklist for future contributors.

Co-Authored-By: Claude Opus 4.6 <[email protected]>
2026-02-15 00:29:47 -08:00
f0a0642e7d feat: multi-provider inference + libSQL onboarding selection (#92)
* feat: add interactive database backend selection during onboarding

Previously the onboarding wizard silently defaulted to PostgreSQL because
libsql wasn't in the default feature set. Now both backends ship by default
and the wizard presents a selection prompt when both are available.

DATABASE_BACKEND env var still bypasses the prompt for headless/CI use.

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* fix: resolve libSQL onboarding crash, keychain double-prompt, and setup audit findings

Three bugs fixed:

1. libSQL onboarding crash ("Missing required setting 'database_url'"):
   DatabaseConfig::resolve() only checked DATABASE_BACKEND env var, falling
   back to Postgres default. Now reads settings.database_backend, plus
   settings.libsql_path and settings.libsql_url as fallbacks.

2. OS keychain prompts twice during startup: Config::from_env() and
   Config::from_db() both called get_master_key(). Now caches the key in
   SECRETS_MASTER_KEY env var after first read so from_db() skips keychain.

3. "Path not found: nearai.session" warning: from_db_map() tried to apply
   app-specific DB keys (nearai.session_token) to the Settings struct.
   Now skips keys that don't map to known Settings fields. Also fixed
   bootstrap migration key mismatch (nearai.session -> nearai.session_token).

Setup module audit fixes (14 findings):
- Replace unreachable!() with proper error in provider match
- Extract setup_api_key_provider() to deduplicate setup_anthropic/setup_openai
- Add SAFETY comments to all unsafe std::env::set_var blocks
- Fix .unwrap() calls with proper error handling
- Remove incorrect #[allow(dead_code)] on used TelegramUpdate::update_id
- Log warnings instead of silently discarding HTTP errors in Telegram binding
- Guard select_many against empty options, fix mask_api_key for non-ASCII
- Update stale doc comment in mod.rs, rename misleading variable
- Add 7 new tests (model fetcher fallbacks, channel discovery, secret gen)

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* fix: address PR review feedback (set_var safety, parse warnings, db_map efficiency)

1. Replace unsafe set_var keychain caching with OnceLock<String> in
   SecretsConfig::resolve(). Eliminates the env var write from main.rs
   entirely, using a process-wide OnceLock cache instead.

2. Log tracing::warn when database_backend or llm_backend settings
   fail to parse, instead of silently falling back to defaults.

3. Remove O(K*S) get() pre-check in from_db_map(). Instead, let set()
   run and match on "Path not found" errors to skip unknown keys,
   avoiding full Settings serialization per key.

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* fix: address critical/high audit findings across WASM sub-crates

- Telegram: remove .unwrap() panic on workspace_read (owner_id check)
- WhatsApp: use configured api_version instead of hardcoded v18.0
- WhatsApp: log config parse errors before falling back to defaults
- Slack: log serialization errors in emit_message and json_response
- Google Docs: safe array access for batch update replies
- Google Sheets: safe array access for add_sheet replies
- Google Calendar: fix doc comment secret name mismatch
- Gmail: avoid unnecessary String allocation in UNREAD check

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* fix: address second-round PR review feedback

- Validate custom model ID is non-empty (loop until valid input)
- Warn on unknown DATABASE_BACKEND env var before defaulting to Postgres
- Force re-selection when llm_backend contains unknown provider value
- Use ok_or_else for proper String error type in google-sheets

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* fix: harden setup module error handling and secret safety

- Introduce ChannelSetupError typed enum replacing raw String errors
  across all channel setup functions (setup_telegram, setup_http,
  setup_tunnel, setup_wasm_channel, validate_telegram_token)
- Add From<ChannelSetupError> for SetupError to simplify call sites
- Convert setup_telegram retry from recursion to loop (unbounded stack)
- Stop printing HTTP webhook secret plaintext to terminal
- Use secret_input() for Turso auth token (was visible input())
- Replace dirs::home_dir().unwrap_or_default() with proper error
- Fix UTF-8 panic in model name truncation (byte-index to chars-based)
- Log warning in secret_exists() instead of silently swallowing errors
- Deduplicate generate_webhook_secret() to delegate to shared helper

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* fix: replace unreachable!() with error return in setup wizard

The provider match in step_inference_provider was guarded by
is_known but used unreachable!() as the catch-all. If a new
provider is added to the is_known check without a corresponding
match arm, this would panic at runtime. Return a typed error
instead.

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* fix: remove unsafe set_var, use thread-safe overlay for injected secrets

Address PR #92 review comments:
- Replace all 5 unsafe `std::env::set_var()` calls with safe alternatives
- Add INJECTED_VARS OnceLock<HashMap> overlay in config.rs, checked by
  optional_env() before falling back to std::env::var()
- Cache wizard API key in SetupWizard.llm_api_key field instead of env
- Pass explicit key param to fetch_anthropic_models/fetch_openai_models
- Persist env-provided API keys to secrets store during onboarding

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* fix: address remaining PR review comments (clippy, TODO, secrets backend ordering)

- Fix empty line after doc comment (clippy: empty_line_after_doc_comments)
- Collapse nested if in optional_env overlay check (clippy: collapsible_if)
- Remove dangling TODO(#XX) placeholder issue ref in channels.rs
- Fix init_secrets_context to respect selected database_backend when both
  postgres and libsql features are compiled, preventing wrong-backend
  secrets storage when DATABASE_URL is set but libsql was chosen

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* fix: address latest PR review comments (SecretString, empty env, docs, embeddings)

- Change wizard llm_api_key from String to SecretString to prevent
  accidental logging of API keys
- Fix inject_llm_keys_from_secrets skipping when env var is set but
  empty, matching optional_env's treatment of empty as unset
- Fix inverted doc comment on INJECTED_VARS (env checked first, overlay
  is the fallback, not the other way around)
- Update stale "env vars" comments in main.rs to reflect overlay pattern
- Fix step_embeddings not seeing cached OpenAI key from wizard session

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* fix: OAuth callback listener binds IPv4 first to match redirect URLs

The listener was binding to [::1] (IPv6) first, but NEAR AI and other
OAuth flows redirect to http://127.0.0.1:9876/... (IPv4 explicit).
On macOS and most systems, [::1] and 127.0.0.1 are separate addresses,
so the browser's connection to 127.0.0.1 was refused when the listener
was on [::1]. Reversed the bind order: try 127.0.0.1 first, fall back
to [::1] if IPv4 is unavailable.

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* fix: cache keychain key eagerly to avoid redundant macOS password dialogs

Replace has_master_key() with get_master_key() in step_security() and
immediately build SecretsCrypto from the result. This eliminates redundant
keychain accesses later in init_secrets_context(), each of which triggers
macOS system dialogs (keychain unlock + app authorization).

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* fix: persist DATABASE_BACKEND to ~/.ironclaw/.env for libSQL startup

The wizard saved database_backend only to the database, but
Config::from_env() needs it BEFORE connecting to any database (to
decide which backend to use). Without it, the backend defaults to
Postgres and then fails with "Missing required setting database_url".

Now save all database bootstrap vars (DATABASE_BACKEND, DATABASE_URL,
LIBSQL_PATH, LIBSQL_URL) to ~/.ironclaw/.env via save_bootstrap_env().

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* fix: status command shows libSQL backend and skips keychain probe

The status command only checked DATABASE_URL (postgres), showing
"not configured" for libSQL users. Now detects the DATABASE_BACKEND
env var and reports libSQL path and Turso sync status.

Also remove the keychain probe from status. get_generic_password()
triggers macOS unlock+authorization dialogs which is terrible UX
for a read-only diagnostic command.

Co-Authored-By: Claude Opus 4.6 <[email protected]>

* style: fix rustfmt formatting in bootstrap test

Co-Authored-By: Claude Opus 4.6 <[email protected]>

---------

Co-authored-by: Claude Opus 4.6 <[email protected]>
2026-02-15 08:24:51 +00:00
64 changed files with 6768 additions and 2619 deletions
+97
View File
@@ -0,0 +1,97 @@
---
description: Fetch a GitHub issue, create a branch, research the codebase, plan the fix, implement with tests, and commit
disable-model-invocation: true
allowed-tools: Bash(gh issue view:*), Bash(gh repo view:*), Bash(git fetch:*), Bash(git checkout:*), Bash(git status:*), Bash(git branch:*), Bash(git add:*), Bash(git commit:*), Bash(cargo fmt:*), Bash(cargo clippy:*), Bash(cargo test:*), Read, Edit, Write, Grep, Glob
argument-hint: "<issue-number or github-issue-url>"
---
# Fix GitHub Issue
## Step 1: Resolve the issue
Parse `$ARGUMENTS` to extract the issue number:
- If it's a URL like `https://github.com/owner/repo/issues/42`, extract `42`.
- If it's a bare number, use it directly.
- If empty, stop and ask the user for an issue number.
Fetch the issue:
```
gh issue view {number} --json title,body,labels,assignees,comments,state
```
If the issue is closed, warn the user and ask if they still want to proceed.
## Step 2: Create a branch
Create a fresh branch off the latest main:
1. Fetch latest: `git fetch origin`
2. Detect default branch: `gh repo view --json defaultBranchRef --jq .defaultBranchRef.name`
3. Create and switch to a new branch: `git checkout -b fix/{number}-{short-slug} origin/{default-branch}`
- `{short-slug}` is 3-5 words from the issue title, lowercase, hyphenated (e.g. `fix/42-idor-workspace-check`)
If the working tree has uncommitted changes, warn the user and stop. Do not stash or discard their work.
## Step 3: Understand the issue
Summarize the issue in 2-3 sentences. Identify:
- **What's broken or missing** (the symptom or feature request)
- **Acceptance criteria** (what "done" looks like, from the issue body or comments)
- **Constraints** (mentioned technologies, backward compatibility, performance requirements)
If the issue is unclear or ambiguous, list the open questions. These will be addressed during planning.
## Step 4: Research the codebase
Before planning, gather context:
1. **Find relevant code** - Search for files, functions, types, and patterns mentioned in the issue. Read them in full.
2. **Trace the flow** - If the issue is about a specific behavior, trace the code path from the entry point (route handler, CLI command, etc.) through to the relevant logic.
3. **Check existing tests** - Find tests related to the affected code. Understand what's already covered.
4. **Check for prior art** - Look for similar patterns in the codebase that solve analogous problems. Prefer consistency with existing patterns.
## Step 5: Enter planning mode
Enter planning mode to design the implementation. The plan MUST cover:
1. **Root cause** (for bugs) or **design approach** (for features)
2. **Files to modify** with specific descriptions of what changes in each
3. **New files** (if any) with justification for why they're needed
4. **Tests to add** - every code path introduced or changed needs a test:
- Happy path (expected input produces expected output)
- Error paths (invalid input, missing data, permission denied)
- Edge cases (empty collections, boundary values, concurrent access)
5. **IronClaw-specific concerns**:
- If the change touches persistence, both database backends must be updated (`postgres.rs` and `libsql_backend.rs`)
- New `Database` trait methods need implementations in both backends
- No `.unwrap()` or `.expect()` in production code
- Use `crate::` imports, not `super::`
- Error types via `thiserror` in `error.rs`
6. **Migration or compatibility concerns** (if any)
Follow the project's CLAUDE.md guidance for architecture decisions.
Wait for user approval before implementing.
## Step 6: Implement
After the plan is approved:
1. Implement each change from the plan.
2. Write all planned tests.
3. Run IronClaw's full quality gate:
- `cargo fmt`
- `cargo clippy --all --benches --tests --examples --all-features` (zero warnings)
- `cargo test --lib` (all tests pass)
4. If any check fails, fix it before proceeding.
Note: Integration tests (`--test workspace_integration`) require PostgreSQL and are expected to fail locally. Only `--lib` test failures are blocking.
## Step 7: Commit and summarize
1. Commit with a descriptive message referencing the issue (e.g. `fix: prevent IDOR in function call outputs (#42)`).
2. Summarize what was done:
- Files changed with line references
- Tests added and what they cover
- Any follow-up work or open questions
+81
View File
@@ -0,0 +1,81 @@
---
description: Respond to PR review comments — triage, plan fixes, implement after confirmation, push, and reply to reviewers
disable-model-invocation: true
allowed-tools: Bash(gh pr list:*), Bash(gh pr comment:*), Bash(gh api:*), Bash(gh repo view:*), Bash(git branch:*), Bash(git status:*), Bash(git add:*), Bash(git commit:*), Bash(git push:*), Bash(cargo fmt:*), Bash(cargo clippy:*), Bash(cargo test:*), Read, Edit, Write, Grep, Glob
argument-hint: "[pr-number (optional, auto-detects from branch)]"
---
# Review and Address PR Comments
## Step 1: Find the PR
If `$ARGUMENTS` is provided, use that as the PR number. Otherwise, detect the PR for the current branch:
```
gh pr list --head $(git branch --show-current) --json number,title,url --jq '.[0]'
```
If no PR is found, tell the user and stop.
## Step 2: Fetch all review comments
Resolve the repo owner and name:
```
gh repo view --json owner,name --jq '"\(.owner.login)/\(.name)"'
```
Fetch the full set of review comments (not issue-level comments):
```
gh api --paginate repos/{owner}/{repo}/pulls/{number}/comments
```
Also fetch the review summaries:
```
gh api --paginate repos/{owner}/{repo}/pulls/{number}/reviews
```
Deduplicate comments that appear multiple times (bots sometimes post the same finding under different IDs). Group by the actual issue being raised, not by comment ID.
## Step 3: Triage and plan
For each unique issue raised in the comments:
1. **Check if already addressed** - Read the current code at the referenced location. If a prior commit already fixed it, note it as "already resolved".
2. **Assess validity** - Determine if the comment identifies a real problem or is a false positive. Be honest about false positives but explain why.
3. **Classify severity** - Critical (security/data loss), High (bugs/broken behavior), Medium (correctness/robustness), Low (style/naming/nits).
4. **Plan the fix** - For each valid unresolved issue, describe the specific code change needed.
Present the plan as a table to the user:
| # | Issue | File:Line | Severity | Status | Planned Fix |
|---|-------|-----------|----------|--------|-------------|
Wait for user confirmation before proceeding to implementation.
## Step 4: Implement fixes
After user confirms:
1. Implement each fix in the plan.
2. Run IronClaw's quality gate to verify nothing breaks:
- `cargo fmt`
- `cargo clippy --all --benches --tests --examples --all-features`
- `cargo test --lib`
3. Commit with a descriptive message referencing the PR review.
4. Push to the branch.
## Step 5: Reply to comments
For each comment addressed, reply on the PR with a short message stating what was fixed and the commit SHA. For false positives or already-resolved items, reply explaining why no change was needed.
## Rules
- Never guess at code you haven't read. Always read the referenced file and line before assessing a comment.
- Group duplicate comments (same issue reported by multiple bots) and reply to all of them.
- Do not make changes beyond what the review comments ask for. Stay focused.
- If a comment suggests a change you disagree with, present your reasoning to the user during the planning phase rather than silently ignoring it.
- Follow IronClaw conventions: no `.unwrap()` in production code, use `crate::` imports, `thiserror` errors.
- If changes touch persistence, verify both database backends are updated.
+245
View File
@@ -0,0 +1,245 @@
---
description: Deep audit of the IronClaw crate for vulnerabilities, bugs, unfinished work, inconsistencies, and oversights
disable-model-invocation: true
allowed-tools: Bash(cargo fmt:*), Bash(cargo clippy:*), Bash(cargo test:*), Bash(cargo audit:*), Bash(git diff:*), Bash(git log:*), Bash(git show:*), Bash(wc:*), Read, Grep, Glob, Task
argument-hint: "[path/to/crate]"
---
# Rust Crate Audit
You are performing a thorough audit of a Rust crate. Your goal is to find every vulnerability, bug, unfinished piece of work, inconsistency, and oversight before it ships. Leave no stone unturned.
## Step 1: Locate the crate
Parse `$ARGUMENTS`:
- If a path is provided, use it as the crate root.
- If empty, use the current working directory.
Verify it's a valid Rust crate by checking for `Cargo.toml`. If not found, stop and ask the user.
## Step 2: Understand the crate
Read `Cargo.toml` to understand:
- Crate name, version, edition
- Dependencies (look for outdated, unmaintained, or suspicious crates)
- Feature flags and their implications
- Build scripts (`build.rs`) if any
Read `CLAUDE.md`, `README.md`, or top-level documentation if present to understand intent and architecture.
Read `src/lib.rs` or `src/main.rs` to get the module tree. Then read each module's `mod.rs` or top-level file to build a mental map of the crate's structure before diving into details.
Read all Rust files (`src/*.rs`) to make sure everything is in context when you are reasoning.
## Step 3: Run the compiler's checks
Run these commands and capture output. Do NOT fix anything, just collect findings:
```
cargo fmt --check 2>&1
```
```
cargo clippy --all --benches --tests --examples --all-features -- -W clippy::all -W clippy::pedantic -W clippy::nursery 2>&1
```
```
cargo test --lib 2>&1
```
If any of these fail, record the failures as findings. If `cargo test` has ignored tests, note which ones and why.
Note: Integration tests (`--test workspace_integration`) require a PostgreSQL database and are expected to fail locally. Only report `--lib` test failures as blocking.
## Step 4: Scan for unfinished work
Search the entire `src/` tree for:
```
todo!
unimplemented!
fixme
FIXME
TODO
HACK
XXX
SAFETY:
stub
placeholder
temporary
```
For each match:
- Is it in production code or test code?
- Is it a genuine incomplete feature or a deliberate placeholder?
- Is there a tracking issue referenced?
- Could this panic at runtime?
Any `todo!()` or `unimplemented!()` in non-test code is **High severity** (runtime panic).
## Step 5: Audit for vulnerabilities and unsafe code
### 5a. Unsafe code
Search for all `unsafe` blocks. For each one:
- Is the safety invariant documented with a `// SAFETY:` comment?
- Is the invariant actually upheld by the surrounding code?
- Could the unsafe block be replaced with a safe alternative?
- Are there any pointer dereferences, transmutes, or FFI calls?
### 5b. Unwrap and panic paths
Search for `.unwrap()`, `.expect(`, `panic!`, `unreachable!` in non-test code. For each:
- Can this actually panic in production?
- Is there a code path that reaches this with None/Err?
- Should it be replaced with proper error handling (`?`, `.ok()`, `.unwrap_or_default()`)?
IronClaw convention: `.unwrap()` and `.expect()` are banned in production code. Any occurrence outside `#[cfg(test)]` blocks is a **High severity** finding.
### 5c. SQL and injection vectors
Search for string formatting used in SQL queries, shell commands, or HTML:
- `format!` used near `.execute(`, `.query(`, `Command::new(`
- String interpolation in query construction vs parameterized queries
- User input flowing into file paths (`Path::new`, `std::fs::`)
IronClaw has two database backends (PostgreSQL and libSQL). Check both for injection vectors.
### 5d. Cryptographic issues
If the crate uses crypto:
- Are comparisons constant-time? (look for `==` on secrets/hashes vs `subtle::ConstantTimeEq`)
- Is randomness from `OsRng` / `thread_rng` and not a fixed seed?
- Are keys/secrets zeroized after use? (`secrecy`, `zeroize` crates)
- Are deprecated algorithms used? (MD5, SHA1 for security, RC4, DES)
### 5e. Resource exhaustion
- Are there unbounded allocations? (`Vec` growing from user input without limits)
- Are there unbounded loops? (retry loops without max attempts)
- Are file reads bounded? (`std::fs::read_to_string` on user-provided paths)
- Are timeouts set on all network operations?
- Are there connection/resource leaks? (opened but never closed, missing `Drop`)
### 5f. Error handling
- Are errors swallowed silently? (`let _ = ...`, `.ok()` discarding errors that matter)
- Do error types carry enough context to debug in production?
- Are there error type mismatches? (returning generic `anyhow::Error` where a typed error would prevent confusion)
- Is `thiserror` used consistently for error types (IronClaw convention)?
## Step 6: Check for inconsistencies
### 6a. Naming conventions
- Are types, functions, modules named consistently? (e.g., mixing `get_` and `fetch_`, `create_` and `new_`)
- Do similar operations follow the same patterns?
### 6b. Duplicate or near-duplicate code
Look for:
- Functions that do nearly the same thing with minor variations (candidates for generics or shared helpers)
- Repeated error mapping patterns that should be extracted
- Copy-pasted SQL queries or string templates with slight differences
- Identical struct definitions or conversion logic in different modules
### 6c. API consistency
- Do similar functions take arguments in the same order?
- Are return types consistent? (e.g., some functions return `Option<T>`, similar ones return `Result<T, E>`)
- Are visibility modifiers consistent? (`pub` where it should be `pub(crate)`, or vice versa)
### 6d. Dead code and unused items
- Are there functions, structs, or modules that nothing references?
- Are there `#[allow(dead_code)]` annotations that should be investigated?
- Are there feature-gated items where the feature is never enabled?
### 6e. Import style
IronClaw convention: use `crate::` imports, not `super::`. Flag any `super::` imports in non-test code.
## Step 7: Inspect for change oversights
### 7a. Partial refactors
- Are there old patterns coexisting with new patterns?
- Are there renamed types/functions where some call sites still use the old name via a compatibility alias?
- Are there comments referencing behavior that no longer exists?
### 7b. Trait implementation gaps
- If a trait is defined, do all intended types implement it?
- Are there `impl` blocks that look incomplete?
- Are `Default` implementations sensible?
IronClaw key traits: `Database` (~60 methods), `Channel`, `Tool`, `LlmProvider`, `SuccessEvaluator`, `EmbeddingProvider`. If any new methods were added to `Database`, verify both `postgres.rs` and `libsql_backend.rs` implement them.
### 7c. Test coverage gaps
- Are there public functions without any test?
- Are there error paths without tests?
- Are there recently-changed functions where the tests still assert old behavior?
### 7d. Documentation drift
- Do doc comments match actual function behavior?
- Are examples in doc comments still valid and compilable?
## Step 8: Dependency audit
Review `Cargo.toml` and `Cargo.lock`:
- Are there duplicate versions of the same crate in the lock file? (potential version conflicts)
- Are there dependencies with known security advisories? Run `cargo audit` to check (install with `cargo install cargo-audit` if not present).
- Are there heavy dependencies used for trivial functionality?
- Are dependency features minimal?
## Step 9: Present findings
Compile all findings into a structured report. Group by severity, then by category.
### Format
For each finding:
```
### [Severity] Category: One-line summary
**Location:** `file_path:line_number`
**Category:** Vulnerability | Bug | Unfinished | Inconsistency | Duplicate | Oversight | Style
**Description:**
Detailed explanation of the issue, why it matters, and how it could manifest.
**Suggested fix:**
Concrete suggestion with code if applicable.
```
### Severity levels
- **Critical**: Security vulnerability, data loss, or crash in production
- **High**: Bug that causes incorrect behavior, `todo!()`/`unimplemented!()` in prod code, or missing validation on trust boundaries
- **Medium**: Inconsistency, duplicate code, incomplete error handling, missing tests for important paths
- **Low**: Naming inconsistency, unnecessary complexity, documentation drift, minor dead code
- **Nit**: Style preference, optional improvement
### Summary table
End with a summary table:
| # | Severity | Category | File:Line | Finding |
|---|----------|----------|-----------|---------|
And a final tally: X Critical, Y High, Z Medium, W Low, V Nit.
## Rules
- Read every file before reporting on it. Never guess about code you haven't seen.
- Be specific. "This might have issues" is worthless. "Line 42 calls `.unwrap()` on a `Result` that returns `Err` when the DB connection is dropped" is useful.
- Distinguish certainty levels: "this IS a bug" vs "this COULD be a bug if X".
- Don't invent problems to look thorough. If the code is solid, say so.
- Focus on substance over style. Don't flag formatting unless it causes real confusion.
- Respect existing project conventions (check CLAUDE.md). Don't flag patterns the project explicitly endorses.
- When in doubt about severity, round up.
- For large crates (>50 files), prioritize: core logic > public API > internal utilities > tests > examples.
- Use the Task tool to parallelize file reading across modules when the crate is large.
- Do NOT fix anything. This is a read-only audit. Report findings for the user to action.
+170
View File
@@ -0,0 +1,170 @@
---
description: Paranoid architect review of a PR — fetches diff, reads changed files, deep review across 6 lenses, posts findings as GitHub comments
disable-model-invocation: true
allowed-tools: Bash(gh pr view:*), Bash(gh pr diff:*), Bash(gh pr comment:*), Bash(gh api:*), Bash(gh repo view:*), Bash(git diff:*), Bash(git log:*), Read, Grep, Glob
argument-hint: "<pr-number or github-pr-url>"
---
# Paranoid Architect Code Review
You are reviewing this PR as a paranoid architect. Your job is to find every bug, vulnerability, race condition, edge case, and undocumented assumption before it ships. Assume adversarial users, concurrent access, and Murphy's law.
## Step 1: Resolve the PR
Parse `$ARGUMENTS` to extract the PR number:
- If it's a URL like `https://github.com/owner/repo/pull/123`, extract `123`.
- If it's a bare number, use it directly.
- If empty, stop and ask the user for a PR number.
Fetch PR metadata (including head commit SHA for posting line comments later):
```
gh pr view {number} --json title,body,baseRefName,headRefName,headRefOid,files,additions,deletions
```
Save the `headRefOid` value, you'll need it as `commit_id` in Step 6.
## Step 2: Load the full diff
```
gh pr diff {number}
```
Also get the list of changed files:
```
gh pr diff {number} --name-only
```
## Step 3: Read every changed file in full
For each changed file, read the ENTIRE current file (not just the diff hunks). You need surrounding context to catch:
- Callers of modified functions that now behave differently
- Trait/interface contracts that the change may violate
- Invariants established elsewhere that the diff breaks
If the PR touches more than 20 files, still read all of them, but process in this priority order: service logic > routes/handlers > models/types > tests > docs. Batch reads in groups of ~20 if needed.
## Step 4: Deep review
Go through the changes with each of these lenses. For every finding, note the file, line range, severity, and a concrete description.
### IronClaw-specific checks
In addition to the general lenses below, check IronClaw conventions (see CLAUDE.md):
- No `.unwrap()` or `.expect()` in production code (tests are fine)
- Use `crate::` imports, not `super::`
- Error types use `thiserror` in `error.rs`
- If the change touches persistence, verify both database backends are updated (PostgreSQL in `postgres.rs` AND libSQL in `libsql_backend.rs`)
- New tools must implement the `Tool` trait correctly and be registered in `registry.rs`
- External tool output must pass through the safety layer
### 4a. Correctness and bugs
- Off-by-one errors, wrong comparison operators, inverted conditions
- Unreachable code, dead branches, impossible match arms
- Type confusion (mixing up IDs, using wrong enum variant)
- Incorrect error propagation (swallowed errors, wrong error type/status code)
- Broken invariants (e.g. uniqueness assumptions violated, ordering assumptions wrong)
- Concurrency issues (TOCTOU, missing locks, race conditions between check and use)
### 4b. Edge cases and failure handling
- What happens with empty input, None/null, zero-length collections?
- What happens when external services fail (DB down, HTTP timeout, malformed response)?
- What happens at integer boundaries (overflow, underflow, i64::MAX)?
- What happens with malformed or adversarial input (invalid UTF-8, huge payloads, deeply nested JSON)?
- Are all error paths tested? Does every `?` propagation make sense?
- Are partial failures handled (e.g. wrote to DB but failed to emit event)?
### 4c. Security (assume a malicious actor)
- **Authentication/Authorization bypass**: Can an unauthenticated user reach this? Can workspace A's user access workspace B's data? Are there IDOR vulnerabilities?
- **Injection**: SQL injection via string interpolation? Command injection? Log injection? Header injection?
- **Data leakage**: Are secrets, PII, or conversation content logged? Returned in error messages? Exposed in API responses?
- **Resource exhaustion / DoS**: Can an attacker send unbounded input? Trigger expensive operations without rate limits? Cause OOM via large allocations?
- **Financial abuse**: Can tokens/credits be consumed without being tracked? Can usage limits be bypassed?
- **Replay / race conditions**: Can the same request be replayed for double-spend? Can concurrent requests bypass limits?
- **Cryptographic issues**: Timing attacks on comparisons? Weak randomness? Missing HMAC verification?
### 4d. Test coverage
- Is every new public function/method tested?
- Are error paths tested (not just happy paths)?
- Are edge cases covered (empty input, boundary values, concurrent access)?
- Do existing tests still make sense with the new changes, or do they assert stale behavior?
- Are there integration/e2e tests for the full flow?
- If a test is missing, describe exactly what test should be written.
### 4e. Documentation and assumptions
- Are new assumptions documented in comments? (e.g. "this field is always non-empty because X")
- Are non-obvious algorithms or business rules explained?
- Are API contracts (request/response shapes, error codes, status codes) documented?
- Are there TODO/FIXME/HACK comments that should be tracked as issues?
### 4f. Architectural concerns
- Does this change follow existing patterns in the codebase, or does it introduce a new one without justification?
- Are there unnecessary abstractions or premature generalizations?
- Is there duplicated logic that should be extracted?
- Are dependencies between modules clean, or does this create circular/tight coupling?
- Will this change make future work harder?
## Step 5: Present findings
Summarize findings to the user as a table:
| # | Severity | Category | File:Line | Finding | Suggested Fix |
|---|----------|----------|-----------|---------|---------------|
Severity levels:
- **Critical**: Security vulnerability, data loss, or financial exploit
- **High**: Bug that will cause incorrect behavior in production
- **Medium**: Robustness issue, missing validation, or incomplete error handling
- **Low**: Style, naming, documentation, or minor improvement
- **Nit**: Optional suggestion, take-it-or-leave-it
Ask the user which findings to post as PR comments. Default: all Critical, High, and Medium.
## Step 6: Post comments on GitHub
Resolve the repo owner and name if not already known:
```
gh repo view --json owner,name --jq '"\(.owner.login)/\(.name)"'
```
For each approved finding, post a review comment on the PR at the specific file and line. Use the `headRefOid` from Step 1 as the `commit_id`:
```
gh api repos/{owner}/{repo}/pulls/{number}/comments \
-f body="..." \
-f path="..." \
-f commit_id="{headRefOid}" \
-F line=... \
-f side="RIGHT"
```
For findings that span multiple locations or are architectural, post as a regular PR comment:
```
gh pr comment {number} --body "..."
```
Format each comment clearly:
- Severity tag (e.g. `**High Severity**`)
- One-line summary
- Detailed explanation of the issue
- Concrete suggestion for the fix (with code if possible)
## Rules
- Read every changed file in full before writing a single finding. Context matters.
- Never post a comment about code you haven't actually read. Verify line numbers against the actual file.
- Be specific. "This might have issues" is useless. "Line 42 returns 404 but should return 400 because X" is useful.
- Distinguish between "this IS a bug" and "this COULD be a bug if X". Be honest about certainty.
- Don't nitpick formatting or style unless it causes actual confusion. Focus on substance.
- If the code is good and you find nothing, say so. Don't invent problems to look thorough.
- Respect the project's CLAUDE.md privacy rules: never include customer data, secrets, or PII in comments.
- When in doubt about severity, round up. It's cheaper to dismiss a false alarm than to miss a real bug.
-1
View File
@@ -39,7 +39,6 @@ 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,6 +1,15 @@
.env
.env.local
.env.*
!.env.example
# Claude Code worktrees
.claude/worktrees/
# Sidecar tool data
.sidecar/
.todos/
target/
+54
View File
@@ -7,6 +7,60 @@ 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
+16
View File
@@ -630,6 +630,22 @@ 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::`
Generated
+20 -174
View File
@@ -352,23 +352,6 @@ dependencies = [
"syn 2.0.114",
]
[[package]]
name = "async-tungstenite"
version = "0.32.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8acc405d38be14342132609f06f02acaf825ddccfe76c4824a69281e0458ebd4"
dependencies = [
"atomic-waker",
"futures-core",
"futures-io",
"futures-task",
"futures-util",
"log",
"pin-project-lite",
"tokio",
"tungstenite 0.28.0",
]
[[package]]
name = "atomic-waker"
version = "1.1.2"
@@ -522,7 +505,7 @@ dependencies = [
"rustc-hash 1.1.0",
"shlex",
"syn 2.0.114",
"which 4.4.2",
"which",
]
[[package]]
@@ -833,72 +816,6 @@ version = "0.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724"
[[package]]
name = "chromiumoxide"
version = "0.8.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6c18200611490f523adb497ddd4744d6d536e243f6add13e7eeeb1c05904fbb1"
dependencies = [
"async-tungstenite",
"base64 0.22.1",
"cfg-if",
"chromiumoxide_cdp",
"chromiumoxide_types",
"dunce",
"fnv",
"futures",
"futures-timer",
"pin-project-lite",
"reqwest",
"serde",
"serde_json",
"thiserror 1.0.69",
"tokio",
"tracing",
"url",
"which 8.0.0",
"windows-registry 0.5.3",
]
[[package]]
name = "chromiumoxide_cdp"
version = "0.8.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b8f78027ced540595dcbaf9e2f3413cbe3708b839ff239d2858acaea73915dcb"
dependencies = [
"chromiumoxide_pdl",
"chromiumoxide_types",
"serde",
"serde_json",
]
[[package]]
name = "chromiumoxide_pdl"
version = "0.8.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "0d2c7b7c6b41a0de36d00a284e619017e0f4aec5c9bc8d90614b9e1687984f20"
dependencies = [
"chromiumoxide_types",
"either",
"heck 0.4.1",
"once_cell",
"proc-macro2",
"quote",
"regex",
"serde",
"serde_json",
]
[[package]]
name = "chromiumoxide_types"
version = "0.8.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "309ba8f378bbc093c93f06beb7bd4c5ceffdf14107ad99cacbbf063709926795"
dependencies = [
"serde",
"serde_json",
]
[[package]]
name = "chrono"
version = "0.4.43"
@@ -910,7 +827,7 @@ dependencies = [
"num-traits",
"serde",
"wasm-bindgen",
"windows-link 0.2.1",
"windows-link",
]
[[package]]
@@ -962,7 +879,7 @@ version = "4.5.55"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a92793da1a46a5f2a02a6f4c46c6496b28c43638adea8306fcb0caa1634f24e5"
dependencies = [
"heck 0.5.0",
"heck",
"proc-macro2",
"quote",
"syn 2.0.114",
@@ -1580,12 +1497,6 @@ version = "0.15.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1aaf95b3e5c8f23aa320147307562d361db0ae0d51242340f558153b4eb2439b"
[[package]]
name = "dunce"
version = "1.0.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "92773504d58c093f6de2459af4af33faa518c13451eb8f2b5698ed3d36e7c813"
[[package]]
name = "dyn-clone"
version = "1.0.20"
@@ -1652,12 +1563,6 @@ dependencies = [
"syn 2.0.114",
]
[[package]]
name = "env_home"
version = "0.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c7f84e12ccf0a7ddc17a6c41c93326024c42920d7ee630d04950e6926645c0fe"
[[package]]
name = "equivalent"
version = "1.0.2"
@@ -2115,12 +2020,6 @@ dependencies = [
"hashbrown 0.14.5",
]
[[package]]
name = "heck"
version = "0.4.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "95505c38b4572b2d910cecb0281560f54b440a19336cbbcb27bf6ce6adc6f5a8"
[[package]]
name = "heck"
version = "0.5.0"
@@ -2368,7 +2267,7 @@ dependencies = [
"tokio",
"tower-service",
"tracing",
"windows-registry 0.6.1",
"windows-registry",
]
[[package]]
@@ -2591,7 +2490,7 @@ dependencies = [
[[package]]
name = "ironclaw"
version = "0.1.3"
version = "0.5.0"
dependencies = [
"aes-gcm",
"aho-corasick",
@@ -2602,7 +2501,6 @@ dependencies = [
"blake3",
"bollard",
"bytes",
"chromiumoxide",
"chrono",
"clap",
"cron",
@@ -2799,7 +2697,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d7c4b02199fee7c5d21a5ae7d8cfa79a6ef5bb2fc834d6e9058e89c825efdc55"
dependencies = [
"cfg-if",
"windows-link 0.2.1",
"windows-link",
]
[[package]]
@@ -3411,7 +3309,7 @@ dependencies = [
"libc",
"redox_syscall 0.5.18",
"smallvec",
"windows-link 0.2.1",
"windows-link",
]
[[package]]
@@ -4048,7 +3946,7 @@ version = "0.8.16"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "72c225407d8e52ef8cf094393781ecda9a99d6544ec28d90a6915751de259264"
dependencies = [
"heck 0.5.0",
"heck",
"proc-macro2",
"quote",
"refinery-core",
@@ -6284,7 +6182,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5f38f7a5eb2f06f53fe943e7fb8bf4197f7cf279f1bc52c0ce56e9d3ffd750a4"
dependencies = [
"anyhow",
"heck 0.5.0",
"heck",
"indexmap 2.13.0",
"wit-parser",
]
@@ -6352,17 +6250,6 @@ dependencies = [
"rustix 0.38.44",
]
[[package]]
name = "which"
version = "8.0.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d3fabb953106c3c8eea8306e4393700d7657561cb43122571b172bbfb7c7ba1d"
dependencies = [
"env_home",
"rustix 1.1.3",
"winsafe",
]
[[package]]
name = "whoami"
version = "2.1.0"
@@ -6396,7 +6283,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8738c5a7ef3a9de0fae10f8b84091a2aa4e059d8fef23de202ab689812b6bc6e"
dependencies = [
"anyhow",
"heck 0.5.0",
"heck",
"proc-macro2",
"quote",
"shellexpand",
@@ -6472,9 +6359,9 @@ checksum = "b8e83a14d34d0623b51dce9581199302a221863196a1dde71a7663a4c2be9deb"
dependencies = [
"windows-implement",
"windows-interface",
"windows-link 0.2.1",
"windows-result 0.4.1",
"windows-strings 0.5.1",
"windows-link",
"windows-result",
"windows-strings",
]
[[package]]
@@ -6499,47 +6386,21 @@ dependencies = [
"syn 2.0.114",
]
[[package]]
name = "windows-link"
version = "0.1.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5e6ad25900d524eaabdbbb96d20b4311e1e7ae1699af4fb28c17ae66c80d798a"
[[package]]
name = "windows-link"
version = "0.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5"
[[package]]
name = "windows-registry"
version = "0.5.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5b8a9ed28765efc97bbc954883f4e6796c33a06546ebafacbabee9696967499e"
dependencies = [
"windows-link 0.1.3",
"windows-result 0.3.4",
"windows-strings 0.4.2",
]
[[package]]
name = "windows-registry"
version = "0.6.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "02752bf7fbdcce7f2a27a742f798510f3e5ad88dbe84871e5168e2120c3d5720"
dependencies = [
"windows-link 0.2.1",
"windows-result 0.4.1",
"windows-strings 0.5.1",
]
[[package]]
name = "windows-result"
version = "0.3.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "56f42bd332cc6c8eac5af113fc0c1fd6a8fd2aa08a0119358686e5160d0586c6"
dependencies = [
"windows-link 0.1.3",
"windows-link",
"windows-result",
"windows-strings",
]
[[package]]
@@ -6548,16 +6409,7 @@ version = "0.4.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7781fa89eaf60850ac3d2da7af8e5242a5ea78d1a11c49bf2910bb5a73853eb5"
dependencies = [
"windows-link 0.2.1",
]
[[package]]
name = "windows-strings"
version = "0.4.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "56e6c93f3a0c3b36176cb1327a4958a0353d5d166c2a35cb268ace15e91d3b57"
dependencies = [
"windows-link 0.1.3",
"windows-link",
]
[[package]]
@@ -6566,7 +6418,7 @@ version = "0.5.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7837d08f69c77cf6b07689544538e017c1bfcf57e34b4c0ff58e6c2cd3b37091"
dependencies = [
"windows-link 0.2.1",
"windows-link",
]
[[package]]
@@ -6611,7 +6463,7 @@ version = "0.61.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "ae137229bcbd6cdf0f7b80a31df61766145077ddf49416a728b02cb3921ff3fc"
dependencies = [
"windows-link 0.2.1",
"windows-link",
]
[[package]]
@@ -6651,7 +6503,7 @@ version = "0.53.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "4945f9f551b88e0d65f3db0bc25c33b8acea4d9e41163edf90dcd0b19f9069f3"
dependencies = [
"windows-link 0.2.1",
"windows-link",
"windows_aarch64_gnullvm 0.53.1",
"windows_aarch64_msvc 0.53.1",
"windows_i686_gnu 0.53.1",
@@ -6809,12 +6661,6 @@ dependencies = [
"memchr",
]
[[package]]
name = "winsafe"
version = "0.0.19"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d135d17ab770252ad95e9a872d365cf3090e3be864a34ab46f48555993efc904"
[[package]]
name = "winx"
version = "0.36.4"
+13 -6
View File
@@ -1,6 +1,14 @@
[workspace]
exclude = [
"channels-src/telegram",
"channels-src/slack",
"channels-src/whatsapp",
"tools-src/gmail",
]
[package]
name = "ironclaw"
version = "0.1.3"
version = "0.5.0"
edition = "2024"
rust-version = "1.92"
description = "Secure personal AI assistant that protects your data and expands its capabilities on the fly"
@@ -122,9 +130,6 @@ bytes = "1"
base64 = "0.22.1"
mime_guess = "2.0.5"
# Headless browser automation via Chrome DevTools Protocol
chromiumoxide = { version = "0.8", default-features = false, features = ["tokio-runtime"] }
# macOS keychain
[target.'cfg(target_os = "macos")'.dependencies]
security-framework = "3"
@@ -142,7 +147,7 @@ pretty_assertions = "1"
tempfile = "3"
[features]
default = ["postgres"]
default = ["postgres", "libsql"]
postgres = [
"dep:deadpool-postgres",
"dep:tokio-postgres",
@@ -186,11 +191,13 @@ 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 = "upload"
pr-run-mode = "skip"
# 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"
+11 -15
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 |
@@ -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) | ✅ | | P2 | Local models |
| Ollama (local) | ✅ | | - | via `rig::providers::ollama` (full support) |
| node-llama-cpp | ✅ | | - | N/A for Rust |
| llama.cpp (native) | ❌ | 🔮 | P3 | Rust bindings |
@@ -174,7 +174,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|---------|----------|----------|-------|
| Auto-discovery | ✅ | ❌ | |
| Failover chains | ✅ | ✅ | `FailoverProvider` with configurable `fallback_model` |
| Cooldown management | ✅ | | Skip failed providers |
| Cooldown management | ✅ | | Lock-free per-provider cooldown in `FailoverProvider` |
| Per-session model override | ✅ | ✅ | Model selector in TUI |
| 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 |
@@ -420,14 +420,10 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
- ✅ Telegram channel (WASM, DM pairing, caption, /start)
- ❌ WhatsApp channel
- ✅ Multi-provider failover (`FailoverProvider` with retryable error classification)
- Hooks system (beforeInbound, beforeToolCall, etc.)
- Hooks system (beforeInbound, beforeToolCall, beforeOutbound, onSessionStart, onSessionEnd, transformResponse)
### P2 - Medium Priority
-Cron job scheduling
- ❌ Web Control UI
- ❌ WebChat channel
- 🚧 Media handling (caption support; no image/PDF processing)
- ❌ CLI subcommands (config, status, memory, doctor)
-Media handling (images, PDFs)
- ❌ Ollama/local model support
- ❌ Configuration hot-reload
- ❌ Webhook trigger endpoint in web gateway
+23
View File
@@ -0,0 +1,23 @@
[package]
name = "discord-channel"
version = "0.1.0"
edition = "2021"
description = "Discord channel for IronClaw"
license = "MIT OR Apache-2.0"
publish = false
[dependencies]
serde = { version = "1.0", features = ["derive"] }
serde_json = "1.0"
wit-bindgen = "0.41.0"
[lib]
crate-type = ["cdylib"]
[profile.release]
strip = true
opt-level = "s"
lto = true
codegen-units = 1
+121
View File
@@ -0,0 +1,121 @@
# Discord Channel for IronClaw
WASM channel for Discord integration - handle slash commands and button interactions via webhooks.
## Features
- **Slash Commands** - Process Discord slash commands
- **Button Interactions** - Handle button clicks
- **Thread Support** - Respond in threads
- **DM Support** - Handle direct messages
## Setup
1. Create a Discord Application at <https://discord.com/developers/applications>
2. Create a Bot and get the token
3. Set up Interactions URL to point to your IronClaw instance
4. Copy the Application ID and Public Key
5. Store in IronClaw secrets:
```bash
ironclaw secret set discord_bot_token YOUR_BOT_TOKEN
```
**Note:** The `discord_bot_token` secret is the only value read directly by this
Discord channel WASM component. The `discord_app_id` and `discord_public_key`
secrets are used by the IronClaw host (for example, to verify Discord
interaction signatures and manage slash command registration) and are not
accessed from the WASM module itself.
## Discord Configuration
### Register Slash Commands
```bash
curl -X POST \
-H "Authorization: Bot YOUR_BOT_TOKEN" \
-H "Content-Type: application/json" \
https://discord.com/api/v10/applications/YOUR_APP_ID/commands \
-d '{
"name": "ask",
"description": "Ask the AI agent",
"options": [{
"name": "question",
"description": "Your question",
"type": 3,
"required": true
}]
}'
```
### Set Interactions Endpoint
In your Discord app settings, set:
- Interactions Endpoint URL: `https://your-ironclaw.com/webhook/discord`
## Usage Examples
### Slash Command
User types: `/ask question: What is the weather?`
The agent receives:
```text
User: @username
Content: /ask question: What is the weather?
```
### Button Click
When a user clicks a button in a message, the agent receives:
```text
User: @username
Content: [Button clicked] Original message content
```
## Error Handling
If an internal error occurs (e.g., metadata serialization failure), the tool attempts to send an ephemeral message to the user:
```text
❌ Internal Error: Failed to process command metadata.
```
Check the host logs for detailed error information.
## Advanced Usage
### Embeds
To send embeds, include an `embeds` array in the `metadata_json` field of the agent's response. The structure should match the Discord API `embed` object.
## Troubleshooting
### "Invalid Signature"
- Check that `discord_public_key` is set correctly in IronClaw secrets.
- This validation happens on the host before reaching the WASM.
### "401 Unauthorized"
- Check that `discord_bot_token` is set correctly in IronClaw secrets.
- Ensure the bot is added to the server.
### "Interaction Failed"
- The interaction might have timed out (Discord requires a response within 3 seconds).
- The `interactions_endpoint_url` might be unreachable.
## Building
```bash
cd channels-src/discord
cargo build --target wasm32-wasi --release
```
## License
MIT/Apache-2.0
@@ -0,0 +1,39 @@
{
"type": "channel",
"name": "discord",
"description": "Discord Gateway/Webhook channel for handling slash commands, buttons, and messages",
"capabilities": {
"http": {
"allowlist": [
{ "host": "discord.com", "path_prefix": "/api/v10" }
],
"credentials": {
"discord_bot_token": {
"secret_name": "discord_bot_token",
"location": { "type": "header", "header_name": "Authorization", "prefix": "Bot " },
"host_patterns": ["discord.com"]
}
},
"rate_limit": {
"requests_per_minute": 60,
"requests_per_hour": 3600
}
},
"secrets": {
"allowed_names": ["discord_bot_token", "discord_*"]
},
"channel": {
"allowed_paths": ["/webhook/discord"],
"allow_polling": false,
"callback_timeout_secs": 45,
"workspace_prefix": "channels/discord/",
"emit_rate_limit": {
"messages_per_minute": 100,
"messages_per_hour": 5000
}
}
},
"config": {
"require_signature_verification": true
}
}
+476
View File
@@ -0,0 +1,476 @@
//! Discord Gateway/Webhook channel for IronClaw.
//!
//! This WASM component implements the channel interface for handling Discord
//! interactions via webhooks and sending messages back to Discord.
//!
//! # Features
//!
//! - URL verification for Discord interactions
//! - Slash command handling
//! - Message event parsing (@mentions, DMs)
//! - Thread support for conversations
//! - Response posting via Discord Web API
//! - Automatic message truncation (> 2000 chars)
//!
//! # Security
//!
//! - Signature validation is handled by the host (webhook secrets)
//! - Bot token is injected by host during HTTP requests
//! - WASM never sees raw credentials
wit_bindgen::generate!({
world: "sandboxed-channel",
path: "../../wit/channel.wit",
});
use serde::{Deserialize, Serialize};
use exports::near::agent::channel::{
AgentResponse, ChannelConfig, Guest, HttpEndpointConfig, IncomingHttpRequest,
OutgoingHttpResponse, StatusUpdate,
};
use near::agent::channel_host::{self, EmittedMessage};
/// Discord interaction wrapper.
#[derive(Debug, Deserialize)]
struct DiscordInteraction {
/// Interaction type (1=Ping, 2=ApplicationCommand, 3=MessageComponent)
#[serde(rename = "type")]
interaction_type: u8,
/// Interaction ID
id: String,
/// Application ID
application_id: String,
/// Guild ID (if in server)
#[allow(dead_code)] // Part of API payload, currently unused
guild_id: Option<String>,
/// Channel ID
channel_id: Option<String>,
/// Member info (if in server)
member: Option<DiscordMember>,
/// User info (if DM)
user: Option<DiscordUser>,
/// Command data (for slash commands)
data: Option<DiscordCommandData>,
/// Message (for component interactions)
message: Option<DiscordMessage>,
/// Token for responding
token: String,
}
#[derive(Debug, Deserialize, Clone)]
struct DiscordMember {
user: DiscordUser,
#[allow(dead_code)] // Part of API payload, currently unused
nick: Option<String>,
}
#[derive(Debug, Deserialize, Clone)]
struct DiscordUser {
id: String,
username: String,
global_name: Option<String>,
}
#[derive(Debug, Deserialize, Clone)]
struct DiscordCommandData {
#[allow(dead_code)] // Part of API payload, currently unused
id: String,
name: String,
options: Option<Vec<DiscordCommandOption>>,
}
#[derive(Debug, Deserialize, Clone)]
struct DiscordCommandOption {
name: String,
value: serde_json::Value,
}
#[derive(Debug, Deserialize, Clone)]
struct DiscordMessage {
#[allow(dead_code)] // Part of API payload, currently unused
id: String,
content: String,
channel_id: String,
#[allow(dead_code)] // Part of API payload, currently unused
author: DiscordUser,
}
/// Metadata stored with emitted messages for response routing.
#[derive(Debug, Serialize, Deserialize)]
struct DiscordMessageMetadata {
/// Discord channel ID
channel_id: String,
/// Interaction ID for followups
interaction_id: String,
/// Interaction token for responding
token: String,
/// Application ID
application_id: String,
/// Thread ID (for forum threads)
thread_id: Option<String>,
}
struct DiscordChannel;
impl Guest for DiscordChannel {
fn on_start(_config_json: String) -> Result<ChannelConfig, String> {
channel_host::log(channel_host::LogLevel::Info, "Discord channel starting");
Ok(ChannelConfig {
display_name: "Discord".to_string(),
http_endpoints: vec![HttpEndpointConfig {
path: "/webhook/discord".to_string(),
methods: vec!["POST".to_string()],
require_secret: true,
}],
poll: None,
})
}
fn on_http_request(req: IncomingHttpRequest) -> OutgoingHttpResponse {
let body_str = match std::str::from_utf8(&req.body) {
Ok(s) => s,
Err(_) => {
return json_response(400, serde_json::json!({"error": "Invalid UTF-8 body"}));
}
};
let interaction: DiscordInteraction = match serde_json::from_str(body_str) {
Ok(i) => i,
Err(e) => {
channel_host::log(
channel_host::LogLevel::Error,
&format!("Failed to parse Discord interaction: {}", e),
);
return json_response(400, serde_json::json!({"error": "Invalid interaction"}));
}
};
match interaction.interaction_type {
// Ping - Discord verification
1 => {
channel_host::log(channel_host::LogLevel::Info, "Responding to Discord ping");
json_response(200, serde_json::json!({"type": 1}))
}
// Application Command (slash command)
2 => {
handle_slash_command(&interaction);
json_response(
200,
serde_json::json!({
"type": 5,
"data": {
"content": "🤔 Thinking..."
}
}),
)
}
// Message Component (buttons, selects)
3 => {
if let Some(ref message) = interaction.message {
handle_message_component(&interaction, message);
}
json_response(200, serde_json::json!({"type": 6}))
}
_ => {
channel_host::log(
channel_host::LogLevel::Warn,
&format!(
"Unknown Discord interaction type: {}",
interaction.interaction_type
),
);
json_response(200, serde_json::json!({"type": 6}))
}
}
}
fn on_poll() {}
fn on_respond(response: AgentResponse) -> Result<(), String> {
let metadata: DiscordMessageMetadata = serde_json::from_str(&response.metadata_json)
.map_err(|e| format!("Failed to parse metadata: {}", e))?;
// Use webhook endpoint for followup
let url = format!(
"https://discord.com/api/v10/webhooks/{}/{}",
metadata.application_id, metadata.token
);
// Truncate content to 2000 characters to comply with Discord limits
let content = truncate_message(&response.content);
let mut payload = serde_json::json!({
"content": content,
});
// Check for embeds in metadata
if let Ok(meta_json) = serde_json::from_str::<serde_json::Value>(&response.metadata_json) {
if let Some(embeds) = meta_json.get("embeds") {
payload["embeds"] = embeds.clone();
}
}
let payload_bytes =
serde_json::to_vec(&payload).map_err(|e| format!("Failed to serialize: {}", e))?;
let headers = serde_json::json!({
"Content-Type": "application/json"
});
let result = channel_host::http_request(
"POST",
&url,
&headers.to_string(),
Some(&payload_bytes),
None,
);
match result {
Ok(http_response) => {
if http_response.status >= 200 && http_response.status < 300 {
channel_host::log(channel_host::LogLevel::Debug, "Posted followup to Discord");
Ok(())
} else {
let body_str = String::from_utf8_lossy(&http_response.body);
Err(format!(
"Discord API error: {} - {}",
http_response.status, body_str
))
}
}
Err(e) => Err(format!("HTTP request failed: {}", e)),
}
}
fn on_status(_update: StatusUpdate) {}
fn on_shutdown() {
channel_host::log(
channel_host::LogLevel::Info,
"Discord channel shutting down",
);
}
}
fn handle_slash_command(interaction: &DiscordInteraction) {
let user = interaction
.member
.as_ref()
.map(|m| &m.user)
.or(interaction.user.as_ref());
let user_id = user.map(|u| u.id.clone()).unwrap_or_default();
let user_name = user
.map(|u| {
u.global_name
.as_ref()
.filter(|s| !s.is_empty())
.unwrap_or(&u.username)
.clone()
})
.unwrap_or_default();
let channel_id = interaction.channel_id.clone().unwrap_or_default();
let command_name = interaction
.data
.as_ref()
.map(|d| d.name.clone())
.unwrap_or_default();
let options = interaction.data.as_ref().and_then(|d| d.options.clone());
let content = if let Some(opts) = options {
let opt_str = opts
.iter()
.map(|o| format!("{}: {}", o.name, o.value))
.collect::<Vec<_>>()
.join(", ");
format!("/{} {}", command_name, opt_str)
} else {
format!("/{}", command_name)
};
let metadata = DiscordMessageMetadata {
channel_id: channel_id.clone(),
interaction_id: interaction.id.clone(),
token: interaction.token.clone(),
application_id: interaction.application_id.clone(),
thread_id: None,
};
let metadata_json = match serde_json::to_string(&metadata) {
Ok(json) => json,
Err(e) => {
channel_host::log(
channel_host::LogLevel::Error,
&format!("Failed to serialize metadata: {}", e),
);
// Attempt to notify user of internal error
let url = format!(
"https://discord.com/api/v10/webhooks/{}/{}",
interaction.application_id, interaction.token
);
let payload = serde_json::json!({
"content": "❌ Internal Error: Failed to process command metadata.",
"flags": 64 // Ephemeral
});
let _ = channel_host::http_request(
"POST",
&url,
&serde_json::json!({"Content-Type": "application/json"}).to_string(),
Some(&serde_json::to_vec(&payload).unwrap_or_default()),
None,
);
return;
}
};
channel_host::emit_message(&EmittedMessage {
user_id,
user_name: Some(user_name),
content,
thread_id: None,
metadata_json,
});
}
fn handle_message_component(interaction: &DiscordInteraction, message: &DiscordMessage) {
// Check member first (for server contexts), then user (for DMs)
let user = interaction
.member
.as_ref()
.map(|m| &m.user)
.or(interaction.user.as_ref());
let user_id = user.map(|u| u.id.clone()).unwrap_or_default();
let user_name = user
.map(|u| {
u.global_name
.as_ref()
.filter(|s| !s.is_empty())
.unwrap_or(&u.username)
.clone()
})
.unwrap_or_default();
let channel_id = message.channel_id.clone();
let metadata = DiscordMessageMetadata {
channel_id: channel_id.clone(),
interaction_id: interaction.id.clone(),
token: interaction.token.clone(),
application_id: interaction.application_id.clone(),
thread_id: None,
};
let metadata_json = match serde_json::to_string(&metadata) {
Ok(json) => json,
Err(e) => {
channel_host::log(
channel_host::LogLevel::Error,
&format!("Failed to serialize metadata: {}", e),
);
return; // Don't emit message if metadata can't be serialized
}
};
channel_host::emit_message(&EmittedMessage {
user_id,
user_name: Some(user_name),
content: format!("[Button clicked] {}", message.content),
thread_id: None,
metadata_json,
});
}
fn json_response(status: u16, value: serde_json::Value) -> OutgoingHttpResponse {
let body = serde_json::to_vec(&value).unwrap_or_default();
let headers = serde_json::json!({"Content-Type": "application/json"});
OutgoingHttpResponse {
status,
headers_json: headers.to_string(),
body,
}
}
export!(DiscordChannel);
fn truncate_message(content: &str) -> String {
if content.len() <= 2000 {
content.to_string()
} else {
let max_bytes = 1990;
let cutoff = content
.char_indices()
.map(|(i, c)| i + c.len_utf8())
.take_while(|&end| end <= max_bytes)
.last()
.unwrap_or(0);
let mut truncated = content[..cutoff].to_string();
truncated.push_str("\n... (truncated)");
truncated
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_truncate_message() {
let short = "Hello world";
assert_eq!(truncate_message(short), short);
let long = "a".repeat(2005);
let truncated = truncate_message(&long);
assert_eq!(truncated.len(), 2006); // 1990 + 16 chars suffix
assert!(truncated.ends_with("\n... (truncated)"));
// Test with multibyte characters (Euro sign is 3 bytes)
// 1000 chars * 3 bytes = 3000 bytes
let multi = "".repeat(1000);
let truncated_multi = truncate_message(&multi);
// 1990 bytes limit. 1990 / 3 = 663 with remainder 1.
// Should truncate at 663 chars (1989 bytes).
// Suffix is 16 bytes. Total: 1989 + 16 = 2005 bytes.
assert!(truncated_multi.len() <= 2006);
assert!(truncated_multi.len() >= 2006 - 4); // Allow for max utf8 char width variance
assert!(truncated_multi.ends_with("\n... (truncated)"));
let content_part = &truncated_multi[..truncated_multi.len() - 16];
assert!(content_part.chars().all(|c| c == '€'));
}
#[test]
fn test_metadata_serialization() {
let metadata = DiscordMessageMetadata {
channel_id: "123".into(),
interaction_id: "456".into(),
token: "abc".into(),
application_id: "789".into(),
thread_id: None,
};
let json = serde_json::to_string(&metadata).unwrap();
let parsed: DiscordMessageMetadata = serde_json::from_str(&json).unwrap();
assert_eq!(parsed.channel_id, "123");
assert_eq!(parsed.interaction_id, "456");
}
}
+14 -2
View File
@@ -338,7 +338,13 @@ fn emit_message(
team_id,
};
let metadata_json = serde_json::to_string(&metadata).unwrap_or_else(|_| "{}".to_string());
let metadata_json = serde_json::to_string(&metadata).unwrap_or_else(|e| {
channel_host::log(
channel_host::LogLevel::Error,
&format!("Failed to serialize Slack metadata: {}", e),
);
"{}".to_string()
});
// Strip @ mentions of the bot from the text for cleaner messages
let cleaned_text = strip_bot_mention(&text);
@@ -366,7 +372,13 @@ 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_default();
let body = serde_json::to_vec(&value).unwrap_or_else(|e| {
channel_host::log(
channel_host::LogLevel::Error,
&format!("Failed to serialize JSON response: {}", e),
);
Vec::new()
});
let headers = serde_json::json!({"Content-Type": "application/json"});
OutgoingHttpResponse {
+9 -19
View File
@@ -285,11 +285,7 @@ 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())
@@ -844,8 +840,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"
@@ -915,15 +911,10 @@ 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_configured = channel_host::workspace_read(OWNER_ID_PATH)
.map(|s| !s.is_empty())
.unwrap_or(false);
let owner_id_str = channel_host::workspace_read(OWNER_ID_PATH).filter(|s| !s.is_empty());
if owner_configured {
if let Ok(owner_id) = channel_host::workspace_read(OWNER_ID_PATH)
.unwrap()
.parse::<i64>()
{
if let Some(ref id_str) = owner_id_str {
if let Ok(owner_id) = id_str.parse::<i64>() {
if from.id != owner_id {
channel_host::log(
channel_host::LogLevel::Debug,
@@ -937,8 +928,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
@@ -1001,8 +992,7 @@ 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 {
+23 -6
View File
@@ -254,10 +254,19 @@ struct WhatsAppChannel;
impl Guest for WhatsAppChannel {
fn on_start(config_json: String) -> Result<ChannelConfig, String> {
let config: WhatsAppConfig = serde_json::from_str(&config_json).unwrap_or(WhatsAppConfig {
api_version: default_api_version(),
reply_to_message: default_reply_to_message(),
});
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(),
}
}
};
channel_host::log(
channel_host::LogLevel::Info,
@@ -267,6 +276,9 @@ impl Guest for WhatsAppChannel {
),
);
// Persist api_version in workspace so on_respond() can read it
let _ = channel_host::workspace_write("channels/whatsapp/api_version", &config.api_version);
// WhatsApp Cloud API is webhook-only, no polling available
Ok(ChannelConfig {
display_name: "WhatsApp".to_string(),
@@ -327,11 +339,16 @@ 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/v18.0/{}/messages",
metadata.phone_number_id
"https://graph.facebook.com/{}/{}/messages",
api_version, metadata.phone_number_id
);
// Build sendMessage payload
+154 -42
View File
@@ -22,6 +22,7 @@ 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;
@@ -67,10 +68,14 @@ 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.
@@ -113,6 +118,7 @@ impl Agent {
deps.safety.clone(),
deps.tools.clone(),
deps.store.clone(),
deps.hooks.clone(),
));
Self {
@@ -138,6 +144,11 @@ 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
}
@@ -150,6 +161,10 @@ 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
@@ -301,7 +316,7 @@ impl Agent {
Some(spawn_heartbeat(
config,
workspace.clone(),
self.llm().clone(),
self.cheap_llm().clone(),
Some(notify_tx),
))
} else {
@@ -417,10 +432,32 @@ impl Agent {
match self.handle_message(&message).await {
Ok(Some(response)) if !response.is_empty() => {
let _ = self
.channels
.respond(&message, OutgoingResponse::text(response))
.await;
// 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;
}
}
}
Ok(Some(_)) => {
// Empty response, nothing to send (e.g. approval handled via send_status)
@@ -466,7 +503,33 @@ impl Agent {
async fn handle_message(&self, message: &IncomingMessage) -> Result<Option<String>, Error> {
// Parse submission type first
let submission = SubmissionParser::parse(&message.content);
let mut submission = SubmissionParser::parse(&message.content);
// Hook: BeforeInbound — allow hooks to modify or reject user input
if let Submission::UserInput { ref content } = submission {
let event = crate::hooks::HookEvent::Inbound {
user_id: message.user_id.clone(),
channel: message.channel.clone(),
content: content.clone(),
thread_id: message.thread_id.clone(),
};
match self.hooks().run(&event).await {
Err(crate::hooks::HookError::Rejected { reason }) => {
return Ok(Some(format!("[Message rejected: {}]", reason)));
}
Err(err) => {
return Ok(Some(format!("[Message blocked by hook policy: {}]", err)));
}
Ok(crate::hooks::HookOutcome::Continue {
modified: Some(new_content),
}) => {
submission = Submission::UserInput {
content: new_content,
};
}
_ => {} // Continue, fail-open errors already logged in registry
}
}
// Hydrate thread from DB if it's a historical thread not in memory
if let Some(ref external_thread_id) = message.thread_id {
@@ -875,6 +938,27 @@ 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
@@ -1152,8 +1236,8 @@ impl Agent {
}
}
// Execute each tool (with approval checking)
for tc in tool_calls {
// Execute each tool (with approval checking and hook interception)
for mut tc in tool_calls {
// Check if tool requires approval
if let Some(tool) = self.tools().get(&tc.name).await
&& tool.requires_approval()
@@ -1164,31 +1248,12 @@ impl Agent {
sess.is_tool_auto_approved(&tc.name)
};
// For shell commands, override auto-approval for
// destructive patterns that should always require
// explicit per-invocation approval.
if is_auto_approved
&& tc.name == "shell"
&& let Some(cmd) = tc
.arguments
.get("command")
.and_then(|c| c.as_str().map(String::from))
.or_else(|| {
tc.arguments
.as_str()
.and_then(|s| {
serde_json::from_str::<serde_json::Value>(s).ok()
})
.and_then(|v| {
v.get("command")
.and_then(|c| c.as_str().map(String::from))
})
})
&& crate::tools::builtin::shell::requires_explicit_approval(&cmd)
{
// 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!(
"Shell command '{}' requires explicit approval despite auto-approve",
cmd.chars().take(80).collect::<String>()
tool = %tc.name,
"Tool requires explicit approval for these parameters despite auto-approve"
);
is_auto_approved = false;
}
@@ -1208,6 +1273,47 @@ impl Agent {
}
}
// Hook: BeforeToolCall — allow hooks to modify or reject tool calls
{
let event = crate::hooks::HookEvent::ToolCall {
tool_name: tc.name.clone(),
parameters: tc.arguments.clone(),
user_id: message.user_id.clone(),
context: "chat".to_string(),
};
match self.hooks().run(&event).await {
Err(crate::hooks::HookError::Rejected { reason }) => {
context_messages.push(ChatMessage::tool_result(
&tc.id,
&tc.name,
format!("Tool call rejected by hook: {}", reason),
));
continue;
}
Err(err) => {
context_messages.push(ChatMessage::tool_result(
&tc.id,
&tc.name,
format!("Tool call blocked by hook policy: {}", err),
));
continue;
}
Ok(crate::hooks::HookOutcome::Continue {
modified: Some(new_params),
}) => match serde_json::from_str(&new_params) {
Ok(parsed) => tc.arguments = parsed,
Err(e) => {
tracing::warn!(
tool = %tc.name,
"Hook returned non-JSON modification for ToolCall, ignoring: {}",
e
);
}
},
_ => {} // Continue, fail-open errors already logged
}
}
let _ = self
.channels
.send_status(
@@ -1473,6 +1579,9 @@ 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;
@@ -1480,7 +1589,6 @@ 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)
@@ -1491,12 +1599,10 @@ 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(messages);
thread.restore_from_messages(checkpoint.messages);
Ok(SubmissionResult::ok_with_message(format!(
"Undone to turn {}. {} undo(s) remaining.",
turn_number, undo_count
@@ -1511,6 +1617,9 @@ 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;
@@ -1518,12 +1627,15 @@ impl Agent {
return Ok(SubmissionResult::ok_with_message("Nothing to redo."));
}
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 }))?;
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) {
thread.restore_from_messages(checkpoint.messages);
Ok(SubmissionResult::ok_with_message(format!(
"Redone to turn {}.",
+5
View File
@@ -14,6 +14,7 @@ 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;
@@ -49,6 +50,7 @@ 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).
@@ -64,6 +66,7 @@ impl Scheduler {
safety: Arc<SafetyLayer>,
tools: Arc<ToolRegistry>,
store: Option<Arc<dyn Database>>,
hooks: Arc<HookRegistry>,
) -> Self {
Self {
config,
@@ -72,6 +75,7 @@ impl Scheduler {
safety,
tools,
store,
hooks,
jobs: Arc::new(RwLock::new(HashMap::new())),
subtasks: Arc::new(RwLock::new(HashMap::new())),
}
@@ -118,6 +122,7 @@ 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,
};
+57 -4
View File
@@ -11,6 +11,7 @@ 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)]
@@ -25,6 +26,7 @@ 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 {
@@ -34,9 +36,16 @@ 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
@@ -54,8 +63,28 @@ impl SessionManager {
return Arc::clone(session);
}
let session = Arc::new(Mutex::new(Session::new(user_id)));
let new_session = Session::new(user_id);
let session_id = new_session.id.to_string();
let session = Arc::new(Mutex::new(new_session));
sessions.insert(user_id.to_string(), Arc::clone(&session));
// 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
}
@@ -173,8 +202,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 session user_ids
let stale_users: Vec<String> = {
// Find stale sessions (user_id + session_id)
let stale_sessions: Vec<(String, String)> = {
let sessions = self.sessions.read().await;
sessions
.iter()
@@ -182,7 +211,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())
Some((user_id.clone(), sess.id.to_string()))
} else {
None
}
@@ -190,6 +219,11 @@ 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;
}
@@ -207,6 +241,25 @@ impl SessionManager {
}
}
// Fire OnSessionEnd hooks for stale sessions (fire-and-forget)
if let Some(ref hooks) = self.hooks {
for (user_id, session_id) in &stale_sessions {
let hooks = hooks.clone();
let uid = user_id.clone();
let sid = session_id.clone();
tokio::spawn(async move {
use crate::hooks::HookEvent;
let event = HookEvent::SessionEnd {
user_id: uid,
session_id: sid,
};
if let Err(e) = hooks.run(&event).await {
tracing::warn!("OnSessionEnd hook error: {}", e);
}
});
}
}
// Remove sessions
let count = {
let mut sessions = self.sessions.write().await;
+136 -16
View File
@@ -43,6 +43,10 @@ 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>,
@@ -68,6 +72,14 @@ 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.
@@ -80,24 +92,23 @@ 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.undo_stack.push_back(checkpoint);
// Trim if over limit
while self.undo_stack.len() > self.max_checkpoints {
self.undo_stack.pop_front();
}
self.push_undo(checkpoint);
}
/// Undo: pop the last checkpoint and return it.
///
/// The current state should be saved to redo stack before calling this.
/// Saves the current state to the redo stack and pops the most recent
/// checkpoint from the undo stack so that repeated undos walk backwards
/// through history.
///
/// Takes ownership of `current_messages`; callers must clone first if
/// they need to retain a copy.
pub fn undo(
&mut self,
current_turn: usize,
current_messages: Vec<ChatMessage>,
) -> Option<&Checkpoint> {
) -> Option<Checkpoint> {
if self.undo_stack.is_empty() {
return None;
}
@@ -110,9 +121,8 @@ impl UndoManager {
);
self.redo_stack.push(current);
// Return the most recent checkpoint without removing it
// (we keep it so multiple undos can work)
self.undo_stack.back()
// Pop and return the most recent checkpoint
self.undo_stack.pop_back()
}
/// Pop the last checkpoint from the undo stack.
@@ -121,7 +131,29 @@ impl UndoManager {
}
/// Redo: restore a previously undone state.
pub fn redo(&mut self) -> Option<Checkpoint> {
///
/// Saves the current state to the undo stack and pops the most recent
/// checkpoint from the redo stack.
///
/// Takes ownership of `current_messages`; callers must clone first if
/// they need to retain a copy.
pub fn redo(
&mut self,
current_turn: usize,
current_messages: Vec<ChatMessage>,
) -> Option<Checkpoint> {
if self.redo_stack.is_empty() {
return None;
}
// Save current state to undo stack
let current = Checkpoint::new(
current_turn,
current_messages,
format!("Turn {}", current_turn),
);
self.push_undo(current);
self.redo_stack.pop()
}
@@ -214,14 +246,16 @@ mod tests {
assert!(manager.can_undo());
assert!(!manager.can_redo());
// Undo
// Undo - returns owned Checkpoint now
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
let restored = manager.redo();
// Redo - now requires current state parameters
let restored = manager.redo(checkpoint.turn_number, checkpoint.messages);
assert!(restored.is_some());
}
@@ -249,4 +283,90 @@ 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);
}
}
+61 -42
View File
@@ -12,6 +12,7 @@ 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,
};
@@ -29,6 +30,7 @@ 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,
}
@@ -352,23 +354,11 @@ 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 tools = self.tools().clone();
let context_manager = self.context_manager().clone();
let safety = self.safety().clone();
let deps = self.deps.clone();
let job_id = self.job_id;
let store = self.deps.store.clone();
async move {
let result = Self::execute_tool_inner(
tools,
context_manager,
safety,
store,
job_id,
&tool_name,
&params,
)
.await;
let result = Self::execute_tool_inner(&deps, job_id, &tool_name, &params).await;
ToolExecResult { result }
}
})
@@ -379,20 +369,18 @@ 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(
tools: Arc<ToolRegistry>,
context_manager: Arc<ContextManager>,
safety: Arc<SafetyLayer>,
store: Option<Arc<dyn Database>>,
deps: &WorkerDeps,
job_id: Uuid,
tool_name: &str,
params: &serde_json::Value,
) -> Result<String, Error> {
let tool = tools
.get(tool_name)
.await
.ok_or_else(|| crate::error::ToolError::NotFound {
name: tool_name.to_string(),
})?;
let tool =
deps.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() {
@@ -402,8 +390,46 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
.into());
}
// Get job context for the tool
let job_ctx = context_manager.get_context(job_id).await?;
// 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(),
}
};
if job_ctx.state == JobState::Cancelled {
return Err(crate::error::ToolError::ExecutionFailed {
name: tool_name.to_string(),
@@ -413,7 +439,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
}
// Validate tool parameters
let validation = safety.validator().validate_tool_params(params);
let validation = deps.safety.validator().validate_tool_params(&params);
if !validation.is_valid {
let details = validation
.errors
@@ -478,8 +504,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| safety.sanitize_tool_output(tool_name, &s).content);
context_manager
.map(|s| deps.safety.sanitize_tool_output(tool_name, &s).content);
deps.context_manager
.update_memory(job_id, |mem| {
let rec = mem.create_action(tool_name, params.clone()).succeed(
output_str.clone(),
@@ -492,7 +518,8 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
.await
.ok()
}
Ok(Err(e)) => context_manager
Ok(Err(e)) => deps
.context_manager
.update_memory(job_id, |mem| {
let rec = mem
.create_action(tool_name, params.clone())
@@ -502,7 +529,8 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
})
.await
.ok(),
Err(_) => context_manager
Err(_) => deps
.context_manager
.update_memory(job_id, |mem| {
let rec = mem
.create_action(tool_name, params.clone())
@@ -515,7 +543,7 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
};
// Persist action to database (fire-and-forget)
if let (Some(action), Some(store)) = (action, store) {
if let (Some(action), Some(store)) = (action, deps.store.clone()) {
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);
@@ -701,16 +729,7 @@ 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.tools().clone(),
self.context_manager().clone(),
self.safety().clone(),
self.deps.store.clone(),
self.job_id,
tool_name,
params,
)
.await
Self::execute_tool_inner(&self.deps, self.job_id, tool_name, params).await
}
async fn mark_completed(&self) -> Result<(), Error> {
+208
View File
@@ -0,0 +1,208 @@
//! Boot screen displayed after all initialization completes.
//!
//! Shows a polished ANSI-styled status panel summarizing the agent's runtime
//! state: model, database, tool count, enabled features, active channels,
//! and the gateway URL.
/// All displayable fields for the boot screen.
pub struct BootInfo {
pub version: String,
pub agent_name: String,
pub llm_backend: String,
pub llm_model: String,
pub cheap_model: Option<String>,
pub db_backend: String,
pub db_connected: bool,
pub tool_count: usize,
pub gateway_url: Option<String>,
pub embeddings_enabled: bool,
pub embeddings_provider: Option<String>,
pub heartbeat_enabled: bool,
pub heartbeat_interval_secs: u64,
pub sandbox_enabled: bool,
pub claude_code_enabled: bool,
pub routines_enabled: bool,
pub channels: Vec<String>,
}
/// Print the boot screen to stdout.
pub fn print_boot_screen(info: &BootInfo) {
// ANSI codes matching existing REPL palette
let bold = "\x1b[1m";
let cyan = "\x1b[36m";
let dim = "\x1b[90m";
let yellow_underline = "\x1b[33;4m";
let reset = "\x1b[0m";
let border = format!(" {dim}{}{reset}", "\u{2576}".repeat(58));
println!();
println!("{border}");
println!();
println!(" {bold}{}{reset} v{}", info.agent_name, info.version);
println!();
// Model line
let model_display = if let Some(ref cheap) = info.cheap_model {
format!(
"{cyan}{}{reset} {dim}cheap{reset} {cyan}{}{reset}",
info.llm_model, cheap
)
} else {
format!("{cyan}{}{reset}", info.llm_model)
};
println!(
" {dim}model{reset} {model_display} {dim}via {}{reset}",
info.llm_backend
);
// Database line
let db_status = if info.db_connected {
"connected"
} else {
"none"
};
println!(
" {dim}database{reset} {cyan}{}{reset} {dim}({db_status}){reset}",
info.db_backend
);
// Tools line
println!(
" {dim}tools{reset} {cyan}{}{reset} {dim}registered{reset}",
info.tool_count
);
// Features line
let mut features = Vec::new();
if info.embeddings_enabled {
if let Some(ref provider) = info.embeddings_provider {
features.push(format!("embeddings ({provider})"));
} else {
features.push("embeddings".to_string());
}
}
if info.heartbeat_enabled {
let mins = info.heartbeat_interval_secs / 60;
features.push(format!("heartbeat ({mins}m)"));
}
if info.sandbox_enabled {
features.push("sandbox".to_string());
}
if info.claude_code_enabled {
features.push("claude-code".to_string());
}
if info.routines_enabled {
features.push("routines".to_string());
}
if !features.is_empty() {
println!(
" {dim}features{reset} {cyan}{}{reset}",
features.join(" ")
);
}
// Channels line
if !info.channels.is_empty() {
println!(
" {dim}channels{reset} {cyan}{}{reset}",
info.channels.join(" ")
);
}
// Gateway URL (highlighted)
if let Some(ref url) = info.gateway_url {
println!();
println!(" {dim}gateway{reset} {yellow_underline}{url}{reset}");
}
println!();
println!("{border}");
println!();
println!(" /help for commands, /quit to exit");
println!();
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_print_boot_screen_full() {
let info = BootInfo {
version: "0.2.0".to_string(),
agent_name: "ironclaw".to_string(),
llm_backend: "nearai".to_string(),
llm_model: "claude-3-5-sonnet-20241022".to_string(),
cheap_model: Some("gpt-4o-mini".to_string()),
db_backend: "libsql".to_string(),
db_connected: true,
tool_count: 24,
gateway_url: Some("http://127.0.0.1:3001/?token=abc123".to_string()),
embeddings_enabled: true,
embeddings_provider: Some("openai".to_string()),
heartbeat_enabled: true,
heartbeat_interval_secs: 1800,
sandbox_enabled: true,
claude_code_enabled: false,
routines_enabled: true,
channels: vec![
"repl".to_string(),
"gateway".to_string(),
"telegram".to_string(),
],
};
// Should not panic
print_boot_screen(&info);
}
#[test]
fn test_print_boot_screen_minimal() {
let info = BootInfo {
version: "0.2.0".to_string(),
agent_name: "ironclaw".to_string(),
llm_backend: "nearai".to_string(),
llm_model: "gpt-4o".to_string(),
cheap_model: None,
db_backend: "none".to_string(),
db_connected: false,
tool_count: 5,
gateway_url: None,
embeddings_enabled: false,
embeddings_provider: None,
heartbeat_enabled: false,
heartbeat_interval_secs: 0,
sandbox_enabled: false,
claude_code_enabled: false,
routines_enabled: false,
channels: vec![],
};
// Should not panic
print_boot_screen(&info);
}
#[test]
fn test_print_boot_screen_no_features() {
let info = BootInfo {
version: "0.1.0".to_string(),
agent_name: "test".to_string(),
llm_backend: "openai".to_string(),
llm_model: "gpt-4o".to_string(),
cheap_model: None,
db_backend: "postgres".to_string(),
db_connected: true,
tool_count: 10,
gateway_url: None,
embeddings_enabled: false,
embeddings_provider: None,
heartbeat_enabled: false,
heartbeat_interval_secs: 0,
sandbox_enabled: false,
claude_code_enabled: false,
routines_enabled: false,
channels: vec!["repl".to_string()],
};
// Should not panic
print_boot_screen(&info);
}
}
+81 -5
View File
@@ -81,17 +81,34 @@ fn migrate_bootstrap_json_to_env(env_path: &std::path::Path) {
}
}
/// Write `DATABASE_URL` to `~/.ironclaw/.env`.
/// Write database bootstrap vars to `~/.ironclaw/.env`.
///
/// These settings form the chicken-and-egg layer: they must be available
/// from the filesystem (env vars) BEFORE any database connection, because
/// they determine which database to connect to. Everything else is stored
/// in the database itself.
///
/// Creates the parent directory if it doesn't exist.
/// The value is double-quoted so that `#` (common in URL-encoded passwords)
/// Values are double-quoted so that `#` (common in URL-encoded passwords)
/// and other shell-special characters are preserved by dotenvy.
pub fn save_database_url(url: &str) -> std::io::Result<()> {
pub fn save_bootstrap_env(vars: &[(&str, &str)]) -> std::io::Result<()> {
let path = ironclaw_env_path();
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
std::fs::write(&path, format!("DATABASE_URL=\"{}\"\n", url))
let mut content = String::new();
for (key, value) in vars {
content.push_str(&format!("{}=\"{}\"\n", key, value));
}
std::fs::write(&path, content)
}
/// Write `DATABASE_URL` to `~/.ironclaw/.env`.
///
/// Convenience wrapper around `save_bootstrap_env` for single-value migration
/// paths. Prefer `save_bootstrap_env` for new code.
pub fn save_database_url(url: &str) -> std::io::Result<()> {
save_bootstrap_env(&[("DATABASE_URL", url)])
}
/// One-time migration of legacy `~/.ironclaw/settings.json` into the database.
@@ -184,7 +201,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", &value)
.set_setting(user_id, "nearai.session_token", &value)
.await
.map_err(|e| {
MigrationError::Database(format!(
@@ -385,4 +402,63 @@ mod tests {
// 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"));
}
}
+14 -2
View File
@@ -184,6 +184,8 @@ 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 {
@@ -193,6 +195,7 @@ 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)),
}
}
@@ -202,9 +205,15 @@ 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)
}
@@ -264,6 +273,7 @@ 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
@@ -298,8 +308,10 @@ impl Channel for ReplChannel {
}
let _ = rl.load_history(&hist_path);
println!("\x1b[1mIronClaw\x1b[0m /help for commands, /quit to exit");
println!();
if !suppress_banner.load(Ordering::Relaxed) {
println!("\x1b[1mIronClaw\x1b[0m /help for commands, /quit to exit");
println!();
}
loop {
let prompt = if debug_mode.load(Ordering::Relaxed) {
+15 -6
View File
@@ -12,6 +12,7 @@ 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 ---
@@ -1001,9 +1002,12 @@ function buildBreadcrumb(path) {
}
function searchMemory(query) {
const normalizedQuery = normalizeSearchQuery(query);
if (!normalizedQuery) return;
apiFetch('/api/memory/search', {
method: 'POST',
body: { query, limit: 20 },
body: { query: normalizedQuery, limit: 20 },
}).then((data) => {
const tree = document.getElementById('memory-tree');
tree.innerHTML = '';
@@ -1014,18 +1018,23 @@ function searchMemory(query) {
for (const result of data.results) {
const item = document.createElement('div');
item.className = 'search-result';
const snippet = snippetAround(result.content, query, 120);
const snippet = snippetAround(result.content, normalizedQuery, 120);
item.innerHTML = '<div class="path">' + escapeHtml(result.path) + '</div>'
+ '<div class="snippet">' + highlightQuery(snippet, query) + '</div>';
+ '<div class="snippet">' + highlightQuery(snippet, normalizedQuery) + '</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(query.toLowerCase());
const idx = lower.indexOf(normalizedQuery.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);
@@ -1038,11 +1047,11 @@ function snippetAround(text, query, len) {
function highlightQuery(text, query) {
if (!query) return escapeHtml(text);
const escaped = escapeHtml(text);
const queryEscaped = query.replace(/[.*+?^${}()|[\]\\]/g, '\\$&');
const normalizedQuery = normalizeSearchQuery(query);
const queryEscaped = normalizedQuery.replace(/[.*+?^${}()|[\]\\]/g, '\\$&');
const re = new RegExp('(' + queryEscaped + ')', 'gi');
return escaped.replace(re, '<mark>$1</mark>');
}
// --- Logs ---
const LOG_MAX_ENTRIES = 2000;
+5 -1
View File
@@ -5,7 +5,11 @@
<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/marked/marked.min.js"></script>
<script
src="https://cdn.jsdelivr.net/npm/[email protected]/lib/marked.umd.min.js"
integrity="sha384-pN9zSKOnTZwXRtYZAu0PBPEgR2B7DOC1aeLxQ33oJ0oy5iN1we6gm57xldM2irDG"
crossorigin="anonymous"
></script>
</head>
<body>
<!-- Auth Screen -->
+8 -9
View File
@@ -80,14 +80,13 @@ pub enum OAuthCallbackError {
/// Bind the OAuth callback listener on the fixed port.
///
/// Tries IPv6 loopback (`[::1]`) first so that `http://localhost:…` redirects
/// work on systems where `localhost` resolves to `::1`. Falls back to IPv4
/// (`127.0.0.1`) only if IPv6 fails for a reason other than `AddrInUse`
/// (e.g., IPv6 not supported on the host). If the port is already occupied
/// on IPv6, the port is occupied period, so we fail immediately.
/// 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 ipv6_addr = format!("[::1]:{}", OAUTH_CALLBACK_PORT);
match TcpListener::bind(&ipv6_addr).await {
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(
@@ -96,10 +95,10 @@ pub async fn bind_callback_listener() -> Result<TcpListener, OAuthCallbackError>
));
}
Err(_) => {
// IPv6 not available on this host, fall back to IPv4
// IPv4 not available, fall back to IPv6
}
}
TcpListener::bind(format!("127.0.0.1:{}", OAUTH_CALLBACK_PORT))
TcpListener::bind(format!("[::1]:{}", OAUTH_CALLBACK_PORT))
.await
.map_err(|e| {
if e.kind() == std::io::ErrorKind::AddrInUse {
+36 -14
View File
@@ -22,15 +22,36 @@ pub async fn run_status_command() -> anyhow::Result<()> {
);
// Database
let db_url_set = std::env::var("DATABASE_URL").is_ok();
print!(" Database: ");
if db_url_set {
match check_database().await {
Ok(()) => println!("connected"),
Err(e) => println!("error ({})", e),
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");
}
}
} else {
println!("not configured");
}
// Session / Auth
@@ -42,16 +63,17 @@ pub async fn run_status_command() -> anyhow::Result<()> {
println!("not found (run `ironclaw onboard`)");
}
// Secrets (auto-detect: env var or keychain)
// Secrets (auto-detect from env only; skip keychain probe to avoid
// triggering macOS system password dialogs on a simple status check)
print!(" Secrets: ");
let has_env_key = std::env::var("SECRETS_MASTER_KEY").is_ok();
let has_keychain = crate::secrets::keychain::has_master_key().await;
if has_env_key {
if std::env::var("SECRETS_MASTER_KEY").is_ok() {
println!("configured (env)");
} else if has_keychain {
println!("configured (keychain)");
} else {
println!("not configured");
// We don't probe the keychain here because get_generic_password()
// triggers macOS unlock+authorization dialogs, which is bad UX for
// a read-only status command. If onboarding completed with keychain
// storage, the key is there; we just can't cheaply verify it.
println!("env not set (keychain may be configured)");
}
// Embeddings
+106 -9
View File
@@ -5,7 +5,9 @@
//! in startup). Everything else comes from env vars, the DB settings
//! table, or auto-detection.
use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::OnceLock;
use std::time::Duration;
use secrecy::{ExposeSecret, SecretString};
@@ -13,6 +15,13 @@ 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 {
@@ -144,6 +153,15 @@ 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;
@@ -379,6 +397,9 @@ 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)
@@ -398,16 +419,35 @@ pub struct NearAiConfig {
/// 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 (default: NearAi)
// Determine backend: env var > settings > 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
};
@@ -433,6 +473,7 @@ 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")?
@@ -444,6 +485,8 @@ impl LlmConfig {
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
@@ -476,6 +519,7 @@ 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 })
@@ -484,8 +528,9 @@ impl LlmConfig {
};
let openai_compatible = if backend == LlmBackend::OpenAiCompatible {
let base_url =
optional_env("LLM_BASE_URL")?.ok_or_else(|| ConfigError::MissingRequired {
let base_url = optional_env("LLM_BASE_URL")?
.or_else(|| settings.openai_compatible_base_url.clone())
.ok_or_else(|| ConfigError::MissingRequired {
key: "LLM_BASE_URL".to_string(),
hint: "Set LLM_BASE_URL when LLM_BACKEND=openai_compatible".to_string(),
})?;
@@ -855,6 +900,11 @@ impl std::fmt::Debug for SecretsConfig {
}
}
/// Process-wide cache for the keychain master key.
///
/// Avoids re-prompting the OS keychain on every `SecretsConfig::resolve()` call
/// (e.g. `Config::from_env()` then `Config::from_db()`). Thread-safe alternative
/// to caching in a process env var.
impl SecretsConfig {
/// Auto-detect secrets master key from env var, then OS keychain.
///
@@ -1338,17 +1388,64 @@ 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(None),
Ok(val) => Ok(Some(val)),
Err(std::env::VarError::NotPresent) => Ok(None),
Err(e) => Err(ConfigError::ParseError(format!(
"failed to read {key}: {e}"
))),
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}"
)));
}
}
// 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>
+3
View File
@@ -40,6 +40,9 @@ 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),
+199
View File
@@ -0,0 +1,199 @@
//! Core hook types and traits.
use std::time::Duration;
use async_trait::async_trait;
/// Points in the agent lifecycle where hooks can be attached.
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum HookPoint {
/// Before processing an inbound user message.
BeforeInbound,
/// Before executing a tool call.
BeforeToolCall,
/// Before sending an outbound response.
BeforeOutbound,
/// When a new session starts.
OnSessionStart,
/// When a session ends (pruned or expired).
OnSessionEnd,
/// Transform the final response before completing a turn.
TransformResponse,
}
/// Contextual data carried with each hook invocation.
#[derive(Debug, Clone)]
pub enum HookEvent {
/// An inbound user message about to be processed.
Inbound {
user_id: String,
channel: String,
content: String,
thread_id: Option<String>,
},
/// A tool call about to be executed.
ToolCall {
tool_name: String,
parameters: serde_json::Value,
user_id: String,
/// "chat" for interactive, or a job ID string for autonomous jobs.
context: String,
},
/// An outbound response about to be sent.
Outbound {
user_id: String,
channel: String,
content: String,
thread_id: Option<String>,
},
/// A new session was created.
SessionStart { user_id: String, session_id: String },
/// A session was ended (pruned).
SessionEnd { user_id: String, session_id: String },
/// The final response is being transformed before completing a turn.
ResponseTransform {
user_id: String,
thread_id: String,
response: String,
},
}
impl HookEvent {
/// Returns the [`HookPoint`] this event corresponds to.
pub fn hook_point(&self) -> HookPoint {
match self {
HookEvent::Inbound { .. } => HookPoint::BeforeInbound,
HookEvent::ToolCall { .. } => HookPoint::BeforeToolCall,
HookEvent::Outbound { .. } => HookPoint::BeforeOutbound,
HookEvent::SessionStart { .. } => HookPoint::OnSessionStart,
HookEvent::SessionEnd { .. } => HookPoint::OnSessionEnd,
HookEvent::ResponseTransform { .. } => HookPoint::TransformResponse,
}
}
/// Apply a modification string to the event's primary content field.
pub fn apply_modification(&mut self, modified: &str) {
match self {
HookEvent::Inbound { content, .. } | HookEvent::Outbound { content, .. } => {
*content = modified.to_string();
}
HookEvent::ToolCall { parameters, .. } => match serde_json::from_str(modified) {
Ok(parsed) => *parameters = parsed,
Err(e) => {
tracing::warn!(
"Hook returned non-JSON modification for ToolCall, ignoring: {}",
e
);
}
},
HookEvent::ResponseTransform { response, .. } => {
*response = modified.to_string();
}
HookEvent::SessionStart { .. } | HookEvent::SessionEnd { .. } => {
// Session events don't have modifiable content
}
}
}
}
/// The result of executing a hook.
#[derive(Debug, Clone)]
pub enum HookOutcome {
/// Continue processing, optionally with modified content.
Continue {
/// If `Some`, replace the event's primary content with this value.
modified: Option<String>,
},
/// Reject the event entirely.
Reject {
/// Human-readable reason for the rejection.
reason: String,
},
}
impl HookOutcome {
/// Shorthand for `Continue { modified: None }`.
pub fn ok() -> Self {
HookOutcome::Continue { modified: None }
}
/// Shorthand for `Continue { modified: Some(value) }`.
pub fn modify(value: String) -> Self {
HookOutcome::Continue {
modified: Some(value),
}
}
/// Shorthand for `Reject { reason }`.
pub fn reject(reason: impl Into<String>) -> Self {
HookOutcome::Reject {
reason: reason.into(),
}
}
}
/// How to handle hook execution failures.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum HookFailureMode {
/// On error/timeout, continue processing as if the hook returned `ok()`.
FailOpen,
/// On error/timeout, reject the event.
FailClosed,
}
/// Hook execution errors.
#[derive(Debug, thiserror::Error)]
pub enum HookError {
#[error("Hook execution failed: {reason}")]
ExecutionFailed { reason: String },
#[error("Hook timed out after {timeout:?}")]
Timeout { timeout: Duration },
#[error("Hook rejected: {reason}")]
Rejected { reason: String },
}
/// Context passed to hooks alongside the event.
pub struct HookContext {
/// Arbitrary metadata hooks can use.
pub metadata: serde_json::Value,
}
impl Default for HookContext {
fn default() -> Self {
Self {
metadata: serde_json::Value::Null,
}
}
}
/// Trait for implementing lifecycle hooks.
///
/// Hooks intercept and can modify agent operations at well-defined points.
#[async_trait]
pub trait Hook: Send + Sync {
/// A unique name for this hook.
fn name(&self) -> &str;
/// The lifecycle points this hook should be called at.
fn hook_points(&self) -> &[HookPoint];
/// How to handle failures in this hook.
///
/// Default: `FailOpen` (continue on error).
fn failure_mode(&self) -> HookFailureMode {
HookFailureMode::FailOpen
}
/// Maximum time this hook is allowed to run.
///
/// Default: 5 seconds.
fn timeout(&self) -> Duration {
Duration::from_secs(5)
}
/// Execute the hook.
async fn execute(&self, event: &HookEvent, ctx: &HookContext)
-> Result<HookOutcome, HookError>;
}
+19
View File
@@ -0,0 +1,19 @@
//! Lifecycle hooks for intercepting and transforming agent operations.
//!
//! The hook system provides 6 well-defined interception points:
//!
//! - **BeforeInbound** — Before processing an inbound user message
//! - **BeforeToolCall** — Before executing a tool call
//! - **BeforeOutbound** — Before sending an outbound response
//! - **OnSessionStart** — When a new session starts
//! - **OnSessionEnd** — When a session ends
//! - **TransformResponse** — Transform the final response before completing a turn
//!
//! Hooks are executed in priority order (lower number = higher priority).
//! Each hook can pass through, modify content, or reject the event.
pub mod hook;
pub mod registry;
pub use hook::{Hook, HookContext, HookError, HookEvent, HookFailureMode, HookOutcome, HookPoint};
pub use registry::HookRegistry;
+555
View File
@@ -0,0 +1,555 @@
//! Hook registry for managing and executing lifecycle hooks.
use std::sync::Arc;
use tokio::sync::RwLock;
use crate::hooks::hook::{Hook, HookContext, HookError, HookEvent, HookFailureMode, HookOutcome};
/// A registered hook with its priority.
struct HookEntry {
hook: Arc<dyn Hook>,
priority: u32,
}
/// Registry that manages hooks and executes them at lifecycle points.
///
/// Hooks are executed in priority order (lower number = higher priority).
/// A `Reject` outcome stops the chain immediately.
/// A `Modify` outcome chains through subsequent hooks.
pub struct HookRegistry {
hooks: RwLock<Vec<HookEntry>>,
}
impl HookRegistry {
/// Create an empty registry.
pub fn new() -> Self {
Self {
hooks: RwLock::new(Vec::new()),
}
}
/// Register a hook with default priority (100).
pub async fn register(&self, hook: Arc<dyn Hook>) {
self.register_with_priority(hook, 100).await;
}
/// Register a hook with a specific priority.
///
/// Lower priority number = runs first.
pub async fn register_with_priority(&self, hook: Arc<dyn Hook>, priority: u32) {
let mut hooks = self.hooks.write().await;
hooks.push(HookEntry { hook, priority });
hooks.sort_by_key(|e| e.priority);
}
/// Unregister a hook by name. Returns `true` if it was found and removed.
pub async fn unregister(&self, name: &str) -> bool {
let mut hooks = self.hooks.write().await;
let before = hooks.len();
hooks.retain(|e| e.hook.name() != name);
hooks.len() < before
}
/// List all registered hook names (in priority order).
pub async fn list(&self) -> Vec<String> {
let hooks = self.hooks.read().await;
hooks.iter().map(|e| e.hook.name().to_string()).collect()
}
/// Run all hooks matching the event's hook point.
///
/// - Hooks run in priority order (lowest first).
/// - `Reject` stops the chain immediately.
/// - `Modify` chains the modification through subsequent hooks.
/// - Timeout/error handling respects each hook's `failure_mode`.
pub async fn run(&self, event: &HookEvent) -> Result<HookOutcome, HookError> {
let point = event.hook_point();
let ctx = HookContext::default();
// Clone matching hooks and drop the read guard before executing.
// Each hook can run up to its timeout, so holding the guard would
// block concurrent register/unregister/run calls.
let matching: Vec<Arc<dyn Hook>> = {
let hooks = self.hooks.read().await;
hooks
.iter()
.filter(|e| e.hook.hook_points().contains(&point))
.map(|e| e.hook.clone())
.collect()
};
if matching.is_empty() {
return Ok(HookOutcome::ok());
}
let mut current_event = event.clone();
for hook in &matching {
let timeout = hook.timeout();
let result = tokio::time::timeout(timeout, hook.execute(&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());
}
}
+2
View File
@@ -39,6 +39,7 @@
//! - **Continuous learning** - Improve estimates from historical data
pub mod agent;
pub mod boot_screen;
pub mod bootstrap;
pub mod channels;
pub mod cli;
@@ -50,6 +51,7 @@ pub mod estimation;
pub mod evaluation;
pub mod extensions;
pub mod history;
pub mod hooks;
pub mod llm;
pub mod orchestrator;
pub mod pairing;
+542 -8
View File
@@ -2,10 +2,15 @@
//!
//! Wraps multiple LlmProvider instances and tries each in sequence
//! until one succeeds. Transparent to callers --- same LlmProvider trait.
//!
//! Providers that fail repeatedly are temporarily placed in cooldown
//! so subsequent requests skip them, reducing latency when a provider
//! is known to be down. Cooldown state is lock-free (atomics only).
use std::future::Future;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::atomic::{AtomicU32, AtomicU64, AtomicUsize, Ordering};
use std::time::{Duration, Instant};
use async_trait::async_trait;
use rust_decimal::Decimal;
@@ -41,61 +46,217 @@ fn is_retryable(err: &LlmError) -> bool {
)
}
/// Configuration for per-provider cooldown behavior.
///
/// When a provider accumulates `failure_threshold` consecutive retryable
/// failures, it enters cooldown for `cooldown_duration`. During cooldown
/// the provider is skipped (unless *all* providers are in cooldown, in
/// which case the oldest-cooled one is tried).
#[derive(Debug, Clone)]
pub struct CooldownConfig {
/// How long a provider stays in cooldown after exceeding the threshold.
pub cooldown_duration: Duration,
/// Number of consecutive retryable failures before cooldown activates.
pub failure_threshold: u32,
}
impl Default for CooldownConfig {
fn default() -> Self {
Self {
cooldown_duration: Duration::from_secs(300),
failure_threshold: 3,
}
}
}
/// Per-provider cooldown state, entirely lock-free.
///
/// All atomic operations use `Relaxed` ordering — consistent with the
/// existing `last_used` field. Stale reads are harmless: the worst case
/// is one extra attempt against a provider that just entered cooldown.
struct ProviderCooldown {
/// Consecutive retryable failures. Reset to 0 on success.
failure_count: AtomicU32,
/// Nanoseconds since `epoch` when cooldown was activated.
/// 0 means the provider is NOT in cooldown.
cooldown_activated_nanos: AtomicU64,
}
impl ProviderCooldown {
fn new() -> Self {
Self {
failure_count: AtomicU32::new(0),
cooldown_activated_nanos: AtomicU64::new(0),
}
}
/// Check whether the provider is currently in cooldown.
fn is_in_cooldown(&self, now_nanos: u64, cooldown_nanos: u64) -> bool {
let activated = self.cooldown_activated_nanos.load(Ordering::Relaxed);
activated != 0 && now_nanos.saturating_sub(activated) < cooldown_nanos
}
/// Record a retryable failure. Returns `true` if the threshold was
/// just reached (caller should activate cooldown).
fn record_failure(&self, threshold: u32) -> bool {
let prev = self.failure_count.fetch_add(1, Ordering::Relaxed);
prev + 1 >= threshold
}
/// Activate cooldown at the given timestamp.
fn activate_cooldown(&self, now_nanos: u64) {
self.cooldown_activated_nanos
.store(now_nanos, Ordering::Relaxed);
}
/// Reset failure count and clear cooldown (called on success).
fn reset(&self) {
self.failure_count.store(0, Ordering::Relaxed);
self.cooldown_activated_nanos.store(0, Ordering::Relaxed);
}
}
/// An LLM provider that wraps multiple providers and tries each in sequence
/// on transient failures.
///
/// The first provider in the list is the primary. If it fails with a retryable
/// error, the next provider is tried, and so on. Non-retryable errors
/// (e.g. `AuthFailed`, `ContextLengthExceeded`) propagate immediately.
///
/// Providers that repeatedly fail with retryable errors are temporarily
/// placed in cooldown and skipped, reducing latency.
pub struct FailoverProvider {
providers: Vec<Arc<dyn LlmProvider>>,
/// Index of the provider that last handled a request successfully.
/// Used by `model_name()` and `cost_per_token()` so downstream cost
/// tracking reflects the provider that actually served the request.
last_used: AtomicUsize,
/// Per-provider cooldown tracking (same length as `providers`).
cooldowns: Vec<ProviderCooldown>,
/// Reference instant for computing elapsed nanos. Shared across all
/// cooldown timestamps so they are comparable.
epoch: Instant,
/// Cooldown configuration.
cooldown_config: CooldownConfig,
}
impl FailoverProvider {
/// Create a new failover provider.
/// Create a new failover provider with default cooldown settings.
///
/// Returns an error if `providers` is empty.
pub fn new(providers: Vec<Arc<dyn LlmProvider>>) -> Result<Self, LlmError> {
Self::with_cooldown(providers, CooldownConfig::default())
}
/// Create a new failover provider with explicit cooldown configuration.
///
/// Returns an error if `providers` is empty.
pub fn with_cooldown(
providers: Vec<Arc<dyn LlmProvider>>,
cooldown_config: CooldownConfig,
) -> Result<Self, LlmError> {
if providers.is_empty() {
return Err(LlmError::RequestFailed {
provider: "failover".to_string(),
reason: "FailoverProvider requires at least one provider".to_string(),
});
}
let cooldowns = (0..providers.len())
.map(|_| ProviderCooldown::new())
.collect();
Ok(Self {
providers,
last_used: AtomicUsize::new(0),
cooldowns,
epoch: Instant::now(),
cooldown_config,
})
}
/// Nanoseconds elapsed since `self.epoch`.
///
/// Truncates `u128` → `u64` (wraps after ~584 years of continuous
/// uptime). Acceptable because `epoch` is set at construction time.
fn now_nanos(&self) -> u64 {
self.epoch.elapsed().as_nanos() as u64
}
/// Try each provider in sequence until one succeeds or all fail.
///
/// Providers in cooldown are skipped unless *all* providers are in
/// cooldown, in which case the one with the oldest cooldown timestamp
/// (most likely to have recovered) is tried.
async fn try_providers<T, F, Fut>(&self, mut call: F) -> Result<T, LlmError>
where
F: FnMut(Arc<dyn LlmProvider>) -> Fut,
Fut: Future<Output = Result<T, LlmError>>,
{
let now_nanos = self.now_nanos();
let cooldown_nanos = self.cooldown_config.cooldown_duration.as_nanos() as u64;
// Partition providers into available and cooled-down.
let (mut available, cooled_down): (Vec<usize>, Vec<usize>) = (0..self.providers.len())
.partition(|&i| !self.cooldowns[i].is_in_cooldown(now_nanos, cooldown_nanos));
// Log skipped providers.
for &i in &cooled_down {
tracing::info!(
provider = %self.providers[i].model_name(),
"Skipping provider (in cooldown)"
);
}
// Never skip ALL providers: if every provider is in cooldown, pick
// the one with the oldest cooldown activation (most likely recovered).
if available.is_empty() {
let oldest = (0..self.providers.len())
.min_by_key(|&i| {
self.cooldowns[i]
.cooldown_activated_nanos
.load(Ordering::Relaxed)
})
.expect("providers list is non-empty");
tracing::info!(
provider = %self.providers[oldest].model_name(),
"All providers in cooldown, trying oldest-cooled provider"
);
available.push(oldest);
}
let mut last_error: Option<LlmError> = None;
for (i, provider) in self.providers.iter().enumerate() {
for (pos, &i) in available.iter().enumerate() {
let provider = &self.providers[i];
let result = call(Arc::clone(provider)).await;
match result {
Ok(response) => {
self.last_used.store(i, Ordering::Relaxed);
self.cooldowns[i].reset();
return Ok(response);
}
Err(err) => {
if !is_retryable(&err) {
return Err(err);
}
if i + 1 < self.providers.len() {
// Increment failure count; activate cooldown if threshold reached.
if self.cooldowns[i].record_failure(self.cooldown_config.failure_threshold) {
let nanos = self.now_nanos();
self.cooldowns[i].activate_cooldown(nanos);
tracing::warn!(
provider = %provider.model_name(),
threshold = self.cooldown_config.failure_threshold,
cooldown_secs = self.cooldown_config.cooldown_duration.as_secs(),
"Provider entered cooldown after repeated failures"
);
}
if pos + 1 < available.len() {
let next_i = available[pos + 1];
tracing::warn!(
provider = %provider.model_name(),
error = %err,
next_provider = %self.providers[i + 1].model_name(),
next_provider = %self.providers[next_i].model_name(),
"Provider failed with retryable error, trying next provider"
);
}
@@ -104,9 +265,9 @@ impl FailoverProvider {
}
}
// SAFETY: providers is non-empty (checked in `new`), so at least one
// SAFETY: `available` is non-empty (guaranteed above), so at least one
// iteration ran and `last_error` is `Some`.
Err(last_error.expect("providers list is non-empty"))
Err(last_error.expect("available providers list is non-empty"))
}
}
@@ -166,7 +327,6 @@ mod tests {
use super::*;
use std::sync::Mutex;
use std::time::Duration;
use crate::llm::provider::{CompletionResponse, FinishReason, ToolCompletionResponse};
@@ -432,6 +592,380 @@ mod tests {
assert!(models.contains(&"model-b".to_string()));
}
// --- MultiCallMockProvider for cooldown tests ---
//
// Unlike `MockProvider` which uses `.take()` (single-use), this mock
// tracks a call counter and returns errors for the first N calls,
// then succeeds.
struct MultiCallMockProvider {
name: String,
/// How many calls should fail before succeeding. 0 = always succeed.
fail_count: u32,
/// Atomically tracks how many times `complete` has been called.
calls: AtomicU32,
/// If true, failures are non-retryable (AuthFailed).
non_retryable: bool,
}
impl MultiCallMockProvider {
/// Always succeeds.
fn always_ok(name: &str) -> Self {
Self {
name: name.to_string(),
fail_count: 0,
calls: AtomicU32::new(0),
non_retryable: false,
}
}
/// Fails with retryable error for the first `n` calls, then succeeds.
fn fail_then_ok(name: &str, n: u32) -> Self {
Self {
name: name.to_string(),
fail_count: n,
calls: AtomicU32::new(0),
non_retryable: false,
}
}
/// Always fails with retryable error.
fn always_fail(name: &str) -> Self {
Self {
name: name.to_string(),
fail_count: u32::MAX,
calls: AtomicU32::new(0),
non_retryable: false,
}
}
/// Always fails with non-retryable error.
fn always_fail_non_retryable(name: &str) -> Self {
Self {
name: name.to_string(),
fail_count: u32::MAX,
calls: AtomicU32::new(0),
non_retryable: true,
}
}
fn call_count(&self) -> u32 {
self.calls.load(Ordering::Relaxed)
}
}
#[async_trait]
impl LlmProvider for MultiCallMockProvider {
fn model_name(&self) -> &str {
&self.name
}
fn cost_per_token(&self) -> (Decimal, Decimal) {
(Decimal::ZERO, Decimal::ZERO)
}
async fn complete(
&self,
_request: CompletionRequest,
) -> Result<CompletionResponse, LlmError> {
let n = self.calls.fetch_add(1, Ordering::Relaxed);
if n < self.fail_count {
if self.non_retryable {
return Err(LlmError::AuthFailed {
provider: self.name.clone(),
});
}
return Err(LlmError::RequestFailed {
provider: self.name.clone(),
reason: format!("call {} failed", n),
});
}
Ok(CompletionResponse {
content: format!("{} ok", self.name),
input_tokens: 10,
output_tokens: 5,
finish_reason: FinishReason::Stop,
response_id: None,
})
}
async fn complete_with_tools(
&self,
_request: ToolCompletionRequest,
) -> Result<ToolCompletionResponse, LlmError> {
let n = self.calls.fetch_add(1, Ordering::Relaxed);
if n < self.fail_count {
if self.non_retryable {
return Err(LlmError::AuthFailed {
provider: self.name.clone(),
});
}
return Err(LlmError::RequestFailed {
provider: self.name.clone(),
reason: format!("call {} failed", n),
});
}
Ok(ToolCompletionResponse {
content: Some(format!("{} ok", self.name)),
tool_calls: vec![],
input_tokens: 10,
output_tokens: 5,
finish_reason: FinishReason::Stop,
response_id: None,
})
}
async fn list_models(&self) -> Result<Vec<String>, LlmError> {
Ok(vec![self.name.clone()])
}
}
// --- Cooldown tests ---
// Cooldown test 1: Provider enters cooldown after `threshold` consecutive failures.
#[tokio::test]
async fn cooldown_activates_after_threshold() {
let config = CooldownConfig {
cooldown_duration: Duration::from_secs(300),
failure_threshold: 2,
};
let p1 = Arc::new(MultiCallMockProvider::always_fail("p1"));
let p2 = Arc::new(MultiCallMockProvider::always_ok("p2"));
let failover =
FailoverProvider::with_cooldown(vec![p1.clone(), p2.clone()], config).unwrap();
// Request 1: p1 fails (count=1, below threshold), p2 succeeds.
let r = failover.complete(make_request()).await.unwrap();
assert_eq!(r.content, "p2 ok");
assert_eq!(p1.call_count(), 1);
// Request 2: p1 fails again (count=2, reaches threshold → cooldown), p2 succeeds.
let r = failover.complete(make_request()).await.unwrap();
assert_eq!(r.content, "p2 ok");
assert_eq!(p1.call_count(), 2);
// Request 3: p1 should be skipped (in cooldown), only p2 called.
let prev_p1_calls = p1.call_count();
let r = failover.complete(make_request()).await.unwrap();
assert_eq!(r.content, "p2 ok");
// p1 was NOT called again.
assert_eq!(p1.call_count(), prev_p1_calls);
}
// Cooldown test 2: Cooldown expires after duration, provider is retried.
#[tokio::test]
async fn cooldown_expires_after_duration() {
let config = CooldownConfig {
cooldown_duration: Duration::from_millis(1),
failure_threshold: 1,
};
// p1 fails once then succeeds (fail_then_ok with n=1 would work,
// but we use always_fail to prove it's skipped, then swap).
let p1 = Arc::new(MultiCallMockProvider::fail_then_ok("p1", 2));
let p2 = Arc::new(MultiCallMockProvider::always_ok("p2"));
let failover =
FailoverProvider::with_cooldown(vec![p1.clone(), p2.clone()], config).unwrap();
// Request 1: p1 fails (threshold=1, enters cooldown immediately), p2 succeeds.
let r = failover.complete(make_request()).await.unwrap();
assert_eq!(r.content, "p2 ok");
assert_eq!(p1.call_count(), 1);
// Request 2: p1 in cooldown, skipped. Only p2 called.
// (But cooldown is 1ms, so wait a bit to let it expire.)
tokio::time::sleep(Duration::from_millis(5)).await;
// After sleep, cooldown should have expired. p1 gets tried again.
// p1 is set to fail 2 times total, so call #2 (index 1) still fails.
// But it proves p1 was attempted again after cooldown expired.
let r = failover.complete(make_request()).await.unwrap();
assert_eq!(p1.call_count(), 2); // p1 was retried
assert_eq!(r.content, "p2 ok"); // p2 handled it
// Wait again for cooldown to expire, p1 call #3 (index 2) succeeds.
tokio::time::sleep(Duration::from_millis(5)).await;
let r = failover.complete(make_request()).await.unwrap();
assert_eq!(r.content, "p1 ok");
assert_eq!(p1.call_count(), 3);
}
// Cooldown test 3: Never skip all providers — oldest-cooled one is tried.
#[tokio::test]
async fn never_skip_all_providers() {
let config = CooldownConfig {
cooldown_duration: Duration::from_secs(300),
failure_threshold: 1,
};
// Both providers always fail.
let p1 = Arc::new(MultiCallMockProvider::always_fail("p1"));
let p2 = Arc::new(MultiCallMockProvider::always_fail("p2"));
let failover =
FailoverProvider::with_cooldown(vec![p1.clone(), p2.clone()], config).unwrap();
// Request 1: both tried, both fail, both enter cooldown.
let _ = failover.complete(make_request()).await;
assert_eq!(p1.call_count(), 1);
assert_eq!(p2.call_count(), 1);
// Request 2: all in cooldown, but the oldest-cooled one (p1, activated
// first) should be tried.
let prev_total = p1.call_count() + p2.call_count();
let _ = failover.complete(make_request()).await;
let new_total = p1.call_count() + p2.call_count();
// Exactly one more call was made (to the oldest-cooled provider).
assert_eq!(new_total, prev_total + 1);
}
// Cooldown test 4: Success resets failure count so it never reaches threshold.
//
// With threshold=3, accumulate 2 failures then succeed. Verify the
// atomic counter is back to 0 and no cooldown was activated. Then
// use a second provider pair to show that without the reset, 3
// consecutive failures DO trigger cooldown (control case).
#[tokio::test]
async fn reset_on_success() {
let config = CooldownConfig {
cooldown_duration: Duration::from_secs(300),
failure_threshold: 3,
};
// p1 fails for calls 0,1 then succeeds on call 2+.
let p1 = Arc::new(MultiCallMockProvider::fail_then_ok("p1", 2));
let p2 = Arc::new(MultiCallMockProvider::always_ok("p2"));
let failover =
FailoverProvider::with_cooldown(vec![p1.clone(), p2.clone()], config.clone()).unwrap();
// Request 1: p1 fails (failure_count=1), p2 succeeds.
let r = failover.complete(make_request()).await.unwrap();
assert_eq!(r.content, "p2 ok");
// Request 2: p1 fails (failure_count=2, still below threshold=3), p2 succeeds.
let r = failover.complete(make_request()).await.unwrap();
assert_eq!(r.content, "p2 ok");
assert_eq!(p1.call_count(), 2);
// Request 3: p1 succeeds (call index 2) → counter resets to 0.
let r = failover.complete(make_request()).await.unwrap();
assert_eq!(r.content, "p1 ok");
assert_eq!(p1.call_count(), 3);
// Verify counter was reset to 0 and no cooldown activated.
let nanos = failover.now_nanos();
let cooldown_nanos = failover.cooldown_config.cooldown_duration.as_nanos() as u64;
assert!(!failover.cooldowns[0].is_in_cooldown(nanos, cooldown_nanos));
assert_eq!(
failover.cooldowns[0].failure_count.load(Ordering::Relaxed),
0
);
// Control: without a success in the middle, 3 failures DO trigger cooldown.
let p3 = Arc::new(MultiCallMockProvider::always_fail("p3"));
let p4 = Arc::new(MultiCallMockProvider::always_ok("p4"));
let control =
FailoverProvider::with_cooldown(vec![p3.clone(), p4.clone()], config).unwrap();
for _ in 0..3 {
let _ = control.complete(make_request()).await.unwrap();
}
let nanos = control.now_nanos();
assert!(control.cooldowns[0].is_in_cooldown(nanos, cooldown_nanos));
}
// Cooldown test 5: threshold-1 failures don't trigger cooldown, threshold does.
#[tokio::test]
async fn threshold_boundary() {
let config = CooldownConfig {
cooldown_duration: Duration::from_secs(300),
failure_threshold: 3,
};
let p1 = Arc::new(MultiCallMockProvider::always_fail("p1"));
let p2 = Arc::new(MultiCallMockProvider::always_ok("p2"));
let failover =
FailoverProvider::with_cooldown(vec![p1.clone(), p2.clone()], config).unwrap();
// 2 requests: p1 fails twice (below threshold of 3), not in cooldown.
for _ in 0..2 {
let r = failover.complete(make_request()).await.unwrap();
assert_eq!(r.content, "p2 ok");
}
assert_eq!(p1.call_count(), 2);
// p1 should still be available (not in cooldown).
let nanos = failover.now_nanos();
let cooldown_nanos = failover.cooldown_config.cooldown_duration.as_nanos() as u64;
assert!(!failover.cooldowns[0].is_in_cooldown(nanos, cooldown_nanos));
// 3rd request: p1 fails → reaches threshold → enters cooldown.
let r = failover.complete(make_request()).await.unwrap();
assert_eq!(r.content, "p2 ok");
assert_eq!(p1.call_count(), 3);
let nanos = failover.now_nanos();
assert!(failover.cooldowns[0].is_in_cooldown(nanos, cooldown_nanos));
// 4th request: p1 should be skipped.
let prev = p1.call_count();
let r = failover.complete(make_request()).await.unwrap();
assert_eq!(r.content, "p2 ok");
assert_eq!(p1.call_count(), prev); // not called
}
// Cooldown test 6: Non-retryable error returns immediately, no failure bump.
#[tokio::test]
async fn non_retryable_does_not_increment_cooldown() {
let config = CooldownConfig {
cooldown_duration: Duration::from_secs(300),
failure_threshold: 1,
};
let p1 = Arc::new(MultiCallMockProvider::always_fail_non_retryable("p1"));
let p2 = Arc::new(MultiCallMockProvider::always_ok("p2"));
let failover =
FailoverProvider::with_cooldown(vec![p1.clone(), p2.clone()], config).unwrap();
// Non-retryable error should return immediately.
let err = failover.complete(make_request()).await.unwrap_err();
assert!(matches!(err, LlmError::AuthFailed { .. }));
assert_eq!(p1.call_count(), 1);
// p2 should NOT have been called (non-retryable = no failover).
assert_eq!(p2.call_count(), 0);
// p1 should NOT be in cooldown (non-retryable doesn't bump count).
let nanos = failover.now_nanos();
let cooldown_nanos = failover.cooldown_config.cooldown_duration.as_nanos() as u64;
assert!(!failover.cooldowns[0].is_in_cooldown(nanos, cooldown_nanos));
}
// Cooldown test 7: Three providers, first in cooldown, second/third available.
#[tokio::test]
async fn three_providers_mixed_cooldown() {
let config = CooldownConfig {
cooldown_duration: Duration::from_secs(300),
failure_threshold: 1,
};
let p1 = Arc::new(MultiCallMockProvider::always_fail("p1"));
let p2 = Arc::new(MultiCallMockProvider::always_ok("p2"));
let p3 = Arc::new(MultiCallMockProvider::always_ok("p3"));
let failover =
FailoverProvider::with_cooldown(vec![p1.clone(), p2.clone(), p3.clone()], config)
.unwrap();
// Request 1: p1 fails → enters cooldown (threshold=1), p2 succeeds.
let r = failover.complete(make_request()).await.unwrap();
assert_eq!(r.content, "p2 ok");
assert_eq!(p1.call_count(), 1);
// Request 2: p1 skipped (cooldown), p2 and p3 available.
let prev = p1.call_count();
let r = failover.complete(make_request()).await.unwrap();
assert_eq!(r.content, "p2 ok");
assert_eq!(p1.call_count(), prev); // p1 skipped
}
// Test: is_retryable correctly classifies errors.
#[test]
fn retryable_classification() {
+106 -1
View File
@@ -17,7 +17,7 @@ mod retry;
mod rig_adapter;
pub mod session;
pub use failover::FailoverProvider;
pub use failover::{CooldownConfig, FailoverProvider};
pub use nearai::{ModelInfo, NearAiProvider};
pub use nearai_chat::NearAiChatProvider;
pub use provider::{
@@ -183,3 +183,108 @@ fn create_openai_compatible_provider(config: &LlmConfig) -> Result<Arc<dyn LlmPr
);
Ok(Arc::new(RigAdapter::new(model, &compat.model)))
}
/// 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());
}
}
+19 -6
View File
@@ -428,17 +428,30 @@ impl SessionManager {
})?;
let user_id = self.user_id.read().await.clone();
let value = store
let value = if let Some(value) = store
.get_setting(&user_id, "nearai.session_token")
.await
.map_err(|e| LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: format!("DB query failed: {}", e),
})?
.ok_or_else(|| LlmError::SessionRenewalFailed {
provider: "nearai".to_string(),
reason: "No session in DB".to_string(),
})?;
})? {
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(),
})?
};
let session: SessionData =
serde_json::from_value(value).map_err(|e| LlmError::SessionRenewalFailed {
+163 -57
View File
@@ -22,9 +22,10 @@ use ironclaw::{
config::Config,
context::ContextManager,
extensions::ExtensionManager,
hooks::HookRegistry,
llm::{
FailoverProvider, LlmProvider, SessionConfig, create_llm_provider,
create_llm_provider_with_config, create_session_manager,
CooldownConfig, FailoverProvider, LlmProvider, SessionConfig, create_cheap_llm_provider,
create_llm_provider, create_llm_provider_with_config, create_session_manager,
},
orchestrator::{
ContainerJobConfig, ContainerJobManager, OrchestratorApi, TokenStore,
@@ -48,7 +49,6 @@ 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();
@@ -308,8 +308,11 @@ async fn main() -> anyhow::Result<()> {
};
let session = create_session_manager(session_config).await;
// Ensure we're authenticated before proceeding (only needed for NEAR AI backend)
if config.llm.backend == ironclaw::config::LlmBackend::NearAi {
// 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()
{
session.ensure_authenticated().await?;
}
@@ -335,7 +338,10 @@ 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 {
Some(ReplChannel::new())
let repl = ReplChannel::new();
// Suppress the one-liner banner; boot screen will be shown instead.
repl.suppress_banner();
Some(repl)
} else {
None
};
@@ -444,6 +450,72 @@ 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());
@@ -464,11 +536,26 @@ async fn main() -> anyhow::Result<()> {
fallback = %fallback.model_name(),
"LLM failover enabled"
);
Arc::new(FailoverProvider::new(vec![llm, fallback])?)
let cooldown_config = CooldownConfig {
cooldown_duration: std::time::Duration::from_secs(
config.llm.nearai.failover_cooldown_secs,
),
failure_threshold: config.llm.nearai.failover_cooldown_threshold,
};
Arc::new(FailoverProvider::with_cooldown(
vec![llm, fallback],
cooldown_config,
)?)
} else {
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");
@@ -542,49 +629,6 @@ 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)
@@ -854,12 +898,14 @@ 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");
}
}
@@ -995,6 +1041,7 @@ async fn main() -> anyhow::Result<()> {
}
}
channel_names.push(channel_name.clone());
channels.add(Box::new(SharedWasmChannel::new(channel_arc)));
}
@@ -1039,6 +1086,7 @@ async fn main() -> anyhow::Result<()> {
.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 {}:{}",
@@ -1101,8 +1149,11 @@ 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());
let session_manager = Arc::new(SessionManager::new().with_hooks(hooks.clone()));
// Register job tools (sandbox deps auto-injected when container_job_manager is available)
tools.register_job_tools(
@@ -1112,6 +1163,7 @@ 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 {
@@ -1144,29 +1196,39 @@ async fn main() -> anyhow::Result<()> {
}
}
gateway_url = Some(format!(
"http://{}:{}/?token={}",
gw_config.host,
gw_config.port,
gw.auth_token()
));
tracing::info!(
"Web gateway enabled on {}:{}",
gw_config.host,
gw_config.port
);
tracing::info!(
"Web UI: http://{}:{}/?token={}",
gw_config.host,
gw_config.port,
gw.auth_token()
);
tracing::info!("Web UI: http://{}:{}/", gw_config.host, gw_config.port);
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(),
@@ -1180,6 +1242,38 @@ 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?;
@@ -1207,6 +1301,18 @@ 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");
}
}
None
}
+57 -11
View File
@@ -40,8 +40,18 @@ pub struct Settings {
#[serde(default)]
pub secrets_master_key_source: KeySource,
// === Step 3: NEAR AI Auth ===
// Session stored separately in session.json
// === 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 4: Model Selection ===
/// Currently selected model.
@@ -504,7 +514,11 @@ impl Settings {
/// 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
// Start with defaults, then overlay each DB setting.
//
// The settings table stores both Settings struct fields and app-specific
// data (e.g. nearai.session_token). Skip keys that don't correspond to
// a known Settings path.
let mut settings = Self::default();
for (key, value) in map {
@@ -513,17 +527,23 @@ 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 => "null".to_string(),
serde_json::Value::Null => continue, // null means default, skip
other => other.to_string(),
};
if let Err(e) = settings.set(key, &value_str) {
tracing::warn!(
"Failed to apply DB setting '{}' = '{}': {}",
key,
value_str,
e
);
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
);
}
}
}
@@ -858,4 +878,30 @@ 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
@@ -0,0 +1,539 @@
# Setup / Onboarding Specification
This document is the authoritative specification for IronClaw's onboarding
wizard. Any code change to `src/setup/` **must** keep this document in sync.
If a future contributor or coding agent modifies setup behavior, update this
file first, then adjust the code to match.
---
## Entry Points
```
ironclaw onboard [--skip-auth] [--channels-only]
```
Explicit invocation. Loads `.env` files, runs the wizard, exits.
```
ironclaw (first run, no database configured)
```
Auto-detection via `check_onboard_needed()` in `main.rs`. Triggers when
none of these are true:
- `DATABASE_URL` env var is set
- `LIBSQL_PATH` env var is set
- `~/.ironclaw/ironclaw.db` exists on disk
The `--no-onboard` CLI flag suppresses auto-detection.
---
## Startup Sequence (main.rs)
```
1. Parse CLI args
2. If Command::Onboard → load .env, run wizard, exit
3. If Command::Run or no command:
a. Load .env files (dotenvy::dotenv() then load_ironclaw_env())
b. check_onboard_needed() → run wizard if needed
c. Config::from_env() → build config from env vars
d. Create SessionManager → load session token
e. ensure_authenticated() → validate session (NEAR AI only)
f. ... rest of agent startup
```
**Critical ordering:** `.env` files must be loaded (step 3a) before
`Config::from_env()` (step 3c) because bootstrap vars like
`DATABASE_BACKEND` live in `~/.ironclaw/.env`.
---
## The 7-Step Wizard
### Overview
```
Step 1: Database Connection
Step 2: Security (master key)
Step 3: Inference Provider ← skipped if --skip-auth
Step 4: Model Selection
Step 5: Embeddings
Step 6: Channel Configuration
Step 7: Background Tasks (heartbeat)
save_and_summarize()
```
`--channels-only` mode runs only Step 6, skipping everything else.
---
### Step 1: Database Connection
**Module:** `wizard.rs``step_database()`
**Goal:** Select backend, establish connection, run migrations.
**Decision tree:**
```
Both features compiled?
├─ Yes → DATABASE_BACKEND env var set?
│ ├─ Yes → use that backend
│ └─ No → interactive selection (PostgreSQL vs libSQL)
├─ Only postgres feature → step_database_postgres()
└─ Only libsql feature → step_database_libsql()
```
**PostgreSQL path** (`step_database_postgres`):
1. Check `DATABASE_URL` from env or settings
2. Test connection (creates `deadpool_postgres::Pool`)
3. Optionally run refinery migrations
4. Store pool in `self.db_pool`
**libSQL path** (`step_database_libsql`):
1. Offer local path (default: `~/.ironclaw/ironclaw.db`)
2. Optional Turso cloud sync (URL + auth token)
3. Test connection (creates `LibSqlBackend`)
4. Always run migrations (idempotent CREATE IF NOT EXISTS)
5. Store backend in `self.db_backend`
**Invariant:** After Step 1, exactly one of `self.db_pool` or
`self.db_backend` is `Some`. This is required for settings persistence
in `save_and_summarize()`.
---
### Step 2: Security (Master Key)
**Module:** `wizard.rs``step_security()`
**Goal:** Configure encryption for API tokens and secrets.
**Decision tree:**
```
SECRETS_MASTER_KEY env var set?
├─ Yes → use env var, done
└─ No → try get_master_key() from OS keychain
├─ Ok(bytes) → cache in self.secrets_crypto, ask "use existing?"
│ ├─ Yes → done (keychain)
│ └─ No → clear cache, fall through to options
└─ Err → fall through to options
├─ OS Keychain: generate + store + build SecretsCrypto
├─ Env variable: generate + print export command
└─ Skip: disable secrets features
```
**CRITICAL CAVEAT: macOS Keychain Dialogs**
On macOS, `security_framework::get_generic_password()` can trigger TWO
system dialogs:
1. "Enter your password to unlock the keychain" (keychain locked)
2. "Allow ironclaw to access this keychain item" (per-app authorization)
This is OS-level behavior we cannot prevent. To minimize pain:
- **Use `get_master_key()` not `has_master_key()`** in step 2. Both call
the same underlying API, but `get_master_key()` returns the key bytes
so we can cache them. `has_master_key()` throws them away, forcing a
second keychain access later.
- **Build `SecretsCrypto` eagerly.** When the keychain key is retrieved,
immediately construct `SecretsCrypto` and store in `self.secrets_crypto`.
Later calls to `init_secrets_context()` check this field first, avoiding
redundant keychain probes.
- **Never probe the keychain in read-only commands** (e.g., `ironclaw status`).
The status command reports "env not set (keychain may be configured)"
rather than triggering system dialogs.
**Invariant:** After Step 2, `self.secrets_crypto` is `Some` if the user
chose Keychain or generated a new key. It may be `None` if the user chose
env-var mode or skipped secrets.
---
### Step 3: Inference Provider
**Module:** `wizard.rs``step_inference_provider()`
**Goal:** Choose LLM backend and authenticate.
**Providers:**
| Provider | Auth Method | Secret Name | Env Var |
|----------|-------------|-------------|---------|
| NEAR AI | Browser OAuth | (session token) | `NEARAI_SESSION_TOKEN` |
| Anthropic | API key | `anthropic_api_key` | `ANTHROPIC_API_KEY` |
| OpenAI | API key | `openai_api_key` | `OPENAI_API_KEY` |
| Ollama | None | - | - |
| OpenAI-compatible | Optional API key | `llm_compatible_api_key` | `LLM_API_KEY` |
**API-key providers** (`setup_api_key_provider`):
1. Check env var → if set, ask to reuse, persist to secrets store
2. Otherwise prompt for key entry via `secret_input()`
3. Store encrypted in secrets via `init_secrets_context()`
4. **Cache key in `self.llm_api_key`** for model fetching in Step 4
**NEAR AI** (`setup_nearai`):
- Calls `session_manager.ensure_authenticated()` which opens browser
- Session token saved to `~/.ironclaw/session.json`
**`self.llm_api_key` caching:** The wizard caches the API key as
`Option<SecretString>` so that Step 4 (model fetching) and Step 5
(embeddings) can use it without re-reading from the secrets store or
mutating environment variables.
---
### Step 4: Model Selection
**Module:** `wizard.rs``step_model_selection()`
**Goal:** Choose which model to use.
**Flow:**
1. If model already set → offer to keep it
2. Fetch models from provider API (5-second timeout)
3. On timeout or error → use static fallback list
4. Present list + "Custom model ID" escape hatch
5. Store in `self.settings.selected_model`
**Model fetchers pass the cached API key explicitly:**
```rust
let cached = self.llm_api_key.as_ref().map(|k| k.expose_secret().to_string());
let models = fetch_anthropic_models(cached.as_deref()).await;
```
This avoids mutating environment variables. The fetcher checks the explicit
key first, then falls back to the standard env var.
---
### Step 5: Embeddings
**Module:** `wizard.rs``step_embeddings()`
**Goal:** Configure semantic search for workspace memory.
**Flow:**
1. Ask "Enable semantic search?" (default: yes)
2. Detect available providers:
- NEAR AI: if backend is `nearai` OR valid session exists
- OpenAI: if `OPENAI_API_KEY` in env OR (backend is `openai` AND cached key)
3. If both available → let user choose
4. If only one → use it
5. If neither → disable embeddings
**Default model:** `text-embedding-3-small` (for both providers)
---
### Step 6: Channel Configuration
**Module:** `wizard.rs``step_channels()`, delegating to `channels.rs`
**Goal:** Enable input channels (TUI, HTTP, Telegram, etc.).
**Sub-steps:**
```
6a. Tunnel setup (if webhook channels needed)
6b. Discover WASM channels from ~/.ironclaw/channels/
6c. Multi-select: CLI/TUI, HTTP, discovered channels, bundled channels
6d. Install missing bundled channels (copy WASM binaries)
6e. Initialize SecretsContext (for token storage)
6f. Setup HTTP webhook (if selected)
6g. Setup each WASM channel (secrets, owner binding)
```
**Tunnel setup** (`setup_tunnel`):
- Options: ngrok, Cloudflare Tunnel, localtunnel, custom URL
- Validates HTTPS requirement
- Stored in `self.settings.tunnel.public_url`
**WASM channel setup** (`setup_wasm_channel`):
- Reads `capabilities.json` for `setup.required_secrets`
- For each secret: check existing, prompt or auto-generate, validate regex
- Save each secret via `SecretsContext`
**Telegram special case** (`setup_telegram`):
- Validates bot token via Telegram `getMe` API
- Owner binding: polls `getUpdates` for 120s to capture sender's user ID
- Optional webhook secret generation
**SecretsContext creation** (`init_secrets_context`):
1. Check `self.secrets_crypto` (set in Step 2) → use if available
2. Else try `SECRETS_MASTER_KEY` env var
3. Else try `get_master_key()` from keychain (only in `channels_only` mode)
4. Create backend-appropriate secrets store (respects selected database backend)
---
### Step 7: Heartbeat
**Module:** `wizard.rs``step_heartbeat()`
**Goal:** Configure periodic background execution.
**Flow:**
1. Ask "Enable heartbeat?" (default: no)
2. If yes: interval in minutes (default: 30), notification channel
3. Store in `self.settings.heartbeat`
---
## Settings Persistence
### Two-Layer Architecture
Settings are persisted in two places:
**Layer 1: `~/.ironclaw/.env`** (bootstrap vars)
Contains only the settings needed BEFORE database connection. Written by
`save_bootstrap_env()` in `bootstrap.rs`.
```env
DATABASE_BACKEND="libsql"
LIBSQL_PATH="/Users/name/.ironclaw/ironclaw.db"
```
Or for PostgreSQL:
```env
DATABASE_BACKEND="postgres"
DATABASE_URL="postgres://user:pass@localhost/ironclaw"
```
**Why separate?** Chicken-and-egg: you need `DATABASE_BACKEND` to know
which database to connect to, so it can't be stored in the database.
**Layer 2: Database settings table** (everything else)
All other settings are stored as key-value pairs in the `settings` table,
keyed by `(user_id, key)`. Written by `set_all_settings()`.
Settings are serialized via `Settings::to_db_map()` as dotted paths:
```
database_backend = "libsql"
llm_backend = "nearai"
selected_model = "anthropic/claude-sonnet-4-5"
embeddings.enabled = "true"
embeddings.provider = "nearai"
channels.http_enabled = "true"
heartbeat.enabled = "true"
heartbeat.interval_secs = "300"
```
### save_and_summarize()
Final step of the wizard:
```
1. Mark onboard_completed = true
2. Write ALL settings to database (try postgres pool, then libSQL backend)
3. Write bootstrap vars to ~/.ironclaw/.env:
- DATABASE_BACKEND (always)
- DATABASE_URL (if postgres)
- LIBSQL_PATH (if libsql)
- LIBSQL_URL (if turso sync)
4. Print configuration summary
```
**Invariant:** Both Layer 1 and Layer 2 must be written. If the database
write fails, the wizard returns an error and the `.env` file is not written.
### Legacy Migration
`bootstrap.rs` handles one-time upgrades from older config formats:
- `bootstrap.json` → extracts `DATABASE_URL`, writes `.env`, renames to `.migrated`
- `settings.json` → migrated to database via `migrate_disk_to_db()`
---
## Settings Struct
**Module:** `settings.rs`
```rust
pub struct Settings {
// Meta
pub onboard_completed: bool,
// Step 1: Database
pub database_backend: Option<String>, // "postgres" | "libsql"
pub database_url: Option<String>,
pub libsql_path: Option<String>,
pub libsql_url: Option<String>,
// Step 2: Security
pub secrets_master_key_source: KeySource, // Keychain | Env | None
// Step 3: Inference
pub llm_backend: Option<String>, // "nearai" | "anthropic" | "openai" | "ollama" | "openai_compatible"
pub ollama_base_url: Option<String>,
pub openai_compatible_base_url: Option<String>,
// Step 4: Model
pub selected_model: Option<String>,
// Step 5: Embeddings
pub embeddings: EmbeddingsSettings, // enabled, provider, model
// Step 6: Channels
pub tunnel: TunnelSettings, // provider, public_url
pub channels: ChannelSettings, // http config, telegram owner, etc.
// Step 7: Heartbeat
pub heartbeat: HeartbeatSettings, // enabled, interval, notify
// Advanced (not in wizard, set via `ironclaw config set`)
pub agent: AgentSettings,
pub wasm: WasmSettings,
pub sandbox: SandboxSettings,
pub safety: SafetySettings,
pub builder: BuilderSettings,
}
```
**KeySource enum:** `Keychain | Env | None`
---
## Secrets Flow
### SecretsContext
Thin wrapper for setup-time secret operations:
```rust
pub struct SecretsContext {
store: Arc<dyn SecretsStore>,
user_id: String,
}
```
Created by `init_secrets_context()` which:
1. Gets `SecretsCrypto` from `self.secrets_crypto` or loads from keychain/env
2. Creates the appropriate backend store:
- If both features compiled: respects `self.settings.database_backend`
- Tries selected backend first, falls back to the other
3. Returns `SecretsContext` wrapping the store
### Secret Storage
Secrets are encrypted with AES-256-GCM using the master key, then stored
in the database `secrets` table. The wizard writes secrets like:
```
telegram_bot_token → encrypted bot token
telegram_webhook_secret → encrypted webhook HMAC secret
anthropic_api_key → encrypted API key
```
---
## Prompt Utilities
**Module:** `prompts.rs`
| Function | Description |
|----------|-------------|
| `select_one(label, options)` | Numbered single-choice menu |
| `select_many(label, options, defaults)` | Checkbox multi-select (raw terminal mode) |
| `input(label)` | Single line text input |
| `optional_input(label, hint)` | Text input that can be empty |
| `secret_input(label)` | Hidden input (shows `*` per char), returns `SecretString` |
| `confirm(label, default)` | `[Y/n]` or `[y/N]` prompt |
| `print_header(text)` | Bold section header with underline |
| `print_step(n, total, text)` | `[1/7] Step Name` |
| `print_success(text)` | Green checkmark prefix |
| `print_error(text)` | Red X prefix |
| `print_info(text)` | Blue info prefix |
`select_many` uses `crossterm` raw mode for arrow key navigation.
Must properly restore terminal state on all exit paths.
---
## Platform Caveats
### macOS Keychain
- `get_generic_password()` triggers system dialogs (unlock + authorize)
- Two dialogs per call is normal, not a bug
- Cache the result after first access to avoid repeat prompts
- Never probe keychain in read-only commands (`status`, `--help`)
- Service name: `"ironclaw"`, account: `"master_key"`
### Linux Secret Service
- Uses GNOME Keyring or KWallet via `secret-service` crate
- May need `gnome-keyring` daemon running
- Collection unlock may prompt for password
### URL Passwords
- `#` is common in URL-encoded passwords (`%23` decoded)
- `.env` values must be double-quoted to preserve `#`
- Display masked: `postgres://user:****@host/db`
### Telegram API
- Bot token format: `123456:ABC-DEF...`
- Token goes in URL path: `https://api.telegram.org/bot{TOKEN}/method`
- Webhook secret header: `X-Telegram-Bot-Api-Secret-Token`
- Owner binding polls `getUpdates` (must delete webhook first)
---
## Testing
Tests live in `mod tests {}` at the bottom of each file.
**What to test when modifying setup:**
- Settings round-trip: `to_db_map()` then `from_db_map()` preserves values
- Bootstrap `.env`: dotenvy can parse what `save_bootstrap_env()` writes
- Model fetchers: static fallback works when API is unreachable
- Channel discovery: handles missing dir, invalid JSON, deduplication
- Prompt functions: not tested (interactive I/O), but ensure error paths
don't panic
**Run setup tests:**
```bash
cargo test --lib -- setup
cargo test --lib -- bootstrap
```
---
## Modification Checklist
When changing the onboarding flow:
1. Update this README first with the intended behavior change
2. If adding a new wizard step:
- Add to the step enum in `run()`, adjust `total_steps`
- Add corresponding settings fields to `Settings`
- Add `to_db_map` / `from_db_map` serialization
- If the setting is needed before DB connection, add to `save_bootstrap_env()`
3. If adding a new provider or channel:
- Add to the selection menu in the appropriate step
- Add authentication flow (API key or OAuth)
- Add model fetcher with static fallback + 5s timeout
4. If touching keychain:
- Cache the result, never call `get_master_key()` twice
- Test on macOS (dialog behavior differs from Linux)
5. If touching secrets:
- Ensure `init_secrets_context()` respects the selected database backend
- Test with both postgres and libsql features
6. Run the full shipping checklist:
```bash
cargo fmt
cargo clippy --all --benches --tests --examples --all-features -- -D warnings
cargo test --lib -- setup bootstrap
```
7. Test a fresh onboarding: `rm -rf ~/.ironclaw && cargo run`
+149 -106
View File
@@ -20,6 +20,22 @@ 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>,
@@ -45,32 +61,39 @@ impl SecretsContext {
}
/// Save a secret to the database.
pub async fn save_secret(&self, name: &str, value: &SecretString) -> Result<(), String> {
pub async fn save_secret(
&self,
name: &str,
value: &SecretString,
) -> Result<(), ChannelSetupError> {
let params = CreateSecretParams::new(name, value.expose_secret());
self.store
.create(&self.user_id, params)
.await
.map_err(|e| format!("Failed to save secret: {}", e))?;
.map_err(|e| ChannelSetupError::Secrets(format!("Failed to save secret: {}", e)))?;
Ok(())
}
/// Check if a secret exists.
pub async fn secret_exists(&self, name: &str) -> bool {
self.store
.exists(&self.user_id, name)
.await
.unwrap_or(false)
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
}
}
}
/// Read a secret from the database (decrypted).
pub async fn get_secret(&self, name: &str) -> Result<SecretString, String> {
pub async fn get_secret(&self, name: &str) -> Result<SecretString, ChannelSetupError> {
let decrypted = self
.store
.get_decrypted(&self.user_id, name)
.await
.map_err(|e| format!("Failed to read secret: {}", e))?;
.map_err(|e| ChannelSetupError::Secrets(format!("Failed to read secret: {}", e)))?;
Ok(SecretString::from(decrypted.expose().to_string()))
}
}
@@ -107,7 +130,6 @@ struct TelegramGetUpdatesResponse {
#[derive(Debug, Deserialize)]
struct TelegramUpdate {
#[allow(dead_code)]
update_id: i64,
message: Option<TelegramUpdateMessage>,
}
@@ -134,7 +156,7 @@ struct TelegramUpdateUser {
pub async fn setup_telegram(
secrets: &SecretsContext,
settings: &Settings,
) -> Result<TelegramSetupResult, String> {
) -> Result<TelegramSetupResult, ChannelSetupError> {
println!("Telegram Setup:");
println!();
print_info("To create a Telegram bot:");
@@ -146,7 +168,7 @@ 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).map_err(|e| e.to_string())? {
if !confirm("Replace existing token?", false)? {
// 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?;
@@ -159,47 +181,48 @@ pub async fn setup_telegram(
}
}
let token = secret_input("Bot token (from @BotFather)").map_err(|e| e.to_string())?;
loop {
let token = secret_input("Bot token (from @BotFather)")?;
// Validate the token
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, &settings.tunnel).await?;
Ok(TelegramSetupResult {
enabled: true,
bot_username: username,
webhook_secret,
owner_id,
})
}
Err(e) => {
print_error(&format!("Token validation failed: {}", e));
return 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).map_err(|e| e.to_string())? {
Box::pin(setup_telegram(secrets, settings)).await
} else {
Ok(TelegramSetupResult {
enabled: false,
bot_username: None,
webhook_secret: None,
owner_id: None,
})
if !confirm("Try again?", true)? {
return Ok(TelegramSetupResult {
enabled: false,
bot_username: None,
webhook_secret: None,
owner_id: None,
});
}
}
}
}
@@ -209,14 +232,14 @@ pub async fn setup_telegram(
///
/// Polls `getUpdates` until a message arrives, then captures the sender's user ID.
/// Returns `None` if the user declines or the flow times out.
async fn bind_telegram_owner(token: &SecretString) -> Result<Option<i64>, String> {
async fn bind_telegram_owner(token: &SecretString) -> Result<Option<i64>, ChannelSetupError> {
println!();
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).map_err(|e| e.to_string())? {
if !confirm("Bind bot to your Telegram account?", true)? {
print_info("Skipping account binding. Bot will accept messages from all users.");
return Ok(None);
}
@@ -227,14 +250,16 @@ async fn bind_telegram_owner(token: &SecretString) -> Result<Option<i64>, String
let client = Client::builder()
.timeout(std::time::Duration::from_secs(35))
.build()
.map_err(|e| format!("Failed to create HTTP client: {}", e))?;
.map_err(|e| ChannelSetupError::Network(format!("Failed to create HTTP client: {}", e)))?;
// Clear any existing webhook so getUpdates works
let delete_url = format!(
"https://api.telegram.org/bot{}/deleteWebhook",
token.expose_secret()
);
let _ = client.post(&delete_url).send().await;
if let Err(e) = client.post(&delete_url).send().await {
tracing::warn!("Failed to delete webhook (getUpdates may not work): {e}");
}
let updates_url = format!(
"https://api.telegram.org/bot{}/getUpdates",
@@ -249,19 +274,23 @@ async fn bind_telegram_owner(token: &SecretString) -> Result<Option<i64>, String
.query(&[("timeout", "30"), ("allowed_updates", "[\"message\"]")])
.send()
.await
.map_err(|e| format!("getUpdates request failed: {}", e))?;
.map_err(|e| ChannelSetupError::Network(format!("getUpdates request failed: {}", e)))?;
if !response.status().is_success() {
return Err(format!("getUpdates returned status {}", response.status()));
return Err(ChannelSetupError::Network(format!(
"getUpdates returned status {}",
response.status()
)));
}
let body: TelegramGetUpdatesResponse = response
.json()
.await
.map_err(|e| format!("Failed to parse getUpdates response: {}", e))?;
let body: TelegramGetUpdatesResponse = response.json().await.map_err(|e| {
ChannelSetupError::Network(format!("Failed to parse getUpdates response: {}", e))
})?;
if !body.ok {
return Err("Telegram API returned error for getUpdates".to_string());
return Err(ChannelSetupError::Network(
"Telegram API returned error for getUpdates".to_string(),
));
}
// Find the first message with a sender
@@ -285,11 +314,14 @@ async fn bind_telegram_owner(token: &SecretString) -> Result<Option<i64>, String
"https://api.telegram.org/bot{}/getUpdates",
token.expose_secret()
);
let _ = client
if let Err(e) = client
.get(&ack_url)
.query(&[("offset", &(update.update_id + 1).to_string())])
.send()
.await;
.await
{
tracing::warn!("Failed to acknowledge Telegram update: {e}");
}
return Ok(Some(from.id));
}
@@ -307,10 +339,10 @@ async fn bind_telegram_owner(token: &SecretString) -> Result<Option<i64>, String
async fn bind_telegram_owner_flow(
secrets: &SecretsContext,
settings: &Settings,
) -> Result<Option<i64>, String> {
) -> Result<Option<i64>, ChannelSetupError> {
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).map_err(|e| e.to_string())? {
if !confirm("Re-bind to a different account?", false)? {
return Ok(settings.channels.telegram_owner_id);
}
}
@@ -325,10 +357,10 @@ async fn bind_telegram_owner_flow(
///
/// This is shared across all channels that need webhook endpoints.
/// Returns the tunnel URL if configured.
pub fn setup_tunnel(settings: &Settings) -> Result<Option<String>, String> {
pub fn setup_tunnel(settings: &Settings) -> Result<Option<String>, ChannelSetupError> {
if let Some(ref url) = settings.tunnel.public_url {
print_info(&format!("Existing tunnel configured: {}", url));
if !confirm("Change tunnel configuration?", false).map_err(|e| e.to_string())? {
if !confirm("Change tunnel configuration?", false)? {
return Ok(Some(url.clone()));
}
}
@@ -348,17 +380,18 @@ pub fn setup_tunnel(settings: &Settings) -> Result<Option<String>, String> {
print_info("Security comes from provider-specific secrets (e.g., Telegram webhook secret).");
println!();
if !confirm("Configure a tunnel?", false).map_err(|e| e.to_string())? {
if !confirm("Configure a tunnel?", false)? {
return Ok(None);
}
let tunnel_url =
input("Tunnel URL (e.g., https://abc123.ngrok.io)").map_err(|e| e.to_string())?;
let tunnel_url = input("Tunnel URL (e.g., https://abc123.ngrok.io)")?;
// Validate URL format
if !tunnel_url.starts_with("https://") {
print_error("URL must start with https:// (webhooks require HTTPS)");
return Err("Invalid tunnel URL: must use HTTPS".to_string());
return Err(ChannelSetupError::Validation(
"Invalid tunnel URL: must use HTTPS".to_string(),
));
}
// Remove trailing slash if present
@@ -378,7 +411,7 @@ pub fn setup_tunnel(settings: &Settings) -> Result<Option<String>, String> {
async fn setup_telegram_webhook_secret(
secrets: &SecretsContext,
tunnel: &TunnelSettings,
) -> Result<Option<String>, String> {
) -> Result<Option<String>, ChannelSetupError> {
if tunnel.public_url.is_none() {
print_info("");
print_info("No tunnel configured. Telegram will use polling mode (30s+ delay).");
@@ -391,7 +424,7 @@ async fn setup_telegram_webhook_secret(
print_info("A webhook secret adds an extra layer of security by validating");
print_info("that requests actually come from Telegram's servers.");
if !confirm("Generate a webhook secret?", true).map_err(|e| e.to_string())? {
if !confirm("Generate a webhook secret?", true)? {
return Ok(None);
}
@@ -410,11 +443,13 @@ 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>, String> {
pub async fn validate_telegram_token(
token: &SecretString,
) -> Result<Option<String>, ChannelSetupError> {
let client = Client::builder()
.timeout(std::time::Duration::from_secs(10))
.build()
.map_err(|e| format!("Failed to create HTTP client: {}", e))?;
.map_err(|e| ChannelSetupError::Network(format!("Failed to create HTTP client: {}", e)))?;
let url = format!(
"https://api.telegram.org/bot{}/getMe",
@@ -425,21 +460,26 @@ pub async fn validate_telegram_token(token: &SecretString) -> Result<Option<Stri
.get(&url)
.send()
.await
.map_err(|e| format!("Request failed: {}", e))?;
.map_err(|e| ChannelSetupError::Network(format!("Request failed: {}", e)))?;
if !response.status().is_success() {
return Err(format!("API returned status {}", response.status()));
return Err(ChannelSetupError::Network(format!(
"API returned status {}",
response.status()
)));
}
let body: TelegramGetMeResponse = response
.json()
.await
.map_err(|e| format!("Failed to parse response: {}", e))?;
.map_err(|e| ChannelSetupError::Network(format!("Failed to parse response: {}", e)))?;
if body.ok {
Ok(body.result.and_then(|u| u.username))
} else {
Err("Telegram API returned error".to_string())
Err(ChannelSetupError::Network(
"Telegram API returned error".to_string(),
))
}
}
@@ -452,38 +492,34 @@ pub struct HttpSetupResult {
}
/// Set up HTTP webhook channel.
pub async fn setup_http(secrets: &SecretsContext) -> Result<HttpSetupResult, String> {
pub async fn setup_http(secrets: &SecretsContext) -> Result<HttpSetupResult, ChannelSetupError> {
println!("HTTP Webhook Setup:");
println!();
print_info("The HTTP webhook allows external services to send messages to the agent.");
println!();
let port_str = optional_input("Port", Some("default: 8080")).map_err(|e| e.to_string())?;
let port_str = optional_input("Port", Some("default: 8080"))?;
let port: u16 = port_str
.as_deref()
.unwrap_or("8080")
.parse()
.map_err(|e| format!("Invalid port: {}", e))?;
.map_err(|e| ChannelSetupError::Validation(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"))
.map_err(|e| e.to_string())?
.unwrap_or_else(|| "0.0.0.0".to_string());
let host =
optional_input("Host", Some("default: 0.0.0.0"))?.unwrap_or_else(|| "0.0.0.0".to_string());
// Generate a webhook secret
if confirm("Generate a webhook secret for authentication?", true).map_err(|e| e.to_string())? {
if confirm("Generate a webhook secret for authentication?", true)? {
let secret = generate_webhook_secret();
secrets
.save_secret("http_webhook_secret", &SecretString::from(secret.clone()))
.save_secret("http_webhook_secret", &SecretString::from(secret))
.await?;
print_success("Webhook secret generated and saved to database");
print_info(&format!(
"Secret: {} (store this for your webhook clients)",
secret
));
print_info("Retrieve it later with: ironclaw secret get http_webhook_secret");
}
print_success(&format!("HTTP webhook will listen on {}:{}", host, port));
@@ -497,11 +533,7 @@ pub async fn setup_http(secrets: &SecretsContext) -> Result<HttpSetupResult, Str
/// Generate a random webhook secret.
pub fn generate_webhook_secret() -> String {
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()
generate_secret_with_length(32)
}
/// Result of WASM channel setup.
@@ -519,7 +551,7 @@ pub async fn setup_wasm_channel(
secrets: &SecretsContext,
channel_name: &str,
setup: &crate::channels::wasm::SetupSchema,
) -> Result<WasmChannelSetupResult, String> {
) -> Result<WasmChannelSetupResult, ChannelSetupError> {
println!("{} Setup:", channel_name);
println!();
@@ -530,7 +562,7 @@ pub async fn setup_wasm_channel(
"Existing {} found in database.",
secret_config.name
));
if !confirm("Replace existing value?", false).map_err(|e| e.to_string())? {
if !confirm("Replace existing value?", false)? {
continue;
}
}
@@ -538,8 +570,7 @@ 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"))
.map_err(|e| e.to_string())?;
optional_input(&secret_config.prompt, Some("leave empty to auto-generate"))?;
if let Some(v) = input_value {
if !v.is_empty() {
@@ -566,18 +597,21 @@ pub async fn setup_wasm_channel(
}
} else {
// Required secret
let input_value = secret_input(&secret_config.prompt).map_err(|e| e.to_string())?;
let input_value = secret_input(&secret_config.prompt)?;
// Validate if pattern is provided
if let Some(ref pattern) = secret_config.validation {
let re = regex::Regex::new(pattern)
.map_err(|e| format!("Invalid validation pattern: {}", e))?;
let re = regex::Regex::new(pattern).map_err(|e| {
ChannelSetupError::Validation(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("Validation failed".to_string());
return Err(ChannelSetupError::Validation(
"Validation failed".to_string(),
));
}
}
@@ -589,14 +623,11 @@ pub async fn setup_wasm_channel(
print_success(&format!("{} saved to database", secret_config.name));
}
// Optionally validate the configuration
// TODO: Substitute secrets into the validation URL and make a
// GET request to verify the configured credentials actually work.
if let Some(ref validation_endpoint) = setup.validation_endpoint {
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 skipped)",
"Validation endpoint configured: {} (validation not yet implemented)",
validation_endpoint
));
}
@@ -620,11 +651,23 @@ fn generate_secret_with_length(length: usize) -> String {
#[cfg(test)]
mod tests {
use super::*;
use crate::setup::channels::generate_webhook_secret;
#[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);
}
}
+3 -2
View File
@@ -3,7 +3,7 @@
//! Provides a guided setup experience for:
//! 1. Database connection
//! 2. Security (secrets master key)
//! 3. NEAR AI authentication
//! 3. Inference provider selection
//! 4. Model selection
//! 5. Embeddings
//! 6. Channel configuration (HTTP, Telegram, etc.)
@@ -24,7 +24,8 @@ mod prompts;
mod wizard;
pub use channels::{
SecretsContext, setup_http, setup_telegram, setup_tunnel, validate_telegram_token,
ChannelSetupError, SecretsContext, setup_http, setup_telegram, setup_tunnel,
validate_telegram_token,
};
pub use prompts::{
confirm, input, optional_input, print_error, print_header, print_info, print_step,
+5
View File
@@ -21,6 +21,7 @@ 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
///
@@ -84,6 +85,10 @@ pub fn select_one(prompt: &str, options: &[&str]) -> io::Result<usize> {
/// ])?;
/// ```
pub fn select_many(prompt: &str, options: &[(&str, bool)]) -> io::Result<Vec<usize>> {
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;
+866 -114
View File
File diff suppressed because it is too large Load Diff
-451
View File
@@ -1,451 +0,0 @@
//! Accessibility tree parsing and element reference generation.
//!
//! Converts Chrome's CDP accessibility tree into a compact, LLM-friendly
//! representation with stable element references (`@e1`, `@e2`, ...).
//!
//! The key insight: sending the full accessibility tree every turn is wasteful.
//! Instead, we assign short IDs to interactive elements and let the LLM
//! reference them by ID for clicks/typing. This is ~93% cheaper in tokens
//! compared to re-sending the full tree each time.
//!
//! ```text
//! Page: https://example.com/login
//! @e1: textbox "Email" [focused]
//! @e2: textbox "Password" [type=password]
//! @e3: button "Sign In"
//! @e4: link "Forgot password?"
//! ```
use std::collections::HashMap;
use std::fmt;
use chromiumoxide::cdp::browser_protocol::accessibility::{AxNode, AxPropertyName};
use chromiumoxide::cdp::browser_protocol::dom::BackendNodeId;
/// A resolved element reference that maps `@eN` back to a DOM target.
#[derive(Debug, Clone)]
pub struct ElementRef {
/// The display label shown to the LLM (e.g., `textbox "Email"`).
#[allow(dead_code)]
pub label: String,
/// CDP backend node ID for targeting this element.
pub backend_node_id: BackendNodeId,
/// CSS selector hint (best-effort, may not be unique).
#[allow(dead_code)]
pub selector_hint: Option<String>,
}
/// Stores the current set of element references for a page snapshot.
#[derive(Debug, Clone, Default)]
pub struct ElementRefMap {
refs: HashMap<String, ElementRef>,
counter: usize,
}
impl ElementRefMap {
pub fn new() -> Self {
Self::default()
}
/// Look up a reference like `@e1` or just `e1`.
pub fn get(&self, ref_id: &str) -> Option<&ElementRef> {
let normalized = ref_id.strip_prefix('@').unwrap_or(ref_id);
self.refs.get(normalized)
}
/// Number of tracked elements.
#[allow(dead_code)]
pub fn len(&self) -> usize {
self.refs.len()
}
pub fn is_empty(&self) -> bool {
self.refs.is_empty()
}
/// Reset all refs. Called before each new `read_page` and when switching tabs.
pub fn reset(&mut self) {
self.refs.clear();
self.counter = 0;
}
/// Allocate the next reference ID and store the element.
fn insert(&mut self, elem: ElementRef) -> String {
self.counter += 1;
let id = format!("e{}", self.counter);
self.refs.insert(id.clone(), elem);
id
}
}
/// Which elements to include when building the tree representation.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ElementFilter {
/// Only interactive elements (buttons, links, inputs, selects, textareas).
Interactive,
/// All elements with meaningful content.
All,
}
impl ElementFilter {
pub fn from_str_opt(s: Option<&str>) -> Self {
match s {
Some("all") => Self::All,
_ => Self::Interactive,
}
}
}
/// Roles that are considered "interactive" for filtering purposes.
const INTERACTIVE_ROLES: &[&str] = &[
"button",
"link",
"textbox",
"searchbox",
"combobox",
"listbox",
"option",
"menuitem",
"menuitemcheckbox",
"menuitemradio",
"radio",
"checkbox",
"switch",
"slider",
"spinbutton",
"tab",
"treeitem",
];
/// Roles to skip entirely (structural noise).
const SKIP_ROLES: &[&str] = &[
"none",
"presentation",
"generic",
"InlineTextBox",
"LineBreak",
];
/// Build a compact page representation from the CDP accessibility tree.
///
/// Returns the text representation and populates `ref_map` with element
/// references the LLM can use for subsequent actions.
pub fn build_page_repr(
url: &str,
title: &str,
nodes: &[AxNode],
filter: ElementFilter,
ref_map: &mut ElementRefMap,
) -> String {
ref_map.reset();
let mut lines = Vec::new();
// Header
lines.push(format!("Page: {}", url));
if !title.is_empty() {
lines.push(format!("Title: {}", title));
}
lines.push(String::new());
// Walk nodes, collecting elements that pass the filter.
for node in nodes {
let role = node_role(node);
if SKIP_ROLES.contains(&role.as_str()) {
continue;
}
// For "interactive" filter, only include interactive roles.
if filter == ElementFilter::Interactive && !INTERACTIVE_ROLES.contains(&role.as_str()) {
continue;
}
// Skip nodes without a name (usually decorative).
let name = node_name(node);
if name.is_empty() && filter == ElementFilter::Interactive {
continue;
}
let backend_id = match node.backend_dom_node_id {
Some(id) => id,
None => continue,
};
// Build display label
let mut label = NodeLabel {
role: role.clone(),
name: truncate_name(&name, 80),
properties: Vec::new(),
};
// Add useful properties
if node_has_property(node, "focused") {
label.properties.push("focused".to_string());
}
if node_has_property(node, "checked") {
label.properties.push("checked".to_string());
}
if node_has_property(node, "disabled") {
label.properties.push("disabled".to_string());
}
if node_has_property(node, "expanded") {
label.properties.push("expanded".to_string());
}
if node_has_property(node, "required") {
label.properties.push("required".to_string());
}
if let Some(val) = node_value(node) {
if !val.is_empty() && val != name {
label
.properties
.push(format!("value=\"{}\"", truncate_name(&val, 40)));
}
}
let display = label.to_string();
let elem_ref = ElementRef {
label: display.clone(),
backend_node_id: backend_id,
selector_hint: guess_selector(node),
};
let ref_id = ref_map.insert(elem_ref);
lines.push(format!("@{}: {}", ref_id, display));
}
if ref_map.is_empty() {
lines.push("(no interactive elements found)".to_string());
}
lines.join("\n")
}
/// Extract the role string from an AX node.
fn node_role(node: &AxNode) -> String {
node.role
.as_ref()
.and_then(|v| v.value.as_ref())
.and_then(|v| v.as_str())
.unwrap_or("unknown")
.to_string()
}
/// Extract the name (accessible label) from an AX node.
fn node_name(node: &AxNode) -> String {
node.name
.as_ref()
.and_then(|v| v.value.as_ref())
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string()
}
/// Extract the value from an AX node (for inputs, etc.).
fn node_value(node: &AxNode) -> Option<String> {
node.value
.as_ref()
.and_then(|v| v.value.as_ref())
.and_then(|v| v.as_str())
.map(|s| s.to_string())
}
/// Map a property name string to the corresponding `AxPropertyName` variant.
fn property_by_name(name: &str) -> Option<AxPropertyName> {
match name {
"focused" => Some(AxPropertyName::Focused),
"checked" => Some(AxPropertyName::Checked),
"disabled" => Some(AxPropertyName::Disabled),
"expanded" => Some(AxPropertyName::Expanded),
"required" => Some(AxPropertyName::Required),
"selected" => Some(AxPropertyName::Selected),
"pressed" => Some(AxPropertyName::Pressed),
"readonly" => Some(AxPropertyName::Readonly),
"hidden" => Some(AxPropertyName::Hidden),
"modal" => Some(AxPropertyName::Modal),
_ => None,
}
}
/// Check if a node has a boolean property set to true.
fn node_has_property(node: &AxNode, prop_name: &str) -> bool {
let Some(props) = &node.properties else {
return false;
};
let Some(target) = property_by_name(prop_name) else {
return false;
};
props.iter().any(|p| {
p.name == target
&& p.value
.value
.as_ref()
.and_then(|v| v.as_bool())
.unwrap_or(false)
})
}
/// Best-effort CSS selector guess from node attributes.
fn guess_selector(node: &AxNode) -> Option<String> {
// We don't have DOM attributes directly from the AX tree,
// so we can only offer role-based hints. The actual targeting
// uses backend_node_id which is precise.
let role = node_role(node);
let name = node_name(node);
if name.is_empty() {
return None;
}
// Build an ARIA selector hint (not used for actual targeting,
// just a human-readable hint in debug output).
Some(format!(
"[role=\"{}\"][name=\"{}\"]",
role,
truncate_name(&name, 30)
))
}
/// Truncate a display name to max chars, adding ellipsis if needed.
fn truncate_name(s: &str, max: usize) -> String {
if s.chars().count() <= max {
s.to_string()
} else {
format!(
"{}...",
s.chars().take(max.saturating_sub(3)).collect::<String>()
)
}
}
/// Helper for formatting a node's display label.
struct NodeLabel {
role: String,
name: String,
properties: Vec<String>,
}
impl fmt::Display for NodeLabel {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.role)?;
if !self.name.is_empty() {
write!(f, " \"{}\"", self.name)?;
}
if !self.properties.is_empty() {
write!(f, " [{}]", self.properties.join(", "))?;
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use crate::tools::builtin::browser::accessibility::{
ElementFilter, ElementRefMap, build_page_repr, truncate_name,
};
use chromiumoxide::cdp::browser_protocol::accessibility::{
AxNode, AxNodeId, AxValue, AxValueType,
};
use chromiumoxide::cdp::browser_protocol::dom::BackendNodeId;
fn make_ax_value(s: &str) -> AxValue {
let mut v = AxValue::new(AxValueType::String);
v.value = Some(serde_json::Value::String(s.to_string()));
v
}
fn make_ax_node(role: &str, name: &str, backend_id: i64) -> AxNode {
let mut node = AxNode::new(AxNodeId::from(format!("node_{}", backend_id)), false);
node.role = Some(make_ax_value(role));
node.name = Some(make_ax_value(name));
node.backend_dom_node_id = Some(BackendNodeId::new(backend_id));
node
}
#[test]
fn test_build_page_repr_interactive_filter() {
let nodes = vec![
make_ax_node("button", "Submit", 1),
make_ax_node("link", "Home", 2),
make_ax_node("textbox", "Email", 3),
make_ax_node("heading", "Welcome", 4), // not interactive
make_ax_node("generic", "", 5), // skip role
];
let mut ref_map = ElementRefMap::new();
let repr = build_page_repr(
"https://example.com",
"Test Page",
&nodes,
ElementFilter::Interactive,
&mut ref_map,
);
assert!(repr.contains("@e1: button \"Submit\""));
assert!(repr.contains("@e2: link \"Home\""));
assert!(repr.contains("@e3: textbox \"Email\""));
assert!(!repr.contains("heading"));
assert!(!repr.contains("generic"));
assert_eq!(ref_map.len(), 3);
}
#[test]
fn test_build_page_repr_all_filter() {
let nodes = vec![
make_ax_node("button", "Submit", 1),
make_ax_node("heading", "Welcome", 2),
];
let mut ref_map = ElementRefMap::new();
let repr = build_page_repr(
"https://example.com",
"",
&nodes,
ElementFilter::All,
&mut ref_map,
);
assert!(repr.contains("button"));
assert!(repr.contains("heading"));
assert_eq!(ref_map.len(), 2);
}
#[test]
fn test_element_ref_lookup() {
let mut ref_map = ElementRefMap::new();
let nodes = vec![make_ax_node("button", "Click me", 1)];
build_page_repr(
"https://x.com",
"",
&nodes,
ElementFilter::Interactive,
&mut ref_map,
);
assert!(ref_map.get("e1").is_some());
assert!(ref_map.get("@e1").is_some()); // with @ prefix
assert!(ref_map.get("e99").is_none());
}
#[test]
fn test_empty_page() {
let mut ref_map = ElementRefMap::new();
let repr = build_page_repr(
"https://empty.com",
"",
&[],
ElementFilter::Interactive,
&mut ref_map,
);
assert!(repr.contains("no interactive elements"));
assert!(ref_map.is_empty());
}
#[test]
fn test_truncate_name() {
assert_eq!(truncate_name("short", 10), "short");
assert_eq!(truncate_name("this is a very long name", 10), "this is...");
}
}
-517
View File
@@ -1,517 +0,0 @@
//! Headless browser tool for web interaction.
//!
//! A single `BrowserTool` that dispatches actions via a tagged enum,
//! keeping the tool registry clean (one tool, not ten). The LLM sends
//! an `action` field to pick the operation:
//!
//! ```json
//! { "action": "navigate", "url": "https://example.com" }
//! { "action": "click", "ref": "@e3" }
//! { "action": "type", "ref": "@e1", "text": "hello" }
//! { "action": "read_page" }
//! { "action": "screenshot" }
//! ```
//!
//! Element references (`@e1`, `@e2`, ...) are assigned by `read_page`
//! and remain valid until the next `read_page` call.
pub mod accessibility;
pub mod session;
pub mod stealth;
use std::time::Duration;
use async_trait::async_trait;
use serde::Deserialize;
use tokio::sync::RwLock;
use crate::context::JobContext;
use crate::tools::builtin::browser::accessibility::ElementFilter;
use crate::tools::builtin::browser::session::BrowserSession;
use crate::tools::tool::{Tool, ToolError, ToolOutput};
/// Actions the LLM can request from the browser tool.
///
/// Uses serde tagged enum: the JSON `"action"` field selects the variant,
/// remaining fields are variant-specific parameters.
#[derive(Debug, Deserialize)]
#[serde(tag = "action", rename_all = "snake_case")]
enum BrowserAction {
/// Navigate to a URL.
Navigate { url: String },
/// Go back in browser history.
Back,
/// Go forward in browser history.
Forward,
/// Read the page's accessibility tree (assigns element refs).
ReadPage {
/// "interactive" (default) or "all"
filter: Option<String>,
},
/// Click an element by reference ID.
Click {
/// Element reference like "@e1" or "e1".
#[serde(alias = "ref")]
ref_id: String,
},
/// Type text into an element by reference ID.
Type {
/// Element reference like "@e1" or "e1".
#[serde(alias = "ref")]
ref_id: String,
text: String,
},
/// Scroll the page.
Scroll {
/// "up", "down", "left", "right"
direction: String,
/// Number of scroll steps (default 3).
amount: Option<u32>,
},
/// Capture a screenshot (returns base64 PNG).
Screenshot {
/// Capture full scrollable page (default false).
full_page: Option<bool>,
},
/// Extract text content from the page or a CSS selector.
Extract {
/// Optional CSS selector. If omitted, extracts all body text.
selector: Option<String>,
},
/// Wait for a CSS selector to appear or a fixed delay.
Wait {
/// CSS selector to wait for. If omitted, just sleeps.
selector: Option<String>,
/// Timeout in milliseconds (default 5000).
timeout_ms: Option<u64>,
},
/// Execute JavaScript (requires user approval).
EvalJs { expression: String },
}
/// Headless browser tool for navigating web pages, interacting with
/// elements, and extracting content.
///
/// Uses Chrome/Chromium via the DevTools Protocol. The browser is launched
/// lazily on first use and includes basic anti-detection patches.
///
/// ## Workflow
///
/// 1. `navigate` to a URL
/// 2. `read_page` to get the accessibility tree with element refs
/// 3. `click` / `type` using the refs
/// 4. `extract` or `screenshot` to get results
///
/// Element refs (`@e1`, `@e2`) are valid until the next `read_page`.
pub struct BrowserTool {
/// Lazily initialized browser session. RwLock because `execute` takes `&self`.
session: RwLock<Option<BrowserSession>>,
}
impl BrowserTool {
pub fn new() -> Self {
Self {
session: RwLock::new(None),
}
}
/// Ensure the browser session is initialized, launching Chrome if needed.
async fn ensure_session(&self) -> Result<(), ToolError> {
let needs_launch = self.session.read().await.is_none();
if needs_launch {
let new_session = BrowserSession::launch().await?;
let mut guard = self.session.write().await;
if guard.is_none() {
*guard = Some(new_session);
}
}
Ok(())
}
}
impl Default for BrowserTool {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl Tool for BrowserTool {
fn name(&self) -> &str {
"browser"
}
fn description(&self) -> &str {
"Control a headless web browser. Navigate pages, read content, click elements, type text, \
take screenshots. Use 'read_page' to get an accessibility tree with element references \
(@e1, @e2...), then use those refs for 'click' and 'type' actions.\n\n\
Actions: navigate, back, forward, read_page, click, type, scroll, screenshot, extract, \
wait, eval_js"
}
fn parameters_schema(&self) -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": {
"action": {
"type": "string",
"enum": [
"navigate", "back", "forward", "read_page", "click",
"type", "scroll", "screenshot", "extract", "wait", "eval_js"
],
"description": "The browser action to perform"
},
"url": {
"type": "string",
"description": "URL to navigate to (for 'navigate' action)"
},
"ref_id": {
"type": "string",
"description": "Element reference like '@e1' (for 'click' and 'type' actions)"
},
"text": {
"type": "string",
"description": "Text to type (for 'type' action)"
},
"direction": {
"type": "string",
"enum": ["up", "down", "left", "right"],
"description": "Scroll direction (for 'scroll' action)"
},
"amount": {
"type": "integer",
"description": "Scroll steps, default 3 (for 'scroll' action)"
},
"full_page": {
"type": "boolean",
"description": "Capture full scrollable page (for 'screenshot' action)"
},
"selector": {
"type": "string",
"description": "CSS selector (for 'extract' and 'wait' actions)"
},
"timeout_ms": {
"type": "integer",
"description": "Timeout in milliseconds (for 'wait' action, default 5000)"
},
"filter": {
"type": "string",
"enum": ["interactive", "all"],
"description": "Element filter for 'read_page' (default: interactive)"
},
"expression": {
"type": "string",
"description": "JavaScript expression (for 'eval_js' action)"
}
},
"required": ["action"]
})
}
async fn execute(
&self,
params: serde_json::Value,
_ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let action: BrowserAction = serde_json::from_value(params)
.map_err(|e| ToolError::InvalidParameters(format!("Invalid browser action: {}", e)))?;
// Launch browser on first use.
self.ensure_session().await?;
match action {
BrowserAction::Navigate { url } => {
let session = self.session.read().await;
let session = session.as_ref().ok_or_else(|| {
ToolError::ExecutionFailed("Browser session not initialized".to_string())
})?;
let title = session.navigate(&url).await?;
let current_url = session.current_url().await?;
Ok(ToolOutput::success(
serde_json::json!({
"url": current_url,
"title": title,
"status": "navigated"
}),
start.elapsed(),
))
}
BrowserAction::Back => {
let session = self.session.read().await;
let session = session.as_ref().ok_or_else(|| {
ToolError::ExecutionFailed("Browser session not initialized".to_string())
})?;
session.go_back().await?;
let url = session.current_url().await?;
Ok(ToolOutput::success(
serde_json::json!({ "url": url, "status": "navigated_back" }),
start.elapsed(),
))
}
BrowserAction::Forward => {
let session = self.session.read().await;
let session = session.as_ref().ok_or_else(|| {
ToolError::ExecutionFailed("Browser session not initialized".to_string())
})?;
session.go_forward().await?;
let url = session.current_url().await?;
Ok(ToolOutput::success(
serde_json::json!({ "url": url, "status": "navigated_forward" }),
start.elapsed(),
))
}
BrowserAction::ReadPage { filter } => {
let session = self.session.read().await;
let session = session.as_ref().ok_or_else(|| {
ToolError::ExecutionFailed("Browser session not initialized".to_string())
})?;
let element_filter = ElementFilter::from_str_opt(filter.as_deref());
let repr = session.read_page(element_filter).await?;
Ok(ToolOutput::text(repr, start.elapsed()))
}
BrowserAction::Click { ref_id } => {
let session = self.session.read().await;
let session = session.as_ref().ok_or_else(|| {
ToolError::ExecutionFailed("Browser session not initialized".to_string())
})?;
session.click_element(&ref_id).await?;
Ok(ToolOutput::success(
serde_json::json!({ "status": "clicked", "ref": ref_id }),
start.elapsed(),
))
}
BrowserAction::Type { ref_id, text } => {
let session = self.session.read().await;
let session = session.as_ref().ok_or_else(|| {
ToolError::ExecutionFailed("Browser session not initialized".to_string())
})?;
session.type_text(&ref_id, &text).await?;
Ok(ToolOutput::success(
serde_json::json!({
"status": "typed",
"ref": ref_id,
"length": text.len()
}),
start.elapsed(),
))
}
BrowserAction::Scroll { direction, amount } => {
let session = self.session.read().await;
let session = session.as_ref().ok_or_else(|| {
ToolError::ExecutionFailed("Browser session not initialized".to_string())
})?;
let steps = amount.unwrap_or(3);
session.scroll(&direction, steps).await?;
Ok(ToolOutput::success(
serde_json::json!({
"status": "scrolled",
"direction": direction,
"amount": steps
}),
start.elapsed(),
))
}
BrowserAction::Screenshot { full_page } => {
let session = self.session.read().await;
let session = session.as_ref().ok_or_else(|| {
ToolError::ExecutionFailed("Browser session not initialized".to_string())
})?;
let b64 = session.screenshot(full_page.unwrap_or(false)).await?;
Ok(ToolOutput::success(
serde_json::json!({
"format": "png",
"encoding": "base64",
"data": b64,
"full_page": full_page.unwrap_or(false)
}),
start.elapsed(),
))
}
BrowserAction::Extract { selector } => {
let session = self.session.read().await;
let session = session.as_ref().ok_or_else(|| {
ToolError::ExecutionFailed("Browser session not initialized".to_string())
})?;
let text = session.extract_text(selector.as_deref()).await?;
// Truncate very long text to avoid blowing up context.
let truncated = if text.len() > 32_000 {
format!(
"{}...\n\n[truncated, {} total chars]",
&text[..32_000],
text.len()
)
} else {
text.clone()
};
Ok(ToolOutput::text(&truncated, start.elapsed()).with_raw(text))
}
BrowserAction::Wait {
selector,
timeout_ms,
} => {
let session = self.session.read().await;
let session = session.as_ref().ok_or_else(|| {
ToolError::ExecutionFailed("Browser session not initialized".to_string())
})?;
let timeout = timeout_ms.unwrap_or(5000);
let found = session.wait(selector.as_deref(), timeout).await?;
Ok(ToolOutput::success(
serde_json::json!({
"found": found,
"selector": selector,
"timeout_ms": timeout
}),
start.elapsed(),
))
}
BrowserAction::EvalJs { expression } => {
let session = self.session.read().await;
let session = session.as_ref().ok_or_else(|| {
ToolError::ExecutionFailed("Browser session not initialized".to_string())
})?;
let result = session.eval_js(&expression).await?;
Ok(ToolOutput::success(
serde_json::json!({ "result": result }),
start.elapsed(),
))
}
}
}
fn estimated_duration(&self, _params: &serde_json::Value) -> Option<Duration> {
Some(Duration::from_secs(10))
}
fn requires_sanitization(&self) -> bool {
true // Page content is untrusted external data
}
fn requires_approval(&self) -> bool {
true // Browser navigates to external sites, executes JS
}
}
#[cfg(test)]
mod tests {
use crate::tools::builtin::browser::BrowserTool;
use crate::tools::tool::Tool;
#[test]
fn test_browser_tool_metadata() {
let tool = BrowserTool::new();
assert_eq!(tool.name(), "browser");
assert!(tool.requires_approval());
assert!(tool.requires_sanitization());
}
#[test]
fn test_schema_has_action_enum() {
let tool = BrowserTool::new();
let schema = tool.parameters_schema();
let action_prop = schema.get("properties").and_then(|p| p.get("action"));
assert!(action_prop.is_some());
let action_enum = action_prop.and_then(|a| a.get("enum"));
assert!(action_enum.is_some());
let actions: Vec<&str> = action_enum
.and_then(|e| e.as_array())
.map(|arr| arr.iter().filter_map(|v| v.as_str()).collect())
.unwrap_or_default();
assert!(actions.contains(&"navigate"));
assert!(actions.contains(&"click"));
assert!(actions.contains(&"type"));
assert!(actions.contains(&"read_page"));
assert!(actions.contains(&"screenshot"));
assert!(actions.contains(&"eval_js"));
}
#[test]
fn test_action_deserialization() {
use super::BrowserAction;
// Navigate
let action: BrowserAction = serde_json::from_value(
serde_json::json!({"action": "navigate", "url": "https://x.com"}),
)
.unwrap();
assert!(matches!(action, BrowserAction::Navigate { url } if url == "https://x.com"));
// Click with "ref" alias
let action: BrowserAction =
serde_json::from_value(serde_json::json!({"action": "click", "ref": "@e1"})).unwrap();
assert!(matches!(action, BrowserAction::Click { ref_id } if ref_id == "@e1"));
// Click with "ref_id"
let action: BrowserAction =
serde_json::from_value(serde_json::json!({"action": "click", "ref_id": "e2"})).unwrap();
assert!(matches!(action, BrowserAction::Click { ref_id } if ref_id == "e2"));
// Type
let action: BrowserAction = serde_json::from_value(
serde_json::json!({"action": "type", "ref": "@e1", "text": "hello"}),
)
.unwrap();
assert!(
matches!(action, BrowserAction::Type { ref_id, text } if ref_id == "@e1" && text == "hello")
);
// ReadPage with default filter
let action: BrowserAction =
serde_json::from_value(serde_json::json!({"action": "read_page"})).unwrap();
assert!(matches!(action, BrowserAction::ReadPage { filter: None }));
// Screenshot
let action: BrowserAction =
serde_json::from_value(serde_json::json!({"action": "screenshot", "full_page": true}))
.unwrap();
assert!(matches!(
action,
BrowserAction::Screenshot {
full_page: Some(true)
}
));
// Invalid action
let result: Result<BrowserAction, _> =
serde_json::from_value(serde_json::json!({"action": "fly_to_moon"}));
assert!(result.is_err());
}
}
-587
View File
@@ -1,587 +0,0 @@
//! Browser session management.
//!
//! Owns the Chrome process lifecycle and per-tab state. Sessions are spawned
//! lazily on first browser action and torn down when dropped.
//!
//! ```text
//! BrowserSession
//! ├── Browser (chromiumoxide, owns Chrome child process)
//! ├── handler_task (JoinHandle polling CDP WebSocket)
//! ├── tabs: HashMap<tab_id, Page>
//! ├── active_tab: current tab id
//! └── element_refs: ElementRefMap (valid until next read_page)
//! ```
use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::Arc;
use chromiumoxide::Page;
use chromiumoxide::browser::{Browser, BrowserConfig};
use chromiumoxide::cdp::browser_protocol::accessibility::GetFullAxTreeParams;
use chromiumoxide::cdp::browser_protocol::dom::{GetBoxModelParams, ScrollIntoViewIfNeededParams};
use chromiumoxide::cdp::browser_protocol::input::{
DispatchMouseEventParams, DispatchMouseEventType, InsertTextParams, MouseButton,
};
use chromiumoxide::cdp::browser_protocol::page::CaptureScreenshotFormat;
use chromiumoxide::page::ScreenshotParams;
use futures::StreamExt;
use tokio::sync::RwLock;
use tokio::task::JoinHandle;
use crate::tools::builtin::browser::accessibility::{
ElementFilter, ElementRefMap, build_page_repr,
};
use crate::tools::builtin::browser::stealth;
use crate::tools::tool::ToolError;
/// Manages a Chrome browser instance and its tabs.
pub struct BrowserSession {
#[allow(dead_code)] // Used by new_tab() which is reserved for tab management actions
browser: Browser,
_handler_task: JoinHandle<()>,
tabs: HashMap<String, Page>,
active_tab: String,
element_refs: Arc<RwLock<ElementRefMap>>,
#[allow(dead_code)] // Used by new_tab() which is reserved for tab management actions
stealth_js: String,
}
impl BrowserSession {
/// Launch a new Chrome browser session.
///
/// Locates Chrome on the system, applies stealth patches, and opens
/// an initial blank tab.
pub async fn launch() -> Result<Self, ToolError> {
let chrome_path = find_chrome().ok_or_else(|| {
ToolError::ExecutionFailed(
"Chrome/Chromium not found. Install Chrome or set CHROME_PATH.".to_string(),
)
})?;
// Shared profile so the agent accumulates useful state across sessions
// (logged-in sessions, dismissed cookie banners, local storage).
// Delete ~/.ironclaw/browser/profile/ to reset.
let profile_dir = browser_profile_dir();
std::fs::create_dir_all(&profile_dir).map_err(|e| {
ToolError::ExecutionFailed(format!("Failed to create browser profile dir: {}", e))
})?;
let mut config_builder = BrowserConfig::builder()
.chrome_executable(&chrome_path)
.user_data_dir(&profile_dir)
.window_size(1920, 1080)
.no_sandbox();
for arg in stealth::stealth_args() {
config_builder = config_builder.arg(arg);
}
let config = config_builder.build().map_err(|e| {
ToolError::ExecutionFailed(format!("Failed to build browser config: {}", e))
})?;
let (browser, mut handler) = Browser::launch(config)
.await
.map_err(|e| ToolError::ExecutionFailed(format!("Failed to launch Chrome: {}", e)))?;
// The handler must be polled continuously or the CDP connection dies.
let handler_task = tokio::spawn(async move {
while let Some(event) = handler.next().await {
if event.is_err() {
tracing::warn!("Browser handler error: {:?}", event);
break;
}
}
});
// Open initial tab.
let page = browser.new_page("about:blank").await.map_err(|e| {
ToolError::ExecutionFailed(format!("Failed to open initial tab: {}", e))
})?;
// Inject stealth JS on every new document load for this page.
let stealth_js = stealth::stealth_js().to_string();
page.evaluate_on_new_document(stealth_js.clone())
.await
.map_err(|e| {
ToolError::ExecutionFailed(format!("Failed to inject stealth JS: {}", e))
})?;
let tab_id = "tab0".to_string();
let mut tabs = HashMap::new();
tabs.insert(tab_id.clone(), page);
Ok(Self {
browser,
_handler_task: handler_task,
tabs,
active_tab: tab_id,
element_refs: Arc::new(RwLock::new(ElementRefMap::new())),
stealth_js,
})
}
/// Get the active page, or error if session is broken.
fn active_page(&self) -> Result<&Page, ToolError> {
self.tabs.get(&self.active_tab).ok_or_else(|| {
ToolError::ExecutionFailed(format!("No active tab: {}", self.active_tab))
})
}
// --- Navigation ---
pub async fn navigate(&self, url: &str) -> Result<String, ToolError> {
let page = self.active_page()?;
page.goto(url)
.await
.map_err(|e| ToolError::ExternalService(format!("Navigation failed: {}", e)))?;
let title = page
.get_title()
.await
.map_err(|e| ToolError::ExecutionFailed(format!("Failed to get page title: {}", e)))?
.unwrap_or_default();
Ok(title)
}
pub async fn go_back(&self) -> Result<(), ToolError> {
let page = self.active_page()?;
page.evaluate("window.history.back()")
.await
.map_err(|e| ToolError::ExecutionFailed(format!("Failed to go back: {}", e)))?;
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
Ok(())
}
pub async fn go_forward(&self) -> Result<(), ToolError> {
let page = self.active_page()?;
page.evaluate("window.history.forward()")
.await
.map_err(|e| ToolError::ExecutionFailed(format!("Failed to go forward: {}", e)))?;
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
Ok(())
}
// --- Page reading ---
/// Build accessibility tree representation and update element refs.
pub async fn read_page(&self, filter: ElementFilter) -> Result<String, ToolError> {
let page = self.active_page()?;
let url = page
.url()
.await
.map_err(|e| ToolError::ExecutionFailed(format!("Failed to get URL: {}", e)))?
.unwrap_or_else(|| "about:blank".to_string());
let title = page
.get_title()
.await
.map_err(|e| ToolError::ExecutionFailed(format!("Failed to get title: {}", e)))?
.unwrap_or_default();
// Fetch full accessibility tree via CDP.
let ax_result = page
.execute(GetFullAxTreeParams::default())
.await
.map_err(|e| {
ToolError::ExecutionFailed(format!("Failed to get accessibility tree: {}", e))
})?;
let nodes = ax_result.result.nodes;
let mut ref_map = self.element_refs.write().await;
let repr = build_page_repr(&url, &title, &nodes, filter, &mut ref_map);
Ok(repr)
}
/// Extract text content from the page or a CSS selector.
pub async fn extract_text(&self, selector: Option<&str>) -> Result<String, ToolError> {
let page = self.active_page()?;
let js = match selector {
Some(sel) => {
let escaped = serde_json::to_string(sel).map_err(|e| {
ToolError::InvalidParameters(format!("Invalid selector: {}", e))
})?;
format!(
"(() => {{ const el = document.querySelector({}); return el ? el.innerText : null; }})()",
escaped
)
}
None => "document.body.innerText".to_string(),
};
let result: Option<String> = page
.evaluate(js.as_str())
.await
.map_err(|e| ToolError::ExecutionFailed(format!("Failed to extract text: {}", e)))?
.into_value()
.map_err(|e| {
ToolError::ExecutionFailed(format!("Failed to deserialize text: {}", e))
})?;
Ok(result.unwrap_or_default())
}
// --- Interaction ---
/// Click an element by reference ID (e.g., "e1" or "@e1").
///
/// Uses DOM.scrollIntoViewIfNeeded + DOM.getBoxModel to find the element's
/// center coordinates, then dispatches mouse press + release at that point.
pub async fn click_element(&self, ref_id: &str) -> Result<(), ToolError> {
let page = self.active_page()?;
let refs = self.element_refs.read().await;
let elem_ref = refs.get(ref_id).ok_or_else(|| {
ToolError::InvalidParameters(format!(
"Unknown element reference '{}'. Call browser with action 'read_page' first.",
ref_id
))
})?;
let backend_node_id = elem_ref.backend_node_id;
drop(refs);
// Scroll the element into the viewport.
page.execute(
ScrollIntoViewIfNeededParams::builder()
.backend_node_id(backend_node_id)
.build(),
)
.await
.map_err(|e| {
ToolError::ExecutionFailed(format!("Failed to scroll element into view: {}", e))
})?;
// Get element's bounding box via DOM.getBoxModel.
let box_result = page
.execute(
GetBoxModelParams::builder()
.backend_node_id(backend_node_id)
.build(),
)
.await
.map_err(|e| {
ToolError::ExecutionFailed(format!("Failed to get element box model: {}", e))
})?;
// Content quad is [x1,y1, x2,y2, x3,y3, x4,y4]. Center = average of 4 corners.
let content = box_result.result.model.content.inner();
if content.len() < 8 {
return Err(ToolError::ExecutionFailed(
"Element has no valid bounding box".to_string(),
));
}
let x = (content[0] + content[2] + content[4] + content[6]) / 4.0;
let y = (content[1] + content[3] + content[5] + content[7]) / 4.0;
// Dispatch mouse press + release at center of element.
page.execute(
DispatchMouseEventParams::builder()
.r#type(DispatchMouseEventType::MousePressed)
.x(x)
.y(y)
.button(MouseButton::Left)
.click_count(1)
.build()
.map_err(|e| {
ToolError::ExecutionFailed(format!("Failed to build mouse event: {}", e))
})?,
)
.await
.map_err(|e| ToolError::ExecutionFailed(format!("Mouse press failed: {}", e)))?;
page.execute(
DispatchMouseEventParams::builder()
.r#type(DispatchMouseEventType::MouseReleased)
.x(x)
.y(y)
.button(MouseButton::Left)
.click_count(1)
.build()
.map_err(|e| {
ToolError::ExecutionFailed(format!("Failed to build mouse event: {}", e))
})?,
)
.await
.map_err(|e| ToolError::ExecutionFailed(format!("Mouse release failed: {}", e)))?;
Ok(())
}
/// Type text into an element by reference ID.
pub async fn type_text(&self, ref_id: &str, text: &str) -> Result<(), ToolError> {
// First click to focus the element.
self.click_element(ref_id).await?;
// Brief delay to let focus settle.
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let page = self.active_page()?;
// Use CDP insertText for reliable IME-style text entry.
page.execute(InsertTextParams::new(text))
.await
.map_err(|e| ToolError::ExecutionFailed(format!("Failed to type text: {}", e)))?;
Ok(())
}
/// Scroll the page.
pub async fn scroll(&self, direction: &str, amount: u32) -> Result<(), ToolError> {
let page = self.active_page()?;
let (dx, dy) = match direction {
"up" => (0, -(amount as i32 * 100)),
"down" => (0, amount as i32 * 100),
"left" => (-(amount as i32 * 100), 0),
"right" => (amount as i32 * 100, 0),
_ => {
return Err(ToolError::InvalidParameters(format!(
"Invalid scroll direction '{}'. Use: up, down, left, right",
direction
)));
}
};
let js = format!("window.scrollBy({}, {})", dx, dy);
page.evaluate(js.as_str())
.await
.map_err(|e| ToolError::ExecutionFailed(format!("Scroll failed: {}", e)))?;
Ok(())
}
/// Wait for a CSS selector to appear, or a fixed timeout.
pub async fn wait(&self, selector: Option<&str>, timeout_ms: u64) -> Result<bool, ToolError> {
let page = self.active_page()?;
let timeout = std::time::Duration::from_millis(timeout_ms);
match selector {
Some(sel) => {
let poll_interval = std::time::Duration::from_millis(100);
let start = std::time::Instant::now();
let escaped = serde_json::to_string(sel).map_err(|e| {
ToolError::InvalidParameters(format!("Invalid selector: {}", e))
})?;
loop {
let js = format!("!!document.querySelector({})", escaped);
let found: bool = page
.evaluate(js.as_str())
.await
.map_err(|e| {
ToolError::ExecutionFailed(format!("Wait poll failed: {}", e))
})?
.into_value()
.unwrap_or(false);
if found {
return Ok(true);
}
if start.elapsed() >= timeout {
return Ok(false);
}
tokio::time::sleep(poll_interval).await;
}
}
None => {
tokio::time::sleep(timeout).await;
Ok(true)
}
}
}
// --- Screenshots ---
/// Capture a screenshot as base64-encoded PNG.
pub async fn screenshot(&self, full_page: bool) -> Result<String, ToolError> {
let page = self.active_page()?;
let params = ScreenshotParams::builder()
.format(CaptureScreenshotFormat::Png)
.full_page(full_page)
.build();
let bytes = page
.screenshot(params)
.await
.map_err(|e| ToolError::ExecutionFailed(format!("Screenshot failed: {}", e)))?;
use base64::Engine;
Ok(base64::engine::general_purpose::STANDARD.encode(&bytes))
}
// --- JavaScript ---
/// Execute arbitrary JavaScript and return the result.
pub async fn eval_js(&self, expression: &str) -> Result<serde_json::Value, ToolError> {
let page = self.active_page()?;
let result = page
.evaluate(expression)
.await
.map_err(|e| ToolError::ExecutionFailed(format!("JS evaluation failed: {}", e)))?;
let value: serde_json::Value = result.into_value().unwrap_or(serde_json::Value::Null);
Ok(value)
}
// --- Tab management ---
/// Open a new tab and make it active.
#[allow(dead_code)] // Reserved for tab management actions
pub async fn new_tab(&mut self, url: &str) -> Result<String, ToolError> {
let page =
self.browser.new_page(url).await.map_err(|e| {
ToolError::ExecutionFailed(format!("Failed to open new tab: {}", e))
})?;
// Inject stealth JS on the new page too.
page.evaluate_on_new_document(self.stealth_js.clone())
.await
.map_err(|e| {
ToolError::ExecutionFailed(format!("Failed to inject stealth JS on new tab: {}", e))
})?;
let tab_id = format!("tab{}", self.tabs.len());
self.tabs.insert(tab_id.clone(), page);
self.active_tab = tab_id.clone();
// Clear element refs since we're on a new page.
self.element_refs.write().await.reset();
Ok(tab_id)
}
/// List open tabs.
#[allow(dead_code)] // Reserved for tab management actions
pub fn list_tabs(&self) -> Vec<String> {
self.tabs.keys().cloned().collect()
}
/// Switch to a different tab.
#[allow(dead_code)] // Reserved for tab management actions
pub async fn switch_tab(&mut self, tab_id: &str) -> Result<(), ToolError> {
if !self.tabs.contains_key(tab_id) {
return Err(ToolError::InvalidParameters(format!(
"Unknown tab '{}'. Open tabs: {:?}",
tab_id,
self.list_tabs()
)));
}
self.active_tab = tab_id.to_string();
// Clear element refs when switching tabs.
self.element_refs.write().await.reset();
Ok(())
}
/// Get current page URL.
pub async fn current_url(&self) -> Result<String, ToolError> {
let page = self.active_page()?;
page.url()
.await
.map_err(|e| ToolError::ExecutionFailed(format!("Failed to get URL: {}", e)))
.map(|u| u.unwrap_or_else(|| "about:blank".to_string()))
}
}
impl Drop for BrowserSession {
fn drop(&mut self) {
tracing::debug!("Browser session dropping, Chrome process will be cleaned up");
}
}
/// Returns `~/.ironclaw/browser/profile/`.
fn browser_profile_dir() -> PathBuf {
dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".ironclaw")
.join("browser")
.join("profile")
}
/// Search common locations for a Chrome/Chromium binary.
pub fn find_chrome() -> Option<PathBuf> {
// Environment variable override.
if let Ok(path) = std::env::var("CHROME_PATH") {
let p = PathBuf::from(&path);
if p.exists() {
return Some(p);
}
}
let candidates = if cfg!(target_os = "macos") {
vec![
"/Applications/Google Chrome.app/Contents/MacOS/Google Chrome",
"/Applications/Chromium.app/Contents/MacOS/Chromium",
"/Applications/Google Chrome Canary.app/Contents/MacOS/Google Chrome Canary",
"/Applications/Brave Browser.app/Contents/MacOS/Brave Browser",
]
} else if cfg!(target_os = "linux") {
vec![
"/usr/bin/google-chrome",
"/usr/bin/google-chrome-stable",
"/usr/bin/chromium",
"/usr/bin/chromium-browser",
"/snap/bin/chromium",
]
} else {
// Windows paths.
vec![
r"C:\Program Files\Google\Chrome\Application\chrome.exe",
r"C:\Program Files (x86)\Google\Chrome\Application\chrome.exe",
]
};
for candidate in candidates {
let p = PathBuf::from(candidate);
if p.exists() {
return Some(p);
}
}
which_chrome_in_path()
}
/// Check if chrome/chromium is available in PATH.
fn which_chrome_in_path() -> Option<PathBuf> {
let path_var = std::env::var("PATH").ok()?;
let separator = if cfg!(windows) { ';' } else { ':' };
for name in &["google-chrome", "chromium", "chromium-browser", "chrome"] {
for dir in path_var.split(separator) {
let candidate = PathBuf::from(dir).join(name);
if candidate.exists() {
return Some(candidate);
}
}
}
None
}
#[cfg(test)]
mod tests {
use crate::tools::builtin::browser::session::find_chrome;
#[test]
fn test_find_chrome_returns_path_or_none() {
let result = find_chrome();
if let Some(path) = &result {
assert!(
path.exists(),
"find_chrome returned non-existent path: {:?}",
path
);
}
}
}
-158
View File
@@ -1,158 +0,0 @@
//! Anti-detection JavaScript patches for headless Chrome.
//!
//! Injects scripts via `Page.addScriptToEvaluateOnNewDocument` to suppress
//! common bot-detection signals. Handles ~80% of detection for legitimate
//! browsing (not adversarial scraping against Cloudflare Enterprise).
//!
//! What we patch:
//! - `navigator.webdriver` (trivial but still checked)
//! - `navigator.plugins` (headless has empty plugin list)
//! - `navigator.languages` (match system locale)
//! - `chrome.runtime` (looks like a real extension API)
//! - `HeadlessChrome` user-agent substring (suppressed via launch flags)
/// Chrome launch arguments that reduce detection surface.
pub fn stealth_args() -> Vec<&'static str> {
vec![
"--disable-blink-features=AutomationControlled",
"--no-first-run",
"--no-default-browser-check",
"--disable-infobars",
"--disable-background-networking",
"--disable-prompt-on-repost",
"--disable-hang-monitor",
"--disable-sync",
"--metrics-recording-only",
"--no-service-autorun",
]
}
/// JavaScript injected before any page scripts run.
///
/// This covers the most common fingerprinting checks. Each patch is
/// a self-contained IIFE so failures in one don't break the others.
pub fn stealth_js() -> &'static str {
r#"
// --- navigator.webdriver ---
// CDP sets this to true; real browsers have it undefined or false.
(() => {
Object.defineProperty(navigator, 'webdriver', {
get: () => undefined,
configurable: true,
});
})();
// --- navigator.plugins ---
// Headless Chrome reports an empty plugin array. Real Chrome on desktop
// always has at least these two. We fake the array shape.
(() => {
const pluginData = [
{ name: 'Chrome PDF Plugin', filename: 'internal-pdf-viewer',
description: 'Portable Document Format' },
{ name: 'Chrome PDF Viewer', filename: 'mhjfbmdgcfjbbpaeojofohoefgiehjai',
description: '' },
];
const makeMimeType = (type_, suffixes, desc, plugin) => {
const mt = Object.create(MimeType.prototype);
Object.defineProperties(mt, {
type: { get: () => type_ },
suffixes: { get: () => suffixes },
description: { get: () => desc },
enabledPlugin: { get: () => plugin },
});
return mt;
};
const makePlugin = (data) => {
const p = Object.create(Plugin.prototype);
const mimes = [makeMimeType('application/pdf', 'pdf', 'Portable Document Format', p)];
Object.defineProperties(p, {
name: { get: () => data.name },
filename: { get: () => data.filename },
description: { get: () => data.description },
length: { get: () => mimes.length },
0: { get: () => mimes[0] },
});
p.item = (i) => mimes[i] || null;
p.namedItem = (name) => mimes.find(m => m.type === name) || null;
return p;
};
const plugins = pluginData.map(makePlugin);
const pluginArray = Object.create(PluginArray.prototype);
Object.defineProperties(pluginArray, {
length: { get: () => plugins.length },
0: { get: () => plugins[0] },
1: { get: () => plugins[1] },
});
pluginArray.item = (i) => plugins[i] || null;
pluginArray.namedItem = (name) => plugins.find(p => p.name === name) || null;
pluginArray.refresh = () => {};
pluginArray[Symbol.iterator] = function* () { yield* plugins; };
Object.defineProperty(navigator, 'plugins', {
get: () => pluginArray,
configurable: true,
});
})();
// --- navigator.languages ---
// Headless sometimes reports just ['en'] instead of a realistic list.
(() => {
Object.defineProperty(navigator, 'languages', {
get: () => ['en-US', 'en'],
configurable: true,
});
})();
// --- chrome.runtime ---
// Bot detectors check for chrome.runtime to see if it's a real Chrome
// extension environment. CDP-controlled Chrome has a broken stub.
(() => {
if (!window.chrome) window.chrome = {};
if (!window.chrome.runtime) {
window.chrome.runtime = {
connect: () => {},
sendMessage: () => {},
id: undefined,
};
}
})();
// --- Permissions API ---
// Headless reports 'denied' for notification permissions by default,
// which is a known fingerprinting signal.
(() => {
const originalQuery = window.Permissions?.prototype?.query;
if (originalQuery) {
window.Permissions.prototype.query = function(params) {
if (params?.name === 'notifications') {
return Promise.resolve({ state: 'prompt', onchange: null });
}
return originalQuery.call(this, params);
};
}
})();
"#
}
#[cfg(test)]
mod tests {
use crate::tools::builtin::browser::stealth;
#[test]
fn stealth_js_is_not_empty() {
let js = stealth::stealth_js();
assert!(js.len() > 100);
assert!(js.contains("navigator"));
assert!(js.contains("webdriver"));
}
#[test]
fn stealth_args_are_valid_flags() {
for arg in stealth::stealth_args() {
assert!(arg.starts_with("--"), "arg should start with --: {}", arg);
}
}
}
+45 -10
View File
@@ -5,13 +5,18 @@ 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};
/// Maximum response body size (5 MB). Prevents OOM from unbounded responses.
/// Maximum response body size (5 MB).
///
/// 5 MB is large enough for typical JSON API responses and moderate HTML pages,
/// but small enough to prevent OOM from malicious or runaway servers. The WASM
/// HTTP wrapper uses the same limit for consistency.
const MAX_RESPONSE_SIZE: usize = 5 * 1024 * 1024;
/// Tool for making HTTP requests.
@@ -230,19 +235,43 @@ impl Tool for HttpTool {
.filter_map(|(k, v)| v.to_str().ok().map(|v| (k.to_string(), v.to_string())))
.collect();
// 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 {
// 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"
);
return Err(ToolError::ExecutionFailed(format!(
"Response body too large ({} bytes, max {})",
body_bytes.len(),
MAX_RESPONSE_SIZE
"Response Content-Length ({} bytes) exceeds maximum allowed size ({} bytes)",
len, MAX_RESPONSE_SIZE
)));
}
// Stream the response body with a hard size cap. Even if Content-Length was
// absent or lied about the size, we stop reading once we exceed the limit.
let mut body = Vec::new();
let mut stream = response.bytes_stream();
while let Some(chunk) = StreamExt::next(&mut stream).await {
let chunk = chunk.map_err(|e| {
ToolError::ExternalService(format!("failed to read response body: {}", e))
})?;
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
@@ -328,4 +357,10 @@ 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);
}
}
-3
View File
@@ -1,6 +1,5 @@
//! Built-in tools that come with the agent.
mod browser;
mod echo;
pub mod extension_tools;
mod file;
@@ -12,8 +11,6 @@ pub mod routine;
pub(crate) mod shell;
mod time;
pub use browser::BrowserTool;
pub use browser::session::find_chrome;
pub use echo::EchoTool;
pub use extension_tools::{
ToolActivateTool, ToolAuthTool, ToolInstallTool, ToolListTool, ToolRemoveTool, ToolSearchTool,
+48
View File
@@ -426,6 +426,26 @@ 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
}
@@ -566,6 +586,34 @@ mod tests {
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()
+5 -6
View File
@@ -14,10 +14,10 @@ use crate::safety::SafetyLayer;
use crate::secrets::SecretsStore;
use crate::tools::builder::{BuildSoftwareTool, BuilderConfig, LlmSoftwareBuilder};
use crate::tools::builtin::{
ApplyPatchTool, BrowserTool, CancelJobTool, CreateJobTool, EchoTool, HttpTool, JobStatusTool,
JsonTool, ListDirTool, ListJobsTool, MemoryReadTool, MemorySearchTool, MemoryTreeTool,
MemoryWriteTool, ReadFileTool, ShellTool, TimeTool, ToolActivateTool, ToolAuthTool,
ToolInstallTool, ToolListTool, ToolRemoveTool, ToolSearchTool, WriteFileTool,
ApplyPatchTool, CancelJobTool, CreateJobTool, EchoTool, HttpTool, JobStatusTool, JsonTool,
ListDirTool, ListJobsTool, MemoryReadTool, MemorySearchTool, MemoryTreeTool, MemoryWriteTool,
ReadFileTool, ShellTool, TimeTool, ToolActivateTool, ToolAuthTool, ToolInstallTool,
ToolListTool, ToolRemoveTool, ToolSearchTool, WriteFileTool,
};
use crate::tools::tool::{Tool, ToolDomain};
use crate::tools::wasm::{
@@ -218,9 +218,8 @@ impl ToolRegistry {
self.register_sync(Arc::new(WriteFileTool::new()));
self.register_sync(Arc::new(ListDirTool::new()));
self.register_sync(Arc::new(ApplyPatchTool::new()));
self.register_sync(Arc::new(BrowserTool::new()));
tracing::info!("Registered 6 development tools (includes browser)");
tracing::info!("Registered 5 development tools");
}
/// Register memory tools with a workspace.
+22
View File
@@ -172,6 +172,21 @@ pub trait Tool: Send + Sync {
false
}
/// Whether this specific invocation should override auto-approval.
///
/// This method is called after checking `requires_approval()` and finding that
/// the tool is auto-approved for this session. Return `true` to force approval
/// for this specific invocation despite auto-approval (for example, for
/// destructive operations like `rm -rf` or `git push --force`).
///
/// Return `false` to allow auto-approval to proceed normally.
///
/// The default returns `false`. Override only if you need parameter-aware
/// approval gating.
fn requires_approval_for(&self, _params: &serde_json::Value) -> bool {
false
}
/// Maximum time this tool is allowed to run before the caller kills it.
/// Override for long-running tools like sandbox execution.
/// Default: 60 seconds.
@@ -330,4 +345,11 @@ mod tests {
let err = require_param(&params, "data").unwrap_err();
assert!(err.to_string().contains("missing 'data'"));
}
#[test]
fn test_requires_approval_for_default() {
let tool = EchoTool;
// Default requires_approval_for() returns false, allowing auto-approval.
assert!(!tool.requires_approval_for(&serde_json::json!({"message": "hi"})));
}
}
-203
View File
@@ -1,203 +0,0 @@
//! Integration test for the browser tool.
//!
//! Requires Chrome installed. Run with:
//! cargo test --test browser_integration -- --nocapture
use ironclaw::context::JobContext;
use ironclaw::tools::Tool;
use ironclaw::tools::builtin::{BrowserTool, find_chrome};
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn test_browser_navigate_and_screenshot() {
// Skip if Chrome/Chromium is not installed (works on macOS, Linux, Windows).
if find_chrome().is_none() {
eprintln!("Skipping: Chrome not found");
return;
}
let tool = BrowserTool::new();
let ctx = JobContext::default();
// 1. Navigate to Wikipedia
eprintln!("=== Navigating to Wikipedia...");
let nav_result = tool
.execute(
serde_json::json!({
"action": "navigate",
"url": "https://en.wikipedia.org/wiki/Mariam_Almheiri"
}),
&ctx,
)
.await;
match &nav_result {
Ok(output) => {
eprintln!(
"Navigation result: {}",
serde_json::to_string_pretty(&output.result).unwrap()
);
let title = output
.result
.get("title")
.and_then(|t| t.as_str())
.unwrap_or("");
assert!(
title.contains("Mariam") || title.contains("Almheiri"),
"Page title should mention Mariam Almheiri, got: {}",
title
);
}
Err(e) => {
eprintln!("Navigation failed: {}", e);
panic!("Navigation should succeed");
}
}
// 2. Read the accessibility tree
eprintln!("\n=== Reading page accessibility tree...");
let read_result = tool
.execute(serde_json::json!({"action": "read_page"}), &ctx)
.await;
match &read_result {
Ok(output) => {
let tree = output.result.as_str().unwrap_or("");
let line_count = tree.lines().count();
eprintln!("Accessibility tree: {} lines", line_count);
// Print first 20 lines
for line in tree.lines().take(20) {
eprintln!(" {}", line);
}
if line_count > 20 {
eprintln!(" ... ({} more lines)", line_count - 20);
}
assert!(line_count > 3, "Should have some elements on the page");
}
Err(e) => {
eprintln!("Read page failed: {}", e);
panic!("Read page should succeed");
}
}
// 3. Get page dimensions via eval_js to compute center
eprintln!("\n=== Getting page dimensions...");
let dims_result = tool
.execute(
serde_json::json!({
"action": "eval_js",
"expression": "JSON.stringify({w: window.innerWidth, h: window.innerHeight, scrollH: document.body.scrollHeight})"
}),
&ctx,
)
.await;
let (viewport_w, viewport_h) = match &dims_result {
Ok(output) => {
let result_str = output
.result
.get("result")
.and_then(|r| r.as_str())
.unwrap_or("{}");
let dims: serde_json::Value = serde_json::from_str(result_str).unwrap_or_default();
let w = dims.get("w").and_then(|v| v.as_f64()).unwrap_or(1920.0);
let h = dims.get("h").and_then(|v| v.as_f64()).unwrap_or(1080.0);
eprintln!("Viewport: {}x{}", w, h);
(w, h)
}
Err(e) => {
eprintln!("eval_js failed: {}", e);
(1920.0, 1080.0)
}
};
// 4. Scroll to middle of page first
eprintln!("\n=== Scrolling to middle of page...");
let _ = tool
.execute(
serde_json::json!({
"action": "eval_js",
"expression": "window.scrollTo(0, document.body.scrollHeight / 2 - window.innerHeight / 2)"
}),
&ctx,
)
.await;
// Brief wait for scroll to settle
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
// 5. Take full viewport screenshot
eprintln!("\n=== Taking viewport screenshot...");
let screenshot_result = tool
.execute(serde_json::json!({"action": "screenshot"}), &ctx)
.await;
match &screenshot_result {
Ok(output) => {
let b64 = output
.result
.get("data")
.and_then(|d| d.as_str())
.unwrap_or("");
eprintln!(
"Screenshot: {} base64 chars ({} bytes decoded)",
b64.len(),
b64.len() * 3 / 4
);
// Save to /tmp for inspection
use base64::Engine;
if let Ok(bytes) = base64::engine::general_purpose::STANDARD.decode(b64) {
let path = "/tmp/ironclaw_browser_test_viewport.png";
if std::fs::write(path, &bytes).is_ok() {
eprintln!("Saved viewport screenshot to {}", path);
}
// Now crop the center 10x10 using raw PNG manipulation
// We'll use eval_js to take a clipped screenshot via CDP directly
}
}
Err(e) => {
eprintln!("Screenshot failed: {}", e);
panic!("Screenshot should succeed");
}
}
// 6. Take a 10x10 screenshot from the center of the viewport using eval_js
// We can't directly use the clip param through the current tool API,
// so we'll take the viewport screenshot and note the center crop coords.
let center_x = (viewport_w / 2.0 - 5.0).max(0.0);
let center_y = (viewport_h / 2.0 - 5.0).max(0.0);
eprintln!(
"\n=== Center 10x10 crop would be at ({}, {}) to ({}, {})",
center_x,
center_y,
center_x + 10.0,
center_y + 10.0
);
// 7. Extract some text to verify content loaded
eprintln!("\n=== Extracting page text...");
let extract_result = tool
.execute(
serde_json::json!({"action": "extract", "selector": "h1"}),
&ctx,
)
.await;
match &extract_result {
Ok(output) => {
let text = output.result.as_str().unwrap_or("");
eprintln!("H1 text: {}", text);
assert!(
text.contains("Mariam") || text.contains("Almheiri"),
"H1 should contain the article subject, got: {}",
text
);
}
Err(e) => {
eprintln!("Extract failed: {}", e);
}
}
eprintln!("\n=== All browser integration tests passed!");
}
+22
View File
@@ -0,0 +1,22 @@
[package]
name = "github-tool"
version = "0.1.0"
edition = "2021"
description = "GitHub integration tool for IronClaw (WASM component)"
license = "MIT OR Apache-2.0"
publish = false
[dependencies]
serde = { version = "1.0", features = ["derive"] }
serde_json = "1.0"
wit-bindgen = "0.41.0"
[lib]
crate-type = ["cdylib"]
[profile.release]
opt-level = "s"
lto = true
strip = true
codegen-units = 1
+189
View File
@@ -0,0 +1,189 @@
# GitHub Tool for IronClaw
WASM tool for GitHub integration - manage repos, issues, PRs, and workflows.
## Features
- **Repository Info** - Get repo details, list user repos
- **Issues** - List, create, and get issue details
- **Pull Requests** - List PRs, get PR details, review files, create reviews
- **File Content** - Read files from repos
- **Workflows** - Trigger GitHub Actions, check run status
## Setup
1. Create a GitHub Personal Access Token at <https://github.com/settings/tokens>
2. Required scopes: `repo`, `workflow`, `read:org`
3. Store the token:
```
ironclaw secret set github_token YOUR_TOKEN
```
## Usage Examples
### Get Repository Info
```json
{
"action": "get_repo",
"owner": "nearai",
"repo": "ironclaw"
}
```
### List Open Issues
```json
{
"action": "list_issues",
"owner": "nearai",
"repo": "ironclaw",
"state": "open",
"limit": 10
}
```
### Create Issue
```json
{
"action": "create_issue",
"owner": "nearai",
"repo": "ironclaw",
"title": "Bug: Something is broken",
"body": "Detailed description...",
"labels": ["bug", "help wanted"]
}
```
### List Pull Requests
```json
{
"action": "list_pull_requests",
"owner": "nearai",
"repo": "ironclaw",
"state": "open",
"limit": 5
}
```
### Review PR
```json
{
"action": "create_pr_review",
"owner": "nearai",
"repo": "ironclaw",
"pr_number": 42,
"body": "LGTM! Great work.",
"event": "APPROVE"
}
```
### Get File Content
```json
{
"action": "get_file_content",
"owner": "nearai",
"repo": "ironclaw",
"path": "README.md",
"ref": "main"
}
```
### Trigger Workflow
```json
{
"action": "trigger_workflow",
"owner": "nearai",
"repo": "ironclaw",
"workflow_id": "ci.yml",
"ref": "main",
"inputs": {
"environment": "staging"
}
}
```
### Check Workflow Runs
```json
{
"action": "get_workflow_runs",
"owner": "nearai",
"repo": "ironclaw",
"limit": 5
}
```
### List Workflow Runs (Pagination)
```json
{
"action": "get_workflow_runs",
"owner": "nearai",
"repo": "ironclaw",
"limit": 5,
"page": 2
}
```
## Error Handling
Errors are returned as strings in the `error` field of the response.
### Rate Limit Exceeded
When the GitHub API rate limit is exceeded (and retries fail), you might see:
```text
GitHub API error 429: { "message": "API rate limit exceeded for user ID ...", ... }
```
The tool automatically logs warnings when the rate limit is low (<10 remaining) and retries on 429/5xx errors.
### Invalid Parameters
```text
Invalid event: 'INVALID'. Must be one of: APPROVE, REQUEST_CHANGES, COMMENT
```
### Missing Token
```text
GitHub token not found in secret store. Set it with: ironclaw secret set github_token <token>...
```
## Troubleshooting
### "GitHub API error 404: Not Found"
- Check that the `owner` and `repo` are correct.
- Ensure the `github_token` has access to the repository (especially for private repos).
- Verify the token scopes include `repo` and `read:org`.
### "GitHub API error 401: Bad credentials"
- The token might be invalid or expired.
- Update the token: `ironclaw secret set github_token NEW_TOKEN`.
### Rate Limiting
- The tool logs a warning when remaining requests drop below 10.
- Check logs for "GitHub API rate limit low".
- If you hit the limit, wait for the reset time (usually 1 hour).
## Building
```bash
cd tools-src/github
cargo build --target wasm32-wasi --release
```
## License
MIT/Apache-2.0
@@ -0,0 +1,41 @@
{
"capabilities": {
"http": {
"allowlist": [
{
"host": "api.github.com",
"path_prefix": "/",
"methods": [
"GET",
"POST"
]
}
],
"credentials": {
"github_token": {
"secret_name": "github_token",
"location": {
"type": "bearer"
},
"host_patterns": [
"api.github.com"
]
}
},
"rate_limit": {
"requests_per_minute": 60,
"requests_per_hour": 3600
}
},
"secrets": {
"allowed_names": [
"github_token",
"github_*"
]
}
},
"config": {
"default_limit": 30,
"max_limit": 100
}
}
+845
View File
@@ -0,0 +1,845 @@
//! GitHub WASM Tool for IronClaw.
//!
//! Provides GitHub integration for reading repos, managing issues,
//! reviewing PRs, and triggering workflows.
//!
//! # Authentication
//!
//! Store your GitHub Personal Access Token:
//! `ironclaw secret set github_token <token>`
//!
//! Token needs these permissions:
//! - repo (for private repos)
//! - workflow (for triggering actions)
//! - read:org (for org repos)
wit_bindgen::generate!({
world: "sandboxed-tool",
path: "../../wit/tool.wit",
});
use serde::Deserialize;
const MAX_TEXT_LENGTH: usize = 65536;
/// Validate input length to prevent oversized payloads.
fn validate_input_length(s: &str, field_name: &str) -> Result<(), String> {
if s.len() > MAX_TEXT_LENGTH {
return Err(format!(
"Input '{}' exceeds maximum length of {} characters",
field_name, MAX_TEXT_LENGTH
));
}
Ok(())
}
/// Percent-encode a string for safe use in URL path segments.
/// Encodes everything except alphanumeric, hyphen, underscore, and dot.
fn url_encode_path(s: &str) -> String {
let mut out = String::with_capacity(s.len() * 2);
for b in s.bytes() {
match b {
b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'_' | b'.' => {
out.push(b as char);
}
_ => {
out.push('%');
out.push(char::from(b"0123456789ABCDEF"[(b >> 4) as usize]));
out.push(char::from(b"0123456789ABCDEF"[(b & 0xf) as usize]));
}
}
}
out
}
/// Percent-encode a string for use as a URL query parameter value.
/// Currently identical to `url_encode_path`.
fn url_encode_query(s: &str) -> String {
url_encode_path(s)
}
/// Validate that a path segment doesn't contain dangerous characters.
/// Returns true if the segment is safe to use.
fn validate_path_segment(s: &str) -> bool {
!s.is_empty() && !s.contains('/') && !s.contains("..") && !s.contains('?') && !s.contains('#')
}
struct GitHubTool;
#[derive(Debug, Deserialize)]
#[serde(tag = "action")]
enum GitHubAction {
#[serde(rename = "get_repo")]
GetRepo { owner: String, repo: String },
#[serde(rename = "list_issues")]
ListIssues {
owner: String,
repo: String,
state: Option<String>,
page: Option<u32>,
limit: Option<u32>,
},
#[serde(rename = "create_issue")]
CreateIssue {
owner: String,
repo: String,
title: String,
body: Option<String>,
labels: Option<Vec<String>>,
},
#[serde(rename = "get_issue")]
GetIssue {
owner: String,
repo: String,
issue_number: u32,
},
#[serde(rename = "list_pull_requests")]
ListPullRequests {
owner: String,
repo: String,
state: Option<String>,
page: Option<u32>,
limit: Option<u32>,
},
#[serde(rename = "get_pull_request")]
GetPullRequest {
owner: String,
repo: String,
pr_number: u32,
},
#[serde(rename = "get_pull_request_files")]
GetPullRequestFiles {
owner: String,
repo: String,
pr_number: u32,
},
#[serde(rename = "create_pr_review")]
CreatePrReview {
owner: String,
repo: String,
pr_number: u32,
body: String,
event: String,
},
#[serde(rename = "list_repos")]
ListRepos {
username: String,
page: Option<u32>,
limit: Option<u32>,
},
#[serde(rename = "get_file_content")]
GetFileContent {
owner: String,
repo: String,
path: String,
r#ref: Option<String>,
},
#[serde(rename = "trigger_workflow")]
TriggerWorkflow {
owner: String,
repo: String,
workflow_id: String,
r#ref: String,
inputs: Option<serde_json::Value>,
},
#[serde(rename = "get_workflow_runs")]
GetWorkflowRuns {
owner: String,
repo: String,
workflow_id: Option<String>,
page: Option<u32>,
limit: Option<u32>,
},
}
impl exports::near::agent::tool::Guest for GitHubTool {
fn execute(req: exports::near::agent::tool::Request) -> exports::near::agent::tool::Response {
match execute_inner(&req.params) {
Ok(result) => exports::near::agent::tool::Response {
output: Some(result),
error: None,
},
Err(e) => exports::near::agent::tool::Response {
output: None,
error: Some(e),
},
}
}
fn schema() -> String {
SCHEMA.to_string()
}
fn description() -> String {
"GitHub integration for managing repositories, issues, pull requests, \
and workflows. Supports reading repo info, listing/creating issues, \
reviewing PRs, and triggering GitHub Actions. \
Authentication is handled via the 'github_token' secret injected by the host."
.to_string()
}
}
fn execute_inner(params: &str) -> Result<String, String> {
let action: GitHubAction =
serde_json::from_str(params).map_err(|e| format!("Invalid parameters: {e}"))?;
// Pre-flight check: ensure token exists in secret store.
// We don't use the returned value because the host injects it into the request.
let _ = get_github_token()?;
match action {
GitHubAction::GetRepo { owner, repo } => get_repo(&owner, &repo),
GitHubAction::ListIssues {
owner,
repo,
state,
page,
limit,
} => list_issues(&owner, &repo, state.as_deref(), page, limit),
GitHubAction::CreateIssue {
owner,
repo,
title,
body,
labels,
} => create_issue(&owner, &repo, &title, body.as_deref(), labels),
GitHubAction::GetIssue {
owner,
repo,
issue_number,
} => get_issue(&owner, &repo, issue_number),
GitHubAction::ListPullRequests {
owner,
repo,
state,
page,
limit,
} => list_pull_requests(&owner, &repo, state.as_deref(), page, limit),
GitHubAction::GetPullRequest {
owner,
repo,
pr_number,
} => get_pull_request(&owner, &repo, pr_number),
GitHubAction::GetPullRequestFiles {
owner,
repo,
pr_number,
} => get_pull_request_files(&owner, &repo, pr_number),
GitHubAction::CreatePrReview {
owner,
repo,
pr_number,
body,
event,
} => create_pr_review(&owner, &repo, pr_number, &body, &event),
GitHubAction::ListRepos {
username,
page,
limit,
} => list_repos(&username, page, limit),
GitHubAction::GetFileContent {
owner,
repo,
path,
r#ref,
} => get_file_content(&owner, &repo, &path, r#ref.as_deref()),
GitHubAction::TriggerWorkflow {
owner,
repo,
workflow_id,
r#ref,
inputs,
} => trigger_workflow(&owner, &repo, &workflow_id, &r#ref, inputs),
GitHubAction::GetWorkflowRuns {
owner,
repo,
workflow_id,
page,
limit,
} => get_workflow_runs(&owner, &repo, workflow_id.as_deref(), page, limit),
}
}
fn get_github_token() -> Result<String, String> {
if near::agent::host::secret_exists("github_token") {
// Return dummy value since we only need to verify existence.
// The actual token is injected by the host.
return Ok("present".to_string());
}
Err("GitHub token not found in secret store. Set it with: ironclaw secret set github_token <token>. \
Token needs 'repo', 'workflow', and 'read:org' scopes.".into())
}
fn github_request(method: &str, path: &str, body: Option<String>) -> Result<String, String> {
let url = format!("https://api.github.com{}", path);
// Authorization header (Bearer <token>) is injected automatically by the host
// via the `http-wrapper` proxy based on the `github_token` secret.
let headers = serde_json::json!({
"Accept": "application/vnd.github+json",
"X-GitHub-Api-Version": "2022-11-28",
"User-Agent": "IronClaw-GitHub-Tool"
});
let body_bytes = body.map(|b| b.into_bytes());
// Simple retry logic for transient errors (max 3 attempts)
let max_retries = 3;
let mut attempt = 0;
loop {
attempt += 1;
let response = near::agent::host::http_request(
method,
&url,
&headers.to_string(),
body_bytes.as_deref(),
None,
);
match response {
Ok(resp) => {
// Log warning if rate limit is low
if let Ok(headers_json) =
serde_json::from_str::<serde_json::Value>(&resp.headers_json)
{
// Header keys are often lowercase in http libs, check case-insensitively if needed,
// but usually standard is lowercase/case-insensitive. Let's try lowercase.
if let Some(remaining) = headers_json
.get("x-ratelimit-remaining")
.and_then(|v| v.as_str())
{
if let Ok(count) = remaining.parse::<u32>() {
if count < 10 {
near::agent::host::log(
near::agent::host::LogLevel::Warn,
&format!("GitHub API rate limit low: {} remaining", count),
);
}
}
}
}
if resp.status >= 200 && resp.status < 300 {
return String::from_utf8(resp.body)
.map_err(|e| format!("Invalid UTF-8: {}", e));
} else if attempt < max_retries && (resp.status == 429 || resp.status >= 500) {
near::agent::host::log(
near::agent::host::LogLevel::Warn,
&format!(
"GitHub API error {} (attempt {}/{}). Retrying...",
resp.status, attempt, max_retries
),
);
// Minimal backoff simulation since we can't block easily in WASM without consuming generic budget?
// actually std::thread::sleep works in WASMtime if configured, but here we might just spin.
// ideally host exposes sleep. For now just retry immediately or rely on host timeout logic?
// Let's assume immediate retry for now as simple strategy.
continue;
} else {
let body_str = String::from_utf8_lossy(&resp.body);
return Err(format!("GitHub API error {}: {}", resp.status, body_str));
}
}
Err(e) => {
if attempt < max_retries {
near::agent::host::log(
near::agent::host::LogLevel::Warn,
&format!(
"HTTP request failed: {} (attempt {}/{}). Retrying...",
e, attempt, max_retries
),
);
continue;
}
return Err(format!(
"HTTP request failed after {} attempts: {}",
max_retries, e
));
}
}
}
}
// === API Functions ===
fn get_repo(owner: &str, repo: &str) -> Result<String, String> {
if !validate_path_segment(owner) || !validate_path_segment(repo) {
return Err("Invalid owner or repo name".into());
}
let encoded_owner = url_encode_path(owner);
let encoded_repo = url_encode_path(repo);
github_request(
"GET",
&format!("/repos/{}/{}", encoded_owner, encoded_repo),
None,
)
}
fn list_issues(
owner: &str,
repo: &str,
state: Option<&str>,
page: Option<u32>,
limit: Option<u32>,
) -> Result<String, String> {
if !validate_path_segment(owner) || !validate_path_segment(repo) {
return Err("Invalid owner or repo name".into());
}
let encoded_owner = url_encode_path(owner);
let encoded_repo = url_encode_path(repo);
let state = state.unwrap_or("open");
let limit = limit.unwrap_or(30).min(100); // Cap at 100
let encoded_state = url_encode_query(state);
let mut path = format!(
"/repos/{}/{}/issues?state={}&per_page={}",
encoded_owner, encoded_repo, encoded_state, limit
);
if let Some(p) = page {
path.push_str(&format!("&page={}", p));
}
github_request("GET", &path, None)
}
fn create_issue(
owner: &str,
repo: &str,
title: &str,
body: Option<&str>,
labels: Option<Vec<String>>,
) -> Result<String, String> {
if !validate_path_segment(owner) || !validate_path_segment(repo) {
return Err("Invalid owner or repo name".into());
}
validate_input_length(title, "title")?;
if let Some(b) = body {
validate_input_length(b, "body")?;
}
let encoded_owner = url_encode_path(owner);
let encoded_repo = url_encode_path(repo);
let path = format!("/repos/{}/{}/issues", encoded_owner, encoded_repo);
let mut req_body = serde_json::json!({
"title": title,
});
if let Some(body) = body {
req_body["body"] = serde_json::json!(body);
}
if let Some(labels) = labels {
req_body["labels"] = serde_json::json!(labels);
}
github_request("POST", &path, Some(req_body.to_string()))
}
fn get_issue(owner: &str, repo: &str, issue_number: u32) -> Result<String, String> {
if !validate_path_segment(owner) || !validate_path_segment(repo) {
return Err("Invalid owner or repo name".into());
}
let encoded_owner = url_encode_path(owner);
let encoded_repo = url_encode_path(repo);
github_request(
"GET",
&format!(
"/repos/{}/{}/issues/{}",
encoded_owner, encoded_repo, issue_number
),
None,
)
}
fn list_pull_requests(
owner: &str,
repo: &str,
state: Option<&str>,
page: Option<u32>,
limit: Option<u32>,
) -> Result<String, String> {
if !validate_path_segment(owner) || !validate_path_segment(repo) {
return Err("Invalid owner or repo name".into());
}
let encoded_owner = url_encode_path(owner);
let encoded_repo = url_encode_path(repo);
let state = state.unwrap_or("open");
let limit = limit.unwrap_or(30).min(100); // Cap at 100
let encoded_state = url_encode_query(state);
let mut path = format!(
"/repos/{}/{}/pulls?state={}&per_page={}",
encoded_owner, encoded_repo, encoded_state, limit
);
if let Some(p) = page {
path.push_str(&format!("&page={}", p));
}
github_request("GET", &path, None)
}
fn get_pull_request(owner: &str, repo: &str, pr_number: u32) -> Result<String, String> {
if !validate_path_segment(owner) || !validate_path_segment(repo) {
return Err("Invalid owner or repo name".into());
}
let encoded_owner = url_encode_path(owner);
let encoded_repo = url_encode_path(repo);
github_request(
"GET",
&format!(
"/repos/{}/{}/pulls/{}",
encoded_owner, encoded_repo, pr_number
),
None,
)
}
fn get_pull_request_files(owner: &str, repo: &str, pr_number: u32) -> Result<String, String> {
if !validate_path_segment(owner) || !validate_path_segment(repo) {
return Err("Invalid owner or repo name".into());
}
let encoded_owner = url_encode_path(owner);
let encoded_repo = url_encode_path(repo);
github_request(
"GET",
&format!(
"/repos/{}/{}/pulls/{}/files",
encoded_owner, encoded_repo, pr_number
),
None,
)
}
fn create_pr_review(
owner: &str,
repo: &str,
pr_number: u32,
body: &str,
event: &str,
) -> Result<String, String> {
if !validate_path_segment(owner) || !validate_path_segment(repo) {
return Err("Invalid owner or repo name".into());
}
validate_input_length(body, "body")?;
let valid_events = ["APPROVE", "REQUEST_CHANGES", "COMMENT"];
if !valid_events.contains(&event) {
return Err(format!(
"Invalid event: '{}'. Must be one of: {}",
event,
valid_events.join(", ")
));
}
let encoded_owner = url_encode_path(owner);
let encoded_repo = url_encode_path(repo);
let path = format!(
"/repos/{}/{}/pulls/{}/reviews",
encoded_owner, encoded_repo, pr_number
);
let req_body = serde_json::json!({
"body": body,
"event": event,
});
github_request("POST", &path, Some(req_body.to_string()))
}
fn list_repos(username: &str, page: Option<u32>, limit: Option<u32>) -> Result<String, String> {
if !validate_path_segment(username) {
return Err("Invalid username".into());
}
let encoded_username = url_encode_path(username);
let limit = limit.unwrap_or(30).min(100); // Cap at 100
let mut path = format!("/users/{}/repos?per_page={}", encoded_username, limit);
if let Some(p) = page {
path.push_str(&format!("&page={}", p));
}
github_request("GET", &path, None)
}
fn get_file_content(
owner: &str,
repo: &str,
path: &str,
r#ref: Option<&str>,
) -> Result<String, String> {
if !validate_path_segment(owner) || !validate_path_segment(repo) {
return Err("Invalid owner or repo name".into());
}
// Validate path segments - reject path traversal attempts and empty segments
for segment in path.split('/') {
if segment == ".." {
return Err("Invalid path: path traversal not allowed".into());
}
if segment.is_empty() {
return Err("Invalid path: empty segment not allowed".into());
}
}
// Validate ref if provided
if let Some(r#ref) = r#ref {
if r#ref.contains("..") || r#ref.contains(':') {
return Err("Invalid ref: must be a valid branch, tag, or commit SHA".into());
}
}
let encoded_owner = url_encode_path(owner);
let encoded_repo = url_encode_path(repo);
// Path can contain slashes, so we encode each segment separately
let encoded_path = path
.split('/')
.map(url_encode_path)
.collect::<Vec<_>>()
.join("/");
let url_path = if let Some(r#ref) = r#ref {
let encoded_ref = url_encode_query(r#ref);
format!(
"/repos/{}/{}/contents/{}?ref={}",
encoded_owner, encoded_repo, encoded_path, encoded_ref
)
} else {
format!(
"/repos/{}/{}/contents/{}",
encoded_owner, encoded_repo, encoded_path
)
};
github_request("GET", &url_path, None)
}
fn trigger_workflow(
owner: &str,
repo: &str,
workflow_id: &str,
r#ref: &str,
inputs: Option<serde_json::Value>,
) -> Result<String, String> {
if !validate_path_segment(owner) || !validate_path_segment(repo) {
return Err("Invalid owner or repo name".into());
}
// Validate inputs size if present
if let Some(valid_inputs) = &inputs {
let inputs_str = valid_inputs.to_string();
validate_input_length(&inputs_str, "inputs")?;
}
// Validate workflow_id - must be a safe filename
if workflow_id.contains('/') || workflow_id.contains("..") || workflow_id.contains(':') {
return Err("Invalid workflow_id: must be a filename or numeric ID".into());
}
// Validate ref - must be a valid git ref
if r#ref.contains("..") || r#ref.contains(':') {
return Err("Invalid ref: must be a valid branch, tag, or commit SHA".into());
}
let encoded_owner = url_encode_path(owner);
let encoded_repo = url_encode_path(repo);
let encoded_workflow_id = url_encode_path(workflow_id);
let path = format!(
"/repos/{}/{}/actions/workflows/{}/dispatches",
encoded_owner, encoded_repo, encoded_workflow_id
);
let mut req_body = serde_json::json!({
"ref": r#ref,
});
if let Some(inputs) = inputs {
req_body["inputs"] = inputs;
}
github_request("POST", &path, Some(req_body.to_string()))
}
fn get_workflow_runs(
owner: &str,
repo: &str,
workflow_id: Option<&str>,
page: Option<u32>,
limit: Option<u32>,
) -> Result<String, String> {
if !validate_path_segment(owner) || !validate_path_segment(repo) {
return Err("Invalid owner or repo name".into());
}
// Validate workflow_id if provided
if let Some(wid) = workflow_id {
if wid.contains('/') || wid.contains("..") || wid.contains(':') {
return Err("Invalid workflow_id: must be a filename or numeric ID".into());
}
}
let encoded_owner = url_encode_path(owner);
let encoded_repo = url_encode_path(repo);
let limit = limit.unwrap_or(30).min(100); // Cap at 100
let mut path = if let Some(workflow_id) = workflow_id {
let encoded_workflow_id = url_encode_path(workflow_id);
format!(
"/repos/{}/{}/actions/workflows/{}/runs?per_page={}",
encoded_owner, encoded_repo, encoded_workflow_id, limit
)
} else {
format!(
"/repos/{}/{}/actions/runs?per_page={}",
encoded_owner, encoded_repo, limit
)
};
if let Some(p) = page {
path.push_str(&format!("&page={}", p));
}
github_request("GET", &path, None)
}
const SCHEMA: &str = r#"{
"type": "object",
"required": ["action"],
"oneOf": [
{
"properties": {
"action": { "const": "get_repo" },
"owner": { "type": "string", "description": "Repository owner (user or org)" },
"repo": { "type": "string", "description": "Repository name" }
},
"required": ["action", "owner", "repo"]
},
{
"properties": {
"action": { "const": "list_issues" },
"owner": { "type": "string" },
"repo": { "type": "string" },
"state": { "type": "string", "enum": ["open", "closed", "all"], "default": "open" },
"limit": { "type": "integer", "default": 30 }
},
"required": ["action", "owner", "repo"]
},
{
"properties": {
"action": { "const": "create_issue" },
"owner": { "type": "string" },
"repo": { "type": "string" },
"title": { "type": "string" },
"body": { "type": "string" },
"labels": { "type": "array", "items": { "type": "string" } }
},
"required": ["action", "owner", "repo", "title"]
},
{
"properties": {
"action": { "const": "get_issue" },
"owner": { "type": "string" },
"repo": { "type": "string" },
"issue_number": { "type": "integer" }
},
"required": ["action", "owner", "repo", "issue_number"]
},
{
"properties": {
"action": { "const": "list_pull_requests" },
"owner": { "type": "string" },
"repo": { "type": "string" },
"state": { "type": "string", "enum": ["open", "closed", "all"], "default": "open" },
"limit": { "type": "integer", "default": 30 }
},
"required": ["action", "owner", "repo"]
},
{
"properties": {
"action": { "const": "get_pull_request" },
"owner": { "type": "string" },
"repo": { "type": "string" },
"pr_number": { "type": "integer" }
},
"required": ["action", "owner", "repo", "pr_number"]
},
{
"properties": {
"action": { "const": "get_pull_request_files" },
"owner": { "type": "string" },
"repo": { "type": "string" },
"pr_number": { "type": "integer" }
},
"required": ["action", "owner", "repo", "pr_number"]
},
{
"properties": {
"action": { "const": "create_pr_review" },
"owner": { "type": "string" },
"repo": { "type": "string" },
"pr_number": { "type": "integer" },
"body": { "type": "string", "description": "Review comment" },
"event": { "type": "string", "enum": ["APPROVE", "REQUEST_CHANGES", "COMMENT"] }
},
"required": ["action", "owner", "repo", "pr_number", "body", "event"]
},
{
"properties": {
"action": { "const": "list_repos" },
"username": { "type": "string" },
"limit": { "type": "integer", "default": 30 }
},
"required": ["action", "username"]
},
{
"properties": {
"action": { "const": "get_file_content" },
"owner": { "type": "string" },
"repo": { "type": "string" },
"path": { "type": "string", "description": "File path in repo" },
"ref": { "type": "string", "description": "Branch/commit (default: default branch)" }
},
"required": ["action", "owner", "repo", "path"]
},
{
"properties": {
"action": { "const": "trigger_workflow" },
"owner": { "type": "string" },
"repo": { "type": "string" },
"workflow_id": { "type": "string", "description": "Workflow filename or ID" },
"ref": { "type": "string", "description": "Branch to run on" },
"inputs": { "type": "object" }
},
"required": ["action", "owner", "repo", "workflow_id", "ref"]
},
{
"properties": {
"action": { "const": "get_workflow_runs" },
"owner": { "type": "string" },
"repo": { "type": "string" },
"workflow_id": { "type": "string" },
"limit": { "type": "integer", "default": 30 }
},
"required": ["action", "owner", "repo"]
}
]
}"#;
export!(GitHubTool);
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_url_encode_path() {
assert_eq!(url_encode_path("foo-bar_123.baz"), "foo-bar_123.baz");
assert_eq!(url_encode_path("foo bar"), "foo%20bar");
assert_eq!(url_encode_path("foo/bar"), "foo%2Fbar");
}
#[test]
fn test_validate_path_segment() {
assert!(validate_path_segment("foo"));
assert!(!validate_path_segment(""));
assert!(!validate_path_segment("foo/bar"));
assert!(!validate_path_segment(".."));
// Empty segments are handled in get_file_content logic, not here
}
#[test]
fn test_validate_event_in_create_pr_review() {
let valid = ["APPROVE", "REQUEST_CHANGES", "COMMENT"];
// Ensure valid inputs are accepted
for event in valid {
assert!(valid.contains(&event));
}
}
#[test]
fn test_input_length_validation() {
assert!(validate_input_length("short", "test").is_ok());
let long = "a".repeat(MAX_TEXT_LENGTH + 1);
assert!(validate_input_length(&long, "test").is_err());
}
}
+2 -2
View File
@@ -136,7 +136,7 @@ fn parse_message(v: &serde_json::Value) -> Message {
date: get_header(payload, "Date"),
body: extract_body(payload),
snippet: v["snippet"].as_str().unwrap_or("").to_string(),
is_unread: label_ids.contains(&"UNREAD".to_string()),
is_unread: label_ids.iter().any(|l| l == "UNREAD"),
label_ids,
}
}
@@ -198,7 +198,7 @@ pub fn list_messages(
to: get_header(payload, "To"),
date: get_header(payload, "Date"),
snippet: msg["snippet"].as_str().unwrap_or("").to_string(),
is_unread: label_ids.contains(&"UNREAD".to_string()),
is_unread: label_ids.iter().any(|l| l == "UNREAD"),
label_ids,
});
}
+1 -1
View File
@@ -6,7 +6,7 @@
//! # Capabilities Required
//!
//! - HTTP: `www.googleapis.com/calendar/v3/*` (GET, POST, PUT, PATCH, DELETE)
//! - Secrets: `google_calendar_token` (OAuth 2.0 token, injected automatically)
//! - Secrets: `google_oauth_token` (OAuth 2.0 token, injected automatically)
//!
//! # Supported Actions
//!
+7 -2
View File
@@ -269,8 +269,13 @@ pub fn replace_text(
let parsed = batch_update_raw(document_id, vec![request])?;
let occurrences = parsed["replies"][0]["replaceAllText"]["occurrencesChanged"]
.as_i64()
let first_reply = parsed["replies"].as_array().and_then(|arr| arr.first());
let occurrences = first_reply
.map(|r| {
r["replaceAllText"]["occurrencesChanged"]
.as_i64()
.unwrap_or(0)
})
.unwrap_or(0);
Ok(ReplaceResult {
+7 -1
View File
@@ -330,7 +330,13 @@ pub fn add_sheet(spreadsheet_id: &str, title: &str) -> Result<AddSheetResult, St
let parsed = batch_update(spreadsheet_id, requests)?;
let reply = &parsed["replies"][0]["addSheet"]["properties"];
let reply = parsed["replies"]
.as_array()
.and_then(|arr| arr.first())
.map(|r| &r["addSheet"]["properties"]);
let reply = reply.ok_or_else(|| "No reply from batch update".to_string())?;
Ok(AddSheetResult {
sheet: SheetInfo {
sheet_id: reply["sheetId"].as_i64().unwrap_or(0),